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);
}
}