shoc-backend/SeaHavenIndustries.Tests/PasswordResetRaceTests.cs

118 lines
5.6 KiB
C#

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;
/// <summary>
/// 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.
/// </summary>
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<Task>? _beforeNextConsume;
private Func<Task>? _afterNextConsume;
public Func<Task>? BeforeNextConsume { set => _beforeNextConsume = value; }
public Func<Task>? AfterNextConsume { set => _afterNextConsume = value; }
public void Register(IServiceCollection services) =>
services.Replace(ServiceDescriptor.Scoped<IForgetPasswordDataService>(provider =>
new RacingForgetPasswordDataService(new ForgetPasswordDataService(provider.GetRequiredService<ApplicationDbContext>()), 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, CancellationToken cancellationToken) =>
_inner.ReplaceCodeAsync(email, userId, codeHash, codeSalt, expiresAtUtc, cancellationToken);
public Task<ForgetPasswordCode?> GetByEmailAsync(string email, CancellationToken cancellationToken) =>
_inner.GetByEmailAsync(email, cancellationToken);
public async Task<bool> 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);
}
}