const 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).?); }