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