shoc-backend/SeaHavenIndustries.Tests/PasswordResetTestHost.cs

321 lines
13 KiB
C#

using System.Collections.Concurrent;
using System.Net.Http.Json;
using System.Text.RegularExpressions;
using Api.SeaHavenIndustries.Controllers;
using Api.SeaHavenIndustries.HostedServices;
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
{
public static readonly string JwtSecret = new('k', 64);
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"] = JwtSecret,
["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();
builder.Services.AddSingleton<Sentry.IHub>(Sentry.Extensibility.HubAdapter.Instance);
builder.Services.AddPasswordResetEmailDelivery();
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, bool keyedHash = true)
{
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 = keyedHash
? SeaHaven.Services.Helpers.PasswordResetCodeSecrets.Hash(SeaHaven.Services.Helpers.PasswordResetCodeSecrets.DeriveKey(JwtSecret), salt, code)
: Convert.ToHexString(System.Security.Cryptography.SHA256.HashData(System.Text.Encoding.UTF8.GetBytes(salt + ":" + code))).ToLowerInvariant(),
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;
}
/// <summary>Requests a code and waits until any queued email has been handed to the sender.</summary>
public async Task<HttpResponseMessage> ForgetPasswordAsync(string email, string? clientIp = null)
{
var response = await Client.SendAsync(Post("api/Authentication/ForgetPassword", new { email }, clientIp));
await WaitForEmailDrainAsync();
return response;
}
public async Task WaitForEmailDrainAsync()
{
var channel = _app.Services.GetRequiredService<PasswordResetEmailChannel>();
var deadline = DateTime.UtcNow.AddSeconds(10);
while (channel.Pending > 0)
{
if (DateTime.UtcNow > deadline)
throw new TimeoutException("Queued password reset emails were not sent.");
await Task.Delay(10);
}
}
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}");
}
}
}