diff --git a/src/atproto/oauth.zig b/src/atproto/oauth.zig index a0e4227..8631387 100644 --- a/src/atproto/oauth.zig +++ b/src/atproto/oauth.zig @@ -445,10 +445,14 @@ fn authorizationCodeToken(request: *http_api.Request, allocator: std.mem.Allocat const client_id = try requireParam(request, params, allocator, "client_id"); const client_metadata = try fetchClientMetadataForAuth(request, allocator, client_id); try requireClientAuth(request, allocator, params, client_id, client_metadata.value); - const oauth_request = (try store.consumeOAuthCode(allocator, code)) orelse { + const oauth_request = (try store.getOAuthRequestByCode(allocator, code)) orelse { log.err("oauth code grant invalid_code client={s} code_len={d}\n", .{ client_id, code.len }); return oauthError(request, .bad_request, "invalid_grant", "Invalid code"); }; + if (oauth_request.expires_at < now()) { + log.err("oauth code grant expired request={s} client={s} expires_at={d}\n", .{ oauth_request.request_id, client_id, oauth_request.expires_at }); + return oauthError(request, .bad_request, "invalid_grant", "Expired code"); + } if (!std.mem.eql(u8, redirect_uri, oauth_request.redirect_uri) or !std.mem.eql(u8, client_id, oauth_request.client_id)) { log.err("oauth code grant mismatch request={s} row_client={s} token_client={s} row_redirect={s} token_redirect={s}\n", .{ oauth_request.request_id, oauth_request.client_id, client_id, oauth_request.redirect_uri, redirect_uri }); return oauthError(request, .bad_request, "invalid_grant", "Mismatched token request"); @@ -473,6 +477,10 @@ fn authorizationCodeToken(request: *http_api.Request, allocator: std.mem.Allocat log.err("oauth code grant dpop_reject request={s} client={s} did={s} err={s}\n", .{ oauth_request.request_id, client_id, did, @errorName(err) }); return handleAuthorizationDpopError(request, allocator, err); }; + if (!try store.consumeOAuthCode(oauth_request.request_id, code)) { + log.err("oauth code grant replay request={s} client={s} did={s}\n", .{ oauth_request.request_id, client_id, did }); + return oauthError(request, .bad_request, "invalid_grant", "Invalid code"); + } try issueTokenResponse(request, allocator, account, oauth_request.client_id, oauth_request.scope, null, oauth_request.dpop_jkt, oauth_request.auth_method); log.info("oauth code grant issued request={s} client={s} did={s} auth_method={s}\n", .{ oauth_request.request_id, client_id, did, oauth_request.auth_method orelse "unknown" }); } diff --git a/src/storage/store.zig b/src/storage/store.zig index e9c8106..23037cf 100644 --- a/src/storage/store.zig +++ b/src/storage/store.zig @@ -1142,7 +1142,7 @@ pub fn authorizeOAuthRequest(request_id: []const u8, did: []const u8, code: []co , .{ did, code, auth_method, request_id }); } -pub fn consumeOAuthCode(allocator: std.mem.Allocator, code: []const u8) !?OAuthRequest { +pub fn getOAuthRequestByCode(allocator: std.mem.Allocator, code: []const u8) !?OAuthRequest { db_mutex.lockUncancelable(store_io); defer db_mutex.unlock(store_io); try requireInitialized(); @@ -1153,9 +1153,21 @@ pub fn consumeOAuthCode(allocator: std.mem.Allocator, code: []const u8) !?OAuthR , .{code}); if (row == null) return null; defer row.?.deinit(); - const request = try oauthRequestFromRow(row.?, allocator); - try conn.exec("UPDATE oauth_requests SET code = NULL WHERE request_id = ?", .{request.request_id}); - return request; + return try oauthRequestFromRow(row.?, allocator); +} + +pub fn consumeOAuthCode(request_id: []const u8, code: []const u8) !bool { + db_mutex.lockUncancelable(store_io); + defer db_mutex.unlock(store_io); + try requireInitialized(); + try conn.exec( + "UPDATE oauth_requests SET code = NULL WHERE request_id = ? AND code = ?", + .{ request_id, code }, + ); + const row = try conn.row("SELECT changes()", .{}); + if (row == null) return false; + defer row.?.deinit(); + return row.?.int(0) == 1; } pub fn putOAuthToken( @@ -7092,6 +7104,41 @@ test "stores discoverable passkey challenges by oauth request" { try std.testing.expect(try getDiscoverableWebAuthnChallenge(allocator, "req-discoverable") == null); } +test "oauth authorization codes remain readable until atomically consumed" { + var arena = std.heap.ArenaAllocator.init(std.testing.allocator); + defer arena.deinit(); + const allocator = arena.allocator(); + + try init(std.Options.debug_io, ":memory:"); + defer close(); + + try putOAuthRequest( + "req-code", + "https://client.example.com/oauth-client-metadata.json", + "https://client.example.com/callback", + "atproto", + "state", + "challenge", + "S256", + "query", + null, + "dpop-thumbprint", + 4102444800, + ); + try authorizeOAuthRequest("req-code", "did:plc:code", "authorization-code", "passkey"); + + const before = (try getOAuthRequestByCode(allocator, "authorization-code")) orelse return error.MissingRecord; + try std.testing.expectEqualStrings("req-code", before.request_id); + try std.testing.expectEqualStrings("authorization-code", before.code.?); + + try std.testing.expect(!try consumeOAuthCode("wrong-request", "authorization-code")); + try std.testing.expect((try getOAuthRequestByCode(allocator, "authorization-code")) != null); + + try std.testing.expect(try consumeOAuthCode("req-code", "authorization-code")); + try std.testing.expect((try getOAuthRequestByCode(allocator, "authorization-code")) == null); + try std.testing.expect(!try consumeOAuthCode("req-code", "authorization-code")); +} + test "oauth app sessions stay active after access token expiry until revoked" { var arena = std.heap.ArenaAllocator.init(std.testing.allocator); defer arena.deinit();