73 lines
2.7 KiB
C#
73 lines
2.7 KiB
C#
using System;
|
|
using System.Threading;
|
|
using System.Threading.Tasks;
|
|
using Microsoft.EntityFrameworkCore;
|
|
using Microsoft.Extensions.DependencyInjection;
|
|
using timetracker.Shared;
|
|
|
|
namespace timetracker.Data;
|
|
|
|
public class TimetrackerDbContext : DbContext
|
|
{
|
|
private readonly IServiceProvider _serviceProvider;
|
|
private ITenantProvider? _tenantProvider;
|
|
private ITenantProvider TenantProvider => _tenantProvider ??= _serviceProvider.GetRequiredService<ITenantProvider>();
|
|
|
|
public TimetrackerDbContext(
|
|
DbContextOptions<TimetrackerDbContext> options,
|
|
IServiceProvider serviceProvider) : base(options)
|
|
{
|
|
_serviceProvider = serviceProvider;
|
|
}
|
|
|
|
public DbSet<Tenant> Tenants => Set<Tenant>();
|
|
public DbSet<User> Users => Set<User>();
|
|
public DbSet<WorkDay> WorkDays => Set<WorkDay>();
|
|
public DbSet<BreakEntry> BreakEntries => Set<BreakEntry>();
|
|
public DbSet<AppSettings> AppSettings => Set<AppSettings>();
|
|
public DbSet<VacationDay> VacationDays => Set<VacationDay>();
|
|
public DbSet<PublicHoliday> PublicHolidays => Set<PublicHoliday>();
|
|
|
|
protected override void OnModelCreating(ModelBuilder modelBuilder)
|
|
{
|
|
base.OnModelCreating(modelBuilder);
|
|
|
|
// Subdomain index unique
|
|
modelBuilder.Entity<Tenant>()
|
|
.HasIndex(t => t.Subdomain)
|
|
.IsUnique();
|
|
|
|
// Global Query Filters (scopes database views by resolved TenantId)
|
|
modelBuilder.Entity<User>().HasQueryFilter(u => u.TenantId == TenantProvider.TenantId);
|
|
modelBuilder.Entity<WorkDay>().HasQueryFilter(w => w.TenantId == TenantProvider.TenantId);
|
|
modelBuilder.Entity<VacationDay>().HasQueryFilter(v => v.TenantId == TenantProvider.TenantId);
|
|
modelBuilder.Entity<AppSettings>().HasQueryFilter(s => s.TenantId == TenantProvider.TenantId);
|
|
}
|
|
|
|
public override Task<int> SaveChangesAsync(CancellationToken cancellationToken = default)
|
|
{
|
|
var tenantId = TenantProvider.TenantId;
|
|
|
|
if (tenantId.HasValue)
|
|
{
|
|
foreach (var entry in ChangeTracker.Entries())
|
|
{
|
|
if (entry.State == EntityState.Added)
|
|
{
|
|
var tenantIdProp = entry.Entity.GetType().GetProperty("TenantId");
|
|
if (tenantIdProp != null && tenantIdProp.CanWrite)
|
|
{
|
|
var currentValue = tenantIdProp.GetValue(entry.Entity);
|
|
if (currentValue == null || (currentValue is int val && val == 0))
|
|
{
|
|
tenantIdProp.SetValue(entry.Entity, tenantId.Value);
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
return base.SaveChangesAsync(cancellationToken);
|
|
}
|
|
}
|