using Elsa.Common.Entities;
using Elsa.Common.Multitenancy;
using Elsa.Extensions;
using Microsoft.EntityFrameworkCore;
using Microsoft.EntityFrameworkCore.ChangeTracking;
using Microsoft.Extensions.DependencyInjection;
namespace Elsa.EntityFrameworkCore;
///
/// An optional base class to implement with some opinions on certain converters to install for certain DB providers.
///
public abstract class ElsaDbContextBase : DbContext, IElsaDbContextSchema
{
private static readonly ISet ModifiedEntityStates = new HashSet
{
EntityState.Added,
EntityState.Modified,
};
protected IServiceProvider ServiceProvider { get; }
private readonly ElsaDbContextOptions? elsaDbContextOptions;
public string? TenantId { get; set; }
///
/// The default schema used by Elsa.
///
public static string ElsaSchema { get; set; } = "Elsa";
///
public string Schema { get; }
///
/// The table used to store the migrations history.
///
public static string MigrationsHistoryTable { get; set; } = "__EFMigrationsHistory";
///
/// Initializes a new instance of the class.
///
protected ElsaDbContextBase(DbContextOptions options, IServiceProvider serviceProvider) : base(options)
{
ServiceProvider = serviceProvider;
elsaDbContextOptions = options.FindExtension()?.Options;
// ReSharper disable once VirtualMemberCallInConstructor
Schema = !string.IsNullOrWhiteSpace(elsaDbContextOptions?.SchemaName) ? elsaDbContextOptions.SchemaName : ElsaSchema;
var tenantAccessor = serviceProvider.GetService();
var tenantId = tenantAccessor?.Tenant?.Id;
if (!string.IsNullOrWhiteSpace(tenantId))
TenantId = tenantId.NullIfEmpty();
}
///
public override async Task SaveChangesAsync(CancellationToken cancellationToken = default)
{
await OnBeforeSavingAsync(cancellationToken);
return await base.SaveChangesAsync(cancellationToken);
}
///
protected override void OnModelCreating(ModelBuilder modelBuilder)
{
if (!string.IsNullOrWhiteSpace(Schema))
modelBuilder.HasDefaultSchema(Schema);
var additionalConfigurations = elsaDbContextOptions?.GetModelConfigurations(this);
additionalConfigurations?.Invoke(modelBuilder);
using var scope = ServiceProvider.CreateScope();
var entityTypeHandlers = scope.ServiceProvider.GetServices().ToList();
foreach (var entityType in modelBuilder.Model.GetEntityTypes().ToList())
{
foreach (var handler in entityTypeHandlers)
handler.Handle(this, modelBuilder, entityType);
}
}
private async Task OnBeforeSavingAsync(CancellationToken cancellationToken)
{
using var scope = ServiceProvider.CreateScope();
var handlers = scope.ServiceProvider.GetServices().ToList();
foreach (var entry in ChangeTracker.Entries().Where(IsModifiedEntity))
{
foreach (var handler in handlers)
await handler.HandleAsync(this, entry, cancellationToken);
}
}
///
/// Determine if an entity was modified.
///
private bool IsModifiedEntity(EntityEntry entityEntry)
{
return ModifiedEntityStates.Contains(entityEntry.State) && entityEntry.Entity is Entity;
}
}