//! Shared vocabulary for every upstream client: endpoint URLs, the four //! disjoint failure groups, the `Client` interface, and response validation. //! //! Everything here except the `Client` vtable is pure. `validateResponse` takes //! two byte slices and returns; `Endpoint.parse` takes text. No socket, no //! clock, no allocator. The transport implementations (DoH, DoT) and the pool //! own all of the `std.Io` work. //! //! The failure classification is the reason this file exists. Health and //! backoff must count only what the peer did wrong: a local `OutOfMemory` says //! nothing about the upstream, `error.Canceled` says nothing at all, and //! `error.BudgetExhausted` says the caller ran out of time before the peer was //! given its interval. The four error sets below are disjoint by construction //! and `group` switches //! over them exhaustively, so a new failure mode cannot silently land in the //! wrong bucket. const std = @import("std"); const packet = @import("../dns/packet.zig"); const name = @import("../dns/name.zig"); /// RFC 1035 §4.2.2: the TCP length prefix is 16-bit, so no DNS message can be /// larger than this on any transport nxdns speaks. pub const max_message_len = 65535; /// RFC 1035 §4.2.2 two-byte big-endian length prefix, shared by every /// stream transport (DoT, plain TCP server, forward-zone TCP fallback). pub const prefix_len = 2; pub fn framePrefix(len: u16) [prefix_len]u8 { var out: [prefix_len]u8 = undefined; std.mem.writeInt(u16, &out, len, .big); return out; } pub fn parsePrefix(bytes: [prefix_len]u8) u16 { return std.mem.readInt(u16, &bytes, .big); } /// The longest DNS name in text form, and so the longest host an endpoint url /// can name. Every consumer of `Endpoint.host` sizes itself from this bound /// rather than re-deriving it: the query log's `upstream` column is built for /// the widest `scheme://host:port` this permits, so a longer host would reach /// storage only as a silently shortened identity. pub const max_host_len = 253; pub const doh_default_port = 443; pub const dot_default_port = 853; // RFC 7858 §3.1 pub const doh_default_path = "/dns-query"; // RFC 8484 §4.1 well-known template pub const Scheme = enum { doh, dot }; /// Borrowed view over the configured URL text; the caller owns the string. pub const Endpoint = struct { scheme: Scheme, /// The original text, for logs and the health API. url: []const u8, /// No brackets, no port, never empty, at most `max_host_len` bytes. Used /// for SNI and certificate verification. host: []const u8, port: u16, /// DoH only; always starts with '/'; `doh_default_path` when absent. A DoT /// endpoint has no request path, so it carries "/" and nothing reads it. path: []const u8, pub const ParseError = error{ UnsupportedScheme, MissingHost, HostTooLong, BadPort, BadUrl }; const doh_prefix = "https://"; const dot_prefix = "tls://"; /// `https://…` => .doh, `tls://…` => .dot (PLAN §9). Accepts `[v6]:port`. /// /// Hand-written rather than a `std.Uri` round trip: `std.Uri` hands back /// percent-encoded components that would need re-decoding into a caller /// buffer, which is a lot of machinery for one household-scale config /// value that is an IP literal or a hostname. pub fn parse(url: []const u8) ParseError!Endpoint { const scheme: Scheme, const rest = if (std.mem.startsWith(u8, url, doh_prefix)) .{ .doh, url[doh_prefix.len..] } else if (std.mem.startsWith(u8, url, dot_prefix)) .{ .dot, url[dot_prefix.len..] } else return error.UnsupportedScheme; const authority, const path = split: { const slash = std.mem.findScalar(u8, rest, '/') orelse break :split .{ rest, "" }; break :split .{ rest[0..slash], rest[slash..] }; }; try rejectDelimiters(authority, "@?#"); try rejectDelimiters(path, "?#"); const host, const port_text = try splitAuthority(authority); if (host.len == 0) return error.MissingHost; // Rejected here rather than tolerated: every identity built from this // endpoint is bounded by `max_host_len`, so a longer host would parse // clean and then be shortened where it is stored or logged. if (host.len > max_host_len) return error.HostTooLong; const port: u16 = if (port_text) |text| blk: { if (text.len == 0) return error.BadPort; break :blk std.fmt.parseInt(u16, text, 10) catch return error.BadPort; } else switch (scheme) { .doh => doh_default_port, .dot => dot_default_port, }; return switch (scheme) { .doh => .{ .scheme = .doh, .url = url, .host = host, .port = port, .path = if (path.len == 0) doh_default_path else path, }, .dot => blk: { // RFC 7858 frames DNS directly on the TLS stream; a path would // be config the transport cannot honour, so it is rejected // rather than ignored. if (path.len != 0 and !std.mem.eql(u8, path, "/")) return error.BadUrl; break :blk .{ .scheme = .dot, .url = url, .host = host, .port = port, .path = "/", }; }, }; } /// `@`, `?` and `#` open a userinfo, query or fragment component. This /// parser implements none of them, so keeping one as literal host or path /// text would let `host` disagree with the authority an RFC 3986 parser /// reads out of the same URL — the name verified against the certificate /// would not be the name dialed. A household config has no use for them, /// so they are a config error rather than something to strip. fn rejectDelimiters(text: []const u8, comptime delimiters: []const u8) ParseError!void { inline for (delimiters) |delimiter| { if (std.mem.findScalar(u8, text, delimiter) != null) return error.BadUrl; } } /// Returns the host without brackets and the port text when one is present. fn splitAuthority(authority: []const u8) ParseError!struct { []const u8, ?[]const u8 } { if (authority.len != 0 and authority[0] == '[') { const close = std.mem.findScalar(u8, authority, ']') orelse return error.BadUrl; const host = authority[1..close]; const tail = authority[close + 1 ..]; if (tail.len == 0) return .{ host, null }; if (tail[0] != ':') return error.BadUrl; return .{ host, tail[1..] }; } const colon = std.mem.findScalar(u8, authority, ':') orelse return .{ authority, null }; return .{ authority[0..colon], authority[colon + 1 ..] }; } }; /// The upstream misbehaved, timed out, or was unreachable. Only these count /// against health. pub const PeerFault = error{ ConnectFailed, TlsFailed, SendFailed, ReceiveFailed, Timeout, /// Unparseable, not a response, or QDCOUNT != 1. BadResponse, /// ID or question does not match the query. ResponseMismatch, /// Does not fit the caller's buffer. ResponseTooLarge, /// DoH: status other than 200. HttpStatus, /// DoH: content-type other than application/dns-message. HttpContentType, }; /// This process ran out of something. Never the upstream's fault, so never /// recorded against an endpoint's health. pub const LocalResource = error{ OutOfMemory, SystemResources, ProcessFdQuotaExceeded, SystemFdQuotaExceeded, /// A caller-supplied buffer cannot hold even a query. BufferTooSmall, Unexpected, }; pub const Cancellation = error{Canceled}; /// The caller's own time ran out before any peer could be given the observation /// interval it was configured to get. Evidence about this process's budget, not /// about any endpoint, so it is never recorded against health — that is the /// whole reason it is not a `PeerFault`. pub const BudgetFault = error{BudgetExhausted}; pub const ExchangeError = PeerFault || LocalResource || Cancellation || BudgetFault; pub const Group = enum { peer_fault, local_resource, cancellation, budget_exhausted }; /// Exhaustive switch over `ExchangeError` — no `else` arm. A new error member /// must break the build here, so no failure can silently land in the wrong /// group. pub fn group(err: ExchangeError) Group { return switch (err) { error.ConnectFailed, error.TlsFailed, error.SendFailed, error.ReceiveFailed, error.Timeout, error.BadResponse, error.ResponseMismatch, error.ResponseTooLarge, error.HttpStatus, error.HttpContentType, => .peer_fault, error.OutOfMemory, error.SystemResources, error.ProcessFdQuotaExceeded, error.SystemFdQuotaExceeded, error.BufferTooSmall, error.Unexpected, => .local_resource, error.Canceled => .cancellation, error.BudgetExhausted => .budget_exhausted, }; } /// Folds a foreign stdlib error into `ExchangeError` when its name matches a /// `LocalResource` or `Cancellation` member; `null` means the caller should /// classify the error as a peer fault from the phase it occurred in. /// /// This is the only place a foreign error set is folded in. Everywhere else /// the call site names the peer fault it means, because the call site is what /// knows whether it was connecting, sending or receiving. pub fn mapLocal(err: anyerror) ?ExchangeError { return switch (err) { error.OutOfMemory => error.OutOfMemory, error.SystemResources => error.SystemResources, error.ProcessFdQuotaExceeded => error.ProcessFdQuotaExceeded, error.SystemFdQuotaExceeded => error.SystemFdQuotaExceeded, error.BufferTooSmall => error.BufferTooSmall, error.Unexpected => error.Unexpected, error.Canceled => error.Canceled, else => null, }; } /// Names the peer fault of the phase the call site is in, unless the error is /// one `mapLocal` claims for this process. Every transport classifies its /// failures through this one function. pub fn mapPhase(err: anyerror, phase: PeerFault) ExchangeError { return mapLocal(err) orelse phase; } /// Closes `target` with cancellation blocked. /// /// A transport's close runs from a `defer` chain that a lost timeout race is /// unwinding. The next cancelable `Io` call in that chain returns /// `error.Canceled` and skips the close, leaking the descriptor, so the close /// swaps cancellation protection for the duration. /// /// `net.Stream` and `net.Socket` close through an `Io`; `tls_client.TlsStream` /// owns the one it was built with and takes none. Both shapes are accepted so /// that one helper covers every close in the transports. pub fn closeBlocked(io: std.Io, target: anytype) void { const prev = io.swapCancelProtection(.blocked); defer _ = io.swapCancelProtection(prev); const Target = @typeInfo(@TypeOf(target)).pointer.child; if (@typeInfo(@TypeOf(Target.close)).@"fn".params.len == 2) { target.close(io); } else { target.close(); } } /// The payload of `f`'s return type, which the race harness requires to be /// `ExchangeError!T`. A raced function with any other error set would let a /// failure reach the pool without passing through `group`. fn RacedPayload(comptime f: anytype) type { const info = @typeInfo(@TypeOf(f)); if (info != .@"fn") @compileError("the race harness needs a function, found " ++ @typeName(@TypeOf(f))); const Return = info.@"fn".return_type orelse @compileError("the race harness needs a function with a concrete return type"); const union_info = switch (@typeInfo(Return)) { .error_union => |u| u, else => @compileError("the race harness needs `ExchangeError!T`, found " ++ @typeName(Return)), }; if (union_info.error_set != ExchangeError) @compileError("the race harness needs `ExchangeError!T`, found " ++ @typeName(Return)); return union_info.payload; } /// Runs `f(args...)` raced against `budget`, and cancels the loser. /// /// No stream read or write in 0.16.0 takes a timeout, so a deadline is a second /// task rather than a socket option. `raceUntilTagged` is the one copy of that /// harness; this is the untagged wrapper for callers that own no deadline and /// only need "a bound on this one operation". /// /// `error.Timeout` means the budget won. A canceled sleep means the whole task /// is being torn down rather than the budget running out, so it stays /// `error.Canceled`. A backend that cannot start a second task is /// `error.SystemResources`, which `group` keeps off the peer's health. pub fn raceWithin( io: std.Io, budget: std.Io.Clock.Duration, comptime f: anytype, args: anytype, ) ExchangeError!RacedPayload(f) { var outcome: RaceOutcome = .completed; return raceUntilTagged(io, .fromNow(io, budget), &outcome, f, args); } /// Which side of the race ended the call. /// /// `completed` means the raced operation itself returned — including when what /// it returned is `error.Timeout`, which is then the peer's own timeout and /// real evidence about that peer. `expired` means the caller's clock ran out /// with the operation still in flight, which is evidence about the budget only. pub const RaceOutcome = enum { completed, expired }; /// `raceWithin` with the two timer origins told apart, and with an ABSOLUTE /// expiry rather than a duration. /// /// The timestamp is the point of the whole function. A duration recomputed from /// a deadline and then slept re-anchors at "now", so every re-race drifts a /// little past the caller's real deadline and a truncated attempt is then /// indistinguishable from a full one. The caller that owns the deadline /// computes the instant once and passes it here. /// /// `outcome` is written before this returns on both racing paths. It is left /// untouched when the race cannot start at all (`error.SystemResources`) or /// when the whole task is being canceled, since neither is an observation about /// this budget; callers initialize it to the value they want in those cases. pub fn raceUntilTagged( io: std.Io, expiry_at: std.Io.Clock.Timestamp, outcome: *RaceOutcome, comptime f: anytype, args: anytype, ) ExchangeError!RacedPayload(f) { const Slot = union(enum) { raced: ExchangeError!RacedPayload(f), expiry: std.Io.Cancelable!void, }; var slots: [2]Slot = undefined; var race: std.Io.Select(Slot) = .init(io, &slots); defer race.cancelDiscard(); race.concurrent(.raced, f, args) catch |err| switch (err) { error.ConcurrencyUnavailable => return error.SystemResources, }; race.concurrent(.expiry, expire, .{ io, expiry_at }) catch |err| switch (err) { error.ConcurrencyUnavailable => return error.SystemResources, }; switch (try race.await()) { .raced => |result| { outcome.* = .completed; return result; }, .expiry => |result| { try result; outcome.* = .expired; return error.Timeout; }, } } fn expire(io: std.Io, expiry_at: std.Io.Clock.Timestamp) std.Io.Cancelable!void { return expiry_at.wait(io); } /// A thing that sends one DNS message and returns one validated DNS message. /// Implemented by DohClient, DotClient, Pool, and test fakes. pub const Client = struct { ptr: *anyopaque, exchangeFn: *const fn ( ptr: *anyopaque, io: std.Io, query: []const u8, response_buf: []u8, selected: *?[]const u8, ) ExchangeError![]u8, /// Returns a prefix of `response_buf`. The returned message has already /// passed `validateResponse` against `query`. /// /// `selected` names the resolver the exchange used. A single-endpoint /// implementation (DoH, DoT, the forward client, test fakes) may write it /// *before* each attempt: it has one resolver and records no health, so /// "the one I tried" is an honest answer even for a failure, and a SERVFAIL /// row without its resolver explains nothing. /// /// `Pool` is stricter, and documents the rule on `Pool.exchange`: it names /// only endpoints whose outcome it recorded, so a query that ran out of /// budget blames nobody and may leave this `null`. Both are within this /// contract — the guarantee here is that a non-null value names a resolver /// this exchange really used, never that a failure leaves one behind. /// /// The slice must outlive the call; every implementation borrows storage it /// owns for at least the query's duration. Callers initialize it to null: a /// `null` after the call means no resolver is being reported. pub fn exchange( self: Client, io: std.Io, query: []const u8, response_buf: []u8, selected: *?[]const u8, ) ExchangeError![]u8 { return self.exchangeFn(self.ptr, io, query, response_buf, selected); } }; pub const ValidateError = error{ BadResponse, ResponseMismatch }; /// RFC 9619: exactly one question on both sides. RFC 4343: names compare /// case-insensitively, so a case-mangling (0x20) upstream still matches. Does /// not inspect the answer section — content policy is not this layer's /// business. pub fn validateResponse(query: []const u8, response: []const u8) ValidateError!void { const resp = packet.parse(response) catch return error.BadResponse; if (!resp.header.flags.qr) return error.BadResponse; if (resp.header.qdcount != 1) return error.BadResponse; // A malformed query here is a bug in this process, not in the upstream. // It still may not crash, so it reports the same structural error. const req = packet.parse(query) catch return error.BadResponse; if (req.header.qdcount != 1) return error.BadResponse; if (resp.header.id != req.header.id) return error.ResponseMismatch; const rq = packet.firstQuestion(resp) orelse return error.BadResponse; const qq = packet.firstQuestion(req) orelse return error.BadResponse; if (rq.qtype != qq.qtype) return error.ResponseMismatch; if (rq.qclass != qq.qclass) return error.ResponseMismatch; if (!name.eqlIgnoreCase(rq.name, qq.name)) return error.ResponseMismatch; } const testing = std.testing; test "framePrefix writes the length big-endian" { try testing.expectEqualSlices(u8, &.{ 0x00, 0x00 }, &framePrefix(0)); try testing.expectEqualSlices(u8, &.{ 0x00, 0x1d }, &framePrefix(29)); try testing.expectEqualSlices(u8, &.{ 0x01, 0x00 }, &framePrefix(256)); try testing.expectEqualSlices(u8, &.{ 0xff, 0xff }, &framePrefix(65535)); } test "parsePrefix reads the length big-endian" { try testing.expectEqual(@as(u16, 0), parsePrefix(.{ 0x00, 0x00 })); try testing.expectEqual(@as(u16, 29), parsePrefix(.{ 0x00, 0x1d })); try testing.expectEqual(@as(u16, 256), parsePrefix(.{ 0x01, 0x00 })); try testing.expectEqual(@as(u16, 65535), parsePrefix(.{ 0xff, 0xff })); } test "framePrefix and parsePrefix round-trip" { for ([_]u16{ 0, 1, 12, 512, 4096, 65534, 65535 }) |len| { try testing.expectEqual(len, parsePrefix(framePrefix(len))); } } test "the prefix ceiling is the message ceiling" { try testing.expectEqual(@as(u16, max_message_len), parsePrefix(.{ 0xff, 0xff })); } test "parse a DoH url with an explicit path" { const e = try Endpoint.parse("https://cloudflare-dns.com/dns-query"); try testing.expectEqual(Scheme.doh, e.scheme); try testing.expectEqualStrings("cloudflare-dns.com", e.host); try testing.expectEqual(@as(u16, 443), e.port); try testing.expectEqualStrings("/dns-query", e.path); try testing.expectEqualStrings("https://cloudflare-dns.com/dns-query", e.url); } test "parse preserves a non-default DoH path" { const e = try Endpoint.parse("https://dns.example/x"); try testing.expectEqualStrings("dns.example", e.host); try testing.expectEqualStrings("/x", e.path); } test "parse defaults the DoH path" { const e = try Endpoint.parse("https://dns.example"); try testing.expectEqualStrings("dns.example", e.host); try testing.expectEqual(@as(u16, 443), e.port); try testing.expectEqualStrings(doh_default_path, e.path); } test "parse a DoT url with an explicit port" { const e = try Endpoint.parse("tls://dns.google:853"); try testing.expectEqual(Scheme.dot, e.scheme); try testing.expectEqualStrings("dns.google", e.host); try testing.expectEqual(@as(u16, 853), e.port); } test "parse defaults the DoT port" { const e = try Endpoint.parse("tls://dns.google"); try testing.expectEqual(Scheme.dot, e.scheme); try testing.expectEqualStrings("dns.google", e.host); try testing.expectEqual(@as(u16, dot_default_port), e.port); } test "parse strips IPv6 brackets and keeps the port" { const e = try Endpoint.parse("https://[2606:4700:4700::1111]:8443/dns-query"); try testing.expectEqualStrings("2606:4700:4700::1111", e.host); try testing.expectEqual(@as(u16, 8443), e.port); try testing.expectEqualStrings("/dns-query", e.path); } test "parse rejects an unsupported scheme" { try testing.expectError(error.UnsupportedScheme, Endpoint.parse("udp://1.1.1.1:53")); try testing.expectError(error.UnsupportedScheme, Endpoint.parse("http://dns.example/")); try testing.expectError(error.UnsupportedScheme, Endpoint.parse("1.1.1.1")); } test "parse rejects an empty host" { try testing.expectError(error.MissingHost, Endpoint.parse("https://")); try testing.expectError(error.MissingHost, Endpoint.parse("https://:443/dns-query")); try testing.expectError(error.MissingHost, Endpoint.parse("tls://")); } test "parse takes a host at the length bound and rejects one past it" { // Four labels, the widest a 253-byte name allows: 3 * (63 + 1) + 61. const at_bound = ("a" ** 63 ++ ".") ** 3 ++ "a" ** 61; comptime std.debug.assert(at_bound.len == max_host_len); const accepted = try Endpoint.parse("https://" ++ at_bound ++ "/dns-query"); try testing.expectEqualStrings(at_bound, accepted.host); // One byte more is one byte no consumer of `host` has room for. try testing.expectError( error.HostTooLong, Endpoint.parse("https://" ++ at_bound ++ "a/dns-query"), ); // The bound is on the host alone, so a port and a path do not spend it, // and the bracketed form is measured with the brackets removed. try testing.expectError( error.HostTooLong, Endpoint.parse("tls://" ++ at_bound ++ "a:853"), ); try testing.expectError( error.HostTooLong, Endpoint.parse("https://[" ++ at_bound ++ "a]:8443/dns-query"), ); } test "parse rejects a bad port" { try testing.expectError(error.BadPort, Endpoint.parse("https://h:99999/")); try testing.expectError(error.BadPort, Endpoint.parse("https://h:/")); try testing.expectError(error.BadPort, Endpoint.parse("https://h:abc/")); try testing.expectError(error.BadPort, Endpoint.parse("tls://[::1]:70000")); } test "parse rejects a malformed url" { try testing.expectError(error.BadUrl, Endpoint.parse("https://[::1")); try testing.expectError(error.BadUrl, Endpoint.parse("https://[::1]x")); // A DoT endpoint has no request path. try testing.expectError(error.BadUrl, Endpoint.parse("tls://dns.google/dns-query")); // A bare trailing slash is not a path, so it is accepted. try testing.expectEqual(Scheme.dot, (try Endpoint.parse("tls://dns.google/")).scheme); } test "parse rejects userinfo" { try testing.expectError(error.BadUrl, Endpoint.parse("https://user@host/")); try testing.expectError(error.BadUrl, Endpoint.parse("https://user:pass@host/dns-query")); try testing.expectError(error.BadUrl, Endpoint.parse("tls://user@dns.google:853")); } test "parse rejects a query string" { try testing.expectError(error.BadUrl, Endpoint.parse("https://host?x")); try testing.expectError(error.BadUrl, Endpoint.parse("https://host/dns-query?x=1")); try testing.expectError(error.BadUrl, Endpoint.parse("tls://dns.google/?x")); } test "parse rejects a fragment" { try testing.expectError(error.BadUrl, Endpoint.parse("https://host#f")); try testing.expectError(error.BadUrl, Endpoint.parse("https://host/dns-query#f")); try testing.expectError(error.BadUrl, Endpoint.parse("tls://dns.google#f")); } test "the four error groups are disjoint" { const sets = .{ PeerFault, LocalResource, Cancellation, BudgetFault }; inline for (sets, 0..) |a, i| { inline for (sets, 0..) |b, j| { if (i >= j) continue; inline for (@typeInfo(a).error_set.?) |member_a| { inline for (@typeInfo(b).error_set.?) |member_b| { try testing.expect(!std.mem.eql(u8, member_a.name, member_b.name)); } } } } // Every member of the union belongs to exactly one group, which `group` // proves by being an exhaustive switch. Assert the counts line up so a // member added to two sets at once cannot pass unnoticed. const total = @typeInfo(PeerFault).error_set.?.len + @typeInfo(LocalResource).error_set.?.len + @typeInfo(Cancellation).error_set.?.len + @typeInfo(BudgetFault).error_set.?.len; try testing.expectEqual(total, @typeInfo(ExchangeError).error_set.?.len); } test "group classifies each member" { try testing.expectEqual(Group.peer_fault, group(error.Timeout)); try testing.expectEqual(Group.peer_fault, group(error.HttpContentType)); try testing.expectEqual(Group.local_resource, group(error.OutOfMemory)); try testing.expectEqual(Group.local_resource, group(error.BufferTooSmall)); try testing.expectEqual(Group.cancellation, group(error.Canceled)); try testing.expectEqual(Group.budget_exhausted, group(error.BudgetExhausted)); } test "mapLocal folds only local and cancellation errors" { try testing.expectEqual(@as(?ExchangeError, error.OutOfMemory), mapLocal(error.OutOfMemory)); try testing.expectEqual(@as(?ExchangeError, error.Canceled), mapLocal(error.Canceled)); try testing.expectEqual(@as(?ExchangeError, error.Unexpected), mapLocal(error.Unexpected)); try testing.expectEqual(@as(?ExchangeError, null), mapLocal(error.ConnectionRefused)); try testing.expectEqual(@as(?ExchangeError, null), mapLocal(error.TlsInitializationFailed)); } test "mapPhase names the phase unless the error is this process's own" { try testing.expectEqual(ExchangeError.Canceled, mapPhase(error.Canceled, error.ReceiveFailed)); try testing.expectEqual( ExchangeError.SystemResources, mapPhase(error.SystemResources, error.ConnectFailed), ); try testing.expectEqual( ExchangeError.ConnectFailed, mapPhase(error.ConnectionRefused, error.ConnectFailed), ); try testing.expectEqual( Group.local_resource, group(mapPhase(error.OutOfMemory, error.SendFailed)), ); } /// Stands in for whatever the pool or the forward client races. The variants /// are the three ways such a task ends: in time, too late, or with a failure of /// its own. fn racedReply(io: std.Io, delay_ms: i64, result: ExchangeError!usize) ExchangeError!usize { if (delay_ms != 0) { const duration: std.Io.Clock.Duration = .{ .raw = .fromMilliseconds(delay_ms), .clock = .awake }; try duration.sleep(io); } return result; } test "raceWithin returns the raced value when it finishes inside the budget" { var threaded: std.Io.Threaded = .init(testing.allocator, .{}); defer threaded.deinit(); const io = threaded.io(); const budget: std.Io.Clock.Duration = .{ .raw = .fromSeconds(30), .clock = .awake }; const len = try raceWithin(io, budget, racedReply, .{ io, 0, @as(ExchangeError!usize, 7) }); try testing.expectEqual(@as(usize, 7), len); } test "raceWithin returns Timeout when the budget wins" { var threaded: std.Io.Threaded = .init(testing.allocator, .{}); defer threaded.deinit(); const io = threaded.io(); const budget: std.Io.Clock.Duration = .{ .raw = .fromMilliseconds(20), .clock = .awake }; const started = std.Io.Clock.awake.now(io); try testing.expectError( error.Timeout, raceWithin(io, budget, racedReply, .{ io, 30_000, @as(ExchangeError!usize, 7) }), ); const elapsed_ns = std.Io.Clock.awake.now(io).nanoseconds - started.nanoseconds; // Far under the raced task's own sleep, so the budget is provably what // ended the call rather than the task finishing on its own. try testing.expect(elapsed_ns < @as(i96, 5) * std.time.ns_per_s); } test "raceWithin passes the raced task's own failure through" { var threaded: std.Io.Threaded = .init(testing.allocator, .{}); defer threaded.deinit(); const io = threaded.io(); const budget: std.Io.Clock.Duration = .{ .raw = .fromSeconds(30), .clock = .awake }; // A local resource error keeps its group: the race must not turn it into // the peer fault the budget would have produced. const failed = raceWithin(io, budget, racedReply, .{ io, 0, @as(ExchangeError!usize, error.OutOfMemory), }); try testing.expectError(error.OutOfMemory, failed); try testing.expectError( error.ConnectFailed, raceWithin(io, budget, racedReply, .{ io, 0, @as(ExchangeError!usize, error.ConnectFailed) }), ); } test "raceUntilTagged tells a leaf Timeout apart from an expiry" { var threaded: std.Io.Threaded = .init(testing.allocator, .{}); defer threaded.deinit(); const io = threaded.io(); // The leaf's own timeout: it returned, so the peer really did time out and // the outcome is `completed` even though the error is the same one an // expiry produces. var outcome: RaceOutcome = .expired; const far: std.Io.Clock.Timestamp = .fromNow(io, .{ .raw = .fromSeconds(30), .clock = .awake }); try testing.expectError( error.Timeout, raceUntilTagged(io, far, &outcome, racedReply, .{ io, 0, @as(ExchangeError!usize, error.Timeout), }), ); try testing.expectEqual(RaceOutcome.completed, outcome); // The expiry side, distinguishable only through the tag. outcome = .completed; const soon: std.Io.Clock.Timestamp = .fromNow(io, .{ .raw = .fromMilliseconds(20), .clock = .awake }); try testing.expectError( error.Timeout, raceUntilTagged(io, soon, &outcome, racedReply, .{ io, 30_000, @as(ExchangeError!usize, 7), }), ); try testing.expectEqual(RaceOutcome.expired, outcome); } test "raceUntilTagged tags a successful completion" { var threaded: std.Io.Threaded = .init(testing.allocator, .{}); defer threaded.deinit(); const io = threaded.io(); var outcome: RaceOutcome = .expired; const far: std.Io.Clock.Timestamp = .fromNow(io, .{ .raw = .fromSeconds(30), .clock = .awake }); const len = try raceUntilTagged(io, far, &outcome, racedReply, .{ io, 0, @as(ExchangeError!usize, 7), }); try testing.expectEqual(@as(usize, 7), len); try testing.expectEqual(RaceOutcome.completed, outcome); } test "raceUntilTagged honours an expiry already in the past" { var threaded: std.Io.Threaded = .init(testing.allocator, .{}); defer threaded.deinit(); const io = threaded.io(); // A deadline the caller has already spent: the expiry wins immediately // rather than re-anchoring a duration at "now" and granting a fresh budget. const past = std.Io.Clock.Timestamp.now(io, .awake) .subDuration(.{ .raw = .fromSeconds(1), .clock = .awake }); var outcome: RaceOutcome = .completed; const started = std.Io.Clock.awake.now(io); try testing.expectError( error.Timeout, raceUntilTagged(io, past, &outcome, racedReply, .{ io, 30_000, @as(ExchangeError!usize, 7), }), ); try testing.expectEqual(RaceOutcome.expired, outcome); const elapsed_ns = std.Io.Clock.awake.now(io).nanoseconds - started.nanoseconds; try testing.expect(elapsed_ns < @as(i96, 5) * std.time.ns_per_s); } test "closeBlocked closes a target of either close shape" { // The two shapes the transports use: a socket or a plain stream, which // closes through the `Io`, and a `TlsStream`, which owns the one it was // built with. Dispatching on the wrong one is a compile error, so // instantiating both is the check. // // That the close runs with cancellation blocked is not asserted here: the // Threaded backend's `swapCancelProtection` is a no-op off one of its own // task threads, so a unit test cannot observe the state it sets. const WithIo = struct { closed: bool = false, fn close(self: *@This(), io: std.Io) void { _ = io; self.closed = true; } }; const WithoutIo = struct { closed: bool = false, fn close(self: *@This()) void { self.closed = true; } }; var threaded: std.Io.Threaded = .init(testing.allocator, .{}); defer threaded.deinit(); const io = threaded.io(); var with_io: WithIo = .{}; closeBlocked(io, &with_io); try testing.expect(with_io.closed); var without_io: WithoutIo = .{}; closeBlocked(io, &without_io); try testing.expect(without_io.closed); } test "a fake client satisfies the Client interface" { const Fake = struct { calls: usize = 0, fn exchangeFn( ptr: *anyopaque, io: std.Io, query: []const u8, response_buf: []u8, selected: *?[]const u8, ) ExchangeError![]u8 { _ = io; selected.* = "fake://echo"; const self: *@This() = @ptrCast(@alignCast(ptr)); self.calls += 1; if (query.len > response_buf.len) return error.ResponseTooLarge; @memcpy(response_buf[0..query.len], query); return response_buf[0..query.len]; } fn client(self: *@This()) Client { return .{ .ptr = self, .exchangeFn = exchangeFn }; } }; var fake: Fake = .{}; var buf: [16]u8 = undefined; var selected: ?[]const u8 = null; const echoed = try fake.client().exchange(undefined, "hello", &buf, &selected); try testing.expectEqualStrings("hello", echoed); try testing.expectEqual(@as(usize, 1), fake.calls); try testing.expectEqualStrings("fake://echo", selected.?); } /// A query for example.com A: id 0x1234, RD set, one question. const query_bytes = "\x12\x34\x01\x00\x00\x01\x00\x00\x00\x00\x00\x00" ++ "\x07example\x03com\x00\x00\x01\x00\x01"; /// The matching response: the question echoed plus one A record. const response_bytes = "\x12\x34\x81\x80\x00\x01\x00\x01\x00\x00\x00\x00" ++ "\x07example\x03com\x00\x00\x01\x00\x01" ++ "\xc0\x0c\x00\x01\x00\x01\x00\x00\x01\x2c\x00\x04\x5d\xb8\xd8\x22"; test "validateResponse accepts a matching pair" { try validateResponse(query_bytes, response_bytes); } test "validateResponse accepts a mixed-case question name" { const mangled = "\x12\x34\x81\x80\x00\x01\x00\x01\x00\x00\x00\x00" ++ "\x07ExAmPlE\x03CoM\x00\x00\x01\x00\x01" ++ "\xc0\x0c\x00\x01\x00\x01\x00\x00\x01\x2c\x00\x04\x5d\xb8\xd8\x22"; try validateResponse(query_bytes, mangled); } test "validateResponse rejects a wrong id" { var bytes: [response_bytes.len]u8 = response_bytes.*; packet.setId(&bytes, 0x4321); try testing.expectError(error.ResponseMismatch, validateResponse(query_bytes, &bytes)); } test "validateResponse rejects a different qtype" { const aaaa = "\x12\x34\x81\x80\x00\x01\x00\x00\x00\x00\x00\x00" ++ "\x07example\x03com\x00\x00\x1c\x00\x01"; try testing.expectError(error.ResponseMismatch, validateResponse(query_bytes, aaaa)); } test "validateResponse rejects a different question name" { const other = "\x12\x34\x81\x80\x00\x01\x00\x00\x00\x00\x00\x00" ++ "\x07example\x03org\x00\x00\x01\x00\x01"; try testing.expectError(error.ResponseMismatch, validateResponse(query_bytes, other)); } test "validateResponse rejects a response with QR clear" { try testing.expectError(error.BadResponse, validateResponse(query_bytes, query_bytes)); } test "validateResponse rejects a response with QDCOUNT 0" { const no_question = "\x12\x34\x81\x80\x00\x00\x00\x00\x00\x00\x00\x00"; try testing.expectError(error.BadResponse, validateResponse(query_bytes, no_question)); } test "validateResponse rejects a response with QDCOUNT 2" { const two = "\x12\x34\x81\x80\x00\x02\x00\x00\x00\x00\x00\x00" ++ "\x07example\x03com\x00\x00\x01\x00\x01" ++ "\x07example\x03com\x00\x00\x01\x00\x01"; try testing.expectError(error.BadResponse, validateResponse(query_bytes, two)); } test "validateResponse rejects truncated garbage" { try testing.expectError(error.BadResponse, validateResponse(query_bytes, "\x12\x34\x81")); try testing.expectError(error.BadResponse, validateResponse(query_bytes, "")); try testing.expectError( error.BadResponse, validateResponse(query_bytes, response_bytes[0 .. response_bytes.len - 3]), ); } test "validateResponse rejects a query that does not carry exactly one question" { const no_question = "\x12\x34\x01\x00\x00\x00\x00\x00\x00\x00\x00\x00"; try testing.expectError(error.BadResponse, validateResponse(no_question, response_bytes)); try testing.expectError(error.BadResponse, validateResponse("\x12\x34", response_bytes)); }