atproto relay in zig zlay.waow.tech
relay zig atproto
Something went wrong. Try again.
123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127//! host_check — liveness & safety checks for a candidate PDS host.//!//! mirrors the Go relay's host_checker.go (describeServer probe), ssrf.go//! (private-IP rejection), and slurper.go (relay-loop detection via the//! Server header). relay policy, so it lives under atproto/.
const std = @import("std");const Io = std.Io;const http = std.http;const Allocator = std.mem.Allocator;const hostname_mod = @import("hostname.zig");
const HostValidationError = hostname_mod.HostValidationError;const log = std.log.scoped(.relay);
/// check that a host is a real PDS by calling describeServer./// also checks the Server header for relay-loop detection./// Go relay: host_checker.go CheckHost + slurper.go Server header check.pub fn checkHost(allocator: Allocator, host: []const u8, io: Io) HostValidationError!void { // SSRF protection: reject private IPs before making any request try rejectPrivateHost(allocator, host); var url_buf: [512]u8 = undefined; const url = std.fmt.bufPrint(&url_buf, "https://{s}/xrpc/com.atproto.server.describeServer", .{host}) catch return error.HostUnreachable;
var client: http.Client = .{ .allocator = allocator, .io = io }; defer client.deinit();
const uri = std.Uri.parse(url) catch return error.HostUnreachable; var req = client.request(.GET, uri, .{}) catch return error.HostUnreachable; defer req.deinit(); req.sendBodiless() catch return error.HostUnreachable;
var redirect_buf: [2048]u8 = undefined; const response = req.receiveHead(&redirect_buf) catch return error.HostUnreachable;
if (response.head.status != .ok) return error.NotAPds;
// relay loop detection: check Server header for "atproto-relay" // Go relay: slurper.go — auto-bans hosts whose Server header contains "atproto-relay" if (findHeaderInRaw(response.head.bytes, "server")) |server_val| { if (std.mem.indexOf(u8, server_val, "atproto-relay") != null) { return error.IsARelay; } }}
/// SSRF protection: resolve hostname and reject private/reserved IP ranges./// Go relay: ssrf.go PublicOnlyTransport — rejects 10/8, 172.16/12, 192.168/16, 127/8, link-local.pub fn rejectPrivateHost(allocator: Allocator, host: []const u8) HostValidationError!void { // null-terminate hostname for getaddrinfo const hostname_z = allocator.dupeSentinel(u8, host, 0) catch return error.HostUnreachable; defer allocator.free(hostname_z);
var hints: std.c.addrinfo = .{ .flags = .{}, .family = std.c.AF.UNSPEC, .socktype = std.c.SOCK.STREAM, .protocol = 0, .addrlen = 0, .addr = null, .canonname = null, .next = null, };
var res: ?*std.c.addrinfo = null; const rc = std.c.getaddrinfo(hostname_z, "443", &hints, &res); if (@backingInt(rc) != 0 or res == null) return error.HostUnreachable; defer std.c.freeaddrinfo(res.?);
// check all resolved addresses — reject if ANY is private var cur = res; var found_any = false; while (cur) |node| : (cur = node.next) { found_any = true; if (node.family == std.c.AF.INET) { const sa: *const std.c.sockaddr.in = @ptrCast(@alignCast(node.addr.?)); const bytes: [4]u8 = @bitCast(sa.addr); if (bytes[0] == 10 or // 10.0.0.0/8 (bytes[0] == 172 and (bytes[1] & 0xf0) == 16) or // 172.16.0.0/12 (bytes[0] == 192 and bytes[1] == 168) or // 192.168.0.0/16 bytes[0] == 127 or // 127.0.0.0/8 bytes[0] == 0 or // 0.0.0.0/8 (bytes[0] == 169 and bytes[1] == 254)) // 169.254.0.0/16 link-local { log.warn("SSRF: {s} resolves to private IP {d}.{d}.{d}.{d}", .{ host, bytes[0], bytes[1], bytes[2], bytes[3] }); return error.HostUnreachable; } } // allow IPv6 for now (could add RFC 4193 check later) } if (!found_any) return error.HostUnreachable;}
/// search raw HTTP headers for a header by name (case-insensitive)./// returns the trimmed value, or null if not found.pub fn findHeaderInRaw(raw: []const u8, name: []const u8) ?[]const u8 { var it = std.mem.splitSequence(u8, raw, "\r\n"); _ = it.next(); // skip status line while (it.next()) |line| { if (line.len == 0) break; const colon = std.mem.indexOfScalar(u8, line, ':') orelse continue; const key = std.mem.trim(u8, line[0..colon], " "); if (key.len != name.len) continue; // case-insensitive compare var match = true; for (key, name) |a, b| { if (std.ascii.toLower(a) != std.ascii.toLower(b)) { match = false; break; } } if (match) { return std.mem.trim(u8, line[colon + 1 ..], " "); } } return null;}
// --- tests ---
test "findHeaderInRaw is case-insensitive and trims" { const raw = "HTTP/1.1 200 OK\r\nServer: atproto-relay/1.0\r\nContent-Type: application/json\r\n\r\n"; try std.testing.expectEqualStrings("atproto-relay/1.0", findHeaderInRaw(raw, "server").?); try std.testing.expectEqualStrings("application/json", findHeaderInRaw(raw, "content-type").?); try std.testing.expect(findHeaderInRaw(raw, "x-missing") == null);}