const 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"), ); }