tests(api.tests): pass cancellation token through delegate

This commit is contained in:
Stevan Freeborn
2026-02-12 20:47:40 -06:00
parent 5fa2d72cee
commit ba66c1396c
3 changed files with 52 additions and 50 deletions
@@ -6,48 +6,47 @@ public abstract class IntegrationTest(TestApi testApi) : IClassFixture<TestApi>,
public async ValueTask InitializeAsync() public async ValueTask InitializeAsync()
{ {
await ExecuteDbContextAsync(static async context => await ExecuteAsync(static async (context, ct) =>
{ {
await context.Database.EnsureCreatedAsync(); await context.Database.EnsureCreatedAsync(ct);
}); }, TestContext.Current.CancellationToken);
} }
protected async Task ExecuteDbContextAsync(Func<DbContext, Task> action) protected async Task ExecuteAsync(Func<DbContext, CancellationToken, Task> action, CancellationToken ct)
{ {
await using var scope = testApi.Services.CreateAsyncScope(); await using var scope = testApi.Services.CreateAsyncScope();
var context = scope.ServiceProvider.GetRequiredService<AppDbContext>(); var context = scope.ServiceProvider.GetRequiredService<AppDbContext>();
await action(context); await action(context, ct);
} }
protected async Task ExecuteDbContextAsync(Func<DbContext, IServiceProvider, Task> action) protected async Task ExecuteAsync(Func<DbContext, CancellationToken, IServiceProvider, Task> action, CancellationToken ct)
{ {
await using var scope = testApi.Services.CreateAsyncScope(); await using var scope = testApi.Services.CreateAsyncScope();
var context = scope.ServiceProvider.GetRequiredService<AppDbContext>(); var context = scope.ServiceProvider.GetRequiredService<AppDbContext>();
await action(context, scope.ServiceProvider); await action(context, ct, scope.ServiceProvider);
} }
protected async Task<T> ExecuteDbContextAsync<T>(Func<DbContext, Task<T>> action) protected async Task<T> ExecuteAsync<T>(Func<DbContext, CancellationToken, Task<T>> action, CancellationToken ct)
{ {
await using var scope = testApi.Services.CreateAsyncScope(); await using var scope = testApi.Services.CreateAsyncScope();
var context = scope.ServiceProvider.GetRequiredService<AppDbContext>(); var context = scope.ServiceProvider.GetRequiredService<AppDbContext>();
return await action(context); return await action(context, ct);
} }
protected async Task<T> ExecuteDbContextAsync<T>(Func<DbContext, IServiceProvider, Task<T>> action) protected async Task<T> ExecuteAsync<T>(Func<DbContext, CancellationToken, IServiceProvider, Task<T>> action, CancellationToken ct)
{ {
await using var scope = testApi.Services.CreateAsyncScope(); await using var scope = testApi.Services.CreateAsyncScope();
var context = scope.ServiceProvider.GetRequiredService<AppDbContext>(); var context = scope.ServiceProvider.GetRequiredService<AppDbContext>();
return await action(context, scope.ServiceProvider); return await action(context, ct, scope.ServiceProvider);
} }
public async ValueTask DisposeAsync() public async ValueTask DisposeAsync()
{ {
await ExecuteDbContextAsync(static async context => await ExecuteAsync(static async (context, ct) =>
{ {
await context.Database.EnsureDeletedAsync(); await context.Database.EnsureDeletedAsync(ct);
}); }, TestContext.Current.CancellationToken);
GC.SuppressFinalize(this); GC.SuppressFinalize(this);
} }
} }
@@ -38,16 +38,16 @@ public class LoginTests(TestApi testApi) : IntegrationTest(testApi)
[Fact] [Fact]
public async Task Login_WhenUserExistsButPasswordIsIncorrect_ItShouldReturn401WithProblemDetails() public async Task Login_WhenUserExistsButPasswordIsIncorrect_ItShouldReturn401WithProblemDetails()
{ {
await ExecuteDbContextAsync(static async (context, sp) => await ExecuteAsync(static async (context, ct, sp) =>
{ {
var passwordHasher = sp.GetRequiredService<IPasswordHasher>(); var passwordHasher = sp.GetRequiredService<IPasswordHasher>();
var encryptor = sp.GetRequiredService<IEncryptor>(); var encryptor = sp.GetRequiredService<IEncryptor>();
var userEncryptionKey = await encryptor.GenerateEncryptedKeyAsync(TestContext.Current.CancellationToken); var userEncryptionKey = await encryptor.GenerateEncryptedKeyAsync(ct);
context.Add(User.From("Stevan", passwordHasher.Hash("@Password1"), userEncryptionKey)); context.Add(User.From("Stevan", passwordHasher.Hash("@Password1"), userEncryptionKey));
await context.SaveChangesAsync(TestContext.Current.CancellationToken); await context.SaveChangesAsync(ct);
}); }, TestContext.Current.CancellationToken);
var req = new var req = new
{ {
@@ -63,16 +63,16 @@ public class LoginTests(TestApi testApi) : IntegrationTest(testApi)
[Fact] [Fact]
public async Task Login_WhenUserExistsAndPasswordIsCorrect_ItShouldReturn200WithJwtTokenAndSetRefreshCookie() public async Task Login_WhenUserExistsAndPasswordIsCorrect_ItShouldReturn200WithJwtTokenAndSetRefreshCookie()
{ {
await ExecuteDbContextAsync(static async (context, sp) => await ExecuteAsync(static async (context, ct, sp) =>
{ {
var passwordHasher = sp.GetRequiredService<IPasswordHasher>(); var passwordHasher = sp.GetRequiredService<IPasswordHasher>();
var encryptor = sp.GetRequiredService<IEncryptor>(); var encryptor = sp.GetRequiredService<IEncryptor>();
var userEncryptionKey = await encryptor.GenerateEncryptedKeyAsync(TestContext.Current.CancellationToken); var userEncryptionKey = await encryptor.GenerateEncryptedKeyAsync(ct);
context.Add(User.From("Stevan", passwordHasher.Hash("@Password1"), userEncryptionKey)); context.Add(User.From("Stevan", passwordHasher.Hash("@Password1"), userEncryptionKey));
await context.SaveChangesAsync(TestContext.Current.CancellationToken); await context.SaveChangesAsync(ct);
}); }, TestContext.Current.CancellationToken);
var req = new var req = new
{ {
@@ -169,4 +169,4 @@ public record LoginValidationTestCase : IXunitSerializable
info.AddValue(nameof(Password), Password); info.AddValue(nameof(Password), Password);
info.AddValue(nameof(ExpectedErrors), ExpectedErrors); info.AddValue(nameof(ExpectedErrors), ExpectedErrors);
} }
} }
@@ -31,14 +31,14 @@ public class RefreshTests(TestApi testApi) : IntegrationTest(testApi)
[Fact] [Fact]
public async Task Refresh_WhenCalledWithTokenBelongingToDifferentUser_ItShouldReturn403WithProblemDetails() public async Task Refresh_WhenCalledWithTokenBelongingToDifferentUser_ItShouldReturn403WithProblemDetails()
{ {
var (users, refreshToken) = await ExecuteDbContextAsync(static async (context, sp) => var (users, refreshToken) = await ExecuteAsync(static async (context, ct, sp) =>
{ {
var passwordHasher = sp.GetRequiredService<IPasswordHasher>(); var passwordHasher = sp.GetRequiredService<IPasswordHasher>();
var tokenGenerator = sp.GetRequiredService<ITokenGenerator>(); var tokenGenerator = sp.GetRequiredService<ITokenGenerator>();
var encryptor = sp.GetRequiredService<IEncryptor>(); var encryptor = sp.GetRequiredService<IEncryptor>();
var user1EncryptionKey = await encryptor.GenerateEncryptedKeyAsync(TestContext.Current.CancellationToken); var user1EncryptionKey = await encryptor.GenerateEncryptedKeyAsync(ct);
var user2EncryptionKey = await encryptor.GenerateEncryptedKeyAsync(TestContext.Current.CancellationToken); var user2EncryptionKey = await encryptor.GenerateEncryptedKeyAsync(ct);
var user1 = User.From("User1", passwordHasher.Hash("@Password1"), user1EncryptionKey); var user1 = User.From("User1", passwordHasher.Hash("@Password1"), user1EncryptionKey);
var user2 = User.From("User2", passwordHasher.Hash("@Password2"), user2EncryptionKey); var user2 = User.From("User2", passwordHasher.Hash("@Password2"), user2EncryptionKey);
var refreshToken1 = tokenGenerator.GenerateRefreshToken(user2); var refreshToken1 = tokenGenerator.GenerateRefreshToken(user2);
@@ -49,9 +49,9 @@ public class RefreshTests(TestApi testApi) : IntegrationTest(testApi)
context.Add(user1); context.Add(user1);
context.Add(user2); context.Add(user2);
await context.SaveChangesAsync(TestContext.Current.CancellationToken); await context.SaveChangesAsync(ct);
return (new User[] { user1, user2 }, refreshToken1); return (new User[] { user1, user2 }, refreshToken1);
}); }, TestContext.Current.CancellationToken);
var jwt = JwtTokenBuilder.New() var jwt = JwtTokenBuilder.New()
.WithClaim(JwtRegisteredClaimNames.Sub, users[0].Id.ToString()) .WithClaim(JwtRegisteredClaimNames.Sub, users[0].Id.ToString())
@@ -65,11 +65,12 @@ public class RefreshTests(TestApi testApi) : IntegrationTest(testApi)
await response.Should().BeProblemDetails(HttpStatusCode.Forbidden); await response.Should().BeProblemDetails(HttpStatusCode.Forbidden);
var unrevokedTokensCountForUser2 = await ExecuteDbContextAsync( var unrevokedTokensCountForUser2 = await ExecuteAsync(
async context => await context.Set<RefreshToken>() async (context, ct) => await context.Set<RefreshToken>()
.Include(t => t.User) .Include(t => t.User)
.Where(t => t.UserId == users[1].Id && t.Revoked == false) .Where(t => t.UserId == users[1].Id && t.Revoked == false)
.CountAsync() .CountAsync(ct),
TestContext.Current.CancellationToken
); );
unrevokedTokensCountForUser2.Should().Be(0); unrevokedTokensCountForUser2.Should().Be(0);
@@ -78,13 +79,13 @@ public class RefreshTests(TestApi testApi) : IntegrationTest(testApi)
[Fact] [Fact]
public async Task Refresh_WhenCalledWithRevokedRefreshToken_ItShouldReturn400WithProblemDetails() public async Task Refresh_WhenCalledWithRevokedRefreshToken_ItShouldReturn400WithProblemDetails()
{ {
var (user, refreshToken) = await ExecuteDbContextAsync(static async (context, sp) => var (user, refreshToken) = await ExecuteAsync(static async (context, ct, sp) =>
{ {
var passwordHasher = sp.GetRequiredService<IPasswordHasher>(); var passwordHasher = sp.GetRequiredService<IPasswordHasher>();
var tokenGenerator = sp.GetRequiredService<ITokenGenerator>(); var tokenGenerator = sp.GetRequiredService<ITokenGenerator>();
var encryptor = sp.GetRequiredService<IEncryptor>(); var encryptor = sp.GetRequiredService<IEncryptor>();
var userEncryptionKey = await encryptor.GenerateEncryptedKeyAsync(TestContext.Current.CancellationToken); var userEncryptionKey = await encryptor.GenerateEncryptedKeyAsync(ct);
var user = User.From("User1", passwordHasher.Hash("@Password1"), userEncryptionKey); var user = User.From("User1", passwordHasher.Hash("@Password1"), userEncryptionKey);
var refreshToken = tokenGenerator.GenerateRefreshToken(user); var refreshToken = tokenGenerator.GenerateRefreshToken(user);
refreshToken.Revoke(); refreshToken.Revoke();
@@ -93,9 +94,9 @@ public class RefreshTests(TestApi testApi) : IntegrationTest(testApi)
context.Add(user); context.Add(user);
await context.SaveChangesAsync(TestContext.Current.CancellationToken); await context.SaveChangesAsync(ct);
return (user, refreshToken); return (user, refreshToken);
}); }, TestContext.Current.CancellationToken);
var jwt = JwtTokenBuilder.New() var jwt = JwtTokenBuilder.New()
.WithClaim(JwtRegisteredClaimNames.Sub, user.Id.ToString()) .WithClaim(JwtRegisteredClaimNames.Sub, user.Id.ToString())
@@ -113,23 +114,23 @@ public class RefreshTests(TestApi testApi) : IntegrationTest(testApi)
[Fact] [Fact]
public async Task Refresh_WhenCalledWithExpiredRefreshToken_ItShouldReturn400WithProblemDetails() public async Task Refresh_WhenCalledWithExpiredRefreshToken_ItShouldReturn400WithProblemDetails()
{ {
var (user, refreshToken) = await ExecuteDbContextAsync(static async (context, sp) => var (user, refreshToken) = await ExecuteAsync(static async (context, ct, sp) =>
{ {
var passwordHasher = sp.GetRequiredService<IPasswordHasher>(); var passwordHasher = sp.GetRequiredService<IPasswordHasher>();
var tokenGenerator = sp.GetRequiredService<ITokenGenerator>(); var tokenGenerator = sp.GetRequiredService<ITokenGenerator>();
var timeProvider = sp.GetRequiredService<TimeProvider>(); var timeProvider = sp.GetRequiredService<TimeProvider>();
var encryptor = sp.GetRequiredService<IEncryptor>(); var encryptor = sp.GetRequiredService<IEncryptor>();
var userEncryptionKey = await encryptor.GenerateEncryptedKeyAsync(TestContext.Current.CancellationToken); var userEncryptionKey = await encryptor.GenerateEncryptedKeyAsync(ct);
var user = User.From("User1", passwordHasher.Hash("@Password1"), userEncryptionKey); var user = User.From("User1", passwordHasher.Hash("@Password1"), userEncryptionKey);
var refreshToken = RefreshToken.From(user.Id, "expiredtoken", timeProvider.GetUtcNow().AddHours(-1)); var refreshToken = RefreshToken.From(user.Id, "expiredtoken", timeProvider.GetUtcNow().AddHours(-1));
user.AddRefreshToken(refreshToken); user.AddRefreshToken(refreshToken);
context.Add(user); context.Add(user);
await context.SaveChangesAsync(TestContext.Current.CancellationToken); await context.SaveChangesAsync(ct);
return (user, refreshToken); return (user, refreshToken);
}); }, TestContext.Current.CancellationToken);
var jwt = JwtTokenBuilder.New() var jwt = JwtTokenBuilder.New()
.WithClaim(JwtRegisteredClaimNames.Sub, user.Id.ToString()) .WithClaim(JwtRegisteredClaimNames.Sub, user.Id.ToString())
@@ -149,23 +150,23 @@ public class RefreshTests(TestApi testApi) : IntegrationTest(testApi)
[InlineData(-5)] [InlineData(-5)]
public async Task Refresh_WhenCalledWithValidRefreshTokenAndExpiredOrNotExpiredAccessToken_ItShouldReturn200WithNewTokensAndSetRefreshCookie(int accessTokenExpiresAtOffset) public async Task Refresh_WhenCalledWithValidRefreshTokenAndExpiredOrNotExpiredAccessToken_ItShouldReturn200WithNewTokensAndSetRefreshCookie(int accessTokenExpiresAtOffset)
{ {
var (user, refreshToken) = await ExecuteDbContextAsync(static async (context, sp) => var (user, refreshToken) = await ExecuteAsync(static async (context, ct, sp) =>
{ {
var passwordHasher = sp.GetRequiredService<IPasswordHasher>(); var passwordHasher = sp.GetRequiredService<IPasswordHasher>();
var tokenGenerator = sp.GetRequiredService<ITokenGenerator>(); var tokenGenerator = sp.GetRequiredService<ITokenGenerator>();
var timeProvider = sp.GetRequiredService<TimeProvider>(); var timeProvider = sp.GetRequiredService<TimeProvider>();
var encryptor = sp.GetRequiredService<IEncryptor>(); var encryptor = sp.GetRequiredService<IEncryptor>();
var userEncryptionKey = await encryptor.GenerateEncryptedKeyAsync(TestContext.Current.CancellationToken); var userEncryptionKey = await encryptor.GenerateEncryptedKeyAsync(ct);
var user = User.From("User1", passwordHasher.Hash("@Password1"), userEncryptionKey); var user = User.From("User1", passwordHasher.Hash("@Password1"), userEncryptionKey);
var refreshToken = tokenGenerator.GenerateRefreshToken(user); var refreshToken = tokenGenerator.GenerateRefreshToken(user);
user.AddRefreshToken(refreshToken); user.AddRefreshToken(refreshToken);
context.Add(user); context.Add(user);
await context.SaveChangesAsync(TestContext.Current.CancellationToken); await context.SaveChangesAsync(ct);
return (user, refreshToken); return (user, refreshToken);
}); }, TestContext.Current.CancellationToken);
var jwt = JwtTokenBuilder.New() var jwt = JwtTokenBuilder.New()
.WithClaim(JwtRegisteredClaimNames.Sub, user.Id.ToString()) .WithClaim(JwtRegisteredClaimNames.Sub, user.Id.ToString())
@@ -181,20 +182,22 @@ public class RefreshTests(TestApi testApi) : IntegrationTest(testApi)
response.Should().HaveSetCookieHeader("fiscalos_refresh_cookie"); response.Should().HaveSetCookieHeader("fiscalos_refresh_cookie");
await response.Should().BeJsonContentOfType<Refresh.Response>(HttpStatusCode.OK); await response.Should().BeJsonContentOfType<Refresh.Response>(HttpStatusCode.OK);
var oldRefreshTokenInDb = await ExecuteDbContextAsync( var oldRefreshTokenInDb = await ExecuteAsync(
async context => await context.Set<RefreshToken>() async (context, ct) => await context.Set<RefreshToken>()
.Include(t => t.User) .Include(t => t.User)
.Where(t => t.UserId == user.Id && t.Token == refreshToken.Token && t.Revoked == true) .Where(t => t.UserId == user.Id && t.Token == refreshToken.Token && t.Revoked == true)
.SingleOrDefaultAsync() .SingleOrDefaultAsync(ct),
TestContext.Current.CancellationToken
); );
oldRefreshTokenInDb.Should().NotBeNull(); oldRefreshTokenInDb.Should().NotBeNull();
var newRefreshTokenInDb = await ExecuteDbContextAsync( var newRefreshTokenInDb = await ExecuteAsync(
async context => await context.Set<RefreshToken>() async (context, ct) => await context.Set<RefreshToken>()
.Include(t => t.User) .Include(t => t.User)
.Where(t => t.UserId == user.Id && t.Revoked == false && t.Token != refreshToken.Token) .Where(t => t.UserId == user.Id && t.Revoked == false && t.Token != refreshToken.Token)
.SingleOrDefaultAsync() .SingleOrDefaultAsync(ct),
TestContext.Current.CancellationToken
); );
newRefreshTokenInDb.Should().NotBeNull(); newRefreshTokenInDb.Should().NotBeNull();