diff --git a/src/broadcaster.zig b/src/broadcaster.zig index f7907db..3f002ac 100644 --- a/src/broadcaster.zig +++ b/src/broadcaster.zig @@ -573,8 +573,7 @@ pub fn formatPrometheusMetrics(stats: *const Stats, buf: []u8) []const u8 { } pub fn formatStatsResponse(stats: *const Stats, buf: []u8) []const u8 { - var json_buf: [2048]u8 = undefined; - const json = std.fmt.bufPrint(&json_buf, + return std.fmt.bufPrint(buf, \\{{"seq":{d},"relay_seq":{d},"consumers":{d},"connected_inbound":{d},"frames_in":{d},"frames_out":{d},"validated":{d},"failed":{d},"skipped":{d},"decode_errors":{d},"cache_hits":{d},"cache_misses":{d},"slow_consumers":{d},"uptime_seconds":{d}}} , .{ stats.seq.load(.acquire), @@ -591,13 +590,7 @@ pub fn formatStatsResponse(stats: *const Stats, buf: []u8) []const u8 { stats.cache_misses.load(.acquire), stats.slow_consumers.load(.acquire), std.time.timestamp() - stats.start_time, - }) catch return "HTTP/1.1 500 Internal Server Error\r\nContent-Length: 0\r\n\r\n"; - - return std.fmt.bufPrint( - buf, - "HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: {d}\r\nConnection: close\r\n\r\n{s}", - .{ json.len, json }, - ) catch "HTTP/1.1 500 Internal Server Error\r\nContent-Length: 0\r\n\r\n"; + }) catch ""; } // --- tests --- @@ -710,7 +703,7 @@ test "formatStatsResponse produces valid JSON" { var buf: [4096]u8 = undefined; const response = formatStatsResponse(&stats, &buf); - try std.testing.expect(std.mem.startsWith(u8, response, "HTTP/1.1 200 OK")); + try std.testing.expect(std.mem.startsWith(u8, response, "{")); try std.testing.expect(std.mem.indexOf(u8, response, "\"seq\":100") != null); try std.testing.expect(std.mem.indexOf(u8, response, "\"consumers\":5") != null); try std.testing.expect(std.mem.indexOf(u8, response, "\"frames_in\":200") != null); diff --git a/src/main.zig b/src/main.zig index 4b1d9f5..b396d1d 100644 --- a/src/main.zig +++ b/src/main.zig @@ -18,6 +18,7 @@ //! /_health, /_stats, /metrics — health, stats, prometheus const std = @import("std"); +const http = std.http; const websocket = @import("websocket"); const broadcaster = @import("broadcaster.zig"); const validator_mod = @import("validator.zig"); @@ -213,102 +214,52 @@ fn installSignalHandlers() void { fn handleHttpConn(stream: std.net.Stream, stats: *broadcaster.Stats, persist: *event_log_mod.DiskPersist, slurper: *slurper_mod.Slurper, ci: *collection_index_mod.CollectionIndex) void { defer stream.close(); - // read request — may need multiple reads if proxy splits headers/body - var buf: [8192]u8 = undefined; - var total: usize = 0; - total = stream.read(&buf) catch return; - if (total == 0) return; - - // parse first line: "METHOD /path HTTP/1.1" - const line_end = std.mem.indexOfScalar(u8, buf[0..total], '\n') orelse return; - const first_line = buf[0..line_end]; - - const method_end = std.mem.indexOfScalar(u8, first_line, ' ') orelse return; - const method = first_line[0..method_end]; - - const path_start = method_end + 1; - const rest = first_line[path_start..]; - const path_end = std.mem.indexOfScalar(u8, rest, ' ') orelse rest.len; - const path = rest[0..path_end]; - - // find end of headers - const header_end = std.mem.indexOf(u8, buf[0..total], "\r\n\r\n"); - - // for POSTs: if we have headers, parse Content-Length and read remaining body - if (std.mem.eql(u8, method, "POST")) { - if (header_end) |he| { - const headers = buf[0..he]; - const content_length = parseContentLength(headers) orelse 0; - const body_start = he + 4; - const body_needed = body_start + content_length; - - // keep reading until we have the full body (or buffer is full) - while (total < body_needed and total < buf.len) { - const m = stream.read(buf[total..]) catch break; - if (m == 0) break; - total += m; - } - } - } + var recv_buf: [8192]u8 = undefined; + var send_buf: [8192]u8 = undefined; + var connection_reader = stream.reader(&recv_buf); + var connection_writer = stream.writer(&send_buf); + var server = http.Server.init(connection_reader.interface(), &connection_writer.interface); - const request = buf[0..total]; - // re-find header_end in case more data shifted things (it won't, but be safe) - const he = std.mem.indexOf(u8, request, "\r\n\r\n"); - const body: []const u8 = if (he) |h| request[h + 4 ..] else ""; + var request = server.receiveHead() catch return; - if (std.mem.eql(u8, method, "GET")) { - handleGet(stream, path, stats, persist, slurper, ci); - } else if (std.mem.eql(u8, method, "POST")) { - handlePost(stream, path, request[0 .. he orelse total], body, persist, slurper); - } else { - httpRespond(stream, "405 Method Not Allowed", "text/plain", "method not allowed"); - } -} + const target = request.head.target; + // extract path and query before reading body (head strings reference recv_buf) + const qmark = std.mem.indexOfScalar(u8, target, '?'); + const path = target[0..(qmark orelse target.len)]; + const query = if (qmark) |q| target[q + 1 ..] else ""; -fn parseContentLength(headers: []const u8) ?usize { - var iter = std.mem.splitScalar(u8, headers, '\n'); - while (iter.next()) |line| { - const trimmed = std.mem.trimRight(u8, line, "\r"); - const colon = std.mem.indexOfScalar(u8, trimmed, ':') orelse continue; - const key = std.mem.trim(u8, trimmed[0..colon], " "); - if (std.ascii.eqlIgnoreCase(key, "content-length")) { - const val = std.mem.trim(u8, trimmed[colon + 1 ..], " "); - return std.fmt.parseInt(usize, val, 10) catch null; - } + if (request.head.method == .GET) { + handleGet(&request, path, query, stats, persist, slurper, ci); + } else if (request.head.method == .POST) { + handlePost(&request, path, persist, slurper); + } else { + respondText(&request, .method_not_allowed, "method not allowed"); } - return null; } -fn handleGet(stream: std.net.Stream, full_path: []const u8, stats: *broadcaster.Stats, persist: *event_log_mod.DiskPersist, slurper: *slurper_mod.Slurper, ci: *collection_index_mod.CollectionIndex) void { - // split path from query string - const qmark = std.mem.indexOfScalar(u8, full_path, '?'); - const path = full_path[0..(qmark orelse full_path.len)]; - const query = if (qmark) |q| full_path[q + 1 ..] else ""; - +fn handleGet(request: *http.Server.Request, path: []const u8, query: []const u8, stats: *broadcaster.Stats, persist: *event_log_mod.DiskPersist, slurper: *slurper_mod.Slurper, ci: *collection_index_mod.CollectionIndex) void { if (std.mem.eql(u8, path, "/_health") or std.mem.eql(u8, path, "/xrpc/_health")) { - httpRespond(stream, "200 OK", "application/json", "{\"status\":\"ok\"}"); + respondJson(request, .ok, "{\"status\":\"ok\"}"); } else if (std.mem.eql(u8, path, "/_stats")) { var stats_buf: [4096]u8 = undefined; - const response = broadcaster.formatStatsResponse(stats, &stats_buf); - _ = stream.write(response) catch {}; + const body = broadcaster.formatStatsResponse(stats, &stats_buf); + respondJson(request, .ok, body); } else if (std.mem.eql(u8, path, "/metrics")) { var metrics_buf: [4096]u8 = undefined; const body = broadcaster.formatPrometheusMetrics(stats, &metrics_buf); - var resp_buf: [8192]u8 = undefined; - const response = std.fmt.bufPrint(&resp_buf, "HTTP/1.1 200 OK\r\nContent-Type: text/plain; version=0.0.4; charset=utf-8\r\nContent-Length: {d}\r\nConnection: close\r\n\r\n{s}", .{ body.len, body }) catch return; - _ = stream.write(response) catch {}; + request.respond(body, .{ .status = .ok, .keep_alive = false, .extra_headers = &.{.{ .name = "content-type", .value = "text/plain; version=0.0.4; charset=utf-8" }} }) catch {}; } else if (std.mem.eql(u8, path, "/xrpc/com.atproto.sync.listRepos")) { - handleListRepos(stream, query, persist); + handleListRepos(request, query, persist); } else if (std.mem.eql(u8, path, "/xrpc/com.atproto.sync.getRepoStatus")) { - handleGetRepoStatus(stream, query, persist); + handleGetRepoStatus(request, query, persist); } else if (std.mem.eql(u8, path, "/xrpc/com.atproto.sync.getLatestCommit")) { - handleGetLatestCommit(stream, query, persist); + handleGetLatestCommit(request, query, persist); } else if (std.mem.eql(u8, path, "/xrpc/com.atproto.sync.listReposByCollection")) { - handleListReposByCollection(stream, query, ci); + handleListReposByCollection(request, query, ci); } else if (std.mem.eql(u8, path, "/admin/hosts")) { - handleAdminListHosts(stream, persist, slurper); + handleAdminListHosts(request, persist, slurper); } else if (std.mem.eql(u8, path, "/")) { - httpRespond(stream, "200 OK", "text/plain", + respondText(request, .ok, \\ _ \\ ___| | __ _ _ _ \\|_ / |/ _` | | | | @@ -323,37 +274,47 @@ fn handleGet(stream: std.net.Stream, full_path: []const u8, stats: *broadcaster. \\ ); } else if (std.mem.eql(u8, path, "/favicon.svg") or std.mem.eql(u8, path, "/favicon.ico")) { - httpRespond(stream, "200 OK", "image/svg+xml", + request.respond( \\ \\ \\Z \\ - ); + , .{ .status = .ok, .keep_alive = false, .extra_headers = &.{.{ .name = "content-type", .value = "image/svg+xml" }} }) catch {}; } else { - httpRespond(stream, "404 Not Found", "text/plain", "not found"); + respondText(request, .not_found, "not found"); } } -fn handlePost(stream: std.net.Stream, path: []const u8, headers: []const u8, body: []const u8, persist: *event_log_mod.DiskPersist, slurper: *slurper_mod.Slurper) void { +fn handlePost(request: *http.Server.Request, path: []const u8, persist: *event_log_mod.DiskPersist, slurper: *slurper_mod.Slurper) void { if (std.mem.eql(u8, path, "/admin/repo/ban")) { - handleBan(stream, headers, body, persist); + handleBan(request, persist); } else if (std.mem.eql(u8, path, "/xrpc/com.atproto.sync.requestCrawl")) { - handleRequestCrawl(stream, body, slurper); + handleRequestCrawl(request, slurper); } else if (std.mem.eql(u8, path, "/admin/hosts/block")) { - handleAdminBlockHost(stream, headers, body, persist); + handleAdminBlockHost(request, persist); } else if (std.mem.eql(u8, path, "/admin/hosts/unblock")) { - handleAdminUnblockHost(stream, headers, body, persist); + handleAdminUnblockHost(request, persist); } else { - httpRespond(stream, "404 Not Found", "text/plain", "not found"); + respondText(request, .not_found, "not found"); } } -fn handleBan(stream: std.net.Stream, headers: []const u8, body: []const u8, persist: *event_log_mod.DiskPersist) void { - if (!checkAdmin(stream, headers)) return; +fn handleBan(request: *http.Server.Request, persist: *event_log_mod.DiskPersist) void { + if (!checkAdmin(request)) return; + + // read body (after checkAdmin which uses iterateHeaders) + var transfer_buf: [4096]u8 = undefined; + const body_reader = request.readerExpectNone(&transfer_buf); + var body_buf: [4096]u8 = undefined; + const body_len = body_reader.readSliceShort(&body_buf) catch { + respondJson(request, .bad_request, "{\"error\":\"failed to read request body\"}"); + return; + }; + const body = body_buf[0..body_len]; // parse JSON body for "did" field const parsed = std.json.parseFromSlice(struct { did: []const u8 }, persist.allocator, body, .{ .ignore_unknown_fields = true }) catch { - httpRespond(stream, "400 Bad Request", "application/json", "{\"error\":\"invalid JSON, expected {\\\"did\\\":\\\"...\\\"}\"}"); + respondJson(request, .bad_request, "{\"error\":\"invalid JSON, expected {\\\"did\\\":\\\"...\\\"}\"}"); return; }; defer parsed.deinit(); @@ -361,21 +322,30 @@ fn handleBan(stream: std.net.Stream, headers: []const u8, body: []const u8, pers // resolve DID → UID and take down const uid = persist.uidForDid(did) catch { - httpRespond(stream, "500 Internal Server Error", "application/json", "{\"error\":\"failed to resolve DID\"}"); + respondJson(request, .internal_server_error, "{\"error\":\"failed to resolve DID\"}"); return; }; persist.takeDownUser(uid) catch { - httpRespond(stream, "500 Internal Server Error", "application/json", "{\"error\":\"takedown failed\"}"); + respondJson(request, .internal_server_error, "{\"error\":\"takedown failed\"}"); return; }; log.info("admin: banned {s} (uid={d})", .{ did, uid }); - httpRespond(stream, "200 OK", "application/json", "{\"success\":true}"); + respondJson(request, .ok, "{\"success\":true}"); } -fn handleRequestCrawl(stream: std.net.Stream, body: []const u8, slurper: *slurper_mod.Slurper) void { +fn handleRequestCrawl(request: *http.Server.Request, slurper: *slurper_mod.Slurper) void { + var transfer_buf: [4096]u8 = undefined; + const body_reader = request.readerExpectNone(&transfer_buf); + var body_buf: [4096]u8 = undefined; + const body_len = body_reader.readSliceShort(&body_buf) catch { + respondJson(request, .bad_request, "{\"error\":\"failed to read request body\"}"); + return; + }; + const body = body_buf[0..body_len]; + const parsed = std.json.parseFromSlice(struct { hostname: []const u8 }, slurper.allocator, body, .{ .ignore_unknown_fields = true }) catch { - httpRespond(stream, "400 Bad Request", "application/json", "{\"error\":\"invalid JSON, expected {\\\"hostname\\\":\\\"...\\\"}\"}"); + respondJson(request, .bad_request, "{\"error\":\"invalid JSON, expected {\\\"hostname\\\":\\\"...\\\"}\"}"); return; }; defer parsed.deinit(); @@ -383,7 +353,7 @@ fn handleRequestCrawl(stream: std.net.Stream, body: []const u8, slurper: *slurpe // fast validation: hostname format (Go relay does this synchronously in handler) const hostname = slurper_mod.validateHostname(slurper.allocator, parsed.value.hostname) catch |err| { log.warn("requestCrawl rejected '{s}': {s}", .{ parsed.value.hostname, @errorName(err) }); - httpRespond(stream, "400 Bad Request", "application/json", switch (err) { + respondJson(request, .bad_request, switch (err) { error.EmptyHostname => "{\"error\":\"empty hostname\"}", error.InvalidCharacter => "{\"error\":\"hostname contains invalid characters\"}", error.InvalidLabel => "{\"error\":\"hostname has invalid label\"}", @@ -400,25 +370,25 @@ fn handleRequestCrawl(stream: std.net.Stream, body: []const u8, slurper: *slurpe // fast validation: domain ban check if (slurper.persist.isDomainBanned(hostname)) { log.warn("requestCrawl rejected '{s}': domain banned", .{hostname}); - httpRespond(stream, "400 Bad Request", "application/json", "{\"error\":\"domain is banned\"}"); + respondJson(request, .bad_request, "{\"error\":\"domain is banned\"}"); return; } // enqueue for async processing (describeServer check happens in crawl processor) slurper.addCrawlRequest(hostname) catch { - httpRespond(stream, "500 Internal Server Error", "application/json", "{\"error\":\"failed to store crawl request\"}"); + respondJson(request, .internal_server_error, "{\"error\":\"failed to store crawl request\"}"); return; }; log.info("crawl requested: {s}", .{hostname}); - httpRespond(stream, "200 OK", "application/json", "{\"success\":true}"); + respondJson(request, .ok, "{\"success\":true}"); } // --- admin host management --- -fn handleAdminListHosts(stream: std.net.Stream, persist: *event_log_mod.DiskPersist, slurper: *slurper_mod.Slurper) void { +fn handleAdminListHosts(request: *http.Server.Request, persist: *event_log_mod.DiskPersist, slurper: *slurper_mod.Slurper) void { const hosts = persist.listAllHosts(persist.allocator) catch { - httpRespondJson(stream, "500 Internal Server Error", "{\"error\":\"DatabaseError\",\"message\":\"query failed\"}"); + respondJson(request, .internal_server_error, "{\"error\":\"DatabaseError\",\"message\":\"query failed\"}"); return; }; defer { @@ -447,101 +417,123 @@ fn handleAdminListHosts(stream: std.net.Stream, persist: *event_log_mod.DiskPers } std.fmt.format(w, "],\"active_workers\":{d}}}", .{slurper.workerCount()}) catch return; - httpRespondJson(stream, "200 OK", fbs.getWritten()); + respondJson(request, .ok, fbs.getWritten()); } -fn handleAdminBlockHost(stream: std.net.Stream, headers: []const u8, body: []const u8, persist: *event_log_mod.DiskPersist) void { - if (!checkAdmin(stream, headers)) return; +fn handleAdminBlockHost(request: *http.Server.Request, persist: *event_log_mod.DiskPersist) void { + if (!checkAdmin(request)) return; + + var transfer_buf: [4096]u8 = undefined; + const body_reader = request.readerExpectNone(&transfer_buf); + var body_buf: [4096]u8 = undefined; + const body_len = body_reader.readSliceShort(&body_buf) catch { + respondJson(request, .bad_request, "{\"error\":\"failed to read request body\"}"); + return; + }; + const body = body_buf[0..body_len]; const parsed = std.json.parseFromSlice(struct { hostname: []const u8 }, persist.allocator, body, .{ .ignore_unknown_fields = true }) catch { - httpRespondJson(stream, "400 Bad Request", "{\"error\":\"BadRequest\",\"message\":\"invalid JSON\"}"); + respondJson(request, .bad_request, "{\"error\":\"BadRequest\",\"message\":\"invalid JSON\"}"); return; }; defer parsed.deinit(); const host_info = persist.getOrCreateHost(parsed.value.hostname) catch { - httpRespondJson(stream, "500 Internal Server Error", "{\"error\":\"DatabaseError\",\"message\":\"host lookup failed\"}"); + respondJson(request, .internal_server_error, "{\"error\":\"DatabaseError\",\"message\":\"host lookup failed\"}"); return; }; persist.updateHostStatus(host_info.id, "blocked") catch { - httpRespondJson(stream, "500 Internal Server Error", "{\"error\":\"DatabaseError\",\"message\":\"status update failed\"}"); + respondJson(request, .internal_server_error, "{\"error\":\"DatabaseError\",\"message\":\"status update failed\"}"); return; }; log.info("admin: blocked host {s} (id={d})", .{ parsed.value.hostname, host_info.id }); - httpRespondJson(stream, "200 OK", "{\"success\":true}"); + respondJson(request, .ok, "{\"success\":true}"); } -fn handleAdminUnblockHost(stream: std.net.Stream, headers: []const u8, body: []const u8, persist: *event_log_mod.DiskPersist) void { - if (!checkAdmin(stream, headers)) return; +fn handleAdminUnblockHost(request: *http.Server.Request, persist: *event_log_mod.DiskPersist) void { + if (!checkAdmin(request)) return; + + var transfer_buf: [4096]u8 = undefined; + const body_reader = request.readerExpectNone(&transfer_buf); + var body_buf: [4096]u8 = undefined; + const body_len = body_reader.readSliceShort(&body_buf) catch { + respondJson(request, .bad_request, "{\"error\":\"failed to read request body\"}"); + return; + }; + const body = body_buf[0..body_len]; const parsed = std.json.parseFromSlice(struct { hostname: []const u8 }, persist.allocator, body, .{ .ignore_unknown_fields = true }) catch { - httpRespondJson(stream, "400 Bad Request", "{\"error\":\"BadRequest\",\"message\":\"invalid JSON\"}"); + respondJson(request, .bad_request, "{\"error\":\"BadRequest\",\"message\":\"invalid JSON\"}"); return; }; defer parsed.deinit(); const host_info = persist.getOrCreateHost(parsed.value.hostname) catch { - httpRespondJson(stream, "500 Internal Server Error", "{\"error\":\"DatabaseError\",\"message\":\"host lookup failed\"}"); + respondJson(request, .internal_server_error, "{\"error\":\"DatabaseError\",\"message\":\"host lookup failed\"}"); return; }; persist.updateHostStatus(host_info.id, "active") catch { - httpRespondJson(stream, "500 Internal Server Error", "{\"error\":\"DatabaseError\",\"message\":\"status update failed\"}"); + respondJson(request, .internal_server_error, "{\"error\":\"DatabaseError\",\"message\":\"status update failed\"}"); return; }; persist.resetHostFailures(host_info.id) catch {}; log.info("admin: unblocked host {s} (id={d})", .{ parsed.value.hostname, host_info.id }); - httpRespondJson(stream, "200 OK", "{\"success\":true}"); + respondJson(request, .ok, "{\"success\":true}"); } /// check admin auth, send error response if not authorized. returns true if authorized. -fn checkAdmin(stream: std.net.Stream, headers: []const u8) bool { +fn checkAdmin(request: *http.Server.Request) bool { const admin_pw = std.posix.getenv("RELAY_ADMIN_PASSWORD") orelse { - httpRespond(stream, "403 Forbidden", "application/json", "{\"error\":\"admin endpoint not configured\"}"); + respondJson(request, .forbidden, "{\"error\":\"admin endpoint not configured\"}"); return false; }; - const auth_value = findHeader(headers, "authorization") orelse { - httpRespond(stream, "401 Unauthorized", "application/json", "{\"error\":\"missing authorization header\"}"); - return false; - }; - const bearer_prefix = "Bearer "; - if (!std.mem.startsWith(u8, auth_value, bearer_prefix)) { - httpRespond(stream, "401 Unauthorized", "application/json", "{\"error\":\"invalid authorization scheme\"}"); - return false; - } - const token = auth_value[bearer_prefix.len..]; - if (!std.mem.eql(u8, token, admin_pw)) { - httpRespond(stream, "401 Unauthorized", "application/json", "{\"error\":\"invalid token\"}"); - return false; + var iter = request.iterateHeaders(); + while (iter.next()) |header| { + if (std.ascii.eqlIgnoreCase(header.name, "authorization")) { + const bearer_prefix = "Bearer "; + if (!std.mem.startsWith(u8, header.value, bearer_prefix)) { + respondJson(request, .unauthorized, "{\"error\":\"invalid authorization scheme\"}"); + return false; + } + const token = header.value[bearer_prefix.len..]; + if (!std.mem.eql(u8, token, admin_pw)) { + respondJson(request, .unauthorized, "{\"error\":\"invalid token\"}"); + return false; + } + return true; + } } - return true; + + respondJson(request, .unauthorized, "{\"error\":\"missing authorization header\"}"); + return false; } // --- XRPC endpoint handlers --- -fn handleListRepos(stream: std.net.Stream, query: []const u8, persist: *event_log_mod.DiskPersist) void { +fn handleListRepos(request: *http.Server.Request, query: []const u8, persist: *event_log_mod.DiskPersist) void { const cursor_str = queryParam(query, "cursor") orelse "0"; const limit_str = queryParam(query, "limit") orelse "500"; const cursor_val = std.fmt.parseInt(i64, cursor_str, 10) catch { - httpRespondJson(stream, "400 Bad Request", "{\"error\":\"BadRequest\",\"message\":\"invalid cursor\"}"); + respondJson(request, .bad_request, "{\"error\":\"BadRequest\",\"message\":\"invalid cursor\"}"); return; }; if (cursor_val < 0) { - httpRespondJson(stream, "400 Bad Request", "{\"error\":\"BadRequest\",\"message\":\"cursor must be >= 0\"}"); + respondJson(request, .bad_request, "{\"error\":\"BadRequest\",\"message\":\"cursor must be >= 0\"}"); return; } const limit = std.fmt.parseInt(i64, limit_str, 10) catch { - httpRespondJson(stream, "400 Bad Request", "{\"error\":\"BadRequest\",\"message\":\"invalid limit\"}"); + respondJson(request, .bad_request, "{\"error\":\"BadRequest\",\"message\":\"invalid limit\"}"); return; }; if (limit < 1 or limit > 1000) { - httpRespondJson(stream, "400 Bad Request", "{\"error\":\"BadRequest\",\"message\":\"limit must be 1..1000\"}"); + respondJson(request, .bad_request, "{\"error\":\"BadRequest\",\"message\":\"limit must be 1..1000\"}"); return; } @@ -552,7 +544,7 @@ fn handleListRepos(stream: std.net.Stream, query: []const u8, persist: *event_lo \\FROM account a LEFT JOIN account_repo r ON a.uid = r.uid \\WHERE a.uid > $1 ORDER BY a.uid ASC LIMIT $2 , .{ cursor_val, limit }) catch { - httpRespondJson(stream, "500 Internal Server Error", "{\"error\":\"DatabaseError\",\"message\":\"query failed\"}"); + respondJson(request, .internal_server_error, "{\"error\":\"DatabaseError\",\"message\":\"query failed\"}"); return; }; defer result.deinit(); @@ -621,20 +613,19 @@ fn handleListRepos(stream: std.net.Stream, query: []const u8, persist: *event_lo w.writeByte('}') catch return; - const resp_body = fbs.getWritten(); - httpRespondJson(stream, "200 OK", resp_body); + respondJson(request, .ok, fbs.getWritten()); } -fn handleGetRepoStatus(stream: std.net.Stream, query: []const u8, persist: *event_log_mod.DiskPersist) void { +fn handleGetRepoStatus(request: *http.Server.Request, query: []const u8, persist: *event_log_mod.DiskPersist) void { var did_buf: [256]u8 = undefined; const did = queryParamDecoded(query, "did", &did_buf) orelse { - httpRespondJson(stream, "400 Bad Request", "{\"error\":\"BadRequest\",\"message\":\"did parameter required\"}"); + respondJson(request, .bad_request, "{\"error\":\"BadRequest\",\"message\":\"did parameter required\"}"); return; }; // basic DID syntax check if (!std.mem.startsWith(u8, did, "did:")) { - httpRespondJson(stream, "400 Bad Request", "{\"error\":\"BadRequest\",\"message\":\"invalid DID\"}"); + respondJson(request, .bad_request, "{\"error\":\"BadRequest\",\"message\":\"invalid DID\"}"); return; } @@ -643,10 +634,10 @@ fn handleGetRepoStatus(stream: std.net.Stream, query: []const u8, persist: *even "SELECT a.uid, a.status, a.upstream_status, COALESCE(r.rev, '') FROM account a LEFT JOIN account_repo r ON a.uid = r.uid WHERE a.did = $1", .{did}, ) catch { - httpRespondJson(stream, "500 Internal Server Error", "{\"error\":\"DatabaseError\",\"message\":\"query failed\"}"); + respondJson(request, .internal_server_error, "{\"error\":\"DatabaseError\",\"message\":\"query failed\"}"); return; }) orelse { - httpRespondJson(stream, "404 Not Found", "{\"error\":\"RepoNotFound\",\"message\":\"account not found\"}"); + respondJson(request, .not_found, "{\"error\":\"RepoNotFound\",\"message\":\"account not found\"}"); return; }; defer row.deinit() catch {}; @@ -683,18 +674,18 @@ fn handleGetRepoStatus(stream: std.net.Stream, query: []const u8, persist: *even } w.writeByte('}') catch return; - httpRespondJson(stream, "200 OK", fbs.getWritten()); + respondJson(request, .ok, fbs.getWritten()); } -fn handleGetLatestCommit(stream: std.net.Stream, query: []const u8, persist: *event_log_mod.DiskPersist) void { +fn handleGetLatestCommit(request: *http.Server.Request, query: []const u8, persist: *event_log_mod.DiskPersist) void { var did_buf: [256]u8 = undefined; const did = queryParamDecoded(query, "did", &did_buf) orelse { - httpRespondJson(stream, "400 Bad Request", "{\"error\":\"BadRequest\",\"message\":\"did parameter required\"}"); + respondJson(request, .bad_request, "{\"error\":\"BadRequest\",\"message\":\"did parameter required\"}"); return; }; if (!std.mem.startsWith(u8, did, "did:")) { - httpRespondJson(stream, "400 Bad Request", "{\"error\":\"BadRequest\",\"message\":\"invalid DID\"}"); + respondJson(request, .bad_request, "{\"error\":\"BadRequest\",\"message\":\"invalid DID\"}"); return; } @@ -703,10 +694,10 @@ fn handleGetLatestCommit(stream: std.net.Stream, query: []const u8, persist: *ev "SELECT a.status, a.upstream_status, COALESCE(r.rev, ''), COALESCE(r.commit_data_cid, '') FROM account a LEFT JOIN account_repo r ON a.uid = r.uid WHERE a.did = $1", .{did}, ) catch { - httpRespondJson(stream, "500 Internal Server Error", "{\"error\":\"DatabaseError\",\"message\":\"query failed\"}"); + respondJson(request, .internal_server_error, "{\"error\":\"DatabaseError\",\"message\":\"query failed\"}"); return; }) orelse { - httpRespondJson(stream, "404 Not Found", "{\"error\":\"RepoNotFound\",\"message\":\"account not found\"}"); + respondJson(request, .not_found, "{\"error\":\"RepoNotFound\",\"message\":\"account not found\"}"); return; }; defer row.deinit() catch {}; @@ -721,21 +712,21 @@ fn handleGetLatestCommit(stream: std.net.Stream, query: []const u8, persist: *ev // check account status (match Go relay behavior) if (std.mem.eql(u8, status, "takendown") or std.mem.eql(u8, status, "suspended")) { - httpRespondJson(stream, "403 Forbidden", "{\"error\":\"RepoTakendown\",\"message\":\"account has been taken down\"}"); + respondJson(request, .forbidden, "{\"error\":\"RepoTakendown\",\"message\":\"account has been taken down\"}"); return; } else if (std.mem.eql(u8, status, "deactivated")) { - httpRespondJson(stream, "403 Forbidden", "{\"error\":\"RepoDeactivated\",\"message\":\"account is deactivated\"}"); + respondJson(request, .forbidden, "{\"error\":\"RepoDeactivated\",\"message\":\"account is deactivated\"}"); return; } else if (std.mem.eql(u8, status, "deleted")) { - httpRespondJson(stream, "403 Forbidden", "{\"error\":\"RepoDeleted\",\"message\":\"account is deleted\"}"); + respondJson(request, .forbidden, "{\"error\":\"RepoDeleted\",\"message\":\"account is deleted\"}"); return; } else if (!std.mem.eql(u8, status, "active")) { - httpRespondJson(stream, "403 Forbidden", "{\"error\":\"RepoInactive\",\"message\":\"account is not active\"}"); + respondJson(request, .forbidden, "{\"error\":\"RepoInactive\",\"message\":\"account is not active\"}"); return; } if (rev.len == 0 or cid.len == 0) { - httpRespondJson(stream, "404 Not Found", "{\"error\":\"RepoNotSynchronized\",\"message\":\"relay has no repo data for this account\"}"); + respondJson(request, .not_found, "{\"error\":\"RepoNotSynchronized\",\"message\":\"relay has no repo data for this account\"}"); return; } @@ -749,27 +740,27 @@ fn handleGetLatestCommit(stream: std.net.Stream, query: []const u8, persist: *ev w.writeAll(rev) catch return; w.writeAll("\"}") catch return; - httpRespondJson(stream, "200 OK", fbs.getWritten()); + respondJson(request, .ok, fbs.getWritten()); } -fn handleListReposByCollection(stream: std.net.Stream, query: []const u8, ci: *collection_index_mod.CollectionIndex) void { +fn handleListReposByCollection(request: *http.Server.Request, query: []const u8, ci: *collection_index_mod.CollectionIndex) void { const collection = queryParam(query, "collection") orelse { - httpRespondJson(stream, "400 Bad Request", "{\"error\":\"BadRequest\",\"message\":\"collection parameter required\"}"); + respondJson(request, .bad_request, "{\"error\":\"BadRequest\",\"message\":\"collection parameter required\"}"); return; }; if (collection.len == 0 or !std.mem.containsAtLeast(u8, collection, 1, ".")) { - httpRespondJson(stream, "400 Bad Request", "{\"error\":\"BadRequest\",\"message\":\"invalid collection NSID\"}"); + respondJson(request, .bad_request, "{\"error\":\"BadRequest\",\"message\":\"invalid collection NSID\"}"); return; } const limit_str = queryParam(query, "limit") orelse "500"; const limit = std.fmt.parseInt(usize, limit_str, 10) catch { - httpRespondJson(stream, "400 Bad Request", "{\"error\":\"BadRequest\",\"message\":\"invalid limit\"}"); + respondJson(request, .bad_request, "{\"error\":\"BadRequest\",\"message\":\"invalid limit\"}"); return; }; if (limit < 1 or limit > 1000) { - httpRespondJson(stream, "400 Bad Request", "{\"error\":\"BadRequest\",\"message\":\"limit must be 1..1000\"}"); + respondJson(request, .bad_request, "{\"error\":\"BadRequest\",\"message\":\"limit must be 1..1000\"}"); return; } @@ -779,7 +770,7 @@ fn handleListReposByCollection(stream: std.net.Stream, query: []const u8, ci: *c // scan collection index var did_buf: [65536]u8 = undefined; const result = ci.listReposByCollection(collection, limit, cursor_did, &did_buf) catch { - httpRespondJson(stream, "500 Internal Server Error", "{\"error\":\"InternalError\",\"message\":\"index scan failed\"}"); + respondJson(request, .internal_server_error, "{\"error\":\"InternalError\",\"message\":\"index scan failed\"}"); return; }; @@ -806,7 +797,7 @@ fn handleListReposByCollection(stream: std.net.Stream, query: []const u8, ci: *c } w.writeByte('}') catch return; - httpRespondJson(stream, "200 OK", fbs.getWritten()); + respondJson(request, .ok, fbs.getWritten()); } // --- query string helpers --- @@ -873,29 +864,14 @@ fn hexVal(c: u8) ?u4 { }; } -fn httpRespondJson(stream: std.net.Stream, status: []const u8, body: []const u8) void { - httpRespond(stream, status, "application/json", body); -} +// --- response helpers --- -fn findHeader(headers: []const u8, name: []const u8) ?[]const u8 { - var iter = std.mem.splitScalar(u8, headers, '\n'); - while (iter.next()) |line| { - const trimmed = std.mem.trimRight(u8, line, "\r"); - const colon = std.mem.indexOfScalar(u8, trimmed, ':') orelse continue; - const key = std.mem.trim(u8, trimmed[0..colon], " "); - if (std.ascii.eqlIgnoreCase(key, name)) { - return std.mem.trim(u8, trimmed[colon + 1 ..], " "); - } - } - return null; +fn respondJson(request: *http.Server.Request, status: http.Status, body: []const u8) void { + request.respond(body, .{ .status = status, .keep_alive = false, .extra_headers = &.{.{ .name = "content-type", .value = "application/json" }} }) catch {}; } -fn httpRespond(stream: std.net.Stream, status: []const u8, content_type: []const u8, body: []const u8) void { - // write headers first, then body separately (body can be much larger than header buffer) - var hdr_buf: [512]u8 = undefined; - const hdr = std.fmt.bufPrint(&hdr_buf, "HTTP/1.1 {s}\r\nContent-Type: {s}\r\nContent-Length: {d}\r\nServer: zlay (atproto-relay)\r\nConnection: close\r\n\r\n", .{ status, content_type, body.len }) catch return; - _ = stream.write(hdr) catch return; - _ = stream.write(body) catch {}; +fn respondText(request: *http.Server.Request, status: http.Status, body: []const u8) void { + request.respond(body, .{ .status = status, .keep_alive = false, .extra_headers = &.{.{ .name = "content-type", .value = "text/plain" }} }) catch {}; } fn parseEnvInt(comptime T: type, key: []const u8, default: T) T {