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(); 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); services.AddScoped(); services.AddScoped(provider => new ThrowingAfterPersistPermissionDataService( provider.GetRequiredService())); await using var serviceProvider = services.BuildServiceProvider(); await using (var createScope = serviceProvider.CreateAsyncScope()) { var service = 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); } 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()); } 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; public ThrowingAfterPersistPermissionDataService(ITeamPermissionOverrideDataService inner) { _inner = inner; } 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); throw new InvalidOperationException("injected mid-sequence failure"); } public Task ClearOverridesAsync(string userId, CancellationToken cancellationToken) => _inner.ClearOverridesAsync(userId, cancellationToken); } }