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