shoc-backend/SeaHavenIndustries.Tests/PasswordResetTestHost.cs
Alexandre Brandizzi c841e130be fix(auth): harden password reset codes against guessing and email enumeration
Forgot Password answers every address the same way and emails a code only
to an active account. Codes are stored as salted SHA-256 hashes, expire 15
minutes after issue, are replaced by a newer request, and are checked only
against the email they were issued to. Five failed checks delete the code;
attempts are reserved with one conditional UPDATE so concurrent guesses
cannot exceed the budget. VerificationCode requires the email, and email and
code are accepted in the JSON body so they stay out of URLs.

The three anonymous endpoints are rate limited to 10 requests per 15
minutes per client IP. Forwarded headers are trusted only through loopback
and private hops, since the API sits behind the EB load balancer and nginx.
The migration adds hash, salt, expiry and attempt columns and deletes the
old plaintext rows.
2026-09-25 12:21:29 -03:00

275 lines
11 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();
}
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}");
}
}
}