diff --git a/Api.SeaHavenIndustries.Tests/AuthenticationServiceTests.cs b/Api.SeaHavenIndustries.Tests/AuthenticationServiceTests.cs index 2fd6382..4f6a138 100644 --- a/Api.SeaHavenIndustries.Tests/AuthenticationServiceTests.cs +++ b/Api.SeaHavenIndustries.Tests/AuthenticationServiceTests.cs @@ -267,7 +267,7 @@ public class AuthenticationServiceTests var result = await service.ResetPasswordAsync("a@b.com", "123456", "New@67890", CancellationToken.None); result.Should().BeFalse(); - forget.Verify(f => f.RemoveByEmailAsync("a@b.com", It.IsAny()), Times.Once); + forget.Verify(f => f.RemoveIssuedThroughAsync("a@b.com", 7, It.IsAny()), Times.Once); store.Verify(s => s.FindByIdAsync(It.IsAny(), It.IsAny()), Times.Never); } diff --git a/SeaHaven.DataServices/Implementation/ForgetPasswordDataService.cs b/SeaHaven.DataServices/Implementation/ForgetPasswordDataService.cs index c824fc9..4d625c6 100644 --- a/SeaHaven.DataServices/Implementation/ForgetPasswordDataService.cs +++ b/SeaHaven.DataServices/Implementation/ForgetPasswordDataService.cs @@ -86,6 +86,14 @@ namespace SeaHaven.DataServices.Implementation .ExecuteDeleteAsync(cancellationToken); } + public async Task RemoveIssuedThroughAsync(string email, int throughId, CancellationToken cancellationToken) + { + var normalizedEmail = Normalize(email); + await _context.ForgetPasswordCodes + .Where(code => code.Email.ToLower().Trim() == normalizedEmail && code.Id <= throughId) + .ExecuteDeleteAsync(cancellationToken); + } + private static string Normalize(string email) => email.ToLower().Trim(); } } diff --git a/SeaHaven.DataServices/Interfaces/IForgetPasswordDataService.cs b/SeaHaven.DataServices/Interfaces/IForgetPasswordDataService.cs index c34e9ba..f1035f2 100644 --- a/SeaHaven.DataServices/Interfaces/IForgetPasswordDataService.cs +++ b/SeaHaven.DataServices/Interfaces/IForgetPasswordDataService.cs @@ -25,5 +25,11 @@ namespace SeaHaven.DataServices.Interfaces Task RefundAttemptAsync(int id, CancellationToken cancellationToken); Task RemoveByEmailAsync(string email, CancellationToken cancellationToken); + + /// + /// Deletes the email's codes issued up to and including , + /// leaving any code issued after it by a concurrent request. + /// + Task RemoveIssuedThroughAsync(string email, int throughId, CancellationToken cancellationToken); } } diff --git a/SeaHaven.Services/Implementation/AuthenticationService.cs b/SeaHaven.Services/Implementation/AuthenticationService.cs index de56193..c00c3c2 100644 --- a/SeaHaven.Services/Implementation/AuthenticationService.cs +++ b/SeaHaven.Services/Implementation/AuthenticationService.cs @@ -208,7 +208,7 @@ namespace SeaHaven.Services.Implementation var user = await _userManager.FindByIdAsync(pending.UserId); if (user == null || user.IsDeleted == true) { - await _forgetPasswordDataService.RemoveByEmailAsync(pending.Email, cancellationToken); + await _forgetPasswordDataService.RemoveIssuedThroughAsync(pending.Email, pending.Id, cancellationToken); return false; } @@ -250,11 +250,13 @@ namespace SeaHaven.Services.Implementation } var nowUtc = _timeProvider.GetUtcNow().UtcDateTime; + // Deletes below stop at this code's id: a code issued by a concurrent request + // after this one was read belongs to that request and must survive. if (!await _forgetPasswordDataService.TryConsumeAttemptAsync(pending.Id, MaxCodeAttempts, nowUtc, cancellationToken)) { // Expired or out of attempts: nothing was compared, so no guess is counted. _resetThrottle.ReleaseCheck(email); - await _forgetPasswordDataService.RemoveByEmailAsync(pending.Email, cancellationToken); + await _forgetPasswordDataService.RemoveIssuedThroughAsync(pending.Email, pending.Id, cancellationToken); return null; } @@ -265,7 +267,7 @@ namespace SeaHaven.Services.Implementation } if (pending.FailedAttempts + 1 >= MaxCodeAttempts) - await _forgetPasswordDataService.RemoveByEmailAsync(pending.Email, cancellationToken); + await _forgetPasswordDataService.RemoveIssuedThroughAsync(pending.Email, pending.Id, cancellationToken); return null; } diff --git a/SeaHavenIndustries.Tests/PasswordResetRaceTests.cs b/SeaHavenIndustries.Tests/PasswordResetRaceTests.cs new file mode 100644 index 0000000..1ae3937 --- /dev/null +++ b/SeaHavenIndustries.Tests/PasswordResetRaceTests.cs @@ -0,0 +1,121 @@ +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); + } +} diff --git a/SeaHavenIndustries.Tests/PasswordResetTestHost.cs b/SeaHavenIndustries.Tests/PasswordResetTestHost.cs index 1199080..58ea020 100644 --- a/SeaHavenIndustries.Tests/PasswordResetTestHost.cs +++ b/SeaHavenIndustries.Tests/PasswordResetTestHost.cs @@ -51,7 +51,7 @@ internal sealed class PasswordResetTestHost : IAsyncDisposable public CapturingLoggerProvider Logged { get; private init; } = null!; public IServiceProvider Services => _app.Services; - public static async Task StartAsync(IEmailSender? emailSender = null) + public static async Task StartAsync(IEmailSender? emailSender = null, Action? configureServices = null) { var databasePath = Path.Combine(Path.GetTempPath(), $"password-reset-{Guid.NewGuid():N}.db"); var connectionString = new SqliteConnectionStringBuilder { DataSource = databasePath, DefaultTimeout = 30 }.ToString(); @@ -84,6 +84,7 @@ internal sealed class PasswordResetTestHost : IAsyncDisposable builder.Services.AddSingleton(time); builder.Services.AddSingleton(emailSender ?? sent); builder.Services.AddDataServices(); + configureServices?.Invoke(builder.Services); builder.Services.AddBusinessServices(builder.Configuration); builder.Services.AddControllers() .AddApplicationPart(typeof(AuthenticationController).Assembly)