using System.Net; using Data.SeaHavenIndustries; using Microsoft.Extensions.DependencyInjection; using Microsoft.Extensions.DependencyInjection.Extensions; using SeaHaven.DataServices.Implementation; using SeaHaven.DataServices.Interfaces; namespace SeaHavenIndustries.Tests; /// /// A new code requested while a check of the previous code is in flight. The /// request runs for real, between the check's steps, through the same host. /// public sealed class PasswordResetRaceTests { private const string Alice = "alice@example.com"; private const string OldPassword = "Old@12345"; private const string NewPassword = "New@67890"; private static int _nextClient; private static string NextClient() { var n = Interlocked.Increment(ref _nextClient); return $"192.0.{n / 250 % 250}.{n % 250 + 1}"; } [Fact] public async Task A_code_requested_while_the_previous_one_is_being_checked_still_resets() { var race = new CheckRace(); await using var host = await PasswordResetTestHost.StartAsync(configureServices: race.Register); await host.AddUserAsync(Alice, OldPassword); await host.ForgetPasswordAsync(Alice, NextClient()); var first = host.Sent.LatestCodeFor(Alice); // The check has read the first code; the new request replaces it before the attempt is counted. race.BeforeNextConsume = () => host.ForgetPasswordAsync(Alice, NextClient()); Assert.Equal(HttpStatusCode.BadRequest, (await host.VerifyAsync(Alice, first, NextClient())).StatusCode); Assert.Equal(2, host.Sent.Messages.Count); var second = host.Sent.LatestCodeFor(Alice); Assert.Equal(HttpStatusCode.OK, (await host.ResetAsync(Alice, second, NewPassword, NextClient())).StatusCode); Assert.Equal(HttpStatusCode.OK, (await host.LoginAsync(Alice, NewPassword)).StatusCode); } [Fact] public async Task A_code_requested_while_the_last_wrong_attempt_is_being_checked_still_resets() { var race = new CheckRace(); await using var host = await PasswordResetTestHost.StartAsync(configureServices: race.Register); await host.AddUserAsync(Alice, OldPassword); await host.ForgetPasswordAsync(Alice, NextClient()); var wrong = host.Sent.LatestCodeFor(Alice) == "000000" ? "111111" : "000000"; for (var attempt = 1; attempt < 5; attempt++) await host.VerifyAsync(Alice, wrong, NextClient()); // The fifth wrong attempt is counted, then the new request lands before the used-up code is deleted. race.AfterNextConsume = () => host.ForgetPasswordAsync(Alice, NextClient()); Assert.Equal(HttpStatusCode.BadRequest, (await host.VerifyAsync(Alice, wrong, NextClient())).StatusCode); Assert.Equal(2, host.Sent.Messages.Count); var second = host.Sent.LatestCodeFor(Alice); Assert.Equal(HttpStatusCode.OK, (await host.ResetAsync(Alice, second, NewPassword, NextClient())).StatusCode); Assert.Equal(HttpStatusCode.OK, (await host.LoginAsync(Alice, NewPassword)).StatusCode); } private sealed class CheckRace { private Func? _beforeNextConsume; private Func? _afterNextConsume; public Func? BeforeNextConsume { set => _beforeNextConsume = value; } public Func? AfterNextConsume { set => _afterNextConsume = value; } public void Register(IServiceCollection services) => services.Replace(ServiceDescriptor.Scoped(provider => new RacingForgetPasswordDataService(new ForgetPasswordDataService(provider.GetRequiredService()), this))); public Task RunBeforeConsumeAsync() => Interlocked.Exchange(ref _beforeNextConsume, null)?.Invoke() ?? Task.CompletedTask; public Task RunAfterConsumeAsync() => Interlocked.Exchange(ref _afterNextConsume, null)?.Invoke() ?? Task.CompletedTask; } private sealed class RacingForgetPasswordDataService : IForgetPasswordDataService { private readonly IForgetPasswordDataService _inner; private readonly CheckRace _race; public RacingForgetPasswordDataService(IForgetPasswordDataService inner, CheckRace race) { _inner = inner; _race = race; } public Task ReplaceCodeAsync(string email, string userId, string codeHash, string codeSalt, DateTime expiresAtUtc, DateTime nowUtc, CancellationToken cancellationToken) => _inner.ReplaceCodeAsync(email, userId, codeHash, codeSalt, expiresAtUtc, nowUtc, cancellationToken); public Task PurgeExpiredAsync(DateTime nowUtc, CancellationToken cancellationToken) => _inner.PurgeExpiredAsync(nowUtc, cancellationToken); public Task GetByEmailAsync(string email, CancellationToken cancellationToken) => _inner.GetByEmailAsync(email, cancellationToken); public async Task TryConsumeAttemptAsync(int id, int maxAttempts, DateTime nowUtc, CancellationToken cancellationToken) { await _race.RunBeforeConsumeAsync(); var consumed = await _inner.TryConsumeAttemptAsync(id, maxAttempts, nowUtc, cancellationToken); await _race.RunAfterConsumeAsync(); return consumed; } public Task RefundAttemptAsync(int id, CancellationToken cancellationToken) => _inner.RefundAttemptAsync(id, cancellationToken); public Task RemoveByEmailAsync(string email, CancellationToken cancellationToken) => _inner.RemoveByEmailAsync(email, cancellationToken); public Task RemoveIssuedThroughAsync(string email, int throughId, CancellationToken cancellationToken) => _inner.RemoveIssuedThroughAsync(email, throughId, cancellationToken); } }