using Elsa.Common.Entities; using Elsa.Common.Multitenancy; using Elsa.EntityFrameworkCore.Common.Contracts; using Microsoft.EntityFrameworkCore; using Microsoft.EntityFrameworkCore.ChangeTracking; using Microsoft.Extensions.DependencyInjection; namespace Elsa.EntityFrameworkCore.Common; /// /// 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 readonly IServiceProvider ServiceProvider; 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; var elsaDbContextOptions = options.FindExtension()?.Options; // ReSharper disable once VirtualMemberCallInConstructor Schema = !string.IsNullOrWhiteSpace(elsaDbContextOptions?.SchemaName) ? elsaDbContextOptions.SchemaName : ElsaSchema; var tenantAccessor = serviceProvider.GetRequiredService(); TenantId = tenantAccessor.CurrentTenant?.Id; } /// 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)) { if (!Database.IsSqlite()) modelBuilder.HasDefaultSchema(Schema); } var entityTypeHandlers = ServiceProvider.GetServices().ToList(); foreach (var entityType in modelBuilder.Model.GetEntityTypes()) { foreach (var handler in entityTypeHandlers) { handler.Handle(this, modelBuilder, entityType); } } } private async Task OnBeforeSavingAsync(CancellationToken cancellationToken) { var handlers = 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; } }