diff --git a/src/internal/identity/did_resolver.zig b/src/internal/identity/did_resolver.zig index caa802c..f05575c 100644 --- a/src/internal/identity/did_resolver.zig +++ b/src/internal/identity/did_resolver.zig @@ -62,47 +62,17 @@ pub const DidResolver = struct { /// resolve did:web via .well-known fn resolveWeb(self: *DidResolver, did: Did) !DidDocument { - // did:web:example.com -> https://example.com/.well-known/did.json - // did:web:example.com:path:to -> https://example.com/path/to/did.json - const domain_and_path = did.raw["did:web:".len..]; - - // decode percent-encoded colons in path - var url_buf: std.ArrayList(u8) = .empty; - defer url_buf.deinit(self.allocator); - - try url_buf.appendSlice(self.allocator, "https://"); - - var first_segment = true; - var it = std.mem.splitScalar(u8, domain_and_path, ':'); - while (it.next()) |segment| { - if (first_segment) { - // first segment is the domain - try url_buf.appendSlice(self.allocator, segment); - first_segment = false; - } else { - // subsequent segments are path components - try url_buf.append(self.allocator, '/'); - try url_buf.appendSlice(self.allocator, segment); - } - } - - // add .well-known/did.json or /did.json - if (std.mem.indexOf(u8, domain_and_path, ":") == null) { - // no path, use .well-known - try url_buf.appendSlice(self.allocator, "/.well-known/did.json"); - } else { - // has path, append did.json - try url_buf.appendSlice(self.allocator, "/did.json"); - } + const url = try buildDidWebDocumentUrl(self.allocator, did); + defer self.allocator.free(url); var checked_url = try network_safety.resolveIdentityUrl( self.allocator, &self.transport, self.doh_endpoint, - url_buf.items, + url, ); defer checked_url.deinit(self.allocator); - return try self.fetchDidDocument(url_buf.items, checked_url.resolvedConnection()); + return try self.fetchDidDocument(url, checked_url.resolvedConnection()); } /// fetch and parse a did document from url @@ -127,6 +97,44 @@ pub const DidResolver = struct { } }; +fn buildDidWebDocumentUrl(allocator: std.mem.Allocator, did: Did) ![]u8 { + const identifier = did.identifier(); + if (identifier.len == 0 or std.mem.indexOfScalar(u8, identifier, ':') != null) { + return error.UnsupportedDidMethod; + } + + const authority = try percentDecodeAlloc(allocator, identifier); + defer allocator.free(authority); + + const scheme: []const u8 = if (std.mem.eql(u8, authority, "localhost") or std.mem.startsWith(u8, authority, "localhost:")) + "http" + else + "https"; + + return try std.fmt.allocPrint(allocator, "{s}://{s}/.well-known/did.json", .{ scheme, authority }); +} + +fn percentDecodeAlloc(allocator: std.mem.Allocator, input: []const u8) ![]u8 { + var out: std.ArrayList(u8) = .empty; + errdefer out.deinit(allocator); + + var i: usize = 0; + while (i < input.len) { + if (input[i] != '%') { + try out.append(allocator, input[i]); + i += 1; + continue; + } + + if (i + 2 >= input.len) return error.InvalidPercentEncoding; + const byte = std.fmt.parseInt(u8, input[i + 1 .. i + 3], 16) catch return error.InvalidPercentEncoding; + try out.append(allocator, byte); + i += 3; + } + + return try out.toOwnedSlice(allocator); +} + // === tests === test "resolve did:plc - integration" { @@ -171,21 +179,18 @@ test "did:web loopback host is rejected before fetch" { } test "did:web url construction" { - // test url building without network - var resolver = DidResolver.init(std.Options.debug_io, std.testing.allocator); - defer resolver.deinit(); + const simple = try buildDidWebDocumentUrl(std.testing.allocator, Did.parse("did:web:example.com").?); + defer std.testing.allocator.free(simple); + try std.testing.expectEqualStrings("https://example.com/.well-known/did.json", simple); - // simple domain - { - const did = Did.parse("did:web:example.com").?; - _ = did; - // would resolve to https://example.com/.well-known/did.json - } + const with_port = try buildDidWebDocumentUrl(std.testing.allocator, Did.parse("did:web:example.com%3A3000").?); + defer std.testing.allocator.free(with_port); + try std.testing.expectEqualStrings("https://example.com:3000/.well-known/did.json", with_port); - // domain with path - { - const did = Did.parse("did:web:example.com:user:alice").?; - _ = did; - // would resolve to https://example.com/user/alice/did.json - } + const localhost = try buildDidWebDocumentUrl(std.testing.allocator, Did.parse("did:web:localhost%3A3000").?); + defer std.testing.allocator.free(localhost); + try std.testing.expectEqualStrings("http://localhost:3000/.well-known/did.json", localhost); + + try std.testing.expectError(error.UnsupportedDidMethod, buildDidWebDocumentUrl(std.testing.allocator, Did.parse("did:web:example.com:user:alice").?)); + try std.testing.expectError(error.InvalidPercentEncoding, buildDidWebDocumentUrl(std.testing.allocator, Did.parse("did:web:example.com%xx").?)); } diff --git a/src/internal/identity/handle_resolver.zig b/src/internal/identity/handle_resolver.zig index bc70354..b7b2c06 100644 --- a/src/internal/identity/handle_resolver.zig +++ b/src/internal/identity/handle_resolver.zig @@ -119,7 +119,7 @@ pub const HandleResolver = struct { return self.didFromDohBody(result.body); } - /// parse a Cloudflare DoH JSON body and extract the first valid `did=` TXT value. + /// parse a Cloudflare DoH JSON body and extract exactly one valid `did=` TXT value. fn didFromDohBody(self: *HandleResolver, body: []const u8) ![]const u8 { // Cloudflare appends an unknown `Comment` field for DNSSEC zones, so the // parse must tolerate fields `DnsResponse` doesn't declare. @@ -136,16 +136,25 @@ pub const HandleResolver = struct { return error.NoDnsRecordsFound; } + var found: ?[]const u8 = null; + var found_count: usize = 0; for (dns_response.Answer.?) |answer| { const data = answer.data orelse continue; const did_str = extractDidFromTxt(data) orelse continue; + found_count += 1; + found = did_str; + } - if (Did.parse(did_str) != null) { - return try self.allocator.dupe(u8, did_str); - } + if (found_count != 1) { + return error.NoValidDidFound; } - return error.NoValidDidFound; + const did_str = found.?; + if (Did.parse(did_str) == null) { + return error.NoValidDidFound; + } + + return try self.allocator.dupe(u8, did_str); } }; @@ -209,6 +218,37 @@ test "didFromDohBody tolerates Cloudflare DNSSEC Comment field" { try std.testing.expectEqualStrings("did:plc:yk4dd2qkboz2yv6tpubpc6co", did); } +test "didFromDohBody requires exactly one did TXT record" { + var resolver = HandleResolver.init(std.Options.debug_io, std.testing.allocator); + defer resolver.deinit(); + + const single = + \\{"Status":0,"TC":false,"RD":true,"RA":true,"AD":false,"CD":false, + \\ "Answer":[ + \\ {"name":"_atproto.example.com","type":16,"TTL":300,"data":"\"foo=bar\""}, + \\ {"name":"_atproto.example.com","type":16,"TTL":300,"data":"\"did=did:plc:abc123\""} + \\ ]} + ; + const did = try resolver.didFromDohBody(single); + defer std.testing.allocator.free(did); + try std.testing.expectEqualStrings("did:plc:abc123", did); + + const multiple = + \\{"Status":0,"TC":false,"RD":true,"RA":true,"AD":false,"CD":false, + \\ "Answer":[ + \\ {"name":"_atproto.example.com","type":16,"TTL":300,"data":"\"did=did:plc:abc123\""}, + \\ {"name":"_atproto.example.com","type":16,"TTL":300,"data":"\"did=did:plc:def456\""} + \\ ]} + ; + try std.testing.expectError(error.NoValidDidFound, resolver.didFromDohBody(multiple)); + + const invalid = + \\{"Status":0,"TC":false,"RD":true,"RA":true,"AD":false,"CD":false, + \\ "Answer":[{"name":"_atproto.example.com","type":16,"TTL":300,"data":"\"did=not-a-did\""}]} + ; + try std.testing.expectError(error.NoValidDidFound, resolver.didFromDohBody(invalid)); +} + test "resolve handle (http) - integration" { // use arena for http client internals that may leak var arena = std.heap.ArenaAllocator.init(std.testing.allocator); diff --git a/src/internal/xrpc/xrpc.zig b/src/internal/xrpc/xrpc.zig index 832438a..61c9b74 100644 --- a/src/internal/xrpc/xrpc.zig +++ b/src/internal/xrpc/xrpc.zig @@ -45,7 +45,7 @@ pub const XrpcClient = struct { const url = try self.buildUrl(nsid, params); defer self.allocator.free(url); - return try self.doRequest(url, null); + return try self.doRequest(.GET, url, null); } /// call a procedure method (POST) @@ -53,7 +53,7 @@ pub const XrpcClient = struct { const url = try self.buildUrl(nsid, null); defer self.allocator.free(url); - return try self.doRequest(url, body); + return try self.doRequest(.POST, url, body); } pub fn queryChecked( @@ -65,7 +65,7 @@ pub const XrpcClient = struct { const url = try self.buildUrl(nsid, params); defer self.allocator.free(url); - return try self.requestCheckedUrl(url, null, retry_policy); + return try self.requestCheckedUrl(.GET, url, null, retry_policy); } pub fn procedureChecked( @@ -77,7 +77,7 @@ pub const XrpcClient = struct { const url = try self.buildUrl(nsid, null); defer self.allocator.free(url); - return try self.requestCheckedUrl(url, body, retry_policy); + return try self.requestCheckedUrl(.POST, url, body, retry_policy); } fn buildUrl(self: *XrpcClient, nsid: Nsid, params: ?std.StringHashMap([]const u8)) ![]u8 { @@ -110,7 +110,7 @@ pub const XrpcClient = struct { return try url.toOwnedSlice(self.allocator); } - fn doRequest(self: *XrpcClient, url: []const u8, body: ?[]const u8) !Response { + fn doRequest(self: *XrpcClient, method: std.http.Method, url: []const u8, body: ?[]const u8) !Response { var auth_header_buf: [max_auth_header_len]u8 = undefined; const auth_value: ?[]const u8 = if (self.access_token) |token| std.fmt.bufPrint(&auth_header_buf, "Bearer {s}", .{token}) catch null @@ -119,7 +119,7 @@ pub const XrpcClient = struct { const result = try self.transport.fetch(.{ .url = url, - .method = if (body != null) .POST else .GET, + .method = method, .payload = body, .authorization = auth_value, }); @@ -132,13 +132,13 @@ pub const XrpcClient = struct { }; } - fn requestCheckedUrl(self: *XrpcClient, url: []const u8, body: ?[]const u8, retry_policy: RetryPolicy) !Result { + fn requestCheckedUrl(self: *XrpcClient, method: std.http.Method, url: []const u8, body: ?[]const u8, retry_policy: RetryPolicy) !Result { const attempts = @max(@as(u8, 1), retry_policy.max_attempts); var attempt: u8 = 0; while (true) : (attempt += 1) { - var response = self.doRequest(url, body) catch |err| { - if (attempt + 1 >= attempts or !retry_policy.retry_transient_errors or !isRetryableTransportError(err)) { + var response = self.doRequest(method, url, body) catch |err| { + if (attempt + 1 >= attempts or !retry_policy.retry_transient_errors or !isRetryableTransportErrorForMethod(method, err)) { return err; } try retry_policy.sleepBeforeRetry(self.transport.io, attempt, null); @@ -149,7 +149,7 @@ pub const XrpcClient = struct { return .{ .ok = response }; } - if (attempt + 1 < attempts and retry_policy.isRetryableStatus(response.status)) { + if (attempt + 1 < attempts and retry_policy.isRetryableStatusForMethod(method, response.status)) { const rate_limit = response.rate_limit; response.deinit(); try retry_policy.sleepBeforeRetry(self.transport.io, attempt, rate_limit); @@ -248,6 +248,12 @@ pub const XrpcClient = struct { }; } + pub fn isRetryableStatusForMethod(self: RetryPolicy, method: std.http.Method, status: std.http.Status) bool { + if (!self.isRetryableStatus(status)) return false; + if (method == .POST and status != .too_many_requests) return false; + return true; + } + pub fn delayMillis(self: RetryPolicy, attempt: u8, rate_limit: ?HttpTransport.RateLimitHeaders) u64 { return self.delayMillisAt(attempt, rate_limit, null); } @@ -305,7 +311,11 @@ pub const XrpcClient = struct { }; }; -fn isRetryableTransportError(err: anyerror) bool { +fn isRetryableTransportErrorForMethod(method: std.http.Method, err: anyerror) bool { + if (method == .POST) { + return err == error.ConnectionRefused; + } + return switch (err) { error.ConnectionRefused, error.ConnectionResetByPeer, @@ -396,6 +406,15 @@ test "retry policy is conservative and deterministic" { try std.testing.expect(!policy.isRetryableStatus(.unauthorized)); try std.testing.expect(!policy.isRetryableStatus(.not_found)); + try std.testing.expect(policy.isRetryableStatusForMethod(.GET, .service_unavailable)); + try std.testing.expect(policy.isRetryableStatusForMethod(.GET, .too_many_requests)); + try std.testing.expect(!policy.isRetryableStatusForMethod(.POST, .service_unavailable)); + try std.testing.expect(policy.isRetryableStatusForMethod(.POST, .too_many_requests)); + + try std.testing.expect(isRetryableTransportErrorForMethod(.GET, error.ConnectionResetByPeer)); + try std.testing.expect(isRetryableTransportErrorForMethod(.POST, error.ConnectionRefused)); + try std.testing.expect(!isRetryableTransportErrorForMethod(.POST, error.ConnectionResetByPeer)); + try std.testing.expectEqual(@as(u64, 100), policy.delayMillis(0, null)); try std.testing.expectEqual(@as(u64, 200), policy.delayMillis(1, null)); try std.testing.expectEqual(@as(u64, 1000), policy.delayMillis(10, null));