using System.Linq.Expressions; using Elsa.Common.Entities; using Elsa.Common.Models; using Elsa.Common.Multitenancy; using Elsa.Persistence.EFCore.Extensions; using Elsa.Extensions; using JetBrains.Annotations; using Microsoft.EntityFrameworkCore; using Microsoft.Extensions.DependencyInjection; using Open.Linq.AsyncExtensions; namespace Elsa.Persistence.EFCore; /// /// A generic repository class around EF Core for accessing entities. /// /// The type of the database context. /// The type of the entity. [PublicAPI] public class Store(IDbContextFactory dbContextFactory, IServiceProvider serviceProvider) where TDbContext : DbContext where TEntity : class, new() { // ReSharper disable once StaticMemberInGenericType // Justification: This is a static member that is used to ensure that only one thread can access the database for TEntity at a time. private static readonly SemaphoreSlim Semaphore = new(1, 1); /// /// Creates a new instance of the database context. /// /// The cancellation token. /// The database context. public async Task CreateDbContextAsync(CancellationToken cancellationToken = default) => await dbContextFactory.CreateDbContextAsync(cancellationToken); /// /// Adds the specified entity. /// /// The entity to add. /// The cancellation token. public async Task AddAsync(TEntity entity, CancellationToken cancellationToken = default) { await AddAsync(entity, null, cancellationToken); } /// /// Adds the specified entity. /// /// The entity to add. /// The callback to invoke before adding the entity. /// The cancellation token. public async Task AddAsync(TEntity entity, Func? onAdding, CancellationToken cancellationToken = default) { await using var dbContext = await CreateDbContextAsync(cancellationToken); if (onAdding != null) await onAdding(dbContext, entity, cancellationToken); var set = dbContext.Set(); await set.AddAsync(entity, cancellationToken); await dbContext.SaveChangesAsync(cancellationToken); } /// /// Adds the specified entities. /// /// The entities to save. /// The cancellation token. public async Task AddManyAsync( IEnumerable entities, CancellationToken cancellationToken = default) { await AddManyAsync(entities, null, cancellationToken); } /// /// Adds the specified entities. /// /// The entities to save. /// The callback to invoke before saving the entity. /// The cancellation token. public async Task AddManyAsync( IEnumerable entities, Func? onSaving = null, CancellationToken cancellationToken = default) { var entityList = entities.ToList(); if (entityList.Count == 0) return; await using var dbContext = await CreateDbContextAsync(cancellationToken); if (onSaving != null) { var savingTasks = entityList.Select(entity => onSaving(dbContext, entity, cancellationToken).AsTask()).ToList(); await Task.WhenAll(savingTasks); } await dbContext.BulkInsertAsync(entityList, cancellationToken); } /// /// Saves the entity. /// /// The entity to save. /// The key selector to get the primary key property. /// The cancellation token. public async Task SaveAsync(TEntity entity, Expression> keySelector, CancellationToken cancellationToken = default) => await SaveAsync(entity, keySelector, null, cancellationToken); /// /// Saves the entity. /// /// The entity to save. /// The key selector to get the primary key property. /// The callback to invoke before saving the entity. /// The cancellation token. public async Task SaveAsync(TEntity entity, Expression> keySelector, Func? onSaving, CancellationToken cancellationToken = default) { await Semaphore.WaitAsync(cancellationToken); // Asynchronous wait try { await using var dbContext = await CreateDbContextAsync(cancellationToken); if (onSaving != null) await onSaving(dbContext, entity, cancellationToken); var set = dbContext.Set(); var lambda = keySelector.BuildEqualsExpression(entity); var exists = await set.AnyAsync(lambda, cancellationToken); set.Entry(entity).State = exists ? EntityState.Modified : EntityState.Added; await dbContext.SaveChangesAsync(cancellationToken); } catch (Exception ex) { var handler = serviceProvider.GetService(); if (handler != null) { var context = new DbUpdateExceptionContext(ex, cancellationToken); await handler.HandleAsync(context); } throw; } finally { Semaphore.Release(); } } /// /// Saves the specified entities. /// /// The entities to save. /// The key selector to get the primary key property. /// The cancellation token. public async Task SaveManyAsync(IEnumerable entities, Expression> keySelector, CancellationToken cancellationToken = default) => await SaveManyAsync(entities, keySelector, null, cancellationToken); /// /// Saves the specified entities. /// /// The entities to save. /// The key selector to get the primary key property. /// The callback to invoke before saving the entity. /// The cancellation token. public async Task SaveManyAsync( IEnumerable entities, Expression> keySelector, Func? onSaving = null, CancellationToken cancellationToken = default) { var entityList = entities.ToList(); if (entityList.Count == 0) return; await using var dbContext = await CreateDbContextAsync(cancellationToken); if (onSaving != null) { var savingTasks = entityList.Select(entity => onSaving(dbContext, entity, cancellationToken).AsTask()).ToList(); await Task.WhenAll(savingTasks); } // When doing a custom SQL query (Bulk Upsert), none of the installed query filters will be applied. Hence, we are assigning the current tenant ID explicitly. var tenantId = serviceProvider.GetRequiredService().Tenant?.Id.NullIfEmpty(); foreach (var entity in entityList) { if (entity is Entity entityWithTenant) entityWithTenant.TenantId = tenantId; } try { await dbContext.BulkUpsertAsync(entityList, keySelector, cancellationToken); } catch (Exception ex) { var handler = serviceProvider.GetService(); if (handler != null) { var context = new DbUpdateExceptionContext(ex, cancellationToken); await handler.HandleAsync(context); } throw; } } /// /// Updates the entity. /// /// The entity to update. /// The cancellation token. public Task UpdateAsync(TEntity entity, CancellationToken cancellationToken = default) { return UpdateAsync(entity, null, cancellationToken); } /// /// Updates the entity. /// /// The entity to update. /// The callback to invoke before saving the entity. /// The cancellation token. public async Task UpdateAsync(TEntity entity, Func? onSaving, CancellationToken cancellationToken = default) { await using var dbContext = await CreateDbContextAsync(cancellationToken); if (onSaving != null) await onSaving(dbContext, entity, cancellationToken); var set = dbContext.Set(); set.Entry(entity).State = EntityState.Modified; await dbContext.SaveChangesAsync(cancellationToken); } /// /// Updates specific properties of an entity in the database. /// /// The entity to update. /// An array of expressions indicating the properties to update. /// The cancellation token. /// A task that represents the asynchronous operation. public async Task UpdatePartialAsync(TEntity entity, Expression>[] properties, CancellationToken cancellationToken = default) { await using var dbContext = await CreateDbContextAsync(cancellationToken); dbContext.Attach(entity); foreach (var property in properties) dbContext.Entry(entity).Property(property).IsModified = true; await dbContext.SaveChangesAsync(cancellationToken); } /// /// Finds the entity matching the specified predicate. /// /// The predicate to use. /// The cancellation token. /// The entity if found, otherwise null. public async Task FindAsync(Expression> predicate, CancellationToken cancellationToken = default) => await FindAsync(predicate, null, cancellationToken); /// /// Finds the entity matching the specified predicate. /// /// The predicate to use. /// A callback to run after the entity is loaded /// The cancellation token. /// public async Task FindAsync(Expression> predicate, Func? onLoading = null, CancellationToken cancellationToken = default) { await using var dbContext = await CreateDbContextAsync(cancellationToken); var set = dbContext.Set().AsNoTracking(); var entity = await set.FirstOrDefaultAsync(predicate, cancellationToken); if (entity == null) return null; if (onLoading != null) entity = onLoading.Invoke(dbContext, entity); return entity; } /// /// Finds a single entity using a query /// /// The query to use /// A callback to run after the entity is loaded /// The cancellation token /// The entity if found, otherwise null public async Task FindAsync(Func, IQueryable> query, Func? onLoading = null, CancellationToken cancellationToken = default) { return await FindAsync(query, onLoading, false, cancellationToken); } /// /// Finds a single entity using a query /// /// The query to use /// A callback to run after the entity is loaded /// Define is the request should be tenant agnostic or not /// The cancellation token /// The entity if found, otherwise null public async Task FindAsync(Func, IQueryable> query, Func? onLoading = null, bool tenantAgnostic = false, CancellationToken cancellationToken = default) { return await QueryAsync(query, onLoading, tenantAgnostic, cancellationToken).FirstOrDefault(); } /// /// Finds a single entity using a query /// /// The query to use /// The cancellation token /// The entity if found, otherwise null public async Task FindAsync(Func, IQueryable> query, CancellationToken cancellationToken = default) { return await FindAsync(query, false, cancellationToken); } /// /// Finds a single entity using a query /// /// The query to use /// Define is the request should be tenant agnostic or not /// The cancellation token /// The entity if found, otherwise null public async Task FindAsync(Func, IQueryable> query, bool tenantAgnostic = false, CancellationToken cancellationToken = default) { return await QueryAsync(query, tenantAgnostic, cancellationToken).FirstOrDefault(); } /// /// Finds a list of entities using a query /// public async Task> FindManyAsync(Expression> predicate, CancellationToken cancellationToken = default) => await FindManyAsync(predicate, null, cancellationToken); /// /// Finds a list of entities using a query /// public async Task> FindManyAsync(Expression> predicate, Action? onLoading = null, CancellationToken cancellationToken = default) { await using var dbContext = await CreateDbContextAsync(cancellationToken); var set = dbContext.Set().AsNoTracking(); var entities = await set.Where(predicate).ToListAsync(cancellationToken); if (onLoading != null) foreach (var entity in entities) onLoading(dbContext, entity); return entities; } /// /// Finds a list of entities using a query /// public async Task> FindManyAsync( Expression> predicate, Expression> orderBy, OrderDirection orderDirection = OrderDirection.Ascending, PageArgs? pageArgs = null, CancellationToken cancellationToken = default) => await FindManyAsync(predicate, orderBy, orderDirection, pageArgs, null, cancellationToken); /// /// Returns a list of entities using a query /// public async Task> FindManyAsync( Expression>? predicate, Expression>? orderBy, OrderDirection orderDirection = OrderDirection.Ascending, PageArgs? pageArgs = null, Func? onLoading = null, CancellationToken cancellationToken = default) { await using var dbContext = await CreateDbContextAsync(cancellationToken); var set = dbContext.Set().AsNoTracking(); if (predicate != null) set = set.Where(predicate); if (orderBy != null) set = orderDirection switch { OrderDirection.Ascending => set.OrderBy(orderBy), OrderDirection.Descending => set.OrderByDescending(orderBy), _ => set.OrderBy(orderBy) }; var page = await set.PaginateAsync(pageArgs); if (onLoading != null) page = page with { Items = page.Items.Select(x => onLoading(dbContext, x)!).ToList() }; return page; } public Task> ListAsync(CancellationToken cancellationToken = default) { return ListAsync(null, cancellationToken); } public async Task> ListAsync(Action? onLoading = null, CancellationToken cancellationToken = default) { await using var dbContext = await CreateDbContextAsync(cancellationToken); var set = dbContext.Set().AsNoTracking(); var entities = await set.ToListAsync(cancellationToken); if (onLoading != null) foreach (var entity in entities) onLoading(dbContext, entity); return entities; } /// /// Finds a single entity using a query. /// /// True if the entity was found, otherwise false. public async Task DeleteAsync(TEntity entity, CancellationToken cancellationToken = default) { await using var dbContext = await CreateDbContextAsync(cancellationToken); var set = dbContext.Set(); set.Attach(entity).State = EntityState.Deleted; return await dbContext.SaveChangesAsync(cancellationToken) == 1; } /// /// Deletes entities using a predicate. /// /// The number of entities deleted. public async Task DeleteWhereAsync(Expression> predicate, CancellationToken cancellationToken = default) { await using var dbContext = await CreateDbContextAsync(cancellationToken); var set = dbContext.Set().AsNoTracking(); return await set.Where(predicate).ExecuteDeleteAsync(cancellationToken); } /// /// Deletes entities using a query. /// /// The number of entities deleted. public async Task DeleteWhereAsync(Func, IQueryable> query, CancellationToken cancellationToken = default) { await using var dbContext = await CreateDbContextAsync(cancellationToken); var set = dbContext.Set().AsNoTracking(); var queryable = query(set.AsQueryable()); return await queryable.ExecuteDeleteAsync(cancellationToken); } /// /// Queries the database using a query. /// public async Task> QueryAsync(Func, IQueryable> query, CancellationToken cancellationToken = default) { return await QueryAsync(query, null, false, cancellationToken); } /// /// Queries the database using a query. /// public async Task> QueryAsync(Func, IQueryable> query, bool tenantAgnostic, CancellationToken cancellationToken = default) { return await QueryAsync(query, null, tenantAgnostic, cancellationToken); } /// /// Queries the database using a query and a selector. /// public async Task> QueryAsync(Func, IQueryable> query, Func? onLoading = null, CancellationToken cancellationToken = default) { return await QueryAsync(query, onLoading, false, cancellationToken); } /// /// Queries the database using a query and a selector. /// public async Task> QueryAsync(Func, IQueryable> query, Func? onLoading = null, bool ignoreQueryFilters = false, CancellationToken cancellationToken = default) { await using var dbContext = await CreateDbContextAsync(cancellationToken); var asNoTracking = onLoading == null; var set = asNoTracking ? dbContext.Set().AsNoTracking() : dbContext.Set(); var queryable = query(set.AsQueryable()); if (ignoreQueryFilters) queryable = queryable.IgnoreQueryFilters(); var entities = await queryable.ToListAsync(cancellationToken); if (onLoading != null) { var loadingTasks = entities.Select(entity => onLoading(dbContext, entity, cancellationToken).AsTask()).ToList(); await Task.WhenAll(loadingTasks); } return entities; } /// /// Queries the database using a query and a selector. /// public async Task> QueryAsync(Func, IQueryable> query, Expression> selector, CancellationToken cancellationToken = default) { return await QueryAsync(query, selector, false, cancellationToken); } /// /// Queries the database using a query and a selector. /// public async Task> QueryAsync(Func, IQueryable> query, Expression> selector, bool ignoreQueryFilters = false, CancellationToken cancellationToken = default) { await using var dbContext = await CreateDbContextAsync(cancellationToken); var set = dbContext.Set().AsNoTracking(); var queryable = query(set.AsQueryable()); if (ignoreQueryFilters) queryable = queryable.IgnoreQueryFilters(); queryable = query(queryable); return await queryable.Select(selector).ToListAsync(cancellationToken); } /// /// Counts the number of entities matching a query. /// public async Task CountAsync(Func, IQueryable> query, CancellationToken cancellationToken = default) { return await CountAsync(query, false, cancellationToken); } /// /// Counts the number of entities matching a query. /// public async Task CountAsync(Func, IQueryable> query, bool ignoreQueryFilters = false, CancellationToken cancellationToken = default) { await using var dbContext = await CreateDbContextAsync(cancellationToken); var set = dbContext.Set().AsNoTracking(); var queryable = query(set.AsQueryable()); if (ignoreQueryFilters) queryable = queryable.IgnoreQueryFilters(); queryable = query(queryable); return await queryable.LongCountAsync(cancellationToken: cancellationToken); } /// /// Checks if any entities exist. /// public async Task AnyAsync(Expression> predicate, CancellationToken cancellationToken = default) { return await AnyAsync(predicate, false, cancellationToken); } /// /// Checks if any entities exist. /// public async Task AnyAsync(Expression> predicate, bool ignoreQueryFilters = false, CancellationToken cancellationToken = default) { await using var dbContext = await CreateDbContextAsync(cancellationToken); var set = dbContext.Set().AsNoTracking(); return await set.AnyAsync(predicate, cancellationToken); } /// /// Counts the number of entities matching a predicate. /// /// The predicate. /// The cancellation token. public async Task CountAsync(Expression> predicate, CancellationToken cancellationToken = default) { return await CountAsync(predicate, false, cancellationToken); } /// /// Counts the number of entities matching a predicate. /// /// The predicate. /// Whether to ignore query filters. /// The cancellation token. public async Task CountAsync(Expression> predicate, bool ignoreQueryFilters = false, CancellationToken cancellationToken = default) { await using var dbContext = await CreateDbContextAsync(cancellationToken); var queryable = dbContext.Set().AsNoTracking(); if (ignoreQueryFilters) queryable = queryable.IgnoreQueryFilters(); return await queryable.CountAsync(predicate, cancellationToken); } /// /// Counts the distinct number of entities matching a predicate. /// /// The predicate. /// The property selector to distinct by. /// The cancellation token. public async Task CountAsync(Expression> predicate, Expression> propertySelector, CancellationToken cancellationToken = default) { return await CountAsync(predicate, propertySelector, false, cancellationToken); } /// /// Counts the distinct number of entities matching a predicate. /// /// The predicate. /// The property selector to distinct by. /// Whether to ignore query filters. /// The cancellation token. public async Task CountAsync(Expression> predicate, Expression> propertySelector, bool ignoreQueryFilters = false, CancellationToken cancellationToken = default) { await using var dbContext = await CreateDbContextAsync(cancellationToken); var queryable = dbContext.Set().AsNoTracking(); if (ignoreQueryFilters) queryable = queryable.IgnoreQueryFilters(); return await queryable .Where(predicate) .Select(propertySelector) .Distinct() .CountAsync(cancellationToken); } }