diff --git a/.changeset/atproto-service-auth.md b/.changeset/atproto-service-auth.md new file mode 100644 index 0000000..d677108 --- /dev/null +++ b/.changeset/atproto-service-auth.md @@ -0,0 +1,5 @@ +--- +"@atmo-dev/contrail": minor +--- + +Add discoverable AT Protocol service authentication for personalized feeds and authoritative update notifications. diff --git a/packages/contrail/package.json b/packages/contrail/package.json index 5ab9994..f754305 100644 --- a/packages/contrail/package.json +++ b/packages/contrail/package.json @@ -79,12 +79,14 @@ "@atcute/lexicon-doc": "3.0.2", "@atcute/lexicons": "^2.0.3", "@atcute/tid": "1.1.4", + "@atcute/xrpc-server": "^2.0.2", "cac": "^7.0.0", "hono": "^4.13.0", "jiti": "^2.7.0", "valibot": "1.4.2" }, "devDependencies": { + "@atcute/crypto": "2.4.4", "@cloudflare/workers-types": "^5.20260804.1", "@types/node": "^26.1.2", "@types/pg": "^8.20.4", diff --git a/packages/contrail/src/cli/commands/connect.ts b/packages/contrail/src/cli/commands/connect.ts index b0d18b2..e973751 100644 --- a/packages/contrail/src/cli/commands/connect.ts +++ b/packages/contrail/src/cli/commands/connect.ts @@ -24,6 +24,7 @@ import { isPublicServiceManifest, normalizePublicServiceEndpoint, validateManifestContract, + type PublicServiceAuthContract, type PublicServiceManifest, } from "../../public-service.js"; import { generateLexiconTypesWithAtcute } from "../atcute.js"; @@ -47,6 +48,7 @@ export interface ProviderLock { contractDigest: string; lexiconDigest: string; methods: string[]; + serviceAuth: PublicServiceAuthContract | null; lexiconRoot: string; } @@ -257,6 +259,7 @@ export async function connectPublicService(options: { contractDigest: manifest.contract.digest, lexiconDigest: manifest.lexicons.digest, methods: [...manifest.methods].sort(), + serviceAuth: manifest.serviceAuth ?? null, lexiconRoot: relative(projectRoot, providerRoot), }; const lockDirectory = dirname(lockPath); diff --git a/packages/contrail/src/core/router/feed.ts b/packages/contrail/src/core/router/feed.ts index 8ef149e..6dc0974 100644 --- a/packages/contrail/src/core/router/feed.ts +++ b/packages/contrail/src/core/router/feed.ts @@ -1,3 +1,4 @@ +import type { Nsid } from "@atcute/lexicons/syntax"; import type { Context, Hono } from "hono"; import type { ContrailConfig, @@ -16,6 +17,7 @@ import { import { resolveActor } from "../identity"; import { backfillUser } from "../backfill"; import { runPipeline } from "./collection"; +import type { ServiceAuthGate } from "../service-auth"; const BACKFILL_TIMEOUT_MS = 30_000; const BACKFILL_REQUEST_TIMEOUT_MS = 10_000; @@ -223,7 +225,8 @@ async function maybeBackfillFeed( export function registerFeedRoutes( app: Hono, db: Database, - config: ContrailConfig + config: ContrailConfig, + serviceAuth?: ServiceAuthGate | null ): void { if (!config.feeds) return; @@ -243,8 +246,23 @@ export function registerFeedRoutes( return c.json({ error: "Unknown feed" }, 404); } + const method = `${ns}.getFeed` as Nsid; + const authorization = serviceAuth?.protects("getFeed") + ? await serviceAuth.authorize(c.req.raw, method) + : null; + if (authorization?.response) return authorization.response; + const did = await resolveActor(db, actor, config); if (!did) return c.json({ error: "Could not resolve actor" }, 400); + if ( + authorization?.principal?.issuer !== undefined && + authorization.principal.issuer !== did + ) { + return c.json( + { error: "feed actor must match the service-auth issuer" }, + 403 + ); + } await maybeBackfillFeed(c, db, config, did, feedName, feedConfig); diff --git a/packages/contrail/src/core/router/index.ts b/packages/contrail/src/core/router/index.ts index a928954..d7ff777 100644 --- a/packages/contrail/src/core/router/index.ts +++ b/packages/contrail/src/core/router/index.ts @@ -11,6 +11,7 @@ import { registerCollectionRoutes } from "./collection"; import { registerFeedRoutes } from "./feed"; import { registerNotifyRoute } from "./notify"; import { resolveProfiles } from "./profiles"; +import { createServiceAuthGate } from "../service-auth"; import { describePublicService, normalizeLexiconDocuments, @@ -32,7 +33,18 @@ export function createApp( options: CreateAppOptions = {}, ): Hono { const app = new Hono(); - app.use("*", cors()); + app.use( + "*", + cors({ + allowHeaders: [ + "Authorization", + "Content-Type", + "Atproto-Accept-Labelers", + ], + exposeHeaders: ["Atproto-Content-Labelers", "WWW-Authenticate"], + }), + ); + const serviceAuth = createServiceAuthGate(config); app.get("/", (c) => c.json({ status: "ok" })); app.get("/status", async (c) => { @@ -72,6 +84,27 @@ export function createApp( c.header("etag", `\"${manifest.contract.digest}\"`); return c.json(manifest); }); + if ( + serviceAuth && + serviceAuth.audience === + `did:web:${new URL(options.publicService.endpoint).hostname}` + ) { + app.get("/.well-known/did.json", (c) => { + c.header("content-type", "application/did+ld+json; charset=UTF-8"); + c.header("cache-control", "public, max-age=300"); + return c.json({ + "@context": ["https://www.w3.org/ns/did/v1"], + id: serviceAuth.audience, + service: [ + { + id: `${serviceAuth.audience}#contrail`, + type: "ContrailService", + serviceEndpoint: options.publicService!.endpoint, + }, + ], + }); + }); + } app.get("/lexicons", async (c) => { const service = await description; c.header("content-type", "application/json; charset=UTF-8"); @@ -140,8 +173,8 @@ export function createApp( registerCursorRoute(app, db, config); registerCollectionRoutes(app, db, config); - registerFeedRoutes(app, db, config); - registerNotifyRoute(app, db, config); + registerFeedRoutes(app, db, config, serviceAuth); + registerNotifyRoute(app, db, config, serviceAuth); return app; } diff --git a/packages/contrail/src/core/router/notify.ts b/packages/contrail/src/core/router/notify.ts index 9ecf5e5..1f11282 100644 --- a/packages/contrail/src/core/router/notify.ts +++ b/packages/contrail/src/core/router/notify.ts @@ -1,11 +1,12 @@ import type { Hono } from "hono"; import type { Database, ContrailConfig, IngestEvent } from "../types"; +import type { ServiceAuthGate } from "../service-auth"; import { shortNameForNsid, getFeedMutatingNsids } from "../types"; import { lookupExistingRecords } from "../db/records"; import { createIngestEvent, ingestRecords, recordTimeUs } from "../ingest"; import { runGatedFeedPrune } from "../jetstream"; import { getPDS } from "../client"; -import type { Did } from "@atcute/lexicons"; +import type { Did, Nsid } from "@atcute/lexicons"; import { parseCanonicalResourceUri } from "@atcute/lexicons/syntax"; /** Parse a canonical (DID-authority) record AT-URI into its components, or null @@ -286,7 +287,8 @@ export async function processNotifyUris( export function registerNotifyRoute( app: Hono, db: Database, - config: ContrailConfig + config: ContrailConfig, + serviceAuth?: ServiceAuthGate | null ) { // Endpoint is off by default. Set config.notify to true or a secret string to enable. if (!config.notify) return; @@ -295,6 +297,12 @@ export function registerNotifyRoute( const secret = typeof config.notify === "string" ? config.notify : null; app.post(`/xrpc/${ns}.notifyOfUpdate`, async (c) => { + const method = `${ns}.notifyOfUpdate` as Nsid; + const authorization = serviceAuth?.protects("notifyOfUpdate") + ? await serviceAuth.authorize(c.req.raw, method) + : null; + if (authorization?.response) return authorization.response; + if (secret) { const auth = c.req.header("Authorization"); if (auth !== `Bearer ${secret}`) { @@ -316,6 +324,17 @@ export function registerNotifyRoute( if (uris.length > MAX_NOTIFY_URIS) { return c.json({ error: `max ${MAX_NOTIFY_URIS} URIs per request` }, 400); } + if (authorization?.principal) { + for (const uri of uris) { + const parsed = parseAtUri(uri); + if (parsed && parsed.did !== authorization.principal.issuer) { + return c.json( + { error: "notified records must belong to the service-auth issuer" }, + 403 + ); + } + } + } const result = await processNotifyUris(db, config, uris); return c.json(result); diff --git a/packages/contrail/src/core/service-auth.ts b/packages/contrail/src/core/service-auth.ts new file mode 100644 index 0000000..0eec5bc --- /dev/null +++ b/packages/contrail/src/core/service-auth.ts @@ -0,0 +1,68 @@ +import { + CompositeDidDocumentResolver, + PlcDidDocumentResolver, + WebDidDocumentResolver, + type DidDocumentResolver, +} from "@atcute/identity-resolver"; +import type { Did, Nsid } from "@atcute/lexicons/syntax"; +import { ServiceJwtVerifier, type VerifiedJwt } from "@atcute/xrpc-server/auth"; +import { XRPCError } from "@atcute/xrpc-server"; +import type { AtprotoServiceAuthMethod, ContrailConfig } from "./types.js"; + +const AUTH_TIMEOUT_MS = 5_000; + +export interface ServiceAuthResult { + principal?: VerifiedJwt; + response?: Response; +} + +export interface ServiceAuthGate { + readonly audience: Did; + protects(method: AtprotoServiceAuthMethod): boolean; + authorize(request: Request, method: Nsid): Promise; +} + +function defaultResolver(): DidDocumentResolver { + return new CompositeDidDocumentResolver({ + methods: { + plc: new PlcDidDocumentResolver(), + web: new WebDidDocumentResolver(), + }, + }); +} + +/** Create the shared verifier used by protected built-in routes. Tokens remain + * method-bound even when a client obtained permission through one wildcard + * OAuth scope (`rpc?lxm=*&aud=`). */ +export function createServiceAuthGate( + config: ContrailConfig, +): ServiceAuthGate | null { + if (!config.serviceAuth) return null; + const serviceAuth = config.serviceAuth; + const protectedMethods = new Set(serviceAuth.methods); + const audience = serviceAuth.audience as Did; + const verifier = new ServiceJwtVerifier({ + acceptAudiences: [audience], + resolver: serviceAuth.resolver ?? defaultResolver(), + maxAge: serviceAuth.maxTokenAgeSeconds, + }); + + return { + audience, + protects(method) { + return protectedMethods.has(method); + }, + async authorize(request, method) { + try { + const principal = await verifier.verifyRequest(request, { + lxm: method, + signal: AbortSignal.timeout(AUTH_TIMEOUT_MS), + }); + return { principal }; + } catch (error) { + if (error instanceof XRPCError) return { response: error.toResponse() }; + throw error; + } + }, + }; +} diff --git a/packages/contrail/src/core/types.ts b/packages/contrail/src/core/types.ts index e71cdb5..610c683 100644 --- a/packages/contrail/src/core/types.ts +++ b/packages/contrail/src/core/types.ts @@ -1,4 +1,5 @@ import type { LexiconDoc } from "@atcute/lexicon-doc"; +import { isDid } from "@atcute/lexicons/syntax"; import type { SqlDialect } from "./dialect"; // Database interface — D1 implements this natively @@ -245,6 +246,19 @@ export interface OrderedSourceConfig { epoch: string; } +export type AtprotoServiceAuthMethod = "getFeed" | "notifyOfUpdate"; + +export interface AtprotoServiceAuthConfig { + /** Plain service DID used as the exact JWT audience. */ + audience: string; + /** Built-in methods that require a method-bound AT Protocol service token. */ + methods: AtprotoServiceAuthMethod[]; + /** Maximum accepted token lifetime and age. Default: 300 seconds. */ + maxTokenAgeSeconds?: number; + /** Optional DID resolver for private networks or controlled resolution. */ + resolver?: import("@atcute/identity-resolver").DidDocumentResolver; +} + export interface ContrailConfig { namespace: string; /** Collections to index, keyed by short name. Short names become endpoint URL segments @@ -270,8 +284,11 @@ export interface ContrailConfig { feeds?: Record; logger?: Logger; /** Expose the notifyOfUpdate HTTP endpoint. Off by default. - * Set to `true` for open access, or a string to require `Authorization: Bearer `. */ + * Set to `true` for open access, or a string to require `Authorization: Bearer `. + * Prefer `serviceAuth.methods: ["notifyOfUpdate"]` for portable user auth. */ notify?: boolean | string; + /** Verify method-bound AT Protocol service JWTs for selected built-in routes. */ + serviceAuth?: AtprotoServiceAuthConfig; /** Labels module configuration. When set, contrail subscribes to the * configured labelers, indexes their labels into a single `labels` table, * and hydrates `record.labels` onto `listRecords` / `getRecord` / profile @@ -399,6 +416,38 @@ export function resolveConfig(config: ContrailConfig): ResolvedContrailConfig { ) { throw new TypeError("orderedSource requires non-empty source and epoch values"); } + if (config.serviceAuth) { + if (!isDid(config.serviceAuth.audience)) { + throw new TypeError("serviceAuth.audience must be a plain DID"); + } + if ( + !Array.isArray(config.serviceAuth.methods) || + new Set(config.serviceAuth.methods).size !== config.serviceAuth.methods.length + ) { + throw new TypeError("serviceAuth.methods must contain unique methods"); + } + if ( + config.serviceAuth.methods.includes("getFeed") && + (!config.feeds || Object.keys(config.feeds).length === 0) + ) { + throw new TypeError("serviceAuth cannot protect getFeed without configured feeds"); + } + if ( + config.serviceAuth.methods.includes("notifyOfUpdate") && + config.notify !== true + ) { + throw new TypeError( + "serviceAuth notifyOfUpdate requires notify: true and replaces shared-secret auth", + ); + } + if ( + config.serviceAuth.maxTokenAgeSeconds !== undefined && + (!Number.isSafeInteger(config.serviceAuth.maxTokenAgeSeconds) || + config.serviceAuth.maxTokenAgeSeconds <= 0) + ) { + throw new TypeError("serviceAuth.maxTokenAgeSeconds must be a positive integer"); + } + } const profiles = (config.profiles ?? DEFAULT_PROFILES).map( normalizeProfileConfig ); diff --git a/packages/contrail/src/lexicons/generate.ts b/packages/contrail/src/lexicons/generate.ts index 56574c8..4107c9f 100644 --- a/packages/contrail/src/lexicons/generate.ts +++ b/packages/contrail/src/lexicons/generate.ts @@ -807,13 +807,19 @@ export function generateLexicons( } const feed = feedLexicon(config, sourceDirs); if (feed) emit(`${config.namespace}.getFeed`, feed); - if (surface === "full" && config.notify) { + if ( + config.notify && + (surface === "full" || + config.serviceAuth?.methods.includes("notifyOfUpdate") === true) + ) { emit(`${config.namespace}.notifyOfUpdate`, { lexicon: 1, id: `${config.namespace}.notifyOfUpdate`, defs: { main: { type: "procedure", + description: + "Fetch changed records from their authoritative PDS for immediate indexing", input: { encoding: "application/json", schema: { @@ -830,7 +836,18 @@ export function generateLexicons( }, output: { encoding: "application/json", - schema: { type: "unknown" }, + schema: { + type: "object", + required: ["indexed", "deleted"], + properties: { + indexed: { type: "integer", minimum: 0 }, + deleted: { type: "integer", minimum: 0 }, + errors: { + type: "array", + items: { type: "string" }, + }, + }, + }, }, }, }, diff --git a/packages/contrail/src/public-service.ts b/packages/contrail/src/public-service.ts index efc4085..2beef81 100644 --- a/packages/contrail/src/public-service.ts +++ b/packages/contrail/src/public-service.ts @@ -1,4 +1,4 @@ -import { isNsid } from "@atcute/lexicons/syntax"; +import { isDid, isNsid } from "@atcute/lexicons/syntax"; import type { ContrailConfig } from "./core/types.js"; import { getCollectionMethods, @@ -21,12 +21,24 @@ export interface PublicServiceCollection { references: string[]; } +export interface PublicServiceProtectedMethod { + id: string; + type: "query" | "procedure"; +} + +export interface PublicServiceAuthContract { + type: "atproto-service-auth"; + audience: string; + methods: PublicServiceProtectedMethod[]; +} + export interface PublicContract { format: "contrail.contract"; version: 1; namespace: string; collections: PublicServiceCollection[]; methods: string[]; + serviceAuth?: PublicServiceAuthContract | null; lexiconDigest: string; } @@ -40,6 +52,7 @@ export interface PublicServiceManifest { status: { url: string }; collections: PublicServiceCollection[]; methods: string[]; + serviceAuth?: PublicServiceAuthContract | null; } export interface LexiconDocument { @@ -151,22 +164,46 @@ function publicTopLevelMethods(config: ContrailConfig): string[] { return methods; } +function publicServiceAuth( + config: ContrailConfig, +): PublicServiceAuthContract | null { + if (!config.serviceAuth || config.serviceAuth.methods.length === 0) + return null; + const methods = config.serviceAuth.methods + .map((method): PublicServiceProtectedMethod => + method === "getFeed" + ? { id: `${config.namespace}.getFeed`, type: "query" } + : { id: `${config.namespace}.notifyOfUpdate`, type: "procedure" }, + ) + .sort((left, right) => left.id.localeCompare(right.id)); + return { + type: "atproto-service-auth", + audience: config.serviceAuth.audience, + methods, + }; +} + export function createPublicContract( config: ContrailConfig, lexiconDigest: string, ): PublicContract { const resolved = resolveConfig(config); const collections = publicCollections(resolved); + const serviceAuth = publicServiceAuth(resolved); + const protectedMethods = new Set( + serviceAuth?.methods.map((method) => method.id) ?? [], + ); const methods = [ ...publicTopLevelMethods(resolved), ...collections.flatMap((collection) => collection.methods), - ]; + ].filter((method) => !protectedMethods.has(method)); return { format: "contrail.contract", version: 1, namespace: resolved.namespace, collections, methods: [...new Set(methods)].sort(), + serviceAuth, lexiconDigest, }; } @@ -195,6 +232,15 @@ export function validateContractLexicons( ); } } + for (const method of contract.serviceAuth?.methods ?? []) { + const document = byId.get(method.id) as + { defs?: { main?: { type?: unknown } } } | undefined; + if (document?.defs?.main?.type !== method.type) { + throw new Error( + `protected method requires a matching ${method.type} Lexicon: ${method.id}`, + ); + } + } return lexicons; } @@ -261,6 +307,7 @@ export async function describePublicService( status: { url: `${endpoint}/status` }, collections: contract.collections, methods: contract.methods, + serviceAuth: contract.serviceAuth, }; return { endpoint, lexicons, manifest, canonicalLexicons }; } @@ -278,6 +325,7 @@ export function contractFromManifest( namespace: manifest.namespace, collections: manifest.collections, methods: manifest.methods, + serviceAuth: manifest.serviceAuth, lexiconDigest: manifest.lexicons.digest, }; } @@ -286,6 +334,9 @@ export function validateManifestContract( manifest: PublicServiceManifest, values: readonly object[], ): LexiconDocument[] { + if (!isPublicServiceAuthContract(manifest.serviceAuth)) { + throw new Error("service manifest contains invalid service auth"); + } if (!uniqueStrings(manifest.methods)) { throw new Error("service manifest contains duplicate methods"); } @@ -297,6 +348,21 @@ export function validateManifestContract( if (manifest.methods.some((method) => !method.startsWith(prefix))) { throw new Error("service manifest method is outside its namespace"); } + const protectedMethods = manifest.serviceAuth?.methods ?? []; + const protectedIds = protectedMethods.map((method) => method.id); + if (!uniqueStrings(protectedIds)) { + throw new Error("service manifest contains duplicate protected methods"); + } + if (protectedIds.some((method) => !method.startsWith(prefix))) { + throw new Error( + "service manifest protected method is outside its namespace", + ); + } + if (protectedIds.some((method) => manifest.methods.includes(method))) { + throw new Error( + "service manifest method cannot be both anonymous and protected", + ); + } const advertised = new Set(manifest.methods); for (const collection of manifest.collections) { if (!uniqueStrings(collection.methods)) { @@ -315,6 +381,28 @@ export function validateManifestContract( return validateContractLexicons(contractFromManifest(manifest), values); } +function isPublicServiceAuthContract( + value: unknown, +): value is PublicServiceAuthContract | null | undefined { + if (value === null || value === undefined) return true; + if (!value || typeof value !== "object") return false; + const auth = value as Partial; + return ( + auth.type === "atproto-service-auth" && + typeof auth.audience === "string" && + isDid(auth.audience) && + Array.isArray(auth.methods) && + auth.methods.every( + (method) => + !!method && + typeof method === "object" && + typeof method.id === "string" && + isNsid(method.id) && + (method.type === "query" || method.type === "procedure"), + ) + ); +} + export function isPublicServiceManifest( value: unknown, ): value is PublicServiceManifest { @@ -335,6 +423,7 @@ export function isPublicServiceManifest( typeof manifest.status?.url !== "string" || !Array.isArray(manifest.collections) || !Array.isArray(manifest.methods) || + !isPublicServiceAuthContract(manifest.serviceAuth) || !manifest.methods.every( (method) => typeof method === "string" && isNsid(method), ) diff --git a/packages/contrail/tests/connect.test.ts b/packages/contrail/tests/connect.test.ts index 58a4714..3e8e765 100644 --- a/packages/contrail/tests/connect.test.ts +++ b/packages/contrail/tests/connect.test.ts @@ -32,6 +32,12 @@ const sourceLexicon = { id: "community.lexicon.calendar.event", defs: { main: { type: "record" } }, }; +const notifyMethod = "atmo.rsvp.notifyOfUpdate"; +const notifyLexicon = { + lexicon: 1, + id: notifyMethod, + defs: { main: { type: "procedure" } }, +}; function providerLock(): ProviderLock { return { @@ -42,6 +48,7 @@ function providerLock(): ProviderLock { contractDigest: `sha256:${"a".repeat(64)}`, lexiconDigest: `sha256:${"b".repeat(64)}`, methods: [method], + serviceAuth: null, lexiconRoot: "lexicons/pulled/api.atmo.rsvp", }; } @@ -68,6 +75,7 @@ async function serviceFixture(values = [methodLexicon, sourceLexicon]) { }, ], methods: [method], + serviceAuth: null, }; manifest.contract.digest = await digestPublicContract( contractFromManifest(manifest), @@ -185,6 +193,34 @@ describe("contrail connect", () => { expect(fetcher).not.toHaveBeenCalled(); }); + it("locks discoverable service-auth methods separately", async () => { + const root = await temporaryRoot(); + const fixture = await serviceFixture([ + methodLexicon, + sourceLexicon, + notifyLexicon, + ]); + fixture.manifest.serviceAuth = { + type: "atproto-service-auth", + audience: "did:web:api.atmo.rsvp", + methods: [{ id: notifyMethod, type: "procedure" }], + }; + fixture.manifest.contract.digest = await digestPublicContract( + contractFromManifest(fixture.manifest), + ); + + const result = await connectPublicService({ + endpoint, + root, + out: "lexicons/pulled", + lock: "contrail.lock.json", + fetcher: fixture.fetcher, + }); + + expect(result.lock.methods).toEqual([method]); + expect(result.lock.serviceAuth).toEqual(fixture.manifest.serviceAuth); + }); + it("preserves the previous provider and lock when an update fails validation", async () => { const root = await temporaryRoot(); const fixture = await serviceFixture(); @@ -262,6 +298,33 @@ describe("contrail connect", () => { ).rejects.toThrow("matching query Lexicon"); }); + it("rejects protected methods without matching procedure or query Lexicons", async () => { + const root = await temporaryRoot(); + const fixture = await serviceFixture([ + methodLexicon, + sourceLexicon, + { ...notifyLexicon, defs: { main: { type: "query" } } }, + ]); + fixture.manifest.serviceAuth = { + type: "atproto-service-auth", + audience: "did:web:api.atmo.rsvp", + methods: [{ id: notifyMethod, type: "procedure" }], + }; + fixture.manifest.contract.digest = await digestPublicContract( + contractFromManifest(fixture.manifest), + ); + + await expect( + connectPublicService({ + endpoint, + root, + out: "lexicons/pulled", + lock: "contrail.lock.json", + fetcher: fixture.fetcher, + }), + ).rejects.toThrow("matching procedure Lexicon"); + }); + it("rejects inconsistent method namespaces and collection capabilities", async () => { const root = await temporaryRoot(); const outside = await serviceFixture(); diff --git a/packages/contrail/tests/service-auth.test.ts b/packages/contrail/tests/service-auth.test.ts new file mode 100644 index 0000000..8550f30 --- /dev/null +++ b/packages/contrail/tests/service-auth.test.ts @@ -0,0 +1,168 @@ +import { Secp256k1PrivateKeyExportable } from "@atcute/crypto"; +import type { Did, Nsid } from "@atcute/lexicons/syntax"; +import { createServiceJwt } from "@atcute/xrpc-server/auth"; +import { beforeAll, describe, expect, it } from "vitest"; +import { createSqliteDatabase } from "../src/adapters/sqlite"; +import { createApp } from "../src/core/router"; +import { + initSchema, + resolveConfig, + type ContrailConfig, + type Database, +} from "../src/index"; + +const issuer = "did:plc:aaaaaaaaaaaaaaaaaaaaaaaa" as Did; +const other = "did:plc:bbbbbbbbbbbbbbbbbbbbbbbb" as Did; +const audience = "did:web:api.example.com" as Did; +let keypair: Secp256k1PrivateKeyExportable; + +beforeAll(async () => { + keypair = await Secp256k1PrivateKeyExportable.createKeypair(); +}); + +function config(): ContrailConfig { + return { + namespace: "com.example", + profiles: [], + notify: true, + serviceAuth: { + audience, + methods: ["getFeed", "notifyOfUpdate"], + resolver: { + async resolve(did) { + return { + "@context": [], + id: did, + verificationMethod: [ + { + id: `${did}#atproto`, + type: "Multikey", + controller: did, + publicKeyMultibase: await keypair.exportPublicKey("multikey"), + }, + ], + }; + }, + }, + }, + collections: { + event: { collection: "community.example.event" }, + follow: { + collection: "app.bsky.graph.follow", + discover: false, + subjectField: "subject", + methods: [], + }, + }, + feeds: { network: { targets: ["event"] } }, + }; +} + +async function setup(): Promise<{ + db: Database; + app: ReturnType; +}> { + const resolved = resolveConfig(config()); + const db = createSqliteDatabase(":memory:"); + await initSchema(db, resolved); + for (const did of [issuer, other]) { + await db + .prepare( + "INSERT INTO identities (did, handle, pds, resolved_at) VALUES (?, NULL, ?, ?)", + ) + .bind(did, "https://pds.example.com", Date.now()) + .run(); + } + await db + .prepare( + "INSERT INTO feed_backfills (actor, feed, completed) VALUES (?, 'network', 1)", + ) + .bind(issuer) + .run(); + return { db, app: createApp(db, resolved) }; +} + +async function token(lxm: string, options: { aud?: Did; iss?: Did } = {}) { + return createServiceJwt({ + keypair, + issuer: options.iss ?? issuer, + audience: options.aud ?? audience, + lxm: lxm as Nsid, + }); +} + +function authorized(url: string, jwt: string, init: RequestInit = {}) { + const headers = new Headers(init.headers); + headers.set("authorization", `Bearer ${jwt}`); + return new Request(url, { ...init, headers }); +} + +describe("AT Protocol service auth", () => { + it("requires exact audience and method-bound tokens", async () => { + const { app } = await setup(); + const url = `https://api.example.com/xrpc/com.example.getFeed?feed=network&actor=${issuer}`; + + const missing = await app.fetch(new Request(url)); + expect(missing.status).toBe(401); + expect(missing.headers.get("www-authenticate")).toContain("Bearer"); + + const wrongMethod = await app.fetch( + authorized(url, await token("com.example.notifyOfUpdate")), + ); + expect(wrongMethod.status).toBe(401); + expect(wrongMethod.headers.get("www-authenticate")).toContain( + "BadJwtLexiconMethod", + ); + + const wrongAudience = await app.fetch( + authorized( + url, + await token("com.example.getFeed", { + aud: "did:web:other.example.com" as Did, + }), + ), + ); + expect(wrongAudience.status).toBe(401); + }); + + it("binds a personalized feed to the token issuer", async () => { + const { app } = await setup(); + const jwt = await token("com.example.getFeed"); + + const allowed = await app.fetch( + authorized( + `https://api.example.com/xrpc/com.example.getFeed?feed=network&actor=${issuer}`, + jwt, + ), + ); + expect(allowed.status).toBe(200); + + const forbidden = await app.fetch( + authorized( + `https://api.example.com/xrpc/com.example.getFeed?feed=network&actor=${other}`, + jwt, + ), + ); + expect(forbidden.status).toBe(403); + }); + + it("only lets an issuer notify its own record URIs", async () => { + const { app } = await setup(); + const jwt = await token("com.example.notifyOfUpdate"); + const response = await app.fetch( + authorized( + "https://api.example.com/xrpc/com.example.notifyOfUpdate", + jwt, + { + method: "POST", + headers: { "content-type": "application/json" }, + body: JSON.stringify({ + uri: `at://${other}/community.example.event/1`, + }), + }, + ), + ); + + expect(response.status).toBe(403); + }); +}); diff --git a/pnpm-lock.yaml b/pnpm-lock.yaml index 053154e..82c7c18 100644 --- a/pnpm-lock.yaml +++ b/pnpm-lock.yaml @@ -253,6 +253,9 @@ importers: '@atcute/tid': specifier: 1.1.4 version: 1.1.4 + '@atcute/xrpc-server': + specifier: ^2.0.2 + version: 2.0.2(@atcute/cid@2.4.2)(@atcute/lexicons@2.0.3)(typescript@6.0.3) cac: specifier: ^7.0.0 version: 7.0.0 @@ -266,6 +269,9 @@ importers: specifier: 1.4.2 version: 1.4.2(typescript@6.0.3) devDependencies: + '@atcute/crypto': + specifier: 2.4.4 + version: 2.4.4 '@cloudflare/workers-types': specifier: ^5.20260804.1 version: 5.20260804.1 @@ -428,6 +434,11 @@ packages: '@atcute/varint@2.0.2': resolution: {integrity: sha512-/+hS1juMgnmf6eL6lICUkTw7wcGTo3I+Q0L1PI521mUz77rGSC6nXAUNKtvm2wYJpuWdEGq+GILGoYkOArn0TQ==} + '@atcute/xrpc-server@2.0.2': + resolution: {integrity: sha512-tjOAqJfrBOwit6ccKg5s/eosvzoHmRszafS008/VP0yYdZkiBcSgYpy9S32e4aVoKeVurudBZjGD937bMcLg1Q==} + peerDependencies: + '@atcute/lexicons': ^2.0.0 + '@babel/runtime@7.29.7': resolution: {integrity: sha512-Nq8OhGWiZIZGV6hLHoyAKLLcJihP/xFeBMGJoUrxTX2psI8dCifzLhZISFb+VWS3wFMRDmCGw5R+dOySCqPLhw==} engines: {node: '>=6.9.0'} @@ -3062,6 +3073,11 @@ packages: engines: {node: ^18 || >=20} hasBin: true + nanoid@6.0.1: + resolution: {integrity: sha512-3wVS3i51pE2pi1k5FFL/95BGfVS0kSsvDVuGXHOtxox/TywUmtgq+3qiTOTbs9J7KfHaXPiN171k/A6dBnaXFw==} + engines: {node: ^22 || ^24 || >=26} + hasBin: true + natural-compare@1.4.0: resolution: {integrity: sha512-OWND8ei3VtNC9h7V60qff3SVobHr996CTwgxubgyQYEpg290h9J0buyECNNJexkFm5sOajh5G116RYA1c8ZMSw==} @@ -4181,6 +4197,21 @@ snapshots: '@atcute/varint@2.0.2': {} + '@atcute/xrpc-server@2.0.2(@atcute/cid@2.4.2)(@atcute/lexicons@2.0.3)(typescript@6.0.3)': + dependencies: + '@atcute/cbor': 2.3.6(@atcute/cid@2.4.2) + '@atcute/crypto': 2.4.4 + '@atcute/identity': 2.0.2(@atcute/lexicons@2.0.3)(typescript@6.0.3) + '@atcute/identity-resolver': 2.0.1(@atcute/identity@2.0.2(@atcute/lexicons@2.0.3)(typescript@6.0.3))(@atcute/lexicons@2.0.3)(typescript@6.0.3) + '@atcute/lexicons': 2.0.3 + '@atcute/multibase': 1.2.5 + '@atcute/uint8array': 1.1.5 + nanoid: 6.0.1 + valibot: 1.4.2(typescript@6.0.3) + transitivePeerDependencies: + - '@atcute/cid' + - typescript + '@babel/runtime@7.29.7': {} '@changesets/apply-release-plan@7.1.1': @@ -6468,6 +6499,8 @@ snapshots: nanoid@5.1.16: {} + nanoid@6.0.1: {} + natural-compare@1.4.0: {} number-flow@0.6.2: