shoc-backend/SeaHavenIndustries.Tests/PasswordResetTestHost.cs
Alexandre Brandizzi bcc9b6d7a2 fix(auth): invalidate every pending reset code for an email together
Two concurrent first requests can leave two pending codes for one email.
Exhausting or using one now deletes all of them, so a sibling code cannot
become live afterwards.
2026-09-25 12:31:17 -03:00

297 lines
12 KiB
C#

using System.Collections.Concurrent;
using System.Net.Http.Json;
using System.Text.RegularExpressions;
using Api.SeaHavenIndustries.Controllers;
using Api.SeaHavenIndustries.Infrastructure;
using Data.SeaHavenIndustries;
using Microsoft.AspNetCore.Builder;
using Microsoft.AspNetCore.Hosting;
using Microsoft.AspNetCore.Hosting.Server;
using Microsoft.AspNetCore.Hosting.Server.Features;
using Microsoft.AspNetCore.Identity;
using Microsoft.AspNetCore.Mvc.ApplicationParts;
using Microsoft.AspNetCore.Mvc.Controllers;
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 Microsoft.Extensions.Logging;
using SeaHaven.DataServices.DependencyInjection;
using SeaHaven.Services.DependencyInjection;
using SeaHaven.Services.Interfaces;
namespace SeaHavenIndustries.Tests;
/// <summary>
/// Hosts the real <see cref="AuthenticationController"/> on Kestrel over a SQLite file
/// database, with the real authentication service, data services, Identity, and the
/// same rate limiting and forwarded-header registration the API host uses. Email goes
/// to an in-memory sender, the clock is manual, and every log line is captured.
/// </summary>
internal sealed class PasswordResetTestHost : IAsyncDisposable
{
private readonly WebApplication _app;
private readonly string _databasePath;
private PasswordResetTestHost(WebApplication app, string databasePath, HttpClient client)
{
_app = app;
_databasePath = databasePath;
Client = client;
}
public HttpClient Client { get; }
public CapturingEmailSender Sent { get; private init; } = null!;
public ManualTimeProvider Time { get; private init; } = null!;
public CapturingLoggerProvider Logged { get; private init; } = null!;
public IServiceProvider Services => _app.Services;
public static async Task<PasswordResetTestHost> StartAsync(IEmailSender? emailSender = null)
{
var databasePath = Path.Combine(Path.GetTempPath(), $"password-reset-{Guid.NewGuid():N}.db");
var connectionString = new SqliteConnectionStringBuilder { DataSource = databasePath, DefaultTimeout = 30 }.ToString();
var sent = new CapturingEmailSender();
var time = new ManualTimeProvider();
var logged = new CapturingLoggerProvider();
var builder = WebApplication.CreateBuilder(new WebApplicationOptions { EnvironmentName = "Testing" });
builder.WebHost.UseUrls("http://127.0.0.1:0");
builder.Configuration.AddInMemoryCollection(new Dictionary<string, string?>
{
["JWT:Secret"] = new string('k', 64),
["JWT:ValidIssuer"] = "issuer",
["JWT:ValidAudience"] = "audience"
});
builder.Logging.ClearProviders();
builder.Logging.SetMinimumLevel(LogLevel.Trace);
builder.Logging.AddProvider(logged);
builder.Services.AddDbContext<ApplicationDbContext>(options => options.UseSqlite(connectionString));
builder.Services.Replace(ServiceDescriptor.Scoped<ApplicationDbContext>(provider =>
new SqlitePasswordResetDbContext(provider.GetRequiredService<DbContextOptions<ApplicationDbContext>>())));
builder.Services.AddIdentity<ApplicationUser, IdentityRole>(options =>
{
options.User.RequireUniqueEmail = false;
})
.AddEntityFrameworkStores<ApplicationDbContext>()
.AddDefaultTokenProviders();
builder.Services.AddSingleton<TimeProvider>(time);
builder.Services.AddSingleton<IEmailSender>(emailSender ?? sent);
builder.Services.AddDataServices();
builder.Services.AddBusinessServices(builder.Configuration);
builder.Services.AddControllers()
.AddApplicationPart(typeof(AuthenticationController).Assembly)
.ConfigureApplicationPartManager(manager =>
{
manager.FeatureProviders.Clear();
manager.FeatureProviders.Add(new OnlyAuthenticationController());
});
builder.Services.AddPasswordResetRateLimiting();
var app = builder.Build();
app.UseForwardedHeaders();
app.UseRouting();
app.UseRateLimiter();
app.MapControllers();
await using (var scope = app.Services.CreateAsyncScope())
await scope.ServiceProvider.GetRequiredService<ApplicationDbContext>().Database.EnsureCreatedAsync();
await app.StartAsync();
var address = app.Services.GetRequiredService<IServer>().Features
.Get<IServerAddressesFeature>()!.Addresses.Single();
return new PasswordResetTestHost(app, databasePath, new HttpClient { BaseAddress = new Uri(address) })
{
Sent = sent,
Time = time,
Logged = logged
};
}
public async Task<ApplicationUser> AddUserAsync(string email, string password)
{
await using var scope = _app.Services.CreateAsyncScope();
var users = scope.ServiceProvider.GetRequiredService<UserManager<ApplicationUser>>();
var user = new ApplicationUser
{
UserName = email,
Email = email,
FirstName = "Alice",
LastName = "Q",
EmailConfirmed = true,
CreatedDate = DateTime.UtcNow
};
var result = await users.CreateAsync(user, password);
Assert.True(result.Succeeded, string.Join("; ", result.Errors.Select(error => error.Description)));
return user;
}
public async Task<List<ForgetPasswordCode>> PendingCodesAsync()
{
await using var scope = _app.Services.CreateAsyncScope();
return await scope.ServiceProvider.GetRequiredService<ApplicationDbContext>()
.ForgetPasswordCodes.AsNoTracking().OrderBy(code => code.Id).ToListAsync();
}
/// <summary>
/// Stores a second pending code for the email directly, as two concurrent first
/// requests could, and returns it. The new row is the newest one.
/// </summary>
public async Task<string> AddSiblingCodeAsync(string email, string userId)
{
const string code = "424242";
const string salt = "0123456789abcdef0123456789abcdef";
await using var scope = _app.Services.CreateAsyncScope();
var context = scope.ServiceProvider.GetRequiredService<ApplicationDbContext>();
context.ForgetPasswordCodes.Add(new ForgetPasswordCode
{
Email = email,
UserId = userId,
CodeSalt = salt,
CodeHash = SeaHaven.Services.Helpers.PasswordResetCodeSecrets.Hash(salt, code),
ExpiresAtUtc = Time.GetUtcNow().UtcDateTime.AddMinutes(15)
});
await context.SaveChangesAsync();
return code;
}
public static HttpRequestMessage Post(string path, object? json = null, string? clientIp = null)
{
var request = new HttpRequestMessage(HttpMethod.Post, path);
if (json != null)
request.Content = JsonContent.Create(json);
if (clientIp != null)
request.Headers.Add("X-Forwarded-For", clientIp);
return request;
}
public Task<HttpResponseMessage> ForgetPasswordAsync(string email, string? clientIp = null) =>
Client.SendAsync(Post("api/Authentication/ForgetPassword", new { email }, clientIp));
public Task<HttpResponseMessage> VerifyAsync(string email, string code, string? clientIp = null) =>
Client.SendAsync(Post("api/Authentication/VerificationCode", new { email, code }, clientIp));
public Task<HttpResponseMessage> ResetAsync(string email, string code, string password, string? clientIp = null) =>
Client.SendAsync(Post("api/Authentication/ResetPassword", new { email, code, password }, clientIp));
public Task<HttpResponseMessage> LoginAsync(string username, string password) =>
Client.SendAsync(Post("api/Authentication/login", new { username, password }));
public async ValueTask DisposeAsync()
{
Client.Dispose();
await _app.StopAsync();
await _app.DisposeAsync();
SqliteConnection.ClearAllPools();
foreach (var path in new[] { _databasePath, _databasePath + "-wal", _databasePath + "-shm", _databasePath + "-journal" })
{
if (File.Exists(path))
File.Delete(path);
}
}
private sealed class OnlyAuthenticationController : ControllerFeatureProvider
{
protected override bool IsController(System.Reflection.TypeInfo typeInfo) =>
typeInfo.AsType() == typeof(AuthenticationController);
}
private sealed class SqlitePasswordResetDbContext : ApplicationDbContext
{
public SqlitePasswordResetDbContext(DbContextOptions<ApplicationDbContext> options)
: base(options)
{
}
protected override void OnModelCreating(ModelBuilder builder)
{
base.OnModelCreating(builder);
foreach (var index in builder.Model.GetEntityTypes().SelectMany(entity => entity.GetIndexes()))
{
// SQL Server filter syntax does not carry over; a filtered unique index
// without its filter would wrongly reject a second user.
if (index.GetFilter() is not null)
{
index.SetFilter(null);
index.IsUnique = false;
}
}
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;
}
}
}
}
internal sealed class ManualTimeProvider : TimeProvider
{
private DateTimeOffset _now = new(2026, 9, 25, 12, 0, 0, TimeSpan.Zero);
public override DateTimeOffset GetUtcNow() => _now;
public void Advance(TimeSpan by) => _now = _now.Add(by);
}
internal sealed class CapturingEmailSender : IEmailSender
{
private static readonly Regex CodePattern = new(@"Your Password Reset Code is: (\d{6})", RegexOptions.CultureInvariant);
public ConcurrentQueue<(string To, string Subject, string Body)> Messages { get; } = new();
public Task<bool> SendEmailAsync(string emailTo, string subject, string body)
{
Messages.Enqueue((emailTo, subject, body));
return Task.FromResult(true);
}
public string LatestCodeFor(string email)
{
var message = Messages.Last(sent => string.Equals(sent.To, email, StringComparison.OrdinalIgnoreCase));
return CodePattern.Match(message.Body).Groups[1].Value;
}
}
internal sealed class CapturingLoggerProvider : ILoggerProvider
{
public ConcurrentQueue<string> Lines { get; } = new();
public ILogger CreateLogger(string categoryName) => new CapturingLogger(this, categoryName);
public void Dispose()
{
}
private sealed class CapturingLogger : ILogger
{
private readonly CapturingLoggerProvider _owner;
private readonly string _category;
public CapturingLogger(CapturingLoggerProvider owner, string category)
{
_owner = owner;
_category = category;
}
public IDisposable? BeginScope<TState>(TState state) where TState : notnull => null;
public bool IsEnabled(LogLevel logLevel) => true;
public void Log<TState>(LogLevel logLevel, EventId eventId, TState state, Exception? exception, Func<TState, Exception?, string> formatter)
{
var values = state is IEnumerable<KeyValuePair<string, object?>> pairs
? string.Join(" ", pairs.Select(pair => $"{pair.Key}={pair.Value}"))
: string.Empty;
_owner.Lines.Enqueue($"{logLevel} {_category} {formatter(state, exception)} {values} {exception}");
}
}
}