atproto pds in zig
Something went wrong. Try again.
123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439const std = @import("std");const auth = @import("../auth/tokens.zig");const clock = @import("../core/clock.zig");const dpop = @import("../internal/dpop.zig");const log = @import("../core/log.zig");const store = @import("../storage/store.zig");const httpz = @import("httpz");
const http = std.http;pub const Request = httpz.Request;pub const Response = httpz.Response;
threadlocal var active_response: ?*Response = null;
pub fn bindResponse(res: *Response) void { active_response = res;}
pub fn unbindResponse() void { active_response = null;}
pub fn currentStatus() u16 { return response().status;}
fn response() *Response { return active_response orelse @panic("http response not bound");}
pub const BearerAccount = struct { account: auth.Account, oauth_scope: ?[]const u8, oauth_client_id: ?[]const u8 = null,};
const TokenScheme = enum { bearer, dpop,
fn name(self: TokenScheme) []const u8 { return switch (self) { .bearer => "Bearer", .dpop => "DPoP", }; }};
pub fn requireBearerAccount(request: *const Request, allocator: std.mem.Allocator) !auth.Account { return (try requireBearerAccountWithScope(request, allocator, "com.atproto.access")).account;}
pub fn requireBearerAccess(request: *const Request, allocator: std.mem.Allocator) !BearerAccount { return requireBearerAccountWithScope(request, allocator, "com.atproto.access");}
pub fn requireBearerAccountWithScope(request: *const Request, allocator: std.mem.Allocator, scope: []const u8) !BearerAccount { const auth_header = headerValue(request, "authorization") orelse { return error.AuthRequired; };
const token_scheme: TokenScheme = if (std.ascii.startsWithIgnoreCase(auth_header, "bearer ")) .bearer else if (std.ascii.startsWithIgnoreCase(auth_header, "dpop ")) .dpop else return error.AuthRequired; const token_start: usize = if (token_scheme == .bearer) "bearer ".len else "dpop ".len;
const token = std.mem.trim(u8, auth_header[token_start..], " \t"); const claims = auth.claimsFromSessionJwt(allocator, token) orelse { log.err("bearer auth invalid: jwt parse/verify failed token_len={d}\n", .{token.len}); challengeInvalidToken(allocator, token_scheme, "Token is invalid"); return error.InvalidToken; }; defer allocator.free(claims.did); defer allocator.free(claims.scope); defer allocator.free(claims.jti); defer if (claims.cnf_jkt) |jkt| allocator.free(jkt); if (!scopeAllows(claims.scope, scope)) { log.err("bearer auth invalid: session scope reject claims_scope={s} required={s} did={s}\n", .{ claims.scope, scope, claims.did }); return error.InvalidToken; } const account = (store.findAccount(allocator, claims.did) catch |err| { log.err("bearer auth invalid: account lookup failed did={s} err={s}\n", .{ claims.did, @errorName(err) }); return error.InvalidToken; }) orelse { log.err("bearer auth invalid: account not found did={s}\n", .{claims.did}); return error.InvalidToken; };
const oauth_token = store.getOAuthToken(allocator, token) catch |err| { log.err("bearer auth invalid: oauth token lookup failed did={s} err={s}\n", .{ claims.did, @errorName(err) }); return error.InvalidToken; }; if (oauth_token) |row| { const ts = now(); if (row.kind != .access) { log.err("oauth auth invalid: non-access token used for resource auth kind={s} did={s}\n", .{ @tagName(row.kind), row.did }); return error.InvalidToken; } if (row.revoked) { log.err("bearer auth invalid: oauth row revoked=true access_expires_at={d} now={d} did={s}\n", .{ row.access_expires_at, ts, row.did }); challengeInvalidToken(allocator, token_scheme, "Token has been revoked"); return error.InvalidToken; } if (row.access_expires_at < ts) { log.err("bearer auth invalid: oauth row revoked=false access_expires_at={d} now={d} did={s}\n", .{ row.access_expires_at, ts, row.did }); challengeInvalidToken(allocator, token_scheme, "Token has expired"); return error.InvalidToken; } if (!std.mem.eql(u8, row.did, claims.did)) { log.err("bearer auth invalid: oauth row did mismatch row_did={s} claims_did={s}\n", .{ row.did, claims.did }); return error.InvalidToken; } if (row.dpop_jkt) |expected_jkt| { if (token_scheme != .dpop) { log.err("oauth auth invalid: dpop-bound token used without DPoP auth scheme did={s}\n", .{claims.did}); return error.InvalidToken; } if (claims.cnf_jkt == null or !std.mem.eql(u8, claims.cnf_jkt.?, expected_jkt)) { log.err("oauth auth invalid: token cnf mismatch did={s}\n", .{claims.did}); return error.InvalidToken; } _ = dpop.verifyRequest(allocator, request, token, expected_jkt) catch |err| switch (err) { error.UseDpopNonce, error.MissingProof => { dpop.challengeResource(@constCast(request), allocator) catch {}; return error.InvalidToken; }, else => { log.err("oauth auth invalid: dpop proof failed did={s} err={s}\n", .{ claims.did, @errorName(err) }); return error.InvalidToken; }, }; dpop.attachNonce(allocator) catch {}; } else if (token_scheme != .bearer or claims.cnf_jkt != null) { log.err("oauth auth invalid: bearer token/proof binding mismatch did={s}\n", .{claims.did}); return error.InvalidToken; } return .{ .account = account, .oauth_scope = row.scope, .oauth_client_id = row.client_id }; }
if (claims.cnf_jkt != null) { log.err("bearer auth invalid: cnf-bound token had no oauth row did={s}\n", .{claims.did}); return error.InvalidToken; }
const active = store.sessionTokenIsActive(claims.did, claims.jti, claims.scope) catch |err| { log.err("bearer auth invalid: session token active lookup failed did={s} jti={s} scope={s} err={s}\n", .{ claims.did, claims.jti, claims.scope, @errorName(err) }); return error.InvalidToken; }; if (!active) { log.err("bearer auth invalid: session token inactive did={s} jti={s} scope={s}\n", .{ claims.did, claims.jti, claims.scope }); return error.InvalidToken; } return .{ .account = account, .oauth_scope = null };}
fn challengeInvalidToken(allocator: std.mem.Allocator, scheme: TokenScheme, description: []const u8) void { var buf: [160]u8 = undefined; const value = formatInvalidTokenChallenge(&buf, scheme, description) catch return; addResponseHeader("www-authenticate", value) catch return; if (scheme == .dpop) dpop.attachNonce(allocator) catch {};}
fn formatInvalidTokenChallenge(buf: []u8, scheme: TokenScheme, description: []const u8) ![]const u8 { return std.fmt.bufPrint( buf, "{s} error=\"invalid_token\", error_description=\"{s}\"", .{ scheme.name(), description }, );}
fn scopeAllows(actual: []const u8, required: []const u8) bool { if (std.mem.eql(u8, actual, required)) return true; if (!std.mem.eql(u8, required, "com.atproto.access")) return false; return std.mem.eql(u8, actual, "com.atproto.appPass") or std.mem.eql(u8, actual, "com.atproto.appPassPrivileged");}
pub fn optionalBearerAccount(request: *const Request, allocator: std.mem.Allocator) ?auth.Account { return requireBearerAccount(request, allocator) catch null;}
fn now() i64 { return clock.now();}
pub fn parseJsonBody(request: *Request, allocator: std.mem.Allocator, body: []const u8) !std.json.Parsed(std.json.Value) { return std.json.parseFromSlice(std.json.Value, allocator, body, .{}) catch { try xrpcError(request, .bad_request, "InvalidRequest", "Expected JSON body"); return error.InvalidJson; };}
pub fn valueString(value: std.json.Value, key: []const u8) ?[]const u8 { return switch (value) { .object => |object| switch (object.get(key) orelse return null) { .string => |string| string, else => null, }, else => null, };}
pub fn readBody(request: *Request, buf: []u8) ![]const u8 { const body = request.body() orelse ""; if (body.len > buf.len) return error.BodyTooLarge; @memcpy(buf[0..body.len], body); return buf[0..body.len];}
pub fn readBodyAlloc(request: *Request, allocator: std.mem.Allocator, max_len: usize) ![]const u8 { const body = request.body() orelse ""; if (body.len > max_len) return error.BodyTooLarge; return try allocator.dupe(u8, body);}
pub fn headerValue(request: *const Request, name: []const u8) ?[]const u8 { var lower_buf: [128]u8 = undefined; const lower = if (name.len <= lower_buf.len) blk: { for (name, 0..) |c, i| lower_buf[i] = std.ascii.toLower(c); break :blk lower_buf[0..name.len]; } else name; return request.header(lower);}
pub fn queryParam(target: []const u8, name: []const u8, out: []u8) ?[]const u8 { const query_start = std.mem.indexOfScalar(u8, target, '?') orelse return null; var params = std.mem.splitScalar(u8, target[query_start + 1 ..], '&'); while (params.next()) |param| { const eq = std.mem.indexOfScalar(u8, param, '=') orelse continue; if (!std.mem.eql(u8, param[0..eq], name)) continue; return percentDecode(param[eq + 1 ..], out) catch null; } return null;}
pub fn queryLimit(target: []const u8, default: usize) usize { var buf: [16]u8 = undefined; const raw = queryParam(target, "limit", &buf) orelse return default; return std.fmt.parseInt(usize, raw, 10) catch default;}
pub fn percentDecode(input: []const u8, out: []u8) ![]const u8 { var write: usize = 0; var read: usize = 0; while (read < input.len) { if (write >= out.len) return error.NoSpaceLeft; switch (input[read]) { '%' => { if (read + 2 >= input.len) return error.InvalidPercentEncoding; out[write] = try std.fmt.parseInt(u8, input[read + 1 .. read + 3], 16); read += 3; }, '+' => { out[write] = ' '; read += 1; }, else => |c| { out[write] = c; read += 1; }, } write += 1; } return out[0..write];}
pub fn corsPreflight(request: *Request) !void { const res = response(); const allow_headers = if (headerValue(request, "access-control-request-headers")) |requested| blk: { const trimmed = std.mem.trim(u8, requested, " \t"); break :blk if (trimmed.len > 0) trimmed else default_cors_allow_headers; } else default_cors_allow_headers;
const headers = [_]http.Header{ .{ .name = "access-control-allow-origin", .value = "*" }, .{ .name = "access-control-allow-methods", .value = "GET, HEAD, POST, OPTIONS" }, .{ .name = "vary", .value = "Access-Control-Request-Headers" }, .{ .name = "access-control-allow-headers", .value = allow_headers }, .{ .name = "access-control-allow-private-network", .value = "true" }, .{ .name = "access-control-max-age", .value = "600" }, .{ .name = "connection", .value = "close" }, }; try setHeaders(res, &headers); res.setStatus(.no_content);}
pub fn toStdMethod(method: httpz.Method) http.Method { return switch (method) { .GET => .GET, .HEAD => .HEAD, .POST => .POST, .PUT => .PUT, .PATCH => .PATCH, .DELETE => .DELETE, .OPTIONS => .OPTIONS, else => .GET, };}
pub fn upgradeWebsocket(comptime Handler: type, request: *Request, ctx: anytype) !bool { return httpz.upgradeWebsocket(Handler, request, response(), ctx);}
pub fn json(request: *Request, status: http.Status, body: []const u8) !void { _ = request; const res = response(); try setHeaders(res, &json_headers); res.setStatus(status); res.body = try res.arena.dupe(u8, body);}
pub fn empty(request: *Request, status: http.Status) !void { _ = request; const res = response(); try setHeaders(res, &empty_headers); res.setStatus(status); res.body = "";}
pub fn text(request: *Request, status: http.Status, body: []const u8) !void { _ = request; const res = response(); try setHeaders(res, &text_headers); res.setStatus(status); res.body = try res.arena.dupe(u8, body);}
pub fn xrpcError( request: *Request, status: http.Status, error_name: []const u8, message: []const u8,) !void { var buf: [512]u8 = undefined; const body = try std.fmt.bufPrint(&buf, "{{\"error\":\"{s}\",\"message\":\"{s}\"}}", .{ error_name, message }); try json(request, status, body);}
const json_headers = [_]http.Header{ .{ .name = "content-type", .value = "application/json; charset=utf-8" }, .{ .name = "access-control-allow-origin", .value = "*" }, .{ .name = "access-control-expose-headers", .value = "dpop-nonce, www-authenticate" }, .{ .name = "access-control-allow-private-network", .value = "true" }, .{ .name = "connection", .value = "close" },};
const empty_headers = [_]http.Header{ .{ .name = "access-control-allow-origin", .value = "*" }, .{ .name = "access-control-expose-headers", .value = "dpop-nonce, www-authenticate" }, .{ .name = "access-control-allow-private-network", .value = "true" }, .{ .name = "connection", .value = "close" },};
const text_headers = [_]http.Header{ .{ .name = "content-type", .value = "text/plain; charset=utf-8" }, .{ .name = "access-control-allow-origin", .value = "*" }, .{ .name = "access-control-expose-headers", .value = "dpop-nonce, www-authenticate" }, .{ .name = "access-control-allow-private-network", .value = "true" }, .{ .name = "connection", .value = "close" },};
const default_cors_allow_headers = "atproto-accept-labelers, atproto-proxy, authorization, content-type, dpop, x-bsky-topics";
pub fn respond(request: *Request, status: http.Status, body: []const u8, headers: []const http.Header) !void { const res = response(); if (wantsConnectionClose(headers)) request.conn.handover = .close; try setHeaders(res, headers); res.setStatus(status); res.body = try res.arena.dupe(u8, body);}
pub fn respondNowClose(request: *Request, status: http.Status, body: []const u8, headers: []const http.Header) !void { const res = response(); var out: std.Io.Writer.Allocating = .init(res.arena); defer out.deinit();
try out.writer.print("HTTP/1.1 {d} \r\n", .{@intFromEnum(status)}); for (headers) |header| { try out.writer.print("{s}: {s}\r\n", .{ header.name, header.value }); } try out.writer.print("content-length: {d}\r\n\r\n", .{body.len});
request.conn.handover = .close; try request.conn.writeAll(out.written()); if (body.len > 0) try request.conn.writeAll(body); res.written = true;}
fn setHeaders(res: *Response, headers: []const http.Header) !void { for (headers) |header| try addHeader(res, header.name, header.value);}
fn wantsConnectionClose(headers: []const http.Header) bool { for (headers) |header| { if (std.ascii.eqlIgnoreCase(header.name, "connection") and std.ascii.eqlIgnoreCase(header.value, "close")) { return true; } } return false;}
pub fn addResponseHeader(name: []const u8, value: []const u8) !void { try addHeader(response(), name, value);}
fn addHeader(res: *Response, name: []const u8, value: []const u8) !void { for (res.headers.keys[0..res.headers.len], 0..) |existing, i| { if (std.ascii.eqlIgnoreCase(existing, name)) { res.headers.keys[i] = try res.arena.dupe(u8, name); res.headers.values[i] = try res.arena.dupe(u8, value); return; } } try res.headerOpts(name, value, .{ .dupe_name = true, .dupe_value = true });}
test "decodes query params" { var buf: [64]u8 = undefined; try std.testing.expectEqualStrings("did:plc:service", queryParam("/xrpc/foo?aud=did%3Aplc%3Aservice", "aud", &buf).?);}
test "formats OAuth invalid token challenges" { var buf: [160]u8 = undefined; try std.testing.expectEqualStrings( "Bearer error=\"invalid_token\", error_description=\"Token has expired\"", try formatInvalidTokenChallenge(&buf, .bearer, "Token has expired"), ); try std.testing.expectEqualStrings( "DPoP error=\"invalid_token\", error_description=\"Token has been revoked\"", try formatInvalidTokenChallenge(&buf, .dpop, "Token has been revoked"), );}