diff --git a/CHANGELOG.md b/CHANGELOG.md index f85d392..bf3652b 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,11 @@ # changelog +## 0.3.6 + +- **feat**: expand `zat.oauth` from low-level PKCE/DPoP/client-assertion helpers into a framework-neutral ATProto OAuth client toolkit. New helpers cover client metadata JSON, authorization-server discovery, authorization URL construction, PAR, code exchange, refresh-token exchange, DPoP nonce retry, and DPoP-authenticated resource requests while leaving cookies, redirects, sessions, and storage policy to applications. +- **fix**: enforce ATProto OAuth profile requirements for metadata discovery, scopes, DPoP nonces, and DPoP `htu` normalization, with focused tests for the new correctness surface. +- **feat**: `HttpTransport.FetchResult` can capture OAuth-relevant response headers (`Content-Type`, `DPoP-Nonce`, `WWW-Authenticate`) and has a `deinit` helper so OAuth clients can stay on the shared transport path instead of dropping to raw `std.http.Client`. + ## 0.3.5 - **feat**: firehose `#commit` events now expose the raw `blocks` CAR bytes, `prevData`, operation `prev` CIDs, and a `toMstOperations()` helper so consumers can call `verifyCommitDiff`/`verifyCommitCar` without re-decoding raw frames. Added `#sync` event decoding and a named `LoadedCommitCar` return type for `loadCommitFromCAR`. diff --git a/README.md b/README.md index 7bf6589..21dfd15 100644 --- a/README.md +++ b/README.md @@ -162,6 +162,41 @@ ES256 (P-256) and ES256K (secp256k1) with low-S normalization. RFC 6979 determin +
+OAuth client helpers - ATProto OAuth ceremony without app storage policy + +```zig +var transport = zat.HttpTransport.init(io, allocator); +defer transport.deinit(); + +const authserver = try zat.oauth.discoverAuthorizationServer(allocator, &transport, pds_url); +defer allocator.free(authserver); + +var metadata = try zat.oauth.fetchAuthorizationServerMetadata(allocator, &transport, authserver); +defer metadata.deinit(allocator); + +var secrets = try zat.oauth.prepareAuthRequestSecrets(allocator, io); +defer secrets.deinit(allocator); + +var par = try zat.oauth.sendParRequest(allocator, io, &transport, .{ + .par_url = metadata.pushed_authorization_request_endpoint, + .authserver_issuer = metadata.issuer, + .client_id = client_id, + .redirect_uri = redirect_uri, + .scope = "atproto repo:example.app.record", + .state = secrets.state, + .pkce_challenge = secrets.pkce_challenge, + .login_hint = handle, + .client_keypair = &client_keypair, + .dpop_keypair = &secrets.dpop_keypair, +}); +defer par.deinit(allocator); +``` + +also includes client metadata JSON generation, authorization URL formatting, code/refresh token exchange, DPoP nonce retry, and DPoP-authenticated resource requests. cookies, sessions, redirects, and persistence stay with your application or web framework. + +
+
repo verification - full AT Protocol trust chain diff --git a/src/internal/oauth.zig b/src/internal/oauth.zig index cd6535d..2bf054d 100644 --- a/src/internal/oauth.zig +++ b/src/internal/oauth.zig @@ -1,391 +1,41 @@ -//! OAuth client primitives for AT Protocol +//! OAuth helpers for AT Protocol. //! -//! PKCE, DPoP proofs, client assertions, and related helpers -//! for implementing AT Protocol OAuth flows (based on OAuth 2.1). -//! -//! see: https://atproto.com/specs/oauth - -const std = @import("std"); -const crypto = std.crypto; -const Io = std.Io; -const Allocator = std.mem.Allocator; -const Keypair = @import("crypto/keypair.zig").Keypair; -const jwt = @import("crypto/jwt.zig"); - -fn timestamp(io: Io) i64 { - return @intCast(@divFloor(Io.Timestamp.now(io, .real).nanoseconds, std.time.ns_per_s)); -} - -/// create a signed JWT from header and payload JSON strings. -/// caller owns returned slice. -pub fn createJwt(allocator: Allocator, header_json: []const u8, payload_json: []const u8, keypair: *const Keypair) ![]u8 { - const header_b64 = try jwt.base64UrlEncode(allocator, header_json); - defer allocator.free(header_b64); - - const payload_b64 = try jwt.base64UrlEncode(allocator, payload_json); - defer allocator.free(payload_b64); - - const signing_input = try std.fmt.allocPrint(allocator, "{s}.{s}", .{ header_b64, payload_b64 }); - defer allocator.free(signing_input); - - const sig = try keypair.sign(signing_input); - const sig_b64 = try jwt.base64UrlEncode(allocator, &sig.bytes); - defer allocator.free(sig_b64); - - return std.fmt.allocPrint(allocator, "{s}.{s}", .{ signing_input, sig_b64 }); -} - -/// create a DPoP proof JWT per RFC 9449. -/// htm: HTTP method, htu: target URI, nonce: server-provided DPoP-Nonce, -/// ath: optional access token hash (base64url-encoded SHA-256). -pub fn createDpopProof( - allocator: Allocator, - io: std.Io, - keypair: *const Keypair, - htm: []const u8, - htu: []const u8, - nonce: ?[]const u8, - ath: ?[]const u8, -) ![]u8 { - const jwk_json = try keypair.jwk(allocator); - defer allocator.free(jwk_json); - - const jti = try generateJti(allocator, io); - defer allocator.free(jti); - - const alg = @tagName(keypair.algorithm()); - const now = timestamp(io); - - // header: {"typ":"dpop+jwt","alg":"...","jwk":{...}} - const header = try std.fmt.allocPrint(allocator, - \\{{"typ":"dpop+jwt","alg":"{s}","jwk":{s}}} - , .{ alg, jwk_json }); - defer allocator.free(header); - - // payload — build with writer for optional fields - var aw: std.Io.Writer.Allocating = .init(allocator); - defer aw.deinit(); - - try aw.writer.print( - \\{{"jti":"{s}","htm":"{s}","htu":"{s}","iat":{d} - , .{ jti, htm, htu, now }); - - if (nonce) |n| { - try aw.writer.print(",\"nonce\":\"{s}\"", .{n}); - } - if (ath) |a| { - try aw.writer.print(",\"ath\":\"{s}\"", .{a}); - } - - try aw.writer.writeAll("}"); - - return createJwt(allocator, header, aw.written(), keypair); -} - -/// create a private_key_jwt client assertion for token endpoint auth. -/// client_id: the OAuth client ID, aud: the token endpoint URL. -pub fn createClientAssertion( - allocator: Allocator, - io: std.Io, - keypair: *const Keypair, - client_id: []const u8, - aud: []const u8, -) ![]u8 { - const jti = try generateJti(allocator, io); - defer allocator.free(jti); - - const kid = try keypair.jwkThumbprint(allocator); - defer allocator.free(kid); - - const alg = @tagName(keypair.algorithm()); - const now = timestamp(io); - - const header = try std.fmt.allocPrint(allocator, - \\{{"typ":"JWT","alg":"{s}","kid":"{s}"}} - , .{ alg, kid }); - defer allocator.free(header); - - const payload = try std.fmt.allocPrint(allocator, - \\{{"iss":"{s}","sub":"{s}","aud":"{s}","jti":"{s}","iat":{d},"exp":{d}}} - , .{ client_id, client_id, aud, jti, now, now + 120 }); - defer allocator.free(payload); - - return createJwt(allocator, header, payload, keypair); -} - -/// generate a random PKCE code verifier (43 chars, base64url-encoded 32 random bytes). -/// caller owns returned slice. -pub fn generatePkceVerifier(allocator: Allocator, io: std.Io) ![]u8 { - var random_bytes: [32]u8 = undefined; - io.random(&random_bytes); - return jwt.base64UrlEncode(allocator, &random_bytes); -} - -/// generate a PKCE S256 challenge from a verifier. -/// caller owns returned slice. -pub fn generatePkceChallenge(allocator: Allocator, verifier: []const u8) ![]u8 { - var hash: [32]u8 = undefined; - crypto.hash.sha2.Sha256.hash(verifier, &hash, .{}); - return jwt.base64UrlEncode(allocator, &hash); -} - -/// generate a random state parameter (CSRF token). -/// caller owns returned slice. -pub fn generateState(allocator: Allocator, io: std.Io) ![]u8 { - var random_bytes: [16]u8 = undefined; - io.random(&random_bytes); - return jwt.base64UrlEncode(allocator, &random_bytes); -} - -/// compute access token hash for DPoP ath claim: base64url(SHA-256(access_token)). -/// caller owns returned slice. -pub fn accessTokenHash(allocator: Allocator, access_token: []const u8) ![]u8 { - var hash: [32]u8 = undefined; - crypto.hash.sha2.Sha256.hash(access_token, &hash, .{}); - return jwt.base64UrlEncode(allocator, &hash); -} - -/// encode key-value pairs as application/x-www-form-urlencoded. -/// caller owns returned slice. -pub fn formEncode(allocator: Allocator, params: []const [2][]const u8) ![]u8 { - var aw: std.Io.Writer.Allocating = .init(allocator); - errdefer aw.deinit(); - - for (params, 0..) |kv, i| { - if (i > 0) try aw.writer.writeAll("&"); - try percentEncode(&aw.writer, kv[0]); - try aw.writer.writeAll("="); - try percentEncode(&aw.writer, kv[1]); - } - - return try aw.toOwnedSlice(); -} - -/// format a JWKS JSON containing a single public key. -/// caller owns returned slice. -pub fn jwksJson(allocator: Allocator, keypair: *const Keypair) ![]u8 { - const jwk_json = try keypair.jwk(allocator); - defer allocator.free(jwk_json); - - return std.fmt.allocPrint(allocator, - \\{{"keys":[{s}]}} - , .{jwk_json}); -} - -// --- helpers --- - -fn generateJti(allocator: Allocator, io: std.Io) ![]u8 { - var random_bytes: [16]u8 = undefined; - io.random(&random_bytes); - return jwt.base64UrlEncode(allocator, &random_bytes); -} - -fn percentEncode(writer: anytype, input: []const u8) !void { - for (input) |c| { - if (isUnreserved(c)) { - try writer.writeByte(c); - } else { - try writer.print("%{X:0>2}", .{c}); - } - } -} - -fn isUnreserved(c: u8) bool { - return switch (c) { - 'A'...'Z', 'a'...'z', '0'...'9', '-', '_', '.', '~' => true, - else => false, - }; -} - -// === tests === - -test "PKCE S256 challenge - RFC 7636 test vector" { - const allocator = std.testing.allocator; - const verifier = "dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEjXk"; - const challenge = try generatePkceChallenge(allocator, verifier); - defer allocator.free(challenge); - try std.testing.expectEqualStrings("E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM", challenge); -} - -test "PKCE verifier is 43 chars" { - const allocator = std.testing.allocator; - const io = std.Options.debug_io; - const verifier = try generatePkceVerifier(allocator, io); - defer allocator.free(verifier); - try std.testing.expectEqual(@as(usize, 43), verifier.len); -} - -test "form URL encoding" { - const allocator = std.testing.allocator; - - const params = [_][2][]const u8{ - .{ "grant_type", "authorization_code" }, - .{ "code", "abc123" }, - .{ "redirect_uri", "https://example.com/callback" }, - }; - - const encoded = try formEncode(allocator, ¶ms); - defer allocator.free(encoded); - - try std.testing.expectEqualStrings( - "grant_type=authorization_code&code=abc123&redirect_uri=https%3A%2F%2Fexample.com%2Fcallback", - encoded, - ); -} - -test "access token hash" { - const allocator = std.testing.allocator; - const ath = try accessTokenHash(allocator, "test-access-token"); - defer allocator.free(ath); - // base64url(SHA-256) is always 43 chars - try std.testing.expectEqual(@as(usize, 43), ath.len); -} - -test "createJwt sign and verify round-trip" { - const allocator = std.testing.allocator; - const multibase = @import("crypto/multibase.zig"); - const multicodec = @import("crypto/multicodec.zig"); - - const keypair = try Keypair.fromSecretKey(.p256, .{ - 0x21, 0x22, 0x23, 0x24, 0x25, 0x26, 0x27, 0x28, - 0x29, 0x2a, 0x2b, 0x2c, 0x2d, 0x2e, 0x2f, 0x30, - 0x31, 0x32, 0x33, 0x34, 0x35, 0x36, 0x37, 0x38, - 0x39, 0x3a, 0x3b, 0x3c, 0x3d, 0x3e, 0x3f, 0x40, - }); - - const header = - \\{"alg":"ES256","typ":"JWT"} - ; - const payload = - \\{"iss":"did:example:test","aud":"did:example:aud","exp":9999999999} - ; - - const token = try createJwt(allocator, header, payload, &keypair); - defer allocator.free(token); - - // parse and verify with existing JWT infrastructure - var parsed_jwt = try jwt.Jwt.parse(allocator, token); - defer parsed_jwt.deinit(); - - try std.testing.expectEqual(jwt.Algorithm.ES256, parsed_jwt.header.alg); - try std.testing.expectEqualStrings("did:example:test", parsed_jwt.payload.iss); - - // verify signature via multibase key - const pk = try keypair.publicKey(); - const mc_bytes = try multicodec.encodePublicKey(allocator, .p256, &pk); - defer allocator.free(mc_bytes); - const multibase_key = try multibase.encode(allocator, .base58btc, mc_bytes); - defer allocator.free(multibase_key); - - try parsed_jwt.verify(multibase_key); -} - -test "DPoP proof structure" { - const allocator = std.testing.allocator; - const io = std.Options.debug_io; - - const keypair = try Keypair.fromSecretKey(.p256, .{ - 0x21, 0x22, 0x23, 0x24, 0x25, 0x26, 0x27, 0x28, - 0x29, 0x2a, 0x2b, 0x2c, 0x2d, 0x2e, 0x2f, 0x30, - 0x31, 0x32, 0x33, 0x34, 0x35, 0x36, 0x37, 0x38, - 0x39, 0x3a, 0x3b, 0x3c, 0x3d, 0x3e, 0x3f, 0x40, - }); - - const proof = try createDpopProof(allocator, io, &keypair, "POST", "https://auth.example.com/token", "server-nonce", null); - defer allocator.free(proof); - - // decode header - var iter = std.mem.splitScalar(u8, proof, '.'); - const header_b64 = iter.next().?; - const payload_b64 = iter.next().?; - - const header_json = try jwt.base64UrlDecode(allocator, header_b64); - defer allocator.free(header_json); - - const header_parsed = try std.json.parseFromSlice(std.json.Value, allocator, header_json, .{}); - defer header_parsed.deinit(); - - try std.testing.expectEqualStrings("dpop+jwt", header_parsed.value.object.get("typ").?.string); - try std.testing.expectEqualStrings("ES256", header_parsed.value.object.get("alg").?.string); - try std.testing.expect(header_parsed.value.object.get("jwk") != null); - - // decode payload - const payload_json = try jwt.base64UrlDecode(allocator, payload_b64); - defer allocator.free(payload_json); - - const payload_parsed = try std.json.parseFromSlice(std.json.Value, allocator, payload_json, .{}); - defer payload_parsed.deinit(); - - const obj = payload_parsed.value.object; - try std.testing.expect(obj.get("jti") != null); - try std.testing.expectEqualStrings("POST", obj.get("htm").?.string); - try std.testing.expectEqualStrings("https://auth.example.com/token", obj.get("htu").?.string); - try std.testing.expect(obj.get("iat") != null); - try std.testing.expectEqualStrings("server-nonce", obj.get("nonce").?.string); -} - -test "client assertion structure" { - const allocator = std.testing.allocator; - const io = std.Options.debug_io; - - const keypair = try Keypair.fromSecretKey(.p256, .{ - 0x21, 0x22, 0x23, 0x24, 0x25, 0x26, 0x27, 0x28, - 0x29, 0x2a, 0x2b, 0x2c, 0x2d, 0x2e, 0x2f, 0x30, - 0x31, 0x32, 0x33, 0x34, 0x35, 0x36, 0x37, 0x38, - 0x39, 0x3a, 0x3b, 0x3c, 0x3d, 0x3e, 0x3f, 0x40, - }); - - const assertion = try createClientAssertion(allocator, io, &keypair, "https://app.example.com/client-metadata", "https://bsky.social/oauth/token"); - defer allocator.free(assertion); - - // decode header - var iter = std.mem.splitScalar(u8, assertion, '.'); - const header_b64 = iter.next().?; - const payload_b64 = iter.next().?; - - const header_json = try jwt.base64UrlDecode(allocator, header_b64); - defer allocator.free(header_json); - - const header_parsed = try std.json.parseFromSlice(std.json.Value, allocator, header_json, .{}); - defer header_parsed.deinit(); - - try std.testing.expectEqualStrings("JWT", header_parsed.value.object.get("typ").?.string); - try std.testing.expectEqualStrings("ES256", header_parsed.value.object.get("alg").?.string); - try std.testing.expect(header_parsed.value.object.get("kid") != null); - - // decode payload - const payload_json = try jwt.base64UrlDecode(allocator, payload_b64); - defer allocator.free(payload_json); - - const payload_parsed = try std.json.parseFromSlice(std.json.Value, allocator, payload_json, .{}); - defer payload_parsed.deinit(); - - const obj = payload_parsed.value.object; - try std.testing.expectEqualStrings("https://app.example.com/client-metadata", obj.get("iss").?.string); - try std.testing.expectEqualStrings("https://app.example.com/client-metadata", obj.get("sub").?.string); - try std.testing.expectEqualStrings("https://bsky.social/oauth/token", obj.get("aud").?.string); - try std.testing.expect(obj.get("jti") != null); - try std.testing.expect(obj.get("iat") != null); - try std.testing.expect(obj.get("exp") != null); -} - -test "JWKS JSON wraps JWK" { - const allocator = std.testing.allocator; - - const keypair = try Keypair.fromSecretKey(.p256, .{ - 0x21, 0x22, 0x23, 0x24, 0x25, 0x26, 0x27, 0x28, - 0x29, 0x2a, 0x2b, 0x2c, 0x2d, 0x2e, 0x2f, 0x30, - 0x31, 0x32, 0x33, 0x34, 0x35, 0x36, 0x37, 0x38, - 0x39, 0x3a, 0x3b, 0x3c, 0x3d, 0x3e, 0x3f, 0x40, - }); - - const jwks = try jwksJson(allocator, &keypair); - defer allocator.free(jwks); - - const parsed = try std.json.parseFromSlice(std.json.Value, allocator, jwks, .{}); - defer parsed.deinit(); - - const keys = parsed.value.object.get("keys").?.array; - try std.testing.expectEqual(@as(usize, 1), keys.items.len); - try std.testing.expectEqualStrings("EC", keys.items[0].object.get("kty").?.string); -} +//! `primitives` contains PKCE, DPoP, client assertions, and form/JWKS helpers. +//! `client` contains framework-neutral ATProto OAuth client ceremony. + +pub const primitives = @import("oauth/primitives.zig"); +pub const client = @import("oauth/client.zig"); + +pub const createJwt = primitives.createJwt; +pub const createDpopProof = primitives.createDpopProof; +pub const createClientAssertion = primitives.createClientAssertion; +pub const generatePkceVerifier = primitives.generatePkceVerifier; +pub const generatePkceChallenge = primitives.generatePkceChallenge; +pub const generateState = primitives.generateState; +pub const accessTokenHash = primitives.accessTokenHash; +pub const formEncode = primitives.formEncode; +pub const jwksJson = primitives.jwksJson; + +pub const AuthorizationServerMetadata = client.AuthorizationServerMetadata; +pub const AuthRequestSecrets = client.AuthRequestSecrets; +pub const ClientMetadataParams = client.ClientMetadataParams; +pub const ParParams = client.ParParams; +pub const ParResult = client.ParResult; +pub const CodeTokenParams = client.CodeTokenParams; +pub const RefreshTokenParams = client.RefreshTokenParams; +pub const TokenResult = client.TokenResult; +pub const DpopRequest = client.DpopRequest; +pub const DpopResponse = client.DpopResponse; + +pub const prepareAuthRequestSecrets = client.prepareAuthRequestSecrets; +pub const parseTokenResponse = client.parseTokenResponse; +pub const clientMetadataJson = client.clientMetadataJson; +pub const authorizationUrl = client.authorizationUrl; +pub const discoverAuthorizationServer = client.discoverAuthorizationServer; +pub const fetchAuthorizationServerMetadata = client.fetchAuthorizationServerMetadata; +pub const parseAuthorizationServerMetadata = client.parseAuthorizationServerMetadata; +pub const sendParRequest = client.sendParRequest; +pub const exchangeCodeForToken = client.exchangeCodeForToken; +pub const refreshAccessToken = client.refreshAccessToken; +pub const dpopRequest = client.dpopRequest; +pub const isDpopNonceChallenge = client.isDpopNonceChallenge; diff --git a/src/internal/oauth/client.zig b/src/internal/oauth/client.zig new file mode 100644 index 0000000..c861bdb --- /dev/null +++ b/src/internal/oauth/client.zig @@ -0,0 +1,898 @@ +//! Framework-neutral ATProto OAuth client helpers. +//! +//! This module owns protocol ceremony: metadata discovery, client metadata, +//! PAR, token exchange, refresh, DPoP nonce retry, and authenticated resource +//! requests. Applications still own cookies, redirects, sessions, and storage. + +const std = @import("std"); +const Io = std.Io; +const Allocator = std.mem.Allocator; +const Keypair = @import("../crypto/keypair.zig").Keypair; +const zat_json = @import("../xrpc/json.zig"); +const HttpTransport = @import("../xrpc/transport.zig").HttpTransport; +const primitives = @import("primitives.zig"); + +pub const AuthorizationServerMetadata = struct { + issuer: []const u8, + authorization_endpoint: []const u8, + token_endpoint: []const u8, + pushed_authorization_request_endpoint: []const u8, + response_types_supported: []const []const u8 = &.{}, + grant_types_supported: []const []const u8 = &.{}, + code_challenge_methods_supported: []const []const u8 = &.{}, + token_endpoint_auth_methods_supported: []const []const u8 = &.{}, + token_endpoint_auth_signing_alg_values_supported: []const []const u8 = &.{}, + scopes_supported: []const []const u8 = &.{}, + dpop_signing_alg_values_supported: []const []const u8 = &.{}, + authorization_response_iss_parameter_supported: bool = false, + require_pushed_authorization_requests: bool = false, + require_request_uri_registration: bool = true, + client_id_metadata_document_supported: bool = false, + + pub fn deinit(self: *AuthorizationServerMetadata, allocator: Allocator) void { + allocator.free(self.issuer); + allocator.free(self.authorization_endpoint); + allocator.free(self.token_endpoint); + allocator.free(self.pushed_authorization_request_endpoint); + freeStringList(allocator, self.response_types_supported); + freeStringList(allocator, self.grant_types_supported); + freeStringList(allocator, self.code_challenge_methods_supported); + freeStringList(allocator, self.token_endpoint_auth_methods_supported); + freeStringList(allocator, self.token_endpoint_auth_signing_alg_values_supported); + freeStringList(allocator, self.scopes_supported); + freeStringList(allocator, self.dpop_signing_alg_values_supported); + self.* = undefined; + } +}; + +pub const AuthRequestSecrets = struct { + state: []const u8, + pkce_verifier: []const u8, + pkce_challenge: []const u8, + dpop_keypair: Keypair, + + pub fn deinit(self: *AuthRequestSecrets, allocator: Allocator) void { + allocator.free(self.state); + allocator.free(self.pkce_verifier); + allocator.free(self.pkce_challenge); + self.* = undefined; + } +}; + +pub const ClientMetadataParams = struct { + client_id: []const u8, + client_name: []const u8, + client_uri: []const u8, + redirect_uris: []const []const u8, + scope: []const u8, + keypair: *const Keypair, + token_endpoint_auth_method: []const u8 = "private_key_jwt", + token_endpoint_auth_signing_alg: ?[]const u8 = null, + application_type: []const u8 = "web", +}; + +pub const ParParams = struct { + par_url: []const u8, + authserver_issuer: []const u8, + client_id: []const u8, + redirect_uri: []const u8, + scope: []const u8, + state: []const u8, + pkce_challenge: []const u8, + login_hint: ?[]const u8 = null, + client_keypair: *const Keypair, + dpop_keypair: *const Keypair, +}; + +pub const ParResult = struct { + request_uri: []const u8, + dpop_nonce: ?[]const u8 = null, + + pub fn deinit(self: *ParResult, allocator: Allocator) void { + allocator.free(self.request_uri); + if (self.dpop_nonce) |nonce| allocator.free(nonce); + self.* = undefined; + } +}; + +pub const CodeTokenParams = struct { + token_url: []const u8, + authserver_issuer: []const u8, + client_id: []const u8, + redirect_uri: []const u8, + code: []const u8, + pkce_verifier: []const u8, + client_keypair: *const Keypair, + dpop_keypair: *const Keypair, + dpop_nonce: ?[]const u8 = null, +}; + +pub const RefreshTokenParams = struct { + token_url: []const u8, + authserver_issuer: []const u8, + client_id: []const u8, + refresh_token: []const u8, + client_keypair: *const Keypair, + dpop_keypair: *const Keypair, + dpop_nonce: ?[]const u8 = null, +}; + +pub const TokenResult = struct { + access_token: []const u8, + refresh_token: []const u8, + scope: []const u8, + sub: ?[]const u8 = null, + dpop_nonce: ?[]const u8 = null, + + pub fn deinit(self: *TokenResult, allocator: Allocator) void { + allocator.free(self.access_token); + allocator.free(self.refresh_token); + allocator.free(self.scope); + if (self.sub) |sub| allocator.free(sub); + if (self.dpop_nonce) |nonce| allocator.free(nonce); + self.* = undefined; + } +}; + +pub const DpopRequest = struct { + url: []const u8, + method: std.http.Method = .GET, + access_token: []const u8, + dpop_keypair: *const Keypair, + dpop_nonce: ?[]const u8 = null, + payload: ?[]const u8 = null, + content_type: ?[]const u8 = "application/json", + accept: ?[]const u8 = "application/json", + max_response_size: ?usize = null, +}; + +pub const DpopResponse = struct { + status: std.http.Status, + body: []u8, + dpop_nonce: ?[]const u8 = null, + + pub fn deinit(self: *DpopResponse, allocator: Allocator) void { + allocator.free(self.body); + if (self.dpop_nonce) |nonce| allocator.free(nonce); + self.* = undefined; + } +}; + +pub fn parseTokenResponse(allocator: Allocator, value: std.json.Value, dpop_nonce: ?[]const u8) !TokenResult { + const scope = zat_json.getString(value, "scope") orelse return error.MissingScope; + if (!scopeContainsAtproto(scope)) return error.MissingAtprotoScope; + const access_token = try allocator.dupe(u8, zat_json.getString(value, "access_token") orelse return error.MissingAccessToken); + errdefer allocator.free(access_token); + const refresh_token = try allocator.dupe(u8, zat_json.getString(value, "refresh_token") orelse return error.MissingRefreshToken); + errdefer allocator.free(refresh_token); + const scope_copy = try allocator.dupe(u8, scope); + errdefer allocator.free(scope_copy); + const sub_copy = if (zat_json.getString(value, "sub")) |sub| try allocator.dupe(u8, sub) else null; + errdefer if (sub_copy) |sub| allocator.free(sub); + const nonce_copy = if (dpop_nonce) |nonce| try allocator.dupe(u8, nonce) else null; + errdefer if (nonce_copy) |nonce| allocator.free(nonce); + return .{ + .access_token = access_token, + .refresh_token = refresh_token, + .scope = scope_copy, + .sub = sub_copy, + .dpop_nonce = nonce_copy, + }; +} + +pub fn prepareAuthRequestSecrets(allocator: Allocator, io: Io) !AuthRequestSecrets { + const state = try primitives.generateState(allocator, io); + errdefer allocator.free(state); + const pkce_verifier = try primitives.generatePkceVerifier(allocator, io); + errdefer allocator.free(pkce_verifier); + const pkce_challenge = try primitives.generatePkceChallenge(allocator, pkce_verifier); + errdefer allocator.free(pkce_challenge); + var dpop_secret: [32]u8 = undefined; + io.random(&dpop_secret); + return .{ + .state = state, + .pkce_verifier = pkce_verifier, + .pkce_challenge = pkce_challenge, + .dpop_keypair = try Keypair.fromSecretKey(.p256, dpop_secret), + }; +} + +pub fn clientMetadataJson(allocator: Allocator, params: ClientMetadataParams) ![]u8 { + if (!scopeContainsAtproto(params.scope)) return error.MissingAtprotoScope; + const jwk_json = try params.keypair.jwk(allocator); + defer allocator.free(jwk_json); + const signing_alg = params.token_endpoint_auth_signing_alg orelse @tagName(params.keypair.algorithm()); + + var out: std.Io.Writer.Allocating = .init(allocator); + errdefer out.deinit(); + try out.writer.print( + \\{{"client_id":{f},"client_name":{f},"client_uri":{f},"application_type":{f},"grant_types":["authorization_code","refresh_token"],"response_types":["code"],"redirect_uris":[ + , .{ + std.json.fmt(params.client_id, .{}), + std.json.fmt(params.client_name, .{}), + std.json.fmt(params.client_uri, .{}), + std.json.fmt(params.application_type, .{}), + }); + for (params.redirect_uris, 0..) |uri, i| { + if (i > 0) try out.writer.writeAll(","); + try out.writer.print("{f}", .{std.json.fmt(uri, .{})}); + } + try out.writer.print( + \\],"token_endpoint_auth_method":{f},"token_endpoint_auth_signing_alg":{f},"scope":{f},"dpop_bound_access_tokens":true,"jwks":{{"keys":[{s}]}}}} + , .{ + std.json.fmt(params.token_endpoint_auth_method, .{}), + std.json.fmt(signing_alg, .{}), + std.json.fmt(params.scope, .{}), + jwk_json, + }); + return out.toOwnedSlice(); +} + +pub fn authorizationUrl( + allocator: Allocator, + authorization_endpoint: []const u8, + request_uri: []const u8, + client_id: []const u8, + state: []const u8, +) ![]u8 { + const sep: []const u8 = if (std.mem.indexOfScalar(u8, authorization_endpoint, '?') == null) "?" else "&"; + const request_uri_enc = try percentEncodeAlloc(allocator, request_uri); + defer allocator.free(request_uri_enc); + const client_id_enc = try percentEncodeAlloc(allocator, client_id); + defer allocator.free(client_id_enc); + const state_enc = try percentEncodeAlloc(allocator, state); + defer allocator.free(state_enc); + return std.fmt.allocPrint(allocator, "{s}{s}request_uri={s}&client_id={s}&state={s}", .{ + authorization_endpoint, + sep, + request_uri_enc, + client_id_enc, + state_enc, + }); +} + +pub fn discoverAuthorizationServer( + allocator: Allocator, + transport: *HttpTransport, + pds_url: []const u8, +) ![]const u8 { + const url = try joinUrl(allocator, pds_url, "/.well-known/oauth-protected-resource"); + defer allocator.free(url); + var result = try transport.fetch(.{ + .url = url, + .method = .GET, + .accept = "application/json", + .max_response_size = 256 * 1024, + .redirect_behavior = .unhandled, + .capture_response_headers = true, + }); + defer result.deinit(allocator); + if (result.status != .ok) return error.HttpStatus; + if (!contentTypeIsJson(result.oauth.content_type)) return error.InvalidContentType; + const parsed = try std.json.parseFromSlice(std.json.Value, allocator, result.body, .{}); + defer parsed.deinit(); + const servers = zat_json.getArray(parsed.value, "authorization_servers") orelse return error.NoAuthorizationServers; + if (servers.len != 1 or servers[0] != .string) return error.NoAuthorizationServers; + try validateSimpleHttpsOrigin(servers[0].string); + return allocator.dupe(u8, servers[0].string); +} + +pub fn fetchAuthorizationServerMetadata( + allocator: Allocator, + transport: *HttpTransport, + issuer: []const u8, +) !AuthorizationServerMetadata { + const url = try joinUrl(allocator, issuer, "/.well-known/oauth-authorization-server"); + defer allocator.free(url); + var result = try transport.fetch(.{ + .url = url, + .method = .GET, + .accept = "application/json", + .max_response_size = 256 * 1024, + .redirect_behavior = .unhandled, + .capture_response_headers = true, + }); + defer result.deinit(allocator); + if (result.status != .ok) return error.HttpStatus; + if (!contentTypeIsJson(result.oauth.content_type)) return error.InvalidContentType; + const parsed = try std.json.parseFromSlice(std.json.Value, allocator, result.body, .{}); + defer parsed.deinit(); + var metadata = try parseAuthorizationServerMetadata(allocator, parsed.value); + errdefer metadata.deinit(allocator); + const expected_issuer = try originFromUrl(allocator, issuer); + defer allocator.free(expected_issuer); + if (!std.mem.eql(u8, metadata.issuer, expected_issuer)) return error.IssuerMismatch; + return metadata; +} + +pub fn parseAuthorizationServerMetadata(allocator: Allocator, value: std.json.Value) !AuthorizationServerMetadata { + var metadata = AuthorizationServerMetadata{ + .issuer = try allocator.dupe(u8, zat_json.getString(value, "issuer") orelse return error.MissingIssuer), + .authorization_endpoint = "", + .token_endpoint = "", + .pushed_authorization_request_endpoint = "", + }; + errdefer metadata.deinit(allocator); + try validateSimpleHttpsOrigin(metadata.issuer); + + metadata.authorization_endpoint = try allocator.dupe(u8, zat_json.getString(value, "authorization_endpoint") orelse return error.MissingAuthorizationEndpoint); + metadata.token_endpoint = try allocator.dupe(u8, zat_json.getString(value, "token_endpoint") orelse return error.MissingTokenEndpoint); + metadata.pushed_authorization_request_endpoint = try allocator.dupe(u8, zat_json.getString(value, "pushed_authorization_request_endpoint") orelse return error.MissingParEndpoint); + + metadata.response_types_supported = try parseRequiredStringArray(allocator, value, "response_types_supported"); + metadata.grant_types_supported = try parseRequiredStringArray(allocator, value, "grant_types_supported"); + metadata.code_challenge_methods_supported = try parseRequiredStringArray(allocator, value, "code_challenge_methods_supported"); + metadata.token_endpoint_auth_methods_supported = try parseRequiredStringArray(allocator, value, "token_endpoint_auth_methods_supported"); + metadata.token_endpoint_auth_signing_alg_values_supported = try parseRequiredStringArray(allocator, value, "token_endpoint_auth_signing_alg_values_supported"); + metadata.scopes_supported = try parseRequiredStringArray(allocator, value, "scopes_supported"); + metadata.dpop_signing_alg_values_supported = try parseRequiredStringArray(allocator, value, "dpop_signing_alg_values_supported"); + metadata.authorization_response_iss_parameter_supported = zat_json.getBool(value, "authorization_response_iss_parameter_supported") orelse false; + metadata.require_pushed_authorization_requests = zat_json.getBool(value, "require_pushed_authorization_requests") orelse false; + metadata.require_request_uri_registration = zat_json.getBool(value, "require_request_uri_registration") orelse true; + metadata.client_id_metadata_document_supported = zat_json.getBool(value, "client_id_metadata_document_supported") orelse false; + try validateAuthorizationServerMetadata(metadata); + return metadata; +} + +pub fn sendParRequest( + allocator: Allocator, + io: Io, + transport: *HttpTransport, + params: ParParams, +) !ParResult { + if (!scopeContainsAtproto(params.scope)) return error.MissingAtprotoScope; + const client_assertion = try primitives.createClientAssertion(allocator, io, params.client_keypair, params.client_id, params.authserver_issuer); + defer allocator.free(client_assertion); + + var form_params: std.ArrayList([2][]const u8) = .empty; + defer form_params.deinit(allocator); + try form_params.appendSlice(allocator, &.{ + .{ "response_type", "code" }, + .{ "code_challenge", params.pkce_challenge }, + .{ "code_challenge_method", "S256" }, + .{ "redirect_uri", params.redirect_uri }, + .{ "scope", params.scope }, + .{ "state", params.state }, + .{ "client_id", params.client_id }, + .{ "client_assertion_type", "urn:ietf:params:oauth:client-assertion-type:jwt-bearer" }, + .{ "client_assertion", client_assertion }, + }); + if (params.login_hint) |hint| { + try form_params.append(allocator, .{ "login_hint", hint }); + } + + const body = try primitives.formEncode(allocator, form_params.items); + defer allocator.free(body); + + const result = try fetchWithDpopNonceRetry(allocator, io, transport, .{ + .url = params.par_url, + .method = .POST, + .payload = body, + .content_type = "application/x-www-form-urlencoded", + .dpop_keypair = params.dpop_keypair, + }); + defer result.deinit(allocator); + if (result.status != .ok and result.status != .created) return error.ParFailed; + + const parsed = try std.json.parseFromSlice(std.json.Value, allocator, result.body, .{}); + defer parsed.deinit(); + return .{ + .request_uri = try allocator.dupe(u8, zat_json.getString(parsed.value, "request_uri") orelse return error.MissingRequestUri), + .dpop_nonce = if (result.dpop_nonce) |nonce| try allocator.dupe(u8, nonce) else null, + }; +} + +pub fn exchangeCodeForToken( + allocator: Allocator, + io: Io, + transport: *HttpTransport, + params: CodeTokenParams, +) !TokenResult { + const client_assertion = try primitives.createClientAssertion(allocator, io, params.client_keypair, params.client_id, params.authserver_issuer); + defer allocator.free(client_assertion); + const form_params = [_][2][]const u8{ + .{ "grant_type", "authorization_code" }, + .{ "code", params.code }, + .{ "redirect_uri", params.redirect_uri }, + .{ "code_verifier", params.pkce_verifier }, + .{ "client_id", params.client_id }, + .{ "client_assertion_type", "urn:ietf:params:oauth:client-assertion-type:jwt-bearer" }, + .{ "client_assertion", client_assertion }, + }; + return tokenRequest(allocator, io, transport, params.token_url, params.dpop_keypair, params.dpop_nonce, &form_params); +} + +pub fn refreshAccessToken( + allocator: Allocator, + io: Io, + transport: *HttpTransport, + params: RefreshTokenParams, +) !TokenResult { + const client_assertion = try primitives.createClientAssertion(allocator, io, params.client_keypair, params.client_id, params.authserver_issuer); + defer allocator.free(client_assertion); + const form_params = [_][2][]const u8{ + .{ "grant_type", "refresh_token" }, + .{ "refresh_token", params.refresh_token }, + .{ "client_id", params.client_id }, + .{ "client_assertion_type", "urn:ietf:params:oauth:client-assertion-type:jwt-bearer" }, + .{ "client_assertion", client_assertion }, + }; + return tokenRequest(allocator, io, transport, params.token_url, params.dpop_keypair, params.dpop_nonce, &form_params); +} + +pub fn dpopRequest( + allocator: Allocator, + io: Io, + transport: *HttpTransport, + params: DpopRequest, +) !DpopResponse { + const ath = try primitives.accessTokenHash(allocator, params.access_token); + defer allocator.free(ath); + return fetchWithDpopNonceRetry(allocator, io, transport, .{ + .url = params.url, + .method = params.method, + .payload = params.payload, + .content_type = params.content_type, + .accept = params.accept, + .dpop_keypair = params.dpop_keypair, + .dpop_nonce = params.dpop_nonce, + .access_token = params.access_token, + .access_token_hash = ath, + .max_response_size = params.max_response_size, + }); +} + +pub fn isDpopNonceChallenge(status: std.http.Status, body: []const u8, www_authenticate: ?[]const u8) bool { + if (status == .bad_request and std.mem.indexOf(u8, body, "use_dpop_nonce") != null) return true; + if (status == .unauthorized and std.mem.indexOf(u8, body, "use_dpop_nonce") != null) return true; + if (status == .unauthorized) { + const header = www_authenticate orelse return false; + return std.mem.indexOf(u8, header, "use_dpop_nonce") != null; + } + return false; +} + +const DpopFetchParams = struct { + url: []const u8, + method: std.http.Method, + payload: ?[]const u8 = null, + content_type: ?[]const u8 = null, + accept: ?[]const u8 = "application/json", + dpop_keypair: *const Keypair, + dpop_nonce: ?[]const u8 = null, + access_token: ?[]const u8 = null, + access_token_hash: ?[]const u8 = null, + max_response_size: ?usize = 256 * 1024, +}; + +fn tokenRequest( + allocator: Allocator, + io: Io, + transport: *HttpTransport, + token_url: []const u8, + dpop_keypair: *const Keypair, + dpop_nonce: ?[]const u8, + form_params: []const [2][]const u8, +) !TokenResult { + const body = try primitives.formEncode(allocator, form_params); + defer allocator.free(body); + var result = try fetchWithDpopNonceRetry(allocator, io, transport, .{ + .url = token_url, + .method = .POST, + .payload = body, + .content_type = "application/x-www-form-urlencoded", + .dpop_keypair = dpop_keypair, + .dpop_nonce = dpop_nonce, + }); + defer result.deinit(allocator); + if (result.status != .ok) return error.TokenRequestFailed; + const parsed = try std.json.parseFromSlice(std.json.Value, allocator, result.body, .{}); + defer parsed.deinit(); + return parseTokenResponse(allocator, parsed.value, result.dpop_nonce); +} + +fn fetchWithDpopNonceRetry( + allocator: Allocator, + io: Io, + transport: *HttpTransport, + params: DpopFetchParams, +) !DpopResponse { + var nonce = params.dpop_nonce; + var retry_nonce: ?[]const u8 = null; + defer if (retry_nonce) |value| allocator.free(value); + + for (0..2) |_| { + const htu = try dpopHtu(allocator, params.url); + defer allocator.free(htu); + const proof = try primitives.createDpopProof( + allocator, + io, + params.dpop_keypair, + methodString(params.method), + htu, + nonce, + params.access_token_hash, + ); + defer allocator.free(proof); + + var auth_buf: [4096]u8 = undefined; + const auth_header = if (params.access_token) |token| + try std.fmt.bufPrint(&auth_buf, "DPoP {s}", .{token}) + else + null; + + const extra = [_]std.http.Header{.{ .name = "DPoP", .value = proof }}; + var fetch_result = try transport.fetch(.{ + .url = params.url, + .method = params.method, + .payload = params.payload, + .authorization = auth_header, + .accept = params.accept, + .content_type = params.content_type, + .extra_headers = &extra, + .max_response_size = params.max_response_size, + .capture_response_headers = true, + }); + + const new_nonce = if (fetch_result.oauth.dpop_nonce) |value| try allocator.dupe(u8, value) else null; + if (new_nonce != null and isDpopNonceChallenge(fetch_result.status, fetch_result.body, fetch_result.oauth.www_authenticate)) { + fetch_result.deinit(allocator); + if (retry_nonce) |value| allocator.free(value); + retry_nonce = new_nonce.?; + nonce = retry_nonce; + continue; + } + + if (new_nonce == null and retry_nonce == null) { + fetch_result.deinit(allocator); + return error.MissingDpopNonce; + } + + const body = fetch_result.body; + fetch_result.body = &.{}; + const returned_nonce = if (new_nonce) |value| + value + else if (retry_nonce) |value| + try allocator.dupe(u8, value) + else + null; + fetch_result.oauth.deinit(allocator); + return .{ + .status = fetch_result.status, + .body = body, + .dpop_nonce = returned_nonce, + }; + } + return error.DpopNonceRetryExhausted; +} + +pub fn validateAuthorizationServerMetadata(metadata: AuthorizationServerMetadata) !void { + try validateSimpleHttpsOrigin(metadata.issuer); + if (!containsString(metadata.response_types_supported, "code")) return error.InvalidAuthorizationServerMetadata; + if (!containsString(metadata.grant_types_supported, "authorization_code")) return error.InvalidAuthorizationServerMetadata; + if (!containsString(metadata.grant_types_supported, "refresh_token")) return error.InvalidAuthorizationServerMetadata; + if (!containsString(metadata.code_challenge_methods_supported, "S256")) return error.InvalidAuthorizationServerMetadata; + if (!containsString(metadata.token_endpoint_auth_methods_supported, "none")) return error.InvalidAuthorizationServerMetadata; + if (!containsString(metadata.token_endpoint_auth_methods_supported, "private_key_jwt")) return error.InvalidAuthorizationServerMetadata; + if (containsString(metadata.token_endpoint_auth_signing_alg_values_supported, "none")) return error.InvalidAuthorizationServerMetadata; + if (!containsString(metadata.token_endpoint_auth_signing_alg_values_supported, "ES256")) return error.InvalidAuthorizationServerMetadata; + if (!containsString(metadata.scopes_supported, "atproto")) return error.InvalidAuthorizationServerMetadata; + if (!containsString(metadata.dpop_signing_alg_values_supported, "ES256")) return error.InvalidAuthorizationServerMetadata; + if (!metadata.authorization_response_iss_parameter_supported) return error.InvalidAuthorizationServerMetadata; + if (!metadata.require_pushed_authorization_requests) return error.InvalidAuthorizationServerMetadata; + if (!metadata.require_request_uri_registration) return error.InvalidAuthorizationServerMetadata; + if (!metadata.client_id_metadata_document_supported) return error.InvalidAuthorizationServerMetadata; +} + +fn joinUrl(allocator: Allocator, base: []const u8, path: []const u8) ![]u8 { + if (base.len > 0 and base[base.len - 1] == '/' and path.len > 0 and path[0] == '/') { + return std.fmt.allocPrint(allocator, "{s}{s}", .{ base[0 .. base.len - 1], path }); + } + if (base.len > 0 and base[base.len - 1] != '/' and path.len > 0 and path[0] != '/') { + return std.fmt.allocPrint(allocator, "{s}/{s}", .{ base, path }); + } + return std.fmt.allocPrint(allocator, "{s}{s}", .{ base, path }); +} + +fn dpopHtu(allocator: Allocator, url: []const u8) ![]u8 { + const parsed = try std.Uri.parse(url); + const scheme = parsed.scheme; + const host = parsed.host orelse return error.InvalidDpopHtu; + const path = parsed.path.percent_encoded; + const port = parsed.port; + if (port) |p| { + return std.fmt.allocPrint(allocator, "{s}://{s}:{d}{s}", .{ scheme, host, p, if (path.len == 0) "/" else path }); + } + return std.fmt.allocPrint(allocator, "{s}://{s}{s}", .{ scheme, host, if (path.len == 0) "/" else path }); +} + +fn originFromUrl(allocator: Allocator, url: []const u8) ![]const u8 { + const parsed = try std.Uri.parse(url); + const scheme = parsed.scheme; + const host = parsed.host orelse return error.InvalidIssuer; + if (parsed.port) |port| { + return std.fmt.allocPrint(allocator, "{s}://{s}:{d}", .{ scheme, host, port }); + } + return std.fmt.allocPrint(allocator, "{s}://{s}", .{ scheme, host }); +} + +fn validateSimpleHttpsOrigin(url: []const u8) !void { + const parsed = try std.Uri.parse(url); + if (!std.mem.eql(u8, parsed.scheme, "https")) return error.InvalidIssuer; + if (parsed.user != null or parsed.password != null) return error.InvalidIssuer; + _ = parsed.host orelse return error.InvalidIssuer; + if (parsed.path.percent_encoded.len != 0 and !std.mem.eql(u8, parsed.path.percent_encoded, "/")) return error.InvalidIssuer; + if (parsed.query != null or parsed.fragment != null) return error.InvalidIssuer; + if (parsed.port == 443) return error.InvalidIssuer; +} + +fn parseRequiredStringArray(allocator: Allocator, value: std.json.Value, path: []const u8) ![]const []const u8 { + const items = zat_json.getArray(value, path) orelse return error.MissingMetadataField; + var out = try allocator.alloc([]const u8, items.len); + errdefer allocator.free(out); + var filled: usize = 0; + errdefer { + for (out[0..filled]) |item| allocator.free(item); + } + for (items) |item| { + if (item != .string) return error.InvalidAuthorizationServerMetadata; + out[filled] = try allocator.dupe(u8, item.string); + filled += 1; + } + return out; +} + +fn freeStringList(allocator: Allocator, items: []const []const u8) void { + if (items.len == 0) return; + for (items) |item| allocator.free(item); + allocator.free(items); +} + +fn contentTypeIsJson(content_type: ?[]const u8) bool { + const value = content_type orelse return false; + var it = std.mem.splitScalar(u8, value, ';'); + const media_type = std.mem.trim(u8, it.next() orelse value, " \t\r\n"); + return std.ascii.eqlIgnoreCase(media_type, "application/json"); +} + +fn containsString(items: []const []const u8, needle: []const u8) bool { + for (items) |item| { + if (std.mem.eql(u8, item, needle)) return true; + } + return false; +} + +fn scopeContainsAtproto(scope: []const u8) bool { + var it = std.mem.splitScalar(u8, scope, ' '); + while (it.next()) |item| { + if (std.mem.eql(u8, item, "atproto")) return true; + } + return false; +} + +fn methodString(method: std.http.Method) []const u8 { + return @tagName(method); +} + +fn percentEncodeAlloc(allocator: Allocator, input: []const u8) ![]u8 { + var out: std.Io.Writer.Allocating = .init(allocator); + errdefer out.deinit(); + try primitives.percentEncode(&out.writer, input); + return out.toOwnedSlice(); +} + +test "parse authorization server metadata" { + const allocator = std.testing.allocator; + const json_str = + \\{ + \\ "issuer": "https://auth.example.com", + \\ "authorization_endpoint": "https://auth.example.com/oauth/authorize", + \\ "token_endpoint": "https://auth.example.com/oauth/token", + \\ "pushed_authorization_request_endpoint": "https://auth.example.com/oauth/par", + \\ "response_types_supported": ["code"], + \\ "grant_types_supported": ["authorization_code", "refresh_token"], + \\ "code_challenge_methods_supported": ["S256"], + \\ "token_endpoint_auth_methods_supported": ["none", "private_key_jwt"], + \\ "token_endpoint_auth_signing_alg_values_supported": ["ES256"], + \\ "scopes_supported": ["atproto", "repo:*"], + \\ "dpop_signing_alg_values_supported": ["ES256"], + \\ "authorization_response_iss_parameter_supported": true, + \\ "require_pushed_authorization_requests": true, + \\ "client_id_metadata_document_supported": true + \\} + ; + const parsed = try std.json.parseFromSlice(std.json.Value, allocator, json_str, .{}); + defer parsed.deinit(); + + var metadata = try parseAuthorizationServerMetadata(allocator, parsed.value); + defer metadata.deinit(allocator); + + try std.testing.expectEqualStrings("https://auth.example.com", metadata.issuer); + try std.testing.expectEqualStrings("https://auth.example.com/oauth/authorize", metadata.authorization_endpoint); + try std.testing.expectEqualStrings("https://auth.example.com/oauth/token", metadata.token_endpoint); + try std.testing.expectEqualStrings("https://auth.example.com/oauth/par", metadata.pushed_authorization_request_endpoint); + try std.testing.expect(containsString(metadata.scopes_supported, "atproto")); +} + +test "authorization server metadata rejects invalid issuer" { + const allocator = std.testing.allocator; + const json_str = + \\{ + \\ "issuer": "https://auth.example.com/oauth", + \\ "authorization_endpoint": "https://auth.example.com/oauth/authorize", + \\ "token_endpoint": "https://auth.example.com/oauth/token", + \\ "pushed_authorization_request_endpoint": "https://auth.example.com/oauth/par", + \\ "response_types_supported": ["code"], + \\ "grant_types_supported": ["authorization_code", "refresh_token"], + \\ "code_challenge_methods_supported": ["S256"], + \\ "token_endpoint_auth_methods_supported": ["none", "private_key_jwt"], + \\ "token_endpoint_auth_signing_alg_values_supported": ["ES256"], + \\ "scopes_supported": ["atproto"], + \\ "dpop_signing_alg_values_supported": ["ES256"], + \\ "authorization_response_iss_parameter_supported": true, + \\ "require_pushed_authorization_requests": true, + \\ "client_id_metadata_document_supported": true + \\} + ; + const parsed = try std.json.parseFromSlice(std.json.Value, allocator, json_str, .{}); + defer parsed.deinit(); + + try std.testing.expectError(error.InvalidIssuer, parseAuthorizationServerMetadata(allocator, parsed.value)); +} + +test "authorization server metadata requires atproto support" { + const allocator = std.testing.allocator; + const json_str = + \\{ + \\ "issuer": "https://auth.example.com", + \\ "authorization_endpoint": "https://auth.example.com/oauth/authorize", + \\ "token_endpoint": "https://auth.example.com/oauth/token", + \\ "pushed_authorization_request_endpoint": "https://auth.example.com/oauth/par", + \\ "response_types_supported": ["code"], + \\ "grant_types_supported": ["authorization_code", "refresh_token"], + \\ "code_challenge_methods_supported": ["S256"], + \\ "token_endpoint_auth_methods_supported": ["none", "private_key_jwt"], + \\ "token_endpoint_auth_signing_alg_values_supported": ["ES256"], + \\ "scopes_supported": ["repo:*"], + \\ "dpop_signing_alg_values_supported": ["ES256"], + \\ "authorization_response_iss_parameter_supported": true, + \\ "require_pushed_authorization_requests": true, + \\ "client_id_metadata_document_supported": true + \\} + ; + const parsed = try std.json.parseFromSlice(std.json.Value, allocator, json_str, .{}); + defer parsed.deinit(); + + try std.testing.expectError(error.InvalidAuthorizationServerMetadata, parseAuthorizationServerMetadata(allocator, parsed.value)); +} + +test "client metadata JSON" { + const allocator = std.testing.allocator; + const keypair = try Keypair.fromSecretKey(.p256, .{ + 0x21, 0x22, 0x23, 0x24, 0x25, 0x26, 0x27, 0x28, + 0x29, 0x2a, 0x2b, 0x2c, 0x2d, 0x2e, 0x2f, 0x30, + 0x31, 0x32, 0x33, 0x34, 0x35, 0x36, 0x37, 0x38, + 0x39, 0x3a, 0x3b, 0x3c, 0x3d, 0x3e, 0x3f, 0x40, + }); + const redirects = [_][]const u8{"https://app.example.com/oauth/callback"}; + const metadata = try clientMetadataJson(allocator, .{ + .client_id = "https://app.example.com/oauth-client-metadata.json", + .client_name = "example app", + .client_uri = "https://app.example.com", + .redirect_uris = &redirects, + .scope = "atproto repo:example.app.record", + .keypair = &keypair, + }); + defer allocator.free(metadata); + + const parsed = try std.json.parseFromSlice(std.json.Value, allocator, metadata, .{}); + defer parsed.deinit(); + try std.testing.expectEqualStrings("https://app.example.com/oauth-client-metadata.json", zat_json.getString(parsed.value, "client_id").?); + try std.testing.expectEqualStrings("private_key_jwt", zat_json.getString(parsed.value, "token_endpoint_auth_method").?); + try std.testing.expectEqualStrings("atproto repo:example.app.record", zat_json.getString(parsed.value, "scope").?); + try std.testing.expectEqualStrings("https://app.example.com/oauth/callback", zat_json.getArray(parsed.value, "redirect_uris").?[0].string); + try std.testing.expectEqual(@as(usize, 1), zat_json.getArray(parsed.value, "jwks.keys").?.len); +} + +test "client metadata requires atproto scope" { + const allocator = std.testing.allocator; + const keypair = try Keypair.fromSecretKey(.p256, .{ + 0x21, 0x22, 0x23, 0x24, 0x25, 0x26, 0x27, 0x28, + 0x29, 0x2a, 0x2b, 0x2c, 0x2d, 0x2e, 0x2f, 0x30, + 0x31, 0x32, 0x33, 0x34, 0x35, 0x36, 0x37, 0x38, + 0x39, 0x3a, 0x3b, 0x3c, 0x3d, 0x3e, 0x3f, 0x40, + }); + const redirects = [_][]const u8{"https://app.example.com/oauth/callback"}; + try std.testing.expectError(error.MissingAtprotoScope, clientMetadataJson(allocator, .{ + .client_id = "https://app.example.com/oauth-client-metadata.json", + .client_name = "example app", + .client_uri = "https://app.example.com", + .redirect_uris = &redirects, + .scope = "repo:example.app.record", + .keypair = &keypair, + })); +} + +test "authorization URL percent-encodes parameters" { + const allocator = std.testing.allocator; + const url = try authorizationUrl( + allocator, + "https://auth.example.com/oauth/authorize", + "urn:ietf:params:oauth:request_uri:abc/123", + "https://app.example.com/oauth-client-metadata.json", + "state value", + ); + defer allocator.free(url); + + try std.testing.expectEqualStrings( + "https://auth.example.com/oauth/authorize?request_uri=urn%3Aietf%3Aparams%3Aoauth%3Arequest_uri%3Aabc%2F123&client_id=https%3A%2F%2Fapp.example.com%2Foauth-client-metadata.json&state=state%20value", + url, + ); +} + +test "DPoP nonce challenge detection" { + try std.testing.expect(isDpopNonceChallenge(.bad_request, "{\"error\":\"use_dpop_nonce\"}", null)); + try std.testing.expect(isDpopNonceChallenge(.unauthorized, "{}", "DPoP error=\"use_dpop_nonce\"")); + try std.testing.expect(!isDpopNonceChallenge(.ok, "{\"error\":\"use_dpop_nonce\"}", null)); +} + +test "DPoP htu omits query and fragment" { + const allocator = std.testing.allocator; + const htu = try dpopHtu(allocator, "https://pds.example.com/xrpc/com.atproto.repo.getRecord?repo=did%3Aplc%3Aabc#frag"); + defer allocator.free(htu); + try std.testing.expectEqualStrings("https://pds.example.com/xrpc/com.atproto.repo.getRecord", htu); +} + +test "DPoP htu preserves non-default port and path" { + const allocator = std.testing.allocator; + const htu = try dpopHtu(allocator, "https://pds.example.com:8443/xrpc/app.bsky.actor.getProfile?actor=alice.test"); + defer allocator.free(htu); + try std.testing.expectEqualStrings("https://pds.example.com:8443/xrpc/app.bsky.actor.getProfile", htu); +} + +test "OAuth metadata content type must be JSON" { + try std.testing.expect(contentTypeIsJson("application/json")); + try std.testing.expect(contentTypeIsJson("application/json; charset=utf-8")); + try std.testing.expect(!contentTypeIsJson("text/json")); + try std.testing.expect(!contentTypeIsJson(null)); +} + +test "token response requires atproto scope" { + const allocator = std.testing.allocator; + const json_str = + \\{"access_token":"access","refresh_token":"refresh","scope":"repo:example.app.record","sub":"did:plc:test"} + ; + const parsed = try std.json.parseFromSlice(std.json.Value, allocator, json_str, .{}); + defer parsed.deinit(); + + try std.testing.expectError(error.MissingAtprotoScope, parseTokenResponse(allocator, parsed.value, "nonce")); +} + +test "token response parses scope and subject" { + const allocator = std.testing.allocator; + const json_str = + \\{"access_token":"access","refresh_token":"refresh","scope":"atproto repo:example.app.record","sub":"did:plc:test"} + ; + const parsed = try std.json.parseFromSlice(std.json.Value, allocator, json_str, .{}); + defer parsed.deinit(); + + var token = try parseTokenResponse(allocator, parsed.value, "nonce"); + defer token.deinit(allocator); + + try std.testing.expectEqualStrings("access", token.access_token); + try std.testing.expectEqualStrings("refresh", token.refresh_token); + try std.testing.expectEqualStrings("atproto repo:example.app.record", token.scope); + try std.testing.expectEqualStrings("did:plc:test", token.sub.?); + try std.testing.expectEqualStrings("nonce", token.dpop_nonce.?); +} + +test "prepare auth request secrets" { + const allocator = std.testing.allocator; + var secrets = try prepareAuthRequestSecrets(allocator, std.Options.debug_io); + defer secrets.deinit(allocator); + + try std.testing.expectEqual(@as(usize, 22), secrets.state.len); + try std.testing.expectEqual(@as(usize, 43), secrets.pkce_verifier.len); + try std.testing.expectEqual(@as(usize, 43), secrets.pkce_challenge.len); + _ = try secrets.dpop_keypair.publicKey(); +} diff --git a/src/internal/oauth/primitives.zig b/src/internal/oauth/primitives.zig new file mode 100644 index 0000000..c8da004 --- /dev/null +++ b/src/internal/oauth/primitives.zig @@ -0,0 +1,370 @@ +//! OAuth client primitives for AT Protocol. +//! +//! PKCE, DPoP proofs, client assertions, form encoding, and related helpers +//! for implementing AT Protocol OAuth flows. + +const std = @import("std"); +const crypto = std.crypto; +const Io = std.Io; +const Allocator = std.mem.Allocator; +const Keypair = @import("../crypto/keypair.zig").Keypair; +const jwt = @import("../crypto/jwt.zig"); + +fn timestamp(io: Io) i64 { + return @intCast(@divFloor(Io.Timestamp.now(io, .real).nanoseconds, std.time.ns_per_s)); +} + +/// create a signed JWT from header and payload JSON strings. +/// caller owns returned slice. +pub fn createJwt(allocator: Allocator, header_json: []const u8, payload_json: []const u8, keypair: *const Keypair) ![]u8 { + const header_b64 = try jwt.base64UrlEncode(allocator, header_json); + defer allocator.free(header_b64); + + const payload_b64 = try jwt.base64UrlEncode(allocator, payload_json); + defer allocator.free(payload_b64); + + const signing_input = try std.fmt.allocPrint(allocator, "{s}.{s}", .{ header_b64, payload_b64 }); + defer allocator.free(signing_input); + + const sig = try keypair.sign(signing_input); + const sig_b64 = try jwt.base64UrlEncode(allocator, &sig.bytes); + defer allocator.free(sig_b64); + + return std.fmt.allocPrint(allocator, "{s}.{s}", .{ signing_input, sig_b64 }); +} + +/// create a DPoP proof JWT per RFC 9449. +/// htm: HTTP method, htu: target URI, nonce: server-provided DPoP-Nonce, +/// ath: optional access token hash (base64url-encoded SHA-256). +pub fn createDpopProof( + allocator: Allocator, + io: std.Io, + keypair: *const Keypair, + htm: []const u8, + htu: []const u8, + nonce: ?[]const u8, + ath: ?[]const u8, +) ![]u8 { + const jwk_json = try keypair.jwk(allocator); + defer allocator.free(jwk_json); + + const jti = try generateJti(allocator, io); + defer allocator.free(jti); + + const alg = @tagName(keypair.algorithm()); + const now = timestamp(io); + + const header = try std.fmt.allocPrint(allocator, + \\{{"typ":"dpop+jwt","alg":"{s}","jwk":{s}}} + , .{ alg, jwk_json }); + defer allocator.free(header); + + var aw: std.Io.Writer.Allocating = .init(allocator); + defer aw.deinit(); + + try aw.writer.print( + \\{{"jti":"{s}","htm":"{s}","htu":"{s}","iat":{d} + , .{ jti, htm, htu, now }); + + if (nonce) |n| { + try aw.writer.print(",\"nonce\":\"{s}\"", .{n}); + } + if (ath) |a| { + try aw.writer.print(",\"ath\":\"{s}\"", .{a}); + } + + try aw.writer.writeAll("}"); + + return createJwt(allocator, header, aw.written(), keypair); +} + +/// create a private_key_jwt client assertion for token endpoint auth. +/// client_id: the OAuth client ID, aud: the token endpoint URL. +pub fn createClientAssertion( + allocator: Allocator, + io: std.Io, + keypair: *const Keypair, + client_id: []const u8, + aud: []const u8, +) ![]u8 { + const jti = try generateJti(allocator, io); + defer allocator.free(jti); + + const kid = try keypair.jwkThumbprint(allocator); + defer allocator.free(kid); + + const alg = @tagName(keypair.algorithm()); + const now = timestamp(io); + + const header = try std.fmt.allocPrint(allocator, + \\{{"typ":"JWT","alg":"{s}","kid":"{s}"}} + , .{ alg, kid }); + defer allocator.free(header); + + const payload = try std.fmt.allocPrint(allocator, + \\{{"iss":"{s}","sub":"{s}","aud":"{s}","jti":"{s}","iat":{d},"exp":{d}}} + , .{ client_id, client_id, aud, jti, now, now + 120 }); + defer allocator.free(payload); + + return createJwt(allocator, header, payload, keypair); +} + +/// generate a random PKCE code verifier (43 chars, base64url-encoded 32 random bytes). +/// caller owns returned slice. +pub fn generatePkceVerifier(allocator: Allocator, io: std.Io) ![]u8 { + var random_bytes: [32]u8 = undefined; + io.random(&random_bytes); + return jwt.base64UrlEncode(allocator, &random_bytes); +} + +/// generate a PKCE S256 challenge from a verifier. +/// caller owns returned slice. +pub fn generatePkceChallenge(allocator: Allocator, verifier: []const u8) ![]u8 { + var hash: [32]u8 = undefined; + crypto.hash.sha2.Sha256.hash(verifier, &hash, .{}); + return jwt.base64UrlEncode(allocator, &hash); +} + +/// generate a random state parameter (CSRF token). +/// caller owns returned slice. +pub fn generateState(allocator: Allocator, io: std.Io) ![]u8 { + var random_bytes: [16]u8 = undefined; + io.random(&random_bytes); + return jwt.base64UrlEncode(allocator, &random_bytes); +} + +/// compute access token hash for DPoP ath claim: base64url(SHA-256(access_token)). +/// caller owns returned slice. +pub fn accessTokenHash(allocator: Allocator, access_token: []const u8) ![]u8 { + var hash: [32]u8 = undefined; + crypto.hash.sha2.Sha256.hash(access_token, &hash, .{}); + return jwt.base64UrlEncode(allocator, &hash); +} + +/// encode key-value pairs as application/x-www-form-urlencoded. +/// caller owns returned slice. +pub fn formEncode(allocator: Allocator, params: []const [2][]const u8) ![]u8 { + var aw: std.Io.Writer.Allocating = .init(allocator); + errdefer aw.deinit(); + + for (params, 0..) |kv, i| { + if (i > 0) try aw.writer.writeAll("&"); + try percentEncode(&aw.writer, kv[0]); + try aw.writer.writeAll("="); + try percentEncode(&aw.writer, kv[1]); + } + + return try aw.toOwnedSlice(); +} + +/// format a JWKS JSON containing a single public key. +/// caller owns returned slice. +pub fn jwksJson(allocator: Allocator, keypair: *const Keypair) ![]u8 { + const jwk_json = try keypair.jwk(allocator); + defer allocator.free(jwk_json); + + return std.fmt.allocPrint(allocator, + \\{{"keys":[{s}]}} + , .{jwk_json}); +} + +pub fn percentEncode(writer: anytype, input: []const u8) !void { + for (input) |c| { + if (isUnreserved(c)) { + try writer.writeByte(c); + } else { + try writer.print("%{X:0>2}", .{c}); + } + } +} + +fn generateJti(allocator: Allocator, io: std.Io) ![]u8 { + var random_bytes: [16]u8 = undefined; + io.random(&random_bytes); + return jwt.base64UrlEncode(allocator, &random_bytes); +} + +fn isUnreserved(c: u8) bool { + return switch (c) { + 'A'...'Z', 'a'...'z', '0'...'9', '-', '_', '.', '~' => true, + else => false, + }; +} + +test "PKCE S256 challenge - RFC 7636 test vector" { + const allocator = std.testing.allocator; + const verifier = "dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEjXk"; + const challenge = try generatePkceChallenge(allocator, verifier); + defer allocator.free(challenge); + try std.testing.expectEqualStrings("E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM", challenge); +} + +test "PKCE verifier is 43 chars" { + const allocator = std.testing.allocator; + const io = std.Options.debug_io; + const verifier = try generatePkceVerifier(allocator, io); + defer allocator.free(verifier); + try std.testing.expectEqual(@as(usize, 43), verifier.len); +} + +test "form URL encoding" { + const allocator = std.testing.allocator; + const params = [_][2][]const u8{ + .{ "grant_type", "authorization_code" }, + .{ "code", "abc123" }, + .{ "redirect_uri", "https://example.com/callback" }, + }; + + const encoded = try formEncode(allocator, ¶ms); + defer allocator.free(encoded); + + try std.testing.expectEqualStrings( + "grant_type=authorization_code&code=abc123&redirect_uri=https%3A%2F%2Fexample.com%2Fcallback", + encoded, + ); +} + +test "access token hash" { + const allocator = std.testing.allocator; + const ath = try accessTokenHash(allocator, "test-access-token"); + defer allocator.free(ath); + try std.testing.expectEqual(@as(usize, 43), ath.len); +} + +test "createJwt sign and verify round-trip" { + const allocator = std.testing.allocator; + const multibase = @import("../crypto/multibase.zig"); + const multicodec = @import("../crypto/multicodec.zig"); + + const keypair = try Keypair.fromSecretKey(.p256, .{ + 0x21, 0x22, 0x23, 0x24, 0x25, 0x26, 0x27, 0x28, + 0x29, 0x2a, 0x2b, 0x2c, 0x2d, 0x2e, 0x2f, 0x30, + 0x31, 0x32, 0x33, 0x34, 0x35, 0x36, 0x37, 0x38, + 0x39, 0x3a, 0x3b, 0x3c, 0x3d, 0x3e, 0x3f, 0x40, + }); + + const header = + \\{"alg":"ES256","typ":"JWT"} + ; + const payload = + \\{"iss":"did:example:test","aud":"did:example:aud","exp":9999999999} + ; + + const token = try createJwt(allocator, header, payload, &keypair); + defer allocator.free(token); + + var parsed_jwt = try jwt.Jwt.parse(allocator, token); + defer parsed_jwt.deinit(); + + try std.testing.expectEqual(jwt.Algorithm.ES256, parsed_jwt.header.alg); + try std.testing.expectEqualStrings("did:example:test", parsed_jwt.payload.iss); + + const pk = try keypair.publicKey(); + const mc_bytes = try multicodec.encodePublicKey(allocator, .p256, &pk); + defer allocator.free(mc_bytes); + const multibase_key = try multibase.encode(allocator, .base58btc, mc_bytes); + defer allocator.free(multibase_key); + + try parsed_jwt.verify(multibase_key); +} + +test "DPoP proof structure" { + const allocator = std.testing.allocator; + const io = std.Options.debug_io; + + const keypair = try Keypair.fromSecretKey(.p256, .{ + 0x21, 0x22, 0x23, 0x24, 0x25, 0x26, 0x27, 0x28, + 0x29, 0x2a, 0x2b, 0x2c, 0x2d, 0x2e, 0x2f, 0x30, + 0x31, 0x32, 0x33, 0x34, 0x35, 0x36, 0x37, 0x38, + 0x39, 0x3a, 0x3b, 0x3c, 0x3d, 0x3e, 0x3f, 0x40, + }); + + const proof = try createDpopProof(allocator, io, &keypair, "POST", "https://auth.example.com/token", "server-nonce", null); + defer allocator.free(proof); + + var iter = std.mem.splitScalar(u8, proof, '.'); + const header_b64 = iter.next().?; + const payload_b64 = iter.next().?; + + const header_json = try jwt.base64UrlDecode(allocator, header_b64); + defer allocator.free(header_json); + const header_parsed = try std.json.parseFromSlice(std.json.Value, allocator, header_json, .{}); + defer header_parsed.deinit(); + + try std.testing.expectEqualStrings("dpop+jwt", header_parsed.value.object.get("typ").?.string); + try std.testing.expectEqualStrings("ES256", header_parsed.value.object.get("alg").?.string); + try std.testing.expect(header_parsed.value.object.get("jwk") != null); + + const payload_json = try jwt.base64UrlDecode(allocator, payload_b64); + defer allocator.free(payload_json); + const payload_parsed = try std.json.parseFromSlice(std.json.Value, allocator, payload_json, .{}); + defer payload_parsed.deinit(); + + const obj = payload_parsed.value.object; + try std.testing.expect(obj.get("jti") != null); + try std.testing.expectEqualStrings("POST", obj.get("htm").?.string); + try std.testing.expectEqualStrings("https://auth.example.com/token", obj.get("htu").?.string); + try std.testing.expect(obj.get("iat") != null); + try std.testing.expectEqualStrings("server-nonce", obj.get("nonce").?.string); +} + +test "client assertion structure" { + const allocator = std.testing.allocator; + const io = std.Options.debug_io; + + const keypair = try Keypair.fromSecretKey(.p256, .{ + 0x21, 0x22, 0x23, 0x24, 0x25, 0x26, 0x27, 0x28, + 0x29, 0x2a, 0x2b, 0x2c, 0x2d, 0x2e, 0x2f, 0x30, + 0x31, 0x32, 0x33, 0x34, 0x35, 0x36, 0x37, 0x38, + 0x39, 0x3a, 0x3b, 0x3c, 0x3d, 0x3e, 0x3f, 0x40, + }); + + const assertion = try createClientAssertion(allocator, io, &keypair, "https://app.example.com/client-metadata", "https://bsky.social/oauth/token"); + defer allocator.free(assertion); + + var iter = std.mem.splitScalar(u8, assertion, '.'); + const header_b64 = iter.next().?; + const payload_b64 = iter.next().?; + + const header_json = try jwt.base64UrlDecode(allocator, header_b64); + defer allocator.free(header_json); + const header_parsed = try std.json.parseFromSlice(std.json.Value, allocator, header_json, .{}); + defer header_parsed.deinit(); + + try std.testing.expectEqualStrings("JWT", header_parsed.value.object.get("typ").?.string); + try std.testing.expectEqualStrings("ES256", header_parsed.value.object.get("alg").?.string); + try std.testing.expect(header_parsed.value.object.get("kid") != null); + + const payload_json = try jwt.base64UrlDecode(allocator, payload_b64); + defer allocator.free(payload_json); + const payload_parsed = try std.json.parseFromSlice(std.json.Value, allocator, payload_json, .{}); + defer payload_parsed.deinit(); + + const obj = payload_parsed.value.object; + try std.testing.expectEqualStrings("https://app.example.com/client-metadata", obj.get("iss").?.string); + try std.testing.expectEqualStrings("https://app.example.com/client-metadata", obj.get("sub").?.string); + try std.testing.expectEqualStrings("https://bsky.social/oauth/token", obj.get("aud").?.string); + try std.testing.expect(obj.get("jti") != null); + try std.testing.expect(obj.get("iat") != null); + try std.testing.expect(obj.get("exp") != null); +} + +test "JWKS JSON wraps JWK" { + const allocator = std.testing.allocator; + const keypair = try Keypair.fromSecretKey(.p256, .{ + 0x21, 0x22, 0x23, 0x24, 0x25, 0x26, 0x27, 0x28, + 0x29, 0x2a, 0x2b, 0x2c, 0x2d, 0x2e, 0x2f, 0x30, + 0x31, 0x32, 0x33, 0x34, 0x35, 0x36, 0x37, 0x38, + 0x39, 0x3a, 0x3b, 0x3c, 0x3d, 0x3e, 0x3f, 0x40, + }); + + const jwks = try jwksJson(allocator, &keypair); + defer allocator.free(jwks); + + const parsed = try std.json.parseFromSlice(std.json.Value, allocator, jwks, .{}); + defer parsed.deinit(); + + const keys = parsed.value.object.get("keys").?.array; + try std.testing.expectEqual(@as(usize, 1), keys.items.len); + try std.testing.expectEqualStrings("EC", keys.items[0].object.get("kty").?.string); +} diff --git a/src/internal/xrpc/transport.zig b/src/internal/xrpc/transport.zig index d4e6f0b..fd31794 100644 --- a/src/internal/xrpc/transport.zig +++ b/src/internal/xrpc/transport.zig @@ -166,6 +166,11 @@ pub const HttpTransport = struct { response: *std.http.Client.Response, ) !FetchResult { const rate_limit = RateLimitHeaders.fromResponseHead(response.head); + var oauth = if (options.capture_response_headers) + try OAuthHeaders.fromResponseHead(self.allocator, response.head) + else + OAuthHeaders{}; + errdefer oauth.deinit(self.allocator); if (options.max_response_size) |max| { const body_buf = try self.allocator.alloc(u8, max); @@ -176,6 +181,7 @@ pub const HttpTransport = struct { .status = response.head.status, .body = try self.allocator.dupe(u8, writer.buffered()), .rate_limit = rate_limit, + .oauth = oauth, }; } @@ -186,6 +192,7 @@ pub const HttpTransport = struct { .status = response.head.status, .body = try self.allocator.dupe(u8, aw.written()), .rate_limit = rate_limit, + .oauth = oauth, }; } @@ -200,6 +207,7 @@ pub const HttpTransport = struct { max_response_size: ?usize = null, redirect_behavior: ?std.http.Client.Request.RedirectBehavior = null, resolved_connection: ?ResolvedConnection = null, + capture_response_headers: bool = false, }; pub const ResolvedConnection = struct { @@ -213,6 +221,41 @@ pub const HttpTransport = struct { status: std.http.Status, body: []u8, rate_limit: RateLimitHeaders = .{}, + oauth: OAuthHeaders = .{}, + + pub fn deinit(self: *FetchResult, allocator: std.mem.Allocator) void { + allocator.free(self.body); + self.oauth.deinit(allocator); + } + }; + + pub const OAuthHeaders = struct { + content_type: ?[]const u8 = null, + dpop_nonce: ?[]const u8 = null, + www_authenticate: ?[]const u8 = null, + + pub fn fromResponseHead(allocator: std.mem.Allocator, head: std.http.Client.Response.Head) !OAuthHeaders { + var result: OAuthHeaders = .{}; + errdefer result.deinit(allocator); + var it = head.iterateHeaders(); + while (it.next()) |header| { + if (std.ascii.eqlIgnoreCase(header.name, "content-type")) { + result.content_type = try allocator.dupe(u8, header.value); + } else if (std.ascii.eqlIgnoreCase(header.name, "dpop-nonce")) { + result.dpop_nonce = try allocator.dupe(u8, header.value); + } else if (std.ascii.eqlIgnoreCase(header.name, "www-authenticate")) { + result.www_authenticate = try allocator.dupe(u8, header.value); + } + } + return result; + } + + pub fn deinit(self: *OAuthHeaders, allocator: std.mem.Allocator) void { + if (self.content_type) |value| allocator.free(value); + if (self.dpop_nonce) |value| allocator.free(value); + if (self.www_authenticate) |value| allocator.free(value); + self.* = .{}; + } }; pub const RateLimitHeaders = struct { @@ -318,3 +361,27 @@ test "transport parses rate limit headers" { try std.testing.expectEqual(@as(?u64, 2), headers.retry_after); try std.testing.expect(!headers.isEmpty()); } + +test "transport parses oauth headers" { + const response_bytes = "HTTP/1.1 401 Unauthorized\r\n" ++ + "DPoP-Nonce: nonce-123\r\n" ++ + "WWW-Authenticate: DPoP error=\"use_dpop_nonce\"\r\n\r\n"; + + const head = try std.http.Client.Response.Head.parse(response_bytes); + var headers = try HttpTransport.OAuthHeaders.fromResponseHead(std.testing.allocator, head); + defer headers.deinit(std.testing.allocator); + + try std.testing.expectEqualStrings("nonce-123", headers.dpop_nonce.?); + try std.testing.expectEqualStrings("DPoP error=\"use_dpop_nonce\"", headers.www_authenticate.?); +} + +test "transport parses oauth content type header" { + const response_bytes = "HTTP/1.1 200 OK\r\n" ++ + "Content-Type: application/json; charset=utf-8\r\n\r\n"; + + const head = try std.http.Client.Response.Head.parse(response_bytes); + var headers = try HttpTransport.OAuthHeaders.fromResponseHead(std.testing.allocator, head); + defer headers.deinit(std.testing.allocator); + + try std.testing.expectEqualStrings("application/json; charset=utf-8", headers.content_type.?); +}