diff --git a/src/internal/xrpc/transport.zig b/src/internal/xrpc/transport.zig index 0c824c3..d4e6f0b 100644 --- a/src/internal/xrpc/transport.zig +++ b/src/internal/xrpc/transport.zig @@ -60,49 +60,40 @@ pub const HttpTransport = struct { 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(); + return try self.fetchUrl(options, headers, extra_buf[0..extra_count]); + } - const result = try self.http_client.fetch(.{ - .location = .{ .url = options.url }, - .response_writer = &aw.writer, - .method = options.method, - .payload = options.payload, + fn fetchUrl( + self: *HttpTransport, + options: FetchOptions, + headers: std.http.Client.Request.Headers, + extra_headers: []const std.http.Header, + ) !FetchResult { + const uri = try std.Uri.parse(options.url); + const redirect_behavior = redirectBehavior(options); + var request = try self.http_client.request(options.method, uri, .{ .headers = headers, - .extra_headers = extra_buf[0..extra_count], + .extra_headers = extra_headers, .keep_alive = self.keep_alive, - .redirect_behavior = options.redirect_behavior, + .redirect_behavior = redirect_behavior, }); + defer request.deinit(); - return .{ - .status = result.status, - .body = try self.allocator.dupe(u8, aw.written()), - }; + 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(); + } + + const redirect_buffer = try self.allocRedirectBuffer(redirect_behavior); + defer self.freeRedirectBuffer(redirect_buffer, redirect_behavior); + + var response = try request.receiveHead(redirect_buffer); + return try self.readFetchResult(options, &response); } fn fetchResolved( @@ -130,12 +121,13 @@ pub const HttpTransport = struct { .proxied_host = logical_host, }); + const redirect_behavior = redirectBehavior(options); 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, + .redirect_behavior = redirect_behavior, }) catch |err| { self.http_client.connection_pool.release(connection, self.io); return err; @@ -152,24 +144,48 @@ pub const HttpTransport = struct { try request.sendBodiless(); } - var response = try request.receiveHead(&.{}); + const redirect_buffer = try self.allocRedirectBuffer(redirect_behavior); + defer self.freeRedirectBuffer(redirect_buffer, redirect_behavior); + + var response = try request.receiveHead(redirect_buffer); + return try self.readFetchResult(options, &response); + } + + fn allocRedirectBuffer(self: *HttpTransport, behavior: std.http.Client.Request.RedirectBehavior) ![]u8 { + if (behavior == .unhandled) return &.{}; + return try self.allocator.alloc(u8, 8 * 1024); + } + + fn freeRedirectBuffer(self: *HttpTransport, buffer: []u8, behavior: std.http.Client.Request.RedirectBehavior) void { + if (behavior != .unhandled) self.allocator.free(buffer); + } + + fn readFetchResult( + self: *HttpTransport, + options: FetchOptions, + response: *std.http.Client.Response, + ) !FetchResult { + const rate_limit = RateLimitHeaders.fromResponseHead(response.head); + 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); + try streamResponseBody(response, &writer); return .{ .status = response.head.status, .body = try self.allocator.dupe(u8, writer.buffered()), + .rate_limit = rate_limit, }; } var aw: std.Io.Writer.Allocating = .init(self.allocator); defer aw.deinit(); - try streamResponseBody(&response, &aw.writer); + try streamResponseBody(response, &aw.writer); return .{ .status = response.head.status, .body = try self.allocator.dupe(u8, aw.written()), + .rate_limit = rate_limit, }; } @@ -196,9 +212,50 @@ pub const HttpTransport = struct { pub const FetchResult = struct { status: std.http.Status, body: []u8, + rate_limit: RateLimitHeaders = .{}, + }; + + pub const RateLimitHeaders = struct { + limit: ?u64 = null, + remaining: ?u64 = null, + reset: ?u64 = null, + retry_after: ?u64 = null, + + pub fn fromResponseHead(head: std.http.Client.Response.Head) RateLimitHeaders { + var result: RateLimitHeaders = .{}; + var it = head.iterateHeaders(); + while (it.next()) |header| { + if (std.ascii.eqlIgnoreCase(header.name, "ratelimit-limit")) { + result.limit = parseHeaderInt(header.value); + } else if (std.ascii.eqlIgnoreCase(header.name, "ratelimit-remaining")) { + result.remaining = parseHeaderInt(header.value); + } else if (std.ascii.eqlIgnoreCase(header.name, "ratelimit-reset")) { + result.reset = parseHeaderInt(header.value); + } else if (std.ascii.eqlIgnoreCase(header.name, "retry-after")) { + result.retry_after = parseHeaderInt(header.value); + } + } + return result; + } + + pub fn isEmpty(self: RateLimitHeaders) bool { + return self.limit == null and self.remaining == null and self.reset == null and self.retry_after == null; + } }; }; +fn redirectBehavior(options: HttpTransport.FetchOptions) std.http.Client.Request.RedirectBehavior { + return options.redirect_behavior orelse if (options.payload == null) + std.http.Client.Request.RedirectBehavior.init(3) + else + .unhandled; +} + +fn parseHeaderInt(value: []const u8) ?u64 { + const trimmed = std.mem.trim(u8, value, " \t"); + return std.fmt.parseInt(u64, trimmed, 10) catch null; +} + fn streamResponseBody(response: *std.http.Client.Response, writer: *std.Io.Writer) !void { var transfer_buffer: [64]u8 = undefined; var decompress: std.http.Decompress = undefined; @@ -244,3 +301,20 @@ test "resolved transport rejects mismatched logical host" { }, })); } + +test "transport parses rate limit headers" { + const response_bytes = "HTTP/1.1 429 Too Many Requests\r\n" ++ + "RateLimit-Limit: 3000\r\n" ++ + "ratelimit-remaining: 0\r\n" ++ + "RateLimit-Reset: 1710000000\r\n" ++ + "Retry-After: 2\r\n\r\n"; + + const head = try std.http.Client.Response.Head.parse(response_bytes); + const headers = HttpTransport.RateLimitHeaders.fromResponseHead(head); + + try std.testing.expectEqual(@as(?u64, 3000), headers.limit); + try std.testing.expectEqual(@as(?u64, 0), headers.remaining); + try std.testing.expectEqual(@as(?u64, 1710000000), headers.reset); + try std.testing.expectEqual(@as(?u64, 2), headers.retry_after); + try std.testing.expect(!headers.isEmpty()); +} diff --git a/src/internal/xrpc/xrpc.zig b/src/internal/xrpc/xrpc.zig index 8614b1a..832438a 100644 --- a/src/internal/xrpc/xrpc.zig +++ b/src/internal/xrpc/xrpc.zig @@ -8,6 +8,7 @@ const std = @import("std"); const Nsid = @import("../syntax/nsid.zig").Nsid; const HttpTransport = @import("transport.zig").HttpTransport; +const json_helpers = @import("json.zig"); pub const XrpcClient = struct { allocator: std.mem.Allocator, @@ -55,6 +56,30 @@ pub const XrpcClient = struct { return try self.doRequest(url, body); } + pub fn queryChecked( + self: *XrpcClient, + nsid: Nsid, + params: ?std.StringHashMap([]const u8), + retry_policy: RetryPolicy, + ) !Result { + const url = try self.buildUrl(nsid, params); + defer self.allocator.free(url); + + return try self.requestCheckedUrl(url, null, retry_policy); + } + + pub fn procedureChecked( + self: *XrpcClient, + nsid: Nsid, + body: ?[]const u8, + retry_policy: RetryPolicy, + ) !Result { + const url = try self.buildUrl(nsid, null); + defer self.allocator.free(url); + + return try self.requestCheckedUrl(url, body, retry_policy); + } + fn buildUrl(self: *XrpcClient, nsid: Nsid, params: ?std.StringHashMap([]const u8)) ![]u8 { var url: std.ArrayList(u8) = .empty; errdefer url.deinit(self.allocator); @@ -103,13 +128,43 @@ pub const XrpcClient = struct { .allocator = self.allocator, .status = result.status, .body = result.body, + .rate_limit = result.rate_limit, }; } + fn requestCheckedUrl(self: *XrpcClient, 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)) { + return err; + } + try retry_policy.sleepBeforeRetry(self.transport.io, attempt, null); + continue; + }; + + if (response.ok()) { + return .{ .ok = response }; + } + + if (attempt + 1 < attempts and retry_policy.isRetryableStatus(response.status)) { + const rate_limit = response.rate_limit; + response.deinit(); + try retry_policy.sleepBeforeRetry(self.transport.io, attempt, rate_limit); + continue; + } + + return .{ .err = try XrpcError.fromResponse(response) }; + } + } + pub const Response = struct { allocator: std.mem.Allocator, status: std.http.Status, body: []u8, + rate_limit: HttpTransport.RateLimitHeaders = .{}, pub fn deinit(self: *Response) void { self.allocator.free(self.body); @@ -117,7 +172,7 @@ pub const XrpcClient = struct { /// check if request succeeded pub fn ok(self: Response) bool { - return self.status == .ok; + return self.status.class() == .success; } /// parse body as json @@ -125,8 +180,148 @@ pub const XrpcClient = struct { return try std.json.parseFromSlice(std.json.Value, self.allocator, self.body, .{}); } }; + + pub const Result = union(enum) { + ok: Response, + err: XrpcError, + + pub fn deinit(self: *Result) void { + switch (self.*) { + .ok => |*response| response.deinit(), + .err => |*xrpc_error| xrpc_error.deinit(), + } + } + }; + + pub const XrpcError = struct { + allocator: std.mem.Allocator, + status: std.http.Status, + error_name: ?[]u8 = null, + message: ?[]u8 = null, + body: []u8, + rate_limit: HttpTransport.RateLimitHeaders = .{}, + + pub fn fromResponse(response: Response) !XrpcError { + var result: XrpcError = .{ + .allocator = response.allocator, + .status = response.status, + .body = response.body, + .rate_limit = response.rate_limit, + }; + errdefer result.deinit(); + + var parsed = std.json.parseFromSlice(std.json.Value, response.allocator, response.body, .{}) catch return result; + defer parsed.deinit(); + + if (json_helpers.getString(parsed.value, "error")) |name| { + result.error_name = try response.allocator.dupe(u8, name); + } + if (json_helpers.getString(parsed.value, "message")) |message| { + result.message = try response.allocator.dupe(u8, message); + } + + return result; + } + + pub fn deinit(self: *XrpcError) void { + if (self.error_name) |name| self.allocator.free(name); + if (self.message) |message| self.allocator.free(message); + self.allocator.free(self.body); + } + }; + + pub const RetryPolicy = struct { + max_attempts: u8 = 3, + base_delay_ms: u64 = 500, + max_delay_ms: u64 = 30_000, + jitter_percent: u8 = 20, + retry_transient_errors: bool = true, + + pub fn none() RetryPolicy { + return .{ .max_attempts = 1 }; + } + + pub fn isRetryableStatus(_: RetryPolicy, status: std.http.Status) bool { + return switch (@intFromEnum(status)) { + 429, 500, 502, 503, 504 => true, + else => false, + }; + } + + pub fn delayMillis(self: RetryPolicy, attempt: u8, rate_limit: ?HttpTransport.RateLimitHeaders) u64 { + return self.delayMillisAt(attempt, rate_limit, null); + } + + pub fn delayMillisAt( + self: RetryPolicy, + attempt: u8, + rate_limit: ?HttpTransport.RateLimitHeaders, + now_unix_seconds: ?u64, + ) u64 { + if (rate_limit) |headers| { + if (headers.retry_after) |seconds| { + const milliseconds = std.math.mul(u64, seconds, std.time.ms_per_s) catch return self.max_delay_ms; + return @min(milliseconds, self.max_delay_ms); + } + if (now_unix_seconds) |now| { + if (headers.reset) |reset| { + if (reset > now) { + const seconds = reset - now; + const milliseconds = std.math.mul(u64, seconds, std.time.ms_per_s) catch return self.max_delay_ms; + return @min(milliseconds, self.max_delay_ms); + } + } + } + } + + const shift: u6 = @intCast(@min(attempt, 16)); + const base = self.base_delay_ms * (@as(u64, 1) << shift); + return @min(base, self.max_delay_ms); + } + + pub fn sleepBeforeRetry(self: RetryPolicy, io: std.Io, attempt: u8, rate_limit: ?HttpTransport.RateLimitHeaders) !void { + const now_seconds: ?u64 = if (rate_limit != null) + @intCast(@max(@as(i64, 0), std.Io.Clock.real.now(io).toSeconds())) + else + null; + var delay_ms = self.delayMillisAt(attempt, rate_limit, now_seconds); + delay_ms = self.jitteredDelayMillis(io, delay_ms); + if (delay_ms == 0) return; + try io.sleep(std.Io.Duration.fromMilliseconds(@intCast(delay_ms)), .awake); + } + + fn jitteredDelayMillis(self: RetryPolicy, io: std.Io, delay_ms: u64) u64 { + if (delay_ms == 0 or self.jitter_percent == 0) return delay_ms; + + const spread = delay_ms * @as(u64, self.jitter_percent) / 100; + if (spread == 0) return delay_ms; + + var source: std.Random.IoSource = .{ .io = io }; + const random = source.interface(); + const min = delay_ms - spread; + const max = std.math.add(u64, delay_ms, spread) catch std.math.maxInt(u64); + return @min(random.intRangeAtMost(u64, min, max), self.max_delay_ms); + } + }; }; +fn isRetryableTransportError(err: anyerror) bool { + return switch (err) { + error.ConnectionRefused, + error.ConnectionResetByPeer, + error.HostUnreachable, + error.NetworkUnreachable, + error.NetworkDown, + error.Timeout, + error.TlsInitializationFailed, + error.Unexpected, + error.ReadFailed, + error.WriteFailed, + => true, + else => false, + }; +} + // === tests === test "build url without params" { @@ -155,3 +350,56 @@ test "build url with params" { try std.testing.expect(std.mem.startsWith(u8, url, "https://bsky.social/xrpc/app.bsky.actor.getProfile?")); try std.testing.expect(std.mem.indexOf(u8, url, "actor=did%3Aplc%3Atest123") != null); } + +test "xrpc error parses atproto error envelope and rate limits" { + const body = try std.testing.allocator.dupe(u8, + \\{"error":"RateLimitExceeded","message":"slow down"} + ); + + const response: XrpcClient.Response = .{ + .allocator = std.testing.allocator, + .status = .too_many_requests, + .body = body, + .rate_limit = .{ + .limit = 3000, + .remaining = 0, + .reset = 1710000000, + .retry_after = 2, + }, + }; + + var xrpc_error = try XrpcClient.XrpcError.fromResponse(response); + defer xrpc_error.deinit(); + + try std.testing.expectEqual(.too_many_requests, xrpc_error.status); + try std.testing.expectEqualStrings("RateLimitExceeded", xrpc_error.error_name.?); + try std.testing.expectEqualStrings("slow down", xrpc_error.message.?); + try std.testing.expectEqual(@as(?u64, 3000), xrpc_error.rate_limit.limit); + try std.testing.expectEqual(@as(?u64, 0), xrpc_error.rate_limit.remaining); + try std.testing.expectEqual(@as(?u64, 1710000000), xrpc_error.rate_limit.reset); + try std.testing.expectEqual(@as(?u64, 2), xrpc_error.rate_limit.retry_after); +} + +test "retry policy is conservative and deterministic" { + const policy: XrpcClient.RetryPolicy = .{ + .base_delay_ms = 100, + .max_delay_ms = 1000, + .jitter_percent = 0, + }; + + try std.testing.expect(policy.isRetryableStatus(.too_many_requests)); + try std.testing.expect(policy.isRetryableStatus(.internal_server_error)); + try std.testing.expect(policy.isRetryableStatus(.bad_gateway)); + try std.testing.expect(policy.isRetryableStatus(.service_unavailable)); + try std.testing.expect(policy.isRetryableStatus(.gateway_timeout)); + try std.testing.expect(!policy.isRetryableStatus(.bad_request)); + try std.testing.expect(!policy.isRetryableStatus(.unauthorized)); + try std.testing.expect(!policy.isRetryableStatus(.not_found)); + + 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)); + try std.testing.expectEqual(@as(u64, 1000), policy.delayMillis(0, .{ .retry_after = 5 })); + try std.testing.expectEqual(@as(u64, 1000), policy.delayMillisAt(0, .{ .reset = 1005 }, 1000)); + try std.testing.expectEqual(@as(u64, 100), policy.delayMillisAt(0, .{ .reset = 999 }, 1000)); +}