atproto pds in zig
Something went wrong. Try again.
123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430const std = @import("std");const clock = @import("../core/clock.zig");const config = @import("../core/config.zig");const log = @import("../core/log.zig");const http_api = @import("../http/api.zig");const eventlog = @import("../storage/eventlog.zig");const store = @import("../storage/store.zig");const zat = @import("zat");const httpz = @import("httpz");
const http = std.http;
const max_subscribe_repos_connections = 32;const notify_threshold_ms = 20 * 60 * 1000;
var subscribe_repos_connections: usize = 0;var get_repo_connections: usize = 0;var crawler_last_notified_ms: std.atomic.Value(i64) = .init(0);var crawler_notify_in_flight: std.atomic.Value(bool) = .init(false);
pub const SubscribeReposClient = struct { state: *StreamState,
pub const Context = struct { cursor: u64, };
pub fn init(conn: *httpz.websocket.Conn, ctx: *const Context) !SubscribeReposClient { const state = try std.heap.smp_allocator.create(StreamState); state.* = .{ .conn = conn, .cursor = ctx.cursor }; return .{ .state = state }; }
pub fn afterInit(self: *SubscribeReposClient) !void { self.state.retain(); const thread = std.Thread.spawn(.{}, streamEvents, .{self.state}) catch |err| { self.state.release(); log.err("sync subscribeRepos spawn failed cursor={d} err={s}\n", .{ self.state.cursor, @errorName(err) }); return err; }; thread.detach(); }
pub fn clientMessage(_: *SubscribeReposClient, _: []const u8) !void {}
pub fn close(self: *SubscribeReposClient) void { @atomicStore(bool, &self.state.closed, true, .release); self.state.finish(); self.state.release(); }
const StreamState = struct { conn: *httpz.websocket.Conn, cursor: u64, closed: bool = false, counted: bool = true, refs: usize = 1,
fn retain(self: *StreamState) void { _ = @atomicRmw(usize, &self.refs, .Add, 1, .monotonic); }
fn release(self: *StreamState) void { if (@atomicRmw(usize, &self.refs, .Sub, 1, .acq_rel) == 1) { std.heap.smp_allocator.destroy(self); } }
fn finish(self: *StreamState) void { if (@atomicRmw(bool, &self.counted, .Xchg, false, .acq_rel)) { _ = @atomicRmw(usize, &subscribe_repos_connections, .Sub, 1, .monotonic); } }
fn isClosed(self: *const StreamState) bool { return @atomicLoad(bool, &self.closed, .acquire); } };
fn streamEvents(state: *StreamState) void { defer state.release(); defer state.finish();
var observed = eventlog.snapshot().generation; while (!state.isClosed()) { var sent = false; { var arena = std.heap.ArenaAllocator.init(std.heap.page_allocator); defer arena.deinit(); const events = store.listSeqEvents(arena.allocator(), state.cursor, 100) catch |err| { log.err("sync subscribeRepos list failed cursor={d} err={s}\n", .{ state.cursor, @errorName(err) }); return; }; for (events) |event| { if (state.isClosed()) return; writeFirehoseFrame(state.conn, event.frame) catch |err| { switch (err) { error.BrokenPipe, error.ConnectionResetByPeer => return, else => log.err("sync subscribeRepos write failed cursor={d} seq={d} bytes={d} err={s}\n", .{ state.cursor, event.seq, event.frame.len, @errorName(err) }), } return; }; state.cursor = event.seq; sent = true; } } if (sent) { observed = eventlog.snapshot().generation; continue; } const next = eventlog.waitForChange(observed) catch return; if (next.generation != observed or next.latest_seq > state.cursor) { observed = next.generation; } } }
fn writeFirehoseFrame(conn: *httpz.websocket.Conn, payload: []const u8) !void { var header: [10]u8 = undefined; const header_len = websocketBinaryHeader(&header, payload.len);
conn.lock.lockUncancelable(conn.io); defer conn.lock.unlock(conn.io);
const socket = conn.stream.socket.handle; try writeSocketAll(socket, header[0..header_len]); try writeSocketAll(socket, payload); }
fn websocketBinaryHeader(header: *[10]u8, payload_len: usize) usize { header[0] = 0x82; if (payload_len < 126) { header[1] = @intCast(payload_len); return 2; } if (payload_len <= std.math.maxInt(u16)) { header[1] = 126; std.mem.writeInt(u16, header[2..4], @intCast(payload_len), .big); return 4; } header[1] = 127; std.mem.writeInt(u64, header[2..10], @intCast(payload_len), .big); return 10; }
fn writeSocketAll(socket: std.posix.socket_t, data: []const u8) !void { var remaining = data; while (remaining.len > 0) { const written = std.c.write(socket, remaining.ptr, remaining.len); if (written > 0) { remaining = remaining[@intCast(written)..]; continue; } if (written == 0) return error.WriteZero; switch (std.posix.errno(written)) { .AGAIN => { std.Thread.yield() catch {}; continue; }, .INTR => continue, .PIPE => return error.BrokenPipe, .CONNRESET => return error.ConnectionResetByPeer, else => return error.WriteFailed, } } }};
pub fn getBlob(request: *http_api.Request) !void { var arena = std.heap.ArenaAllocator.init(std.heap.page_allocator); defer arena.deinit(); const allocator = arena.allocator();
var did_buf: [256]u8 = undefined; const did = http_api.queryParam(request.url.raw, "did", &did_buf) orelse { return http_api.xrpcError(request, .bad_request, "InvalidRequest", "Missing did"); }; var cid_buf: [256]u8 = undefined; const cid = http_api.queryParam(request.url.raw, "cid", &cid_buf) orelse { return http_api.xrpcError(request, .bad_request, "InvalidRequest", "Missing cid"); }; try requirePublicRepoAvailable(request, did); const blob = store.getPublicBlob(allocator, did, cid) orelse { return http_api.xrpcError(request, .not_found, "BlobNotFound", "Blob not found"); }; const disposition = try std.fmt.allocPrint(allocator, "attachment; filename=\"{s}\"", .{cid}); const headers = [_]http.Header{ .{ .name = "content-type", .value = blob.mime_type }, .{ .name = "cache-control", .value = "public, max-age=31536000, immutable" }, .{ .name = "x-content-type-options", .value = "nosniff" }, .{ .name = "content-disposition", .value = disposition }, .{ .name = "content-security-policy", .value = "default-src 'none'; sandbox" }, .{ .name = "access-control-allow-origin", .value = "*" }, .{ .name = "access-control-allow-private-network", .value = "true" }, .{ .name = "connection", .value = "close" }, }; try http_api.respondNowClose(request, .ok, if (request.method == .HEAD) "" else blob.data, &headers);}
pub fn getRepo(request: *http_api.Request) !void { var arena = std.heap.ArenaAllocator.init(std.heap.page_allocator); defer arena.deinit(); const allocator = arena.allocator();
var did_buf: [256]u8 = undefined; const did = http_api.queryParam(request.url.raw, "did", &did_buf) orelse { return http_api.xrpcError(request, .bad_request, "InvalidRequest", "Missing did"); }; var since_buf: [256]u8 = undefined; const since = http_api.queryParam(request.url.raw, "since", &since_buf); try requirePublicRepoAvailable(request, did);
const full_export = since == null; if (full_export) { const active = @atomicRmw(usize, &get_repo_connections, .Add, 1, .monotonic); const max_connections = config.maxConcurrentRepoExports(); if (active >= max_connections) { _ = @atomicRmw(usize, &get_repo_connections, .Sub, 1, .monotonic); log.debug("sync getRepo rejected too_many_connections active={d} max={d} did={s}\n", .{ active + 1, max_connections, did }); return http_api.xrpcError(request, .too_many_requests, "RateLimitExceeded", "too many full getRepo exports"); } } defer { if (full_export) _ = @atomicRmw(usize, &get_repo_connections, .Sub, 1, .monotonic); }
const body = store.writeRepoCarSince(allocator, did, since) catch { return http_api.xrpcError(request, .not_found, "RepoNotFound", "Repo not found"); }; const headers = [_]http.Header{ .{ .name = "content-type", .value = "application/vnd.ipld.car" }, .{ .name = "access-control-allow-origin", .value = "*" }, .{ .name = "access-control-allow-private-network", .value = "true" }, .{ .name = "connection", .value = "close" }, }; if (request.method == .HEAD) { return http_api.respondNowClose(request, .ok, "", &headers); } try http_api.respond(request, .ok, body, &headers);}
pub fn listBlobs(request: *http_api.Request) !void { var arena = std.heap.ArenaAllocator.init(std.heap.page_allocator); defer arena.deinit(); const allocator = arena.allocator();
var did_buf: [256]u8 = undefined; const did = http_api.queryParam(request.url.raw, "did", &did_buf) orelse { return http_api.xrpcError(request, .bad_request, "InvalidRequest", "Missing did"); }; var since_buf: [256]u8 = undefined; var cursor_buf: [256]u8 = undefined; try requirePublicRepoAvailable(request, did); const body = try store.writeBlobListJson( allocator, did, http_api.queryParam(request.url.raw, "since", &since_buf), http_api.queryParam(request.url.raw, "cursor", &cursor_buf), @min(http_api.queryLimit(request.url.raw, 500), 1000), ); return http_api.json(request, .ok, body);}
pub fn getLatestCommit(request: *http_api.Request) !void { var arena = std.heap.ArenaAllocator.init(std.heap.page_allocator); defer arena.deinit(); const allocator = arena.allocator();
var did_buf: [256]u8 = undefined; const did = http_api.queryParam(request.url.raw, "did", &did_buf) orelse { return http_api.xrpcError(request, .bad_request, "InvalidRequest", "Missing did"); }; try requirePublicRepoAvailable(request, did); const body = store.writeLatestCommitJson(allocator, did) catch |err| switch (err) { error.RepoNotFound => return http_api.xrpcError(request, .not_found, "RepoNotFound", "Repo not found"), else => return err, }; return http_api.json(request, .ok, body);}
pub fn listRepos(request: *http_api.Request) !void { var arena = std.heap.ArenaAllocator.init(std.heap.page_allocator); defer arena.deinit(); const allocator = arena.allocator();
var cursor_buf: [512]u8 = undefined; const body = store.writeRepoListJson(allocator, http_api.queryParam(request.url.raw, "cursor", &cursor_buf), http_api.queryLimit(request.url.raw, 500)) catch |err| switch (err) { error.InvalidCursor => return http_api.xrpcError(request, .bad_request, "InvalidRequest", "Malformed cursor"), else => return err, }; return http_api.json(request, .ok, body);}
pub fn subscribeRepos(request: *http_api.Request) !void { const active = @atomicRmw(usize, &subscribe_repos_connections, .Add, 1, .monotonic); if (active >= max_subscribe_repos_connections) { _ = @atomicRmw(usize, &subscribe_repos_connections, .Sub, 1, .monotonic); log.debug("sync subscribeRepos rejected too_many_connections active={d} max={d}\n", .{ active + 1, max_subscribe_repos_connections }); return http_api.xrpcError(request, .too_many_requests, "RateLimitExceeded", "too many subscribeRepos connections"); } var upgraded = false; defer if (!upgraded) { _ = @atomicRmw(usize, &subscribe_repos_connections, .Sub, 1, .monotonic); };
var cursor_buf: [32]u8 = undefined; const cursor: u64 = if (http_api.queryParam(request.url.raw, "cursor", &cursor_buf)) |raw| std.fmt.parseInt(u64, raw, 10) catch 0 else 0;
const ctx = SubscribeReposClient.Context{ .cursor = cursor }; upgraded = try http_api.upgradeWebsocket(SubscribeReposClient, request, &ctx); if (!upgraded) { return http_api.xrpcError(request, .upgrade_required, "InvalidRequest", "Expected WebSocket upgrade"); }}
pub fn getRepoStatus(request: *http_api.Request) !void { var arena = std.heap.ArenaAllocator.init(std.heap.page_allocator); defer arena.deinit(); const allocator = arena.allocator();
var did_buf: [256]u8 = undefined; const did = http_api.queryParam(request.url.raw, "did", &did_buf) orelse { return http_api.xrpcError(request, .bad_request, "InvalidRequest", "Missing did"); }; const body = store.writeRepoStatusJson(allocator, did) catch |err| switch (err) { error.RepoNotFound => return http_api.xrpcError(request, .not_found, "RepoNotFound", "Repo not found"), else => return err, }; return http_api.json(request, .ok, body);}
pub fn notifyOfUpdate(request: *http_api.Request) !void { return http_api.json(request, .ok, "{}");}
pub fn requestCrawl(request: *http_api.Request) !void { return http_api.json(request, .ok, "{}");}
fn requirePublicRepoAvailable(request: *http_api.Request, did: []const u8) !void { const status = store.accountStatus(did) catch |err| switch (err) { error.RepoNotFound => return http_api.xrpcError(request, .not_found, "RepoNotFound", "Repo not found"), else => return err, }; switch (status) { .active => return, .takendown => return http_api.xrpcError(request, .bad_request, "RepoTakendown", "Repo has been taken down"), .suspended => return http_api.xrpcError(request, .bad_request, "RepoSuspended", "Repo is suspended"), .deactivated => return http_api.xrpcError(request, .bad_request, "RepoDeactivated", "Repo is deactivated"), .deleted => return http_api.xrpcError(request, .not_found, "RepoNotFound", "Repo not found"), }}
pub fn notifyCrawlers(force: bool) void { const now = nowMillis(); if (!force) { const last = crawler_last_notified_ms.load(.acquire); if (last != 0 and now - last < notify_threshold_ms) return; }
if (crawler_notify_in_flight.cmpxchgStrong(false, true, .acq_rel, .acquire) != null) return; crawler_last_notified_ms.store(now, .release); const thread = std.Thread.spawn(.{}, notifyCrawlersThread, .{}) catch |err| { crawler_notify_in_flight.store(false, .release); log.err("sync request_crawl spawn failed err={s}\n", .{@errorName(err)}); return; }; thread.detach();}
fn nowMillis() i64 { return clock.nowMillis();}
fn notifyCrawlersThread() void { defer crawler_notify_in_flight.store(false, .release); var arena = std.heap.ArenaAllocator.init(std.heap.page_allocator); defer arena.deinit(); requestConfiguredCrawls(arena.allocator());}
fn requestConfiguredCrawls(allocator: std.mem.Allocator) void { const host = publicHostname(config.publicUrl()) orelse { log.err("sync request_crawl skipped invalid_public_url url={s}\n", .{config.publicUrl()}); return; }; var crawlers = std.mem.splitScalar(u8, config.crawlers(), ','); while (crawlers.next()) |crawler_raw| { const crawler = std.mem.trim(u8, crawler_raw, " \t\r\n/"); if (crawler.len == 0) continue; requestCrawler(allocator, crawler, host) catch |err| { log.err("sync request_crawl failed crawler={s} host={s} err={s}\n", .{ crawler, host, @errorName(err) }); continue; }; }}
fn requestCrawler(allocator: std.mem.Allocator, crawler: []const u8, host: []const u8) !void { const url = try std.fmt.allocPrint(allocator, "{s}/xrpc/com.atproto.sync.requestCrawl", .{crawler}); const payload = try std.fmt.allocPrint(allocator, "{{\"hostname\":{f}}}", .{std.json.fmt(host, .{})}); var transport = zat.HttpTransport.init(store.currentIo(), allocator); defer transport.deinit(); const result = try transport.fetch(.{ .url = url, .method = .POST, .payload = payload, .content_type = "application/json", .max_response_size = 64 * 1024, }); if (@intFromEnum(result.status) < 200 or @intFromEnum(result.status) >= 300) return error.CrawlerRejected; log.info("sync request_crawl ok crawler={s} host={s}\n", .{ crawler, host });}
fn publicHostname(url: []const u8) ?[]const u8 { const scheme = std.mem.indexOf(u8, url, "://") orelse return null; var rest = url[scheme + 3 ..]; if (std.mem.indexOfScalar(u8, rest, '/')) |slash| rest = rest[0..slash]; if (std.mem.indexOfScalar(u8, rest, '@')) |at| rest = rest[at + 1 ..]; if (rest.len == 0) return null; if (std.mem.startsWith(u8, rest, "[")) { const close = std.mem.indexOfScalar(u8, rest, ']') orelse return null; return rest[0 .. close + 1]; } if (std.mem.indexOfScalar(u8, rest, ':')) |colon| return rest[0..colon]; return rest;}