From 8287ff236f0cffb1d2d89b8888ffbc800b0a9a5e Mon Sep 17 00:00:00 2001 From: zzstoatzz Date: Thu, 23 Apr 2026 22:33:06 -0500 Subject: [PATCH] harden identity network resolution --- src/internal/identity/did_resolver.zig | 58 +++-- src/internal/identity/handle_resolver.zig | 34 ++- src/internal/identity/network_safety.zig | 283 ++++++++++++++++++++++ src/internal/xrpc/transport.zig | 152 +++++++++++- 4 files changed, 488 insertions(+), 39 deletions(-) create mode 100644 src/internal/identity/network_safety.zig diff --git a/src/internal/identity/did_resolver.zig b/src/internal/identity/did_resolver.zig index 3b6cbe5..caa802c 100644 --- a/src/internal/identity/did_resolver.zig +++ b/src/internal/identity/did_resolver.zig @@ -8,6 +8,9 @@ const std = @import("std"); const Did = @import("../syntax/did.zig").Did; const DidDocument = @import("did_document.zig").DidDocument; const HttpTransport = @import("../xrpc/transport.zig").HttpTransport; +const network_safety = @import("network_safety.zig"); + +const max_did_document_size = 1 * 1024 * 1024; pub const DidResolver = struct { allocator: std.mem.Allocator, @@ -15,6 +18,8 @@ pub const DidResolver = struct { /// plc directory url (default: https://plc.directory) plc_url: []const u8 = "https://plc.directory", + /// DoH endpoint for did:web host safety preflight. + doh_endpoint: []const u8 = "https://cloudflare-dns.com/dns-query", pub fn init(io: std.Io, allocator: std.mem.Allocator) DidResolver { return initWithOptions(io, allocator, .{}); @@ -52,7 +57,7 @@ pub const DidResolver = struct { const url = try std.fmt.allocPrint(self.allocator, "{s}/{s}", .{ self.plc_url, did.raw }); defer self.allocator.free(url); - return try self.fetchDidDocument(url); + return try self.fetchDidDocument(url, null); } /// resolve did:web via .well-known @@ -90,12 +95,28 @@ pub const DidResolver = struct { try url_buf.appendSlice(self.allocator, "/did.json"); } - return try self.fetchDidDocument(url_buf.items); + var checked_url = try network_safety.resolveIdentityUrl( + self.allocator, + &self.transport, + self.doh_endpoint, + url_buf.items, + ); + defer checked_url.deinit(self.allocator); + return try self.fetchDidDocument(url_buf.items, checked_url.resolvedConnection()); } /// fetch and parse a did document from url - fn fetchDidDocument(self: *DidResolver, url: []const u8) !DidDocument { - const result = try self.transport.fetch(.{ .url = url }); + fn fetchDidDocument( + self: *DidResolver, + url: []const u8, + resolved_connection: ?HttpTransport.ResolvedConnection, + ) !DidDocument { + const result = try self.transport.fetch(.{ + .url = url, + .max_response_size = max_did_document_size, + .redirect_behavior = .not_allowed, + .resolved_connection = resolved_connection, + }); defer self.allocator.free(result.body); if (result.status != .ok) { @@ -141,37 +162,12 @@ test "resolve did:plc - leak check (no arena)" { try std.testing.expectEqualStrings("did:plc:z72i7hdynmk6r22z27h6tvur", doc.id); } -test "regression: transport errors propagate distinct kinds" { - // before this fix, transport.fetch had `catch return error.RequestFailed` - // and fetchDidDocument had `catch return error.DidResolutionFailed`, so - // every transport-layer failure (DNS, TCP, TLS) collapsed to one - // indistinguishable error and callers had no way to see what was wrong. - // this regression test asserts the underlying error kind survives the - // resolver layer for at least one common transport failure mode. - // - // history: zlay 2026-04-08, where the host_authority pool failed at 100% - // and we had no production telemetry on which transport error fired - // because both layers had been swallowed. see relay docs/zlay-external- - // review-2026-04-09.md. +test "did:web loopback host is rejected before fetch" { var resolver = DidResolver.init(std.Options.debug_io, std.testing.allocator); defer resolver.deinit(); - // 127.0.0.1:443 is almost certainly not listening on a test machine. - // did:web:127.0.0.1 → https://127.0.0.1/.well-known/did.json → connect refused. const did = Did.parse("did:web:127.0.0.1") orelse return error.SkipZigTest; - if (resolver.resolve(did)) |doc| { - // someone is actually serving a DID doc on 127.0.0.1:443 — skip rather - // than fail, since the assertion below assumes a transport failure - var d = doc; - d.deinit(); - return error.SkipZigTest; - } else |err| { - // exact error name varies by platform (ConnectionRefused on linux/darwin, - // possibly different elsewhere). just assert it's not the catch-all that - // the pre-fix code returned for everything. - try std.testing.expect(err != error.DidResolutionFailed); - try std.testing.expect(err != error.RequestFailed); - } + try std.testing.expectError(error.UnsafeIdentityHost, resolver.resolve(did)); } test "did:web url construction" { diff --git a/src/internal/identity/handle_resolver.zig b/src/internal/identity/handle_resolver.zig index 38a75c6..0b9decb 100644 --- a/src/internal/identity/handle_resolver.zig +++ b/src/internal/identity/handle_resolver.zig @@ -3,8 +3,8 @@ //! resolves AT Protocol handles via HTTP: //! https://{handle}/.well-known/atproto-did //! -//! note: DNS TXT resolution (_atproto.{handle}) not yet implemented -//! as zig std doesn't provide TXT record lookup. +//! DNS TXT resolution uses DNS-over-HTTPS because Zig stdlib does not expose +//! direct TXT lookup. //! //! see: https://atproto.com/specs/handle @@ -12,6 +12,10 @@ const std = @import("std"); const Handle = @import("../syntax/handle.zig").Handle; const Did = @import("../syntax/did.zig").Did; const HttpTransport = @import("../xrpc/transport.zig").HttpTransport; +const network_safety = @import("network_safety.zig"); + +const max_handle_response_size = 8 * 1024; +const max_dns_txt_response_size = 1 * 1024 * 1024; pub const HandleResolver = struct { allocator: std.mem.Allocator, @@ -34,8 +38,9 @@ pub const HandleResolver = struct { pub fn resolve(self: *HandleResolver, handle: Handle) ![]const u8 { if (self.resolveHttp(handle)) |did| { return did; - } else |_| { - return try self.resolveDns(handle); + } else |err| switch (err) { + error.UnsafeIdentityHost => return err, + else => return try self.resolveDns(handle), } } @@ -48,7 +53,24 @@ pub const HandleResolver = struct { ); defer self.allocator.free(url); - const result = self.transport.fetch(.{ .url = url }) catch return error.HttpResolutionFailed; + var checked_url = network_safety.resolveIdentityUrl( + self.allocator, + &self.transport, + self.doh_endpoint, + url, + ) catch |err| switch (err) { + error.OutOfMemory => |e| return e, + error.UnsafeIdentityHost => |e| return e, + else => return error.HttpResolutionFailed, + }; + defer checked_url.deinit(self.allocator); + + const result = self.transport.fetch(.{ + .url = url, + .max_response_size = max_handle_response_size, + .redirect_behavior = .not_allowed, + .resolved_connection = checked_url.resolvedConnection(), + }) catch return error.HttpResolutionFailed; defer self.allocator.free(result.body); if (result.status != .ok) { @@ -85,6 +107,8 @@ pub const HandleResolver = struct { const result = self.transport.fetch(.{ .url = url, .accept = "application/dns-json", + .max_response_size = max_dns_txt_response_size, + .redirect_behavior = .not_allowed, }) catch return error.DnsResolutionFailed; defer self.allocator.free(result.body); diff --git a/src/internal/identity/network_safety.zig b/src/internal/identity/network_safety.zig new file mode 100644 index 0000000..9f8003e --- /dev/null +++ b/src/internal/identity/network_safety.zig @@ -0,0 +1,283 @@ +//! Network safety checks for identity resolution. + +const std = @import("std"); +const HttpTransport = @import("../xrpc/transport.zig").HttpTransport; + +pub const NetworkSafetyError = error{ + MissingHost, + UnsafeIdentityHost, + IdentityDnsResolutionFailed, +}; + +const max_dns_response_size = 1 * 1024 * 1024; + +pub const CheckedIdentityUrl = struct { + host: []u8, + dial_host: ?[]u8, + + pub fn deinit(self: *CheckedIdentityUrl, allocator: std.mem.Allocator) void { + allocator.free(self.host); + if (self.dial_host) |dial_host| allocator.free(dial_host); + } + + pub fn resolvedConnection(self: CheckedIdentityUrl) ?HttpTransport.ResolvedConnection { + return .{ + .dial_host = self.dial_host orelse return null, + .logical_host = self.host, + }; + } +}; + +pub fn checkIdentityUrl(url: []const u8) (std.Uri.ParseError || NetworkSafetyError)!void { + const uri = try std.Uri.parse(url); + var host_buf: [std.Io.net.HostName.max_len]u8 = undefined; + const host_name = uri.getHost(&host_buf) catch return error.MissingHost; + try checkIdentityHost(host_name.bytes); +} + +pub fn checkIdentityUrlResolved( + allocator: std.mem.Allocator, + transport: *HttpTransport, + doh_endpoint: []const u8, + url: []const u8, +) (std.Uri.ParseError || NetworkSafetyError || error{ OutOfMemory, ResponseTooLarge })!void { + var checked = try resolveIdentityUrl(allocator, transport, doh_endpoint, url); + checked.deinit(allocator); +} + +pub fn resolveIdentityUrl( + allocator: std.mem.Allocator, + transport: *HttpTransport, + doh_endpoint: []const u8, + url: []const u8, +) (std.Uri.ParseError || NetworkSafetyError || error{ OutOfMemory, ResponseTooLarge })!CheckedIdentityUrl { + const host = try hostFromUrlAlloc(allocator, url); + errdefer allocator.free(host); + try checkIdentityHost(host); + + // Literal IPs were fully checked above. Only DNS names need DoH preflight. + if (isIpLiteral(host)) { + return .{ .host = host, .dial_host = null }; + } + + var saw_address = false; + const dial_host = try checkDnsAnswers(allocator, transport, doh_endpoint, host, "A", &saw_address); + errdefer if (dial_host) |addr| allocator.free(addr); + const maybe_ip6 = try checkDnsAnswers(allocator, transport, doh_endpoint, host, "AAAA", &saw_address); + if (maybe_ip6) |ip6| allocator.free(ip6); + if (!saw_address) return error.IdentityDnsResolutionFailed; + if (dial_host == null) return error.IdentityDnsResolutionFailed; + + return .{ .host = host, .dial_host = dial_host }; +} + +fn hostFromUrlAlloc( + allocator: std.mem.Allocator, + url: []const u8, +) (std.Uri.ParseError || NetworkSafetyError || error{OutOfMemory})![]u8 { + const uri = try std.Uri.parse(url); + var host_buf: [std.Io.net.HostName.max_len]u8 = undefined; + const host_name = uri.getHost(&host_buf) catch |err| switch (err) { + error.UriMissingHost => return error.MissingHost, + }; + return try allocator.dupe(u8, host_name.bytes); +} + +pub fn checkIdentityHost(host: []const u8) NetworkSafetyError!void { + if (host.len == 0) return error.MissingHost; + + const host_without_trailing_dot = stripTrailingDot(host); + if (std.ascii.eqlIgnoreCase(host_without_trailing_dot, "localhost")) { + return error.UnsafeIdentityHost; + } + + if (std.Io.net.Ip4Address.parse(host, 0)) |ip4| { + if (isNonRoutableIp4(ip4.bytes)) return error.UnsafeIdentityHost; + return; + } else |_| {} + + const ip6_text = stripIp6Brackets(host); + if (std.Io.net.Ip6Address.parse(ip6_text, 0)) |ip6| { + if (ip4FromIp6Mapped(ip6.bytes)) |ip4| { + if (isNonRoutableIp4(ip4)) return error.UnsafeIdentityHost; + } + if (isNonRoutableIp6(ip6.bytes)) return error.UnsafeIdentityHost; + return; + } else |_| {} +} + +fn isIpLiteral(host: []const u8) bool { + if (std.Io.net.Ip4Address.parse(host, 0)) |_| return true else |_| {} + + const ip6_text = stripIp6Brackets(host); + if (std.Io.net.Ip6Address.parse(ip6_text, 0)) |_| return true else |_| {} + + return false; +} + +fn checkDnsAnswers( + allocator: std.mem.Allocator, + transport: *HttpTransport, + doh_endpoint: []const u8, + host: []const u8, + record_type: []const u8, + saw_address: *bool, +) (NetworkSafetyError || error{ OutOfMemory, ResponseTooLarge })!?[]u8 { + const url = try std.fmt.allocPrint( + allocator, + "{s}?name={s}&type={s}", + .{ doh_endpoint, host, record_type }, + ); + defer allocator.free(url); + + const result = transport.fetch(.{ + .url = url, + .accept = "application/dns-json", + .max_response_size = max_dns_response_size, + .redirect_behavior = .not_allowed, + }) catch |err| switch (err) { + error.OutOfMemory => |e| return e, + error.ResponseTooLarge => |e| return e, + else => return error.IdentityDnsResolutionFailed, + }; + defer allocator.free(result.body); + + if (result.status != .ok) return error.IdentityDnsResolutionFailed; + + const parsed = std.json.parseFromSlice(DnsResponse, allocator, result.body, .{}) catch + return error.IdentityDnsResolutionFailed; + defer parsed.deinit(); + + const answers = parsed.value.Answer orelse return null; + var first_ip4: ?[]u8 = null; + errdefer if (first_ip4) |ip| allocator.free(ip); + + for (answers) |answer| { + const data = answer.data orelse continue; + switch (answer.type) { + 1 => { + saw_address.* = true; + try checkIdentityHost(data); + if (first_ip4 == null) first_ip4 = try allocator.dupe(u8, data); + }, + 28 => { + saw_address.* = true; + try checkIdentityHost(data); + }, + else => {}, + } + } + return first_ip4; +} + +fn stripIp6Brackets(host: []const u8) []const u8 { + if (host.len >= 2 and host[0] == '[' and host[host.len - 1] == ']') { + return host[1 .. host.len - 1]; + } + return host; +} + +fn stripTrailingDot(host: []const u8) []const u8 { + if (host.len > 0 and host[host.len - 1] == '.') return host[0 .. host.len - 1]; + return host; +} + +fn isNonRoutableIp4(ip: [4]u8) bool { + return ip[0] == 0 or + ip[0] == 10 or + ip[0] == 127 or + (ip[0] == 169 and ip[1] == 254) or + (ip[0] == 172 and ip[1] >= 16 and ip[1] <= 31) or + (ip[0] == 192 and ip[1] == 168); +} + +fn isNonRoutableIp6(ip: [16]u8) bool { + const all_zero = for (ip) |b| { + if (b != 0) break false; + } else true; + + return all_zero or + isIp6Loopback(ip) or + (ip[0] == 0xfe and (ip[1] & 0xc0) == 0x80) or // fe80::/10 link-local + (ip[0] & 0xfe) == 0xfc; // fc00::/7 unique local +} + +fn isIp6Loopback(ip: [16]u8) bool { + for (ip[0..15]) |b| { + if (b != 0) return false; + } + return ip[15] == 1; +} + +fn ip4FromIp6Mapped(ip: [16]u8) ?[4]u8 { + for (ip[0..10]) |b| { + if (b != 0) return null; + } + if (ip[10] != 0xff or ip[11] != 0xff) return null; + return ip[12..16].*; +} + +const DnsResponse = struct { + Status: i32, + TC: bool = false, + RD: bool = false, + RA: bool = false, + AD: bool = false, + CD: bool = false, + Question: ?[]Question = null, + Answer: ?[]Answer = null, +}; + +const Question = struct { + name: []const u8, + type: i32, +}; + +const Answer = struct { + name: []const u8, + type: i32, + TTL: i32, + data: ?[]const u8 = null, +}; + +test "identity host rejects obvious non-routable hosts" { + try std.testing.expectError(error.UnsafeIdentityHost, checkIdentityHost("localhost")); + try std.testing.expectError(error.UnsafeIdentityHost, checkIdentityHost("localhost.")); + try std.testing.expectError(error.UnsafeIdentityHost, checkIdentityHost("127.0.0.1")); + try std.testing.expectError(error.UnsafeIdentityHost, checkIdentityHost("10.1.2.3")); + try std.testing.expectError(error.UnsafeIdentityHost, checkIdentityHost("172.16.0.1")); + try std.testing.expectError(error.UnsafeIdentityHost, checkIdentityHost("192.168.1.1")); + try std.testing.expectError(error.UnsafeIdentityHost, checkIdentityHost("169.254.1.1")); + try std.testing.expectError(error.UnsafeIdentityHost, checkIdentityHost("[::1]")); + try std.testing.expectError(error.UnsafeIdentityHost, checkIdentityHost("::ffff:127.0.0.1")); + try std.testing.expectError(error.UnsafeIdentityHost, checkIdentityHost("fc00::1")); + try checkIdentityHost("example.com"); + try checkIdentityHost("8.8.8.8"); +} + +test "identity url check rejects literal localhost" { + try std.testing.expectError(error.UnsafeIdentityHost, checkIdentityUrl("https://127.0.0.1/.well-known/did.json")); +} + +test "identity dns answers reject non-routable addresses" { + const json = + \\{ + \\ "Status": 0, + \\ "Answer": [ + \\ {"name": "evil.example.", "type": 1, "TTL": 60, "data": "127.0.0.1"} + \\ ] + \\} + ; + + const parsed = try std.json.parseFromSlice(DnsResponse, std.testing.allocator, json, .{}); + defer parsed.deinit(); + + var saw_address = false; + for (parsed.value.Answer.?) |answer| { + if (answer.type == 1 or answer.type == 28) { + saw_address = true; + try std.testing.expectError(error.UnsafeIdentityHost, checkIdentityHost(answer.data.?)); + } + } + try std.testing.expect(saw_address); +} diff --git a/src/internal/xrpc/transport.zig b/src/internal/xrpc/transport.zig index 210b199..0c824c3 100644 --- a/src/internal/xrpc/transport.zig +++ b/src/internal/xrpc/transport.zig @@ -25,9 +25,6 @@ pub const HttpTransport = struct { /// fetch a URL and write response to provided writer pub fn fetch(self: *HttpTransport, options: FetchOptions) !FetchResult { - var aw: std.Io.Writer.Allocating = .init(self.allocator); - defer aw.deinit(); - var headers: std.http.Client.Request.Headers = .{ .accept_encoding = .{ .override = "identity" }, // disable gzip - zig stdlib issue .content_type = if (options.payload != null) .{ .override = "application/json" } else .default, @@ -59,6 +56,38 @@ pub const HttpTransport = struct { } } + if (options.resolved_connection) |resolved| { + return try self.fetchResolved(options, resolved, headers, extra_buf[0..extra_count]); + } + + if (options.max_response_size) |max| { + const body_buf = try self.allocator.alloc(u8, max); + defer self.allocator.free(body_buf); + var writer = std.Io.Writer.fixed(body_buf); + + const result = self.http_client.fetch(.{ + .location = .{ .url = options.url }, + .response_writer = &writer, + .method = options.method, + .payload = options.payload, + .headers = headers, + .extra_headers = extra_buf[0..extra_count], + .keep_alive = self.keep_alive, + .redirect_behavior = options.redirect_behavior, + }) catch |err| switch (err) { + error.WriteFailed => return error.ResponseTooLarge, + else => |e| return e, + }; + + return .{ + .status = result.status, + .body = try self.allocator.dupe(u8, writer.buffered()), + }; + } + + var aw: std.Io.Writer.Allocating = .init(self.allocator); + defer aw.deinit(); + const result = try self.http_client.fetch(.{ .location = .{ .url = options.url }, .response_writer = &aw.writer, @@ -67,6 +96,7 @@ pub const HttpTransport = struct { .headers = headers, .extra_headers = extra_buf[0..extra_count], .keep_alive = self.keep_alive, + .redirect_behavior = options.redirect_behavior, }); return .{ @@ -75,6 +105,74 @@ pub const HttpTransport = struct { }; } + fn fetchResolved( + self: *HttpTransport, + options: FetchOptions, + resolved: ResolvedConnection, + headers: std.http.Client.Request.Headers, + extra_headers: []const std.http.Header, + ) !FetchResult { + const uri = try std.Uri.parse(options.url); + const protocol = std.http.Client.Protocol.fromUri(uri) orelse return error.UnsupportedUriScheme; + + var uri_host_buf: [std.Io.net.HostName.max_len]u8 = undefined; + const uri_host = try uri.getHost(&uri_host_buf); + if (!std.ascii.eqlIgnoreCase(uri_host.bytes, resolved.logical_host)) { + return error.ResolvedHostMismatch; + } + + const dial_host = try std.Io.net.HostName.init(resolved.dial_host); + const logical_host = try std.Io.net.HostName.init(resolved.logical_host); + const connection = try self.http_client.connectTcpOptions(.{ + .host = dial_host, + .port = uri.port orelse defaultPort(protocol), + .protocol = protocol, + .proxied_host = logical_host, + }); + + var request = self.http_client.request(options.method, uri, .{ + .connection = connection, + .headers = headers, + .extra_headers = extra_headers, + .keep_alive = self.keep_alive, + .redirect_behavior = options.redirect_behavior orelse .unhandled, + }) catch |err| { + self.http_client.connection_pool.release(connection, self.io); + return err; + }; + defer request.deinit(); + + if (options.payload) |payload| { + request.transfer_encoding = .{ .content_length = payload.len }; + var body = try request.sendBodyUnflushed(&.{}); + try body.writer.writeAll(payload); + try body.end(); + try request.connection.?.flush(); + } else { + try request.sendBodiless(); + } + + var response = try request.receiveHead(&.{}); + if (options.max_response_size) |max| { + const body_buf = try self.allocator.alloc(u8, max); + defer self.allocator.free(body_buf); + var writer = std.Io.Writer.fixed(body_buf); + try streamResponseBody(&response, &writer); + return .{ + .status = response.head.status, + .body = try self.allocator.dupe(u8, writer.buffered()), + }; + } + + var aw: std.Io.Writer.Allocating = .init(self.allocator); + defer aw.deinit(); + try streamResponseBody(&response, &aw.writer); + return .{ + .status = response.head.status, + .body = try self.allocator.dupe(u8, aw.written()), + }; + } + pub const FetchOptions = struct { url: []const u8, method: std.http.Method = .GET, @@ -83,6 +181,16 @@ pub const HttpTransport = struct { accept: ?[]const u8 = null, content_type: ?[]const u8 = null, extra_headers: ?[]const std.http.Header = null, + max_response_size: ?usize = null, + redirect_behavior: ?std.http.Client.Request.RedirectBehavior = null, + resolved_connection: ?ResolvedConnection = null, + }; + + pub const ResolvedConnection = struct { + /// Checked address to dial. Currently IPv4 text, which std.Io.net.HostName accepts. + dial_host: []const u8, + /// Original URL host, used by std.http for TLS/SNI and connection identity. + logical_host: []const u8, }; pub const FetchResult = struct { @@ -91,6 +199,24 @@ pub const HttpTransport = struct { }; }; +fn streamResponseBody(response: *std.http.Client.Response, writer: *std.Io.Writer) !void { + var transfer_buffer: [64]u8 = undefined; + var decompress: std.http.Decompress = undefined; + const reader = response.readerDecompressing(&transfer_buffer, &decompress, &.{}); + _ = reader.streamRemaining(writer) catch |err| switch (err) { + error.ReadFailed => return response.bodyErr().?, + error.WriteFailed => return error.ResponseTooLarge, + else => |e| return e, + }; +} + +fn defaultPort(protocol: std.http.Client.Protocol) u16 { + return switch (protocol) { + .plain => 80, + .tls => 443, + }; +} + // === tests === test "transport init/deinit" { @@ -98,3 +224,23 @@ test "transport init/deinit" { var transport = HttpTransport.init(io, std.testing.allocator); defer transport.deinit(); } + +test "transport fixed writer maps overflow to ResponseTooLarge" { + var buf: [0]u8 = .{}; + var writer = std.Io.Writer.fixed(&buf); + try std.testing.expectError(error.WriteFailed, writer.writeAll("x")); +} + +test "resolved transport rejects mismatched logical host" { + const io = std.Options.debug_io; + var transport = HttpTransport.init(io, std.testing.allocator); + defer transport.deinit(); + + try std.testing.expectError(error.ResolvedHostMismatch, transport.fetch(.{ + .url = "http://example.com/", + .resolved_connection = .{ + .dial_host = "127.0.0.1", + .logical_host = "other.example", + }, + })); +} -- 2.51.2