using System.Security.Claims; using Data.SeaHavenIndustries; using Data.SeaHavenIndustries.Enums; using Microsoft.AspNetCore.Identity; using Microsoft.AspNetCore.Identity.EntityFrameworkCore; using Microsoft.Data.Sqlite; using Microsoft.EntityFrameworkCore; using Microsoft.EntityFrameworkCore.Metadata; using Microsoft.Extensions.Configuration; using Microsoft.Extensions.DependencyInjection; using Microsoft.Extensions.DependencyInjection.Extensions; using SeaHaven.DataServices.Dto; using SeaHaven.DataServices.Implementation; using SeaHaven.DataServices.DependencyInjection; using SeaHaven.DataServices.Interfaces; using SeaHaven.Services.DTOs; using SeaHaven.Services.DependencyInjection; using SeaHaven.Services.Implementation; using SeaHaven.Services.Interfaces; namespace SeaHavenIndustries.Tests; public sealed class TeamMemberCreateTransactionTests { [Fact] public async Task Create_MidSequencePersistenceFailure_RollsBackAllMemberRows() { await using var connection = new SqliteConnection("Data Source=:memory:;Foreign Keys=True"); await connection.OpenAsync(); var options = new DbContextOptionsBuilder() .UseSqlite(connection) .Options; await using (var setup = new SqliteTeamMemberTestDbContext(options)) await setup.Database.EnsureCreatedAsync(); await using var serviceProvider = BuildServiceProvider(connection, injectFailure: true); ThrowingAfterPersistPermissionDataService failingPermissionData; await using (var createScope = serviceProvider.CreateAsyncScope()) { var service = createScope.ServiceProvider.GetRequiredService(); failingPermissionData = (ThrowingAfterPersistPermissionDataService)createScope.ServiceProvider .GetRequiredService(); var failure = await Record.ExceptionAsync(() => service.CreateAsync( ValidRequest(), Admin(), CancellationToken.None)); Assert.IsType(failure); Assert.Equal("injected mid-sequence failure", failure!.Message); } Assert.NotNull(failingPermissionData.RowsAtFailure); var rowsAtFailure = failingPermissionData.RowsAtFailure!.Value; Assert.Equal(1, rowsAtFailure.Users); Assert.Equal(1, rowsAtFailure.Roles); Assert.Equal(1, rowsAtFailure.UserRoles); Assert.Equal(2, rowsAtFailure.ServiceAreas); Assert.Equal(1, rowsAtFailure.PermissionOverrides); await using var verifyScope = serviceProvider.CreateAsyncScope(); var verify = verifyScope.ServiceProvider.GetRequiredService(); Assert.Empty(await verify.Users.AsNoTracking().ToListAsync()); Assert.Empty(await verify.Roles.AsNoTracking().ToListAsync()); Assert.Empty(await verify.UserRoles.AsNoTracking().ToListAsync()); Assert.Empty(await verify.UserServiceAreas.AsNoTracking().ToListAsync()); Assert.Empty(await verify.UserPermissionOverrides.AsNoTracking().ToListAsync()); } [Fact] public async Task Create_SuccessfulTransaction_CommitsPendingMemberAndAssociations() { await using var connection = new SqliteConnection("Data Source=:memory:;Foreign Keys=True"); await connection.OpenAsync(); var options = new DbContextOptionsBuilder() .UseSqlite(connection) .Options; await using (var setup = new SqliteTeamMemberTestDbContext(options)) await setup.Database.EnsureCreatedAsync(); await using var serviceProvider = BuildServiceProvider(connection); await using (var createScope = serviceProvider.CreateAsyncScope()) { var service = createScope.ServiceProvider.GetRequiredService(); var outcome = await service.CreateAsync(ValidRequest(), Admin(), CancellationToken.None); Assert.True(outcome.Success); Assert.True(outcome.Member!.PendingRegistration); Assert.Equal("taylor@example.com", outcome.Member.Email); } await using var verifyScope = serviceProvider.CreateAsyncScope(); var verify = verifyScope.ServiceProvider.GetRequiredService(); var user = await verify.Users.AsNoTracking().SingleAsync(user => user.Email == "taylor@example.com"); Assert.True(user.PendingRegistration); var role = await verify.Roles.AsNoTracking().SingleAsync(role => role.Name == "Dispatcher"); Assert.Contains(await verify.UserRoles.AsNoTracking().ToListAsync(), row => row.UserId == user.Id && row.RoleId == role.Id); var areas = await verify.UserServiceAreas.AsNoTracking() .Where(area => area.UserId == user.Id) .Select(area => area.Area) .OrderBy(area => area) .ToListAsync(); Assert.Equal(new[] { "East", "West" }, areas); var permissionOverride = await verify.UserPermissionOverrides.AsNoTracking() .SingleAsync(permission => permission.UserId == user.Id); Assert.Equal("deleteSites", permissionOverride.PermissionKey); } private static ServiceProvider BuildServiceProvider(SqliteConnection connection, bool injectFailure = false) { var configuration = new ConfigurationBuilder().Build(); var services = new ServiceCollection(); services.AddLogging(); services.AddDbContext(builder => builder.UseSqlite(connection)); services.Replace(ServiceDescriptor.Scoped(provider => new SqliteTeamMemberTestDbContext( provider.GetRequiredService>()))); services.AddIdentity() .AddEntityFrameworkStores() .AddDefaultTokenProviders(); services.AddDataServices(); services.AddBusinessServices(configuration); if (injectFailure) { services.AddScoped(); services.AddScoped(provider => new ThrowingAfterPersistPermissionDataService( provider.GetRequiredService(), provider.GetRequiredService())); } return services.BuildServiceProvider(); } private static CreateTeamMemberRequestDTO ValidRequest() => new() { Name = "Taylor Dispatcher", Role = "dispatcher", Color = "#F59E0B", Email = "taylor@example.com", Phone = "555-0100", ServiceAreas = new[] { "east", "West" }, PermissionOverrides = new Dictionary { ["deleteSites"] = UserPermissionState.Allow } }; private static ClaimsPrincipal Admin() => new(new ClaimsIdentity(new[] { new Claim(ClaimTypes.Role, "Admin") }, "test")); private sealed class SqliteTeamMemberTestDbContext : ApplicationDbContext { public SqliteTeamMemberTestDbContext(DbContextOptions options) : base(options) { } protected override void OnModelCreating(ModelBuilder builder) { base.OnModelCreating(builder); foreach (var index in builder.Model.GetEntityTypes().SelectMany(entity => entity.GetIndexes())) { if (index.GetFilter() is not null) index.SetFilter(null); } foreach (var property in builder.Model.GetEntityTypes() .SelectMany(entity => entity.GetProperties()) .Where(property => property.Name == "RowVersion" && property.ClrType == typeof(byte[]))) { property.ValueGenerated = ValueGenerated.Never; property.IsConcurrencyToken = false; } } } private sealed class ThrowingAfterPersistPermissionDataService : ITeamPermissionOverrideDataService { private readonly ITeamPermissionOverrideDataService _inner; private readonly ApplicationDbContext _context; public ThrowingAfterPersistPermissionDataService( ITeamPermissionOverrideDataService inner, ApplicationDbContext context) { _inner = inner; _context = context; } public (int Users, int Roles, int UserRoles, int ServiceAreas, int PermissionOverrides)? RowsAtFailure { get; private set; } public Task GetUserAsync(string userId, CancellationToken cancellationToken) => _inner.GetUserAsync(userId, cancellationToken); public Task SetOverrideAsync( string userId, string permissionKey, UserPermissionState state, CancellationToken cancellationToken) => _inner.SetOverrideAsync(userId, permissionKey, state, cancellationToken); public async Task SetOverridesAsync( string userId, IReadOnlyDictionary overrides, CancellationToken cancellationToken) { await _inner.SetOverridesAsync(userId, overrides, cancellationToken); RowsAtFailure = ( await _context.Users.CountAsync(cancellationToken), await _context.Roles.CountAsync(cancellationToken), await _context.UserRoles.CountAsync(cancellationToken), await _context.UserServiceAreas.CountAsync(cancellationToken), await _context.UserPermissionOverrides.CountAsync(cancellationToken)); throw new InvalidOperationException("injected mid-sequence failure"); } public Task ClearOverridesAsync(string userId, CancellationToken cancellationToken) => _inner.ClearOverridesAsync(userId, cancellationToken); } }