using LANCommander.Server.Data; using LANCommander.Server.Data.Enums; using LANCommander.Server.Data.Models; using LANCommander.Server.Services.Models; using Microsoft.EntityFrameworkCore; using Microsoft.Extensions.Logging; using System.Linq.Expressions; using System.Reflection; using AutoMapper; using AutoMapper.QueryableExtensions; using LANCommander.Server.Services.Abstractions; using Microsoft.AspNetCore.Http; using ZiggyCreatures.Caching.Fusion; namespace LANCommander.Server.Services { public abstract class BaseDatabaseService( ILogger logger, SettingsProvider settingsProvider, IFusionCache cache, IMapper mapper, IHttpContextAccessor httpContextAccessor, IDbContextFactory dbContextFactory) : BaseService(logger, settingsProvider), IBaseDatabaseService where T : class, IBaseModel { protected readonly List, IQueryable>> _modifiers = new(); public IBaseDatabaseService AsNoTracking() { return Query((queryable) => { return queryable.AsNoTracking(); }); } public IBaseDatabaseService AsSplitQuery() { return Query((queryable) => { return queryable.AsSplitQuery(); }); } public IBaseDatabaseService Query(Func, IQueryable> modifier) { _modifiers.Add(modifier); return this; } public IBaseDatabaseService Include(params string[] includes) { return Include(includes); } public IBaseDatabaseService Include(IEnumerable includes) { return Query((queryable) => { foreach (var include in includes) { queryable = queryable.Include(include); } return queryable; }); } public IBaseDatabaseService Include(params Expression>[] expressions) { return Query((queryable) => { foreach (var expression in expressions) { queryable = queryable.Include(expression); } return queryable; }); } public IBaseDatabaseService SortBy(Expression> expression, SortDirection direction = SortDirection.Ascending) { switch (direction) { case SortDirection.Descending: return Query((queryable) => { return queryable.OrderByDescending(expression); }); case SortDirection.Ascending: default: return Query((queryable) => { return queryable.OrderBy(expression); }); } } public virtual async Task> GetAsync() { return await GetAsync(x => true); } public virtual async Task> GetAsync() { return await GetAsync(x => true); } public virtual async Task GetAsync(Guid id) { return await FirstOrDefaultAsync(x => x.Id == id); } public virtual async Task GetAsync(Guid id) { return await FirstOrDefaultAsync(x => x.Id == id); } public virtual async Task> GetAsync(Expression> predicate) { try { using var context = await dbContextFactory.CreateDbContextAsync(); var queryable = context.Set().AsQueryable(); foreach (var modifier in _modifiers) queryable = modifier.Invoke(queryable); return await queryable.Where(predicate).ToListAsync(); } finally { Reset(); } } public virtual async Task> GetAsync(Expression> predicate) { try { using var context = await dbContextFactory.CreateDbContextAsync(); var queryable = context.Set().AsQueryable(); foreach (var modifier in _modifiers) queryable = modifier.Invoke(queryable); return await queryable.Where(predicate).ProjectTo(mapper.ConfigurationProvider).ToListAsync(); } finally { Reset(); } } public virtual async Task FirstAsync(Expression> predicate) { try { using var context = await dbContextFactory.CreateDbContextAsync(); var queryable = context.Set().AsQueryable(); foreach (var modifier in _modifiers) queryable = modifier.Invoke(queryable); return await queryable.FirstAsync(predicate); } finally { Reset(); } } public virtual async Task FirstAsync(Expression> predicate) { try { using var context = await dbContextFactory.CreateDbContextAsync(); var queryable = context.Set().AsQueryable(); foreach (var modifier in _modifiers) queryable = modifier.Invoke(queryable); return await queryable.Where(predicate).ProjectTo(mapper.ConfigurationProvider).FirstAsync(); } finally { Reset(); } } public virtual async Task FirstOrDefaultAsync(Expression> predicate) { try { using var context = await dbContextFactory.CreateDbContextAsync(); var queryable = context.Set().AsQueryable(); foreach (var modifier in _modifiers) queryable = modifier.Invoke(queryable); return await queryable.FirstOrDefaultAsync(predicate); } finally { Reset(); } } public virtual async Task FirstOrDefaultAsync(Expression> predicate) { try { using var context = await dbContextFactory.CreateDbContextAsync(); var queryable = context.Set().AsQueryable(); foreach (var modifier in _modifiers) queryable = modifier.Invoke(queryable); var entity = await queryable.Where(predicate).FirstOrDefaultAsync(); return mapper.Map(entity); } catch (Exception ex) { throw ex; } finally { Reset(); } } public virtual async Task ExistsAsync(Guid id) { return (await FirstOrDefaultAsync(x => x.Id == id)) != null; } public virtual async Task ExistsAsync(Expression> predicate) { return (await GetAsync(predicate)).Any(); } public abstract Task AddAsync(T entity); protected async Task AddAsync(T addedEntity, Action> additionalMapping = null) { try { using var context = await dbContextFactory.CreateDbContextAsync(); var currentUser = await GetCurrentUserAsync(context); var newEntity = Activator.CreateInstance(); context.Entry(newEntity).CurrentValues.SetValues(addedEntity); newEntity.CreatedOn = DateTime.UtcNow; newEntity.CreatedById = currentUser?.Id; if (additionalMapping != null) { var updateContext = new UpdateEntityContext(context, newEntity, addedEntity); additionalMapping?.Invoke(updateContext); } newEntity = (await context.AddAsync(newEntity)).Entity; await context.SaveChangesAsync(); return newEntity; } finally { Reset(); } } /// /// Adds an entity to the database if it does exist as dictated by the predicate /// /// Qualifier expressoin /// Entity to add /// Newly created or existing entity public virtual async Task> AddMissingAsync(Expression> predicate, T entity) { var existing = await FirstOrDefaultAsync(predicate); if (existing == null) { await cache.ExpireAsync($"{typeof(T).FullName}"); entity = await AddAsync(entity); return new ExistingEntityResult { Value = entity, Existing = false }; } else { return new ExistingEntityResult { Value = existing, Existing = true }; } } public abstract Task UpdateAsync(T entity); protected async Task UpdateAsync(T updatedEntity, Action> additionalMapping = null) { using var context = await dbContextFactory.CreateDbContextAsync(); if (updatedEntity.CreatedById != null && updatedEntity.CreatedBy == null) updatedEntity.CreatedById = null; var existingEntity = await context.Set().FirstOrDefaultAsync(e => e.Id == updatedEntity.Id); context.Entry(existingEntity).CurrentValues.SetValues(updatedEntity); if (additionalMapping != null) { var updateContext = new UpdateEntityContext(context, existingEntity, updatedEntity); additionalMapping?.Invoke(updateContext); } var currentUser = await GetCurrentUserAsync(context); existingEntity.UpdatedOn = DateTime.UtcNow; existingEntity.UpdatedById = currentUser?.Id; await context.SaveChangesAsync(); return updatedEntity; } public virtual async Task DeleteAsync(T entity) { try { await cache.ExpireAsync($"{typeof(T).FullName}"); using var context = await dbContextFactory.CreateDbContextAsync(); context.Set().Remove(entity); await context.SaveChangesAsync(); } finally { Reset(); } } public virtual async Task DeleteRangeAsync(IEnumerable entities) { try { await cache.ExpireAsync($"{typeof(T).FullName}"); using var context = await dbContextFactory.CreateDbContextAsync(); context.Set().RemoveRange(entities); await context.SaveChangesAsync(); } finally { Reset(); } } public virtual async Task SyncRelatedCollectionAsync( T entity, Expression>> navigationProperty, IEnumerable records, Func>> matchExpression) where TChild : class where T : class { using var context = await dbContextFactory.CreateDbContextAsync(); var entry = context.Entry(entity); var enumerableExpr = Expression.Lambda>>( navigationProperty.Body, navigationProperty.Parameters); var collectionEntry = entry.Collection(enumerableExpr); if (!collectionEntry.IsLoaded) await collectionEntry.LoadAsync(); var collection = navigationProperty.Compile().Invoke(entity); if (collection == null) { collection = new List(); if (navigationProperty.Body is not MemberExpression memberExpression || memberExpression.Member is not PropertyInfo propertyInfo) throw new InvalidOperationException($"Navigation expression '{navigationProperty}' must point to a property."); propertyInfo.SetValue(entity, collection); } var matchedChildren = new HashSet(); foreach (var record in records) { var matchPredicate = matchExpression(record); var existingChild = collection.FirstOrDefault(matchPredicate.Compile()); if (existingChild == null) { existingChild = await context.Set() .FirstOrDefaultAsync(matchPredicate); } if (existingChild != null) { if (!collection.Contains(existingChild)) collection.Add(existingChild); matchedChildren.Add(existingChild); } } var toDelete = collection.Where(child => !matchedChildren.Contains(child)); foreach (var child in toDelete) { collection.Remove(child); context.Remove(child); } try { await context.SaveChangesAsync(); } finally { Reset(); } } private async Task GetCurrentUserAsync(DatabaseContext context) { var httpContext = httpContextAccessor?.HttpContext; if (httpContext != null && httpContext.User != null && httpContext.User.Identity != null && httpContext.User.Identity.IsAuthenticated) { return await GetUserAsync(httpContext.User.Identity?.Name, context); } return null; } private static async Task GetUserAsync(string? username, DatabaseContext context) => await context.Users.FirstOrDefaultAsync(u => u.UserName == username); protected void Reset() { _modifiers.Clear(); } } }