atproto pds in zig pds.zat.dev
pds atproto
Something went wrong. Try again.
123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281const std = @import("std");const auth = @import("../auth/tokens.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;}
fn response() *Response { return active_response orelse @panic("http response not bound");}
pub const BearerAccount = struct { account: auth.Account, oauth_scope: ?[]const u8,};
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_start = if (std.ascii.startsWithIgnoreCase(auth_header, "bearer ")) "bearer ".len else if (std.ascii.startsWithIgnoreCase(auth_header, "dpop ")) "dpop ".len else return error.AuthRequired;
const token = std.mem.trim(u8, auth_header[token_start..], " \t"); const claims = auth.claimsFromSessionJwt(allocator, token) orelse return error.InvalidToken; defer allocator.free(claims.did); defer allocator.free(claims.scope); defer allocator.free(claims.jti); if (!scopeAllows(claims.scope, scope)) return error.InvalidToken; const account = (store.findAccount(allocator, claims.did) catch null) orelse return error.InvalidToken;
const oauth_token = store.getOAuthToken(allocator, token) catch return error.InvalidToken; if (oauth_token) |row| { if (row.revoked or row.expires_at < now()) return error.InvalidToken; if (!std.mem.eql(u8, row.did, claims.did)) return error.InvalidToken; return .{ .account = account, .oauth_scope = row.scope }; }
if (!(store.sessionTokenIsActive(claims.did, claims.jti, claims.scope) catch false)) return error.InvalidToken; return .{ .account = account, .oauth_scope = null };}
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 { var ts: std.posix.timespec = undefined; return switch (std.posix.errno(std.posix.system.clock_gettime(.REALTIME, &ts))) { .SUCCESS => ts.sec, else => 0, };}
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(); if (headerValue(request, "access-control-request-headers")) |requested| { const trimmed = std.mem.trim(u8, requested, " \t"); if (trimmed.len > 0) try addHeader(res, "access-control-allow-headers", trimmed); } try setHeaders(res, &cors_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-allow-private-network", .value = "true" }, .{ .name = "connection", .value = "close" },};
const empty_headers = [_]http.Header{ .{ .name = "access-control-allow-origin", .value = "*" }, .{ .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-allow-private-network", .value = "true" }, .{ .name = "connection", .value = "close" },};
const cors_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 = "atproto-accept-labelers, atproto-proxy, authorization, content-type, dpop, x-bsky-topics" }, .{ .name = "access-control-allow-private-network", .value = "true" }, .{ .name = "access-control-max-age", .value = "600" }, .{ .name = "connection", .value = "close" },};
pub fn respond(request: *Request, status: http.Status, body: []const u8, headers: []const http.Header) !void { _ = request; const res = response(); try setHeaders(res, headers); res.setStatus(status); res.body = try res.arena.dupe(u8, body);}
fn setHeaders(res: *Response, headers: []const http.Header) !void { for (headers) |header| try addHeader(res, header.name, header.value);}
fn addHeader(res: *Response, name: []const u8, value: []const u8) !void { 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).?);}