From 92b0b019ba7fc4b70e689f2481c9e46f12aee492 Mon Sep 17 00:00:00 2001 From: Florian <45694132+flo-bit@users.noreply.github.com> Date: Fri, 7 Aug 2026 23:46:54 +0200 Subject: [PATCH] Simplify public service clients --- packages/contrail/package.json | 6 +- packages/contrail/src/cli/commands/connect.ts | 43 ++++ packages/contrail/src/public-client.ts | 207 ++++++++++++++++++ packages/contrail/tests/built-client.mjs | 11 + packages/contrail/tests/connect.test.ts | 28 +++ packages/contrail/tests/public-client.test.ts | 127 +++++++++++ .../contrail/tests/public-service-e2e.test.ts | 21 +- packages/contrail/tsup.config.ts | 1 + 8 files changed, 431 insertions(+), 13 deletions(-) create mode 100644 packages/contrail/src/public-client.ts create mode 100644 packages/contrail/tests/built-client.mjs create mode 100644 packages/contrail/tests/public-client.test.ts diff --git a/packages/contrail/package.json b/packages/contrail/package.json index f754305..23e19c9 100644 --- a/packages/contrail/package.json +++ b/packages/contrail/package.json @@ -19,6 +19,10 @@ "types": "./dist/server.d.ts", "import": "./dist/server.js" }, + "./client": { + "types": "./dist/public-client.d.ts", + "import": "./dist/public-client.js" + }, "./sqlite": { "types": "./dist/adapters/sqlite.d.ts", "import": "./dist/adapters/sqlite.js" @@ -65,7 +69,7 @@ "clean": "rm -rf dist", "typecheck": "tsc --noEmit", "test": "vitest run", - "test:built": "node tests/built-sqlite.mjs && node tests/built-lexicons.mjs", + "test:built": "node tests/built-sqlite.mjs && node tests/built-lexicons.mjs && node tests/built-client.mjs", "test:watch": "vitest" }, "dependencies": { diff --git a/packages/contrail/src/cli/commands/connect.ts b/packages/contrail/src/cli/commands/connect.ts index e973751..7a6e0da 100644 --- a/packages/contrail/src/cli/commands/connect.ts +++ b/packages/contrail/src/cli/commands/connect.ts @@ -32,6 +32,15 @@ import { generateLexiconTypesWithAtcute } from "../atcute.js"; const MAX_DISCOVERY_BYTES = 10 * 1024 * 1024; const REQUEST_TIMEOUT_MS = 15_000; +const LEXICON_CONFIG_NAMES = [ + "lex.config.js", + "lex.config.mjs", + "lex.config.cjs", + "lex.config.ts", + "lex.config.mts", + "lex.config.cts", +]; + interface ConnectOptions { root: string; out: string; @@ -149,6 +158,33 @@ async function exists(path: string): Promise { } } +/** Create a dependency-free Atcute config for ordinary consumer projects. + * Existing JavaScript or TypeScript configs remain entirely consumer-owned. */ +export async function ensureConsumerLexiconConfig(options: { + root: string; + out: string; +}): Promise<{ path: string; created: boolean }> { + const root = resolve(options.root); + for (const name of LEXICON_CONFIG_NAMES) { + const path = join(root, name); + if (await exists(path)) return { path, created: false }; + } + + const lexiconRoot = resolveInsideRoot(root, options.out); + const patternRoot = relative(root, lexiconRoot).replaceAll("\\", "/"); + const path = join(root, "lex.config.js"); + const source = `// Generated by \`contrail connect\`. Customize as needed.\nexport default {\n generate: {\n files: [${JSON.stringify(`${patternRoot}/**/*.json`)}],\n outdir: "src/lexicons/",\n },\n};\n`; + try { + await writeFile(path, source, { flag: "wx" }); + return { path, created: true }; + } catch (error) { + if ((error as NodeJS.ErrnoException).code === "EEXIST") { + return { path, created: false }; + } + throw error; + } +} + export async function connectPublicService(options: { endpoint: string; root: string; @@ -335,6 +371,13 @@ export function registerConnect(cli: CAC): void { `connected ${result.lock.endpoint}: ${result.written} Lexicons, contract ${result.lock.contractDigest}`, ); if (options.generate !== false) { + const config = await ensureConsumerLexiconConfig({ + root: options.root, + out: options.out, + }); + if (config.created) { + console.log(`created ${relative(resolve(options.root), config.path)}`); + } generateLexiconTypesWithAtcute(resolve(options.root)); } }); diff --git a/packages/contrail/src/public-client.ts b/packages/contrail/src/public-client.ts new file mode 100644 index 0000000..8979cc3 --- /dev/null +++ b/packages/contrail/src/public-client.ts @@ -0,0 +1,207 @@ +import type {} from "@atcute/atproto"; +import { + Client, + simpleFetchHandler, + type FetchHandler, +} from "@atcute/client"; +import type { Did, Nsid } from "@atcute/lexicons/syntax"; +import { + isPublicServiceManifest, + normalizePublicServiceEndpoint, + type PublicServiceAuthContract, +} from "./public-service.js"; + +const DISCOVERY_TIMEOUT_MS = 15_000; +const TOKEN_EXPIRY_SKEW_MS = 5_000; + +interface CachedToken { + value: string; + expiresAt: number; +} + +export interface PublicServiceClientOptions { + /** Canonical public Contrail HTTPS origin. */ + endpoint: string; + /** Existing authenticated PDS client used to mint service tokens. Omit when + * the consumer only needs anonymous methods. */ + authenticatedPds?: Client; + /** Optional contract pin from `contrail.lock.json`. */ + contractDigest?: string; + /** Browser, test, or instrumented fetch implementation. */ + fetch?: typeof globalThis.fetch; +} + +function xrpcMethod(pathname: string): Nsid | null { + const path = pathname.startsWith("http") + ? new URL(pathname).pathname + : new URL(pathname, "https://contrail.invalid").pathname; + const prefix = "/xrpc/"; + if (!path.startsWith(prefix)) return null; + try { + const method = decodeURIComponent(path.slice(prefix.length)); + return method.includes("/") ? null : (method as Nsid); + } catch { + return null; + } +} + +function tokenExpiration(token: string): number { + try { + const part = token.split(".")[1]; + if (!part) return Date.now() + 30_000; + const base64 = part.replace(/-/g, "+").replace(/_/g, "/"); + const padded = base64.padEnd(Math.ceil(base64.length / 4) * 4, "="); + const bytes = Uint8Array.from(atob(padded), (character) => + character.charCodeAt(0), + ); + const payload = JSON.parse(new TextDecoder().decode(bytes)) as { + exp?: unknown; + }; + return typeof payload.exp === "number" && Number.isSafeInteger(payload.exp) + ? payload.exp * 1_000 + : Date.now() + 30_000; + } catch { + return Date.now() + 30_000; + } +} + +function withBearer(init: RequestInit, token: string): RequestInit { + const headers = new Headers(init.headers); + headers.set("authorization", `Bearer ${token}`); + return { ...init, headers }; +} + +/** Fetch handler that keeps anonymous reads cheap while automatically minting, + * caching, and attaching method-bound AT Protocol service tokens after a + * protected route challenges the first request. */ +export function publicServiceFetchHandler( + options: PublicServiceClientOptions, +): FetchHandler { + const endpoint = normalizePublicServiceEndpoint(options.endpoint); + const fetcher = options.fetch ?? fetch; + const base = simpleFetchHandler({ service: endpoint, fetch: fetcher }); + const tokens = new Map(); + const pendingTokens = new Map>(); + let serviceAuthPromise: Promise | null = null; + + const discoverServiceAuth = () => { + if (serviceAuthPromise) return serviceAuthPromise; + serviceAuthPromise = (async () => { + const response = await fetcher(`${endpoint}/.well-known/contrail`, { + signal: AbortSignal.timeout(DISCOVERY_TIMEOUT_MS), + }); + if (!response.ok) { + throw new Error(`Contrail discovery failed: ${response.status}`); + } + if (response.url && new URL(response.url).origin !== endpoint) { + throw new Error("Contrail discovery redirected to a different origin"); + } + const value: unknown = await response.json(); + if (!isPublicServiceManifest(value)) { + throw new Error("response is not a supported Contrail service manifest"); + } + if (normalizePublicServiceEndpoint(value.endpoint) !== endpoint) { + throw new Error("Contrail manifest endpoint mismatch"); + } + if ( + options.contractDigest && + value.contract.digest !== options.contractDigest + ) { + throw new Error( + `Contrail contract digest mismatch: expected ${options.contractDigest}, received ${value.contract.digest}`, + ); + } + return value.serviceAuth ?? null; + })(); + return serviceAuthPromise; + }; + + const protectedMethod = async (method: string) => { + const auth = await discoverServiceAuth(); + return auth?.methods.some((candidate) => candidate.id === method) + ? auth + : null; + }; + + const tokenFor = async ( + method: Nsid, + auth: PublicServiceAuthContract, + force = false, + ): Promise => { + if (!options.authenticatedPds) { + throw new Error( + `Contrail method ${method} requires an authenticated PDS client`, + ); + } + const cached = tokens.get(method); + if (!force && cached && cached.expiresAt > Date.now() + TOKEN_EXPIRY_SKEW_MS) { + return cached.value; + } + if (!force) { + const pending = pendingTokens.get(method); + if (pending) return pending; + } + + const pending = (async () => { + const response = await options.authenticatedPds!.get( + "com.atproto.server.getServiceAuth", + { + params: { + aud: auth.audience as Did, + lxm: method, + }, + }, + ); + if (!response.ok) { + throw new Error( + `Could not mint service token for ${method}: ${response.status}`, + ); + } + const token = response.data.token; + tokens.set(method, { value: token, expiresAt: tokenExpiration(token) }); + return token; + })(); + pendingTokens.set(method, pending); + try { + return await pending; + } finally { + if (pendingTokens.get(method) === pending) pendingTokens.delete(method); + } + }; + + return async (pathname, init) => { + const method = xrpcMethod(pathname); + if (!method || !options.authenticatedPds) return base(pathname, init); + + // Once discovery has been loaded, avoid the initial challenge on subsequent + // protected calls. Anonymous calls never wait for discovery. + if (serviceAuthPromise) { + const auth = await protectedMethod(method); + if (auth) { + const token = await tokenFor(method, auth); + const response = await base(pathname, withBearer(init, token)); + if (response.status !== 401) return response; + await response.body?.cancel(); + tokens.delete(method); + const refreshed = await tokenFor(method, auth, true); + return base(pathname, withBearer(init, refreshed)); + } + } + + const response = await base(pathname, init); + if (response.status !== 401) return response; + const auth = await protectedMethod(method); + if (!auth) return response; + await response.body?.cancel(); + const token = await tokenFor(method, auth); + return base(pathname, withBearer(init, token)); + }; +} + +/** Create a typed Atcute client for anonymous and service-auth Contrail methods. + * Generated Lexicon imports still supply the method-specific TypeScript API. */ +export function createPublicServiceClient( + options: PublicServiceClientOptions, +): Client { + return new Client({ handler: publicServiceFetchHandler(options) }); +} diff --git a/packages/contrail/tests/built-client.mjs b/packages/contrail/tests/built-client.mjs new file mode 100644 index 0000000..1bc59d4 --- /dev/null +++ b/packages/contrail/tests/built-client.mjs @@ -0,0 +1,11 @@ +import assert from "node:assert/strict"; +import { createPublicServiceClient } from "../dist/public-client.js"; + +const client = createPublicServiceClient({ + endpoint: "https://api.example.com", + fetch: async () => Response.json({ records: [] }), +}); +const response = await client.get("com.example.listRecords"); +assert.equal(response.ok, true); +assert.deepEqual(response.data, { records: [] }); +console.log("built public client passed"); diff --git a/packages/contrail/tests/connect.test.ts b/packages/contrail/tests/connect.test.ts index 3e8e765..ee73699 100644 --- a/packages/contrail/tests/connect.test.ts +++ b/packages/contrail/tests/connect.test.ts @@ -4,6 +4,7 @@ import { join } from "node:path"; import { afterEach, describe, expect, it, vi } from "vitest"; import { connectPublicService, + ensureConsumerLexiconConfig, type ProviderLock, } from "../src/cli/commands/connect"; import { @@ -98,6 +99,33 @@ async function temporaryRoot() { } describe("contrail connect", () => { + it("creates a default Atcute config without replacing consumer config", async () => { + const root = await temporaryRoot(); + const generated = await ensureConsumerLexiconConfig({ + root, + out: "lexicons/providers", + }); + expect(generated.created).toBe(true); + expect(await readFile(generated.path, "utf8")).toContain( + 'files: ["lexicons/providers/**/*.json"]', + ); + expect(await readFile(generated.path, "utf8")).toContain( + 'outdir: "src/lexicons/"', + ); + + await writeFile(join(root, "lex.config.ts"), "export default { mine: true }"); + await rm(generated.path); + const existing = await ensureConsumerLexiconConfig({ + root, + out: "lexicons/other", + }); + expect(existing.created).toBe(false); + expect(existing.path).toBe(join(root, "lex.config.ts")); + expect(await readFile(existing.path, "utf8")).toBe( + "export default { mine: true }", + ); + }); + it("verifies and atomically locks a discovered service", async () => { const root = await temporaryRoot(); const { fetcher, manifest, values } = await serviceFixture(); diff --git a/packages/contrail/tests/public-client.test.ts b/packages/contrail/tests/public-client.test.ts new file mode 100644 index 0000000..97187c5 --- /dev/null +++ b/packages/contrail/tests/public-client.test.ts @@ -0,0 +1,127 @@ +import type {} from "@atcute/atproto"; +import { Client } from "@atcute/client"; +import { describe, expect, it, vi } from "vitest"; +import { + createPublicServiceClient, + publicServiceFetchHandler, +} from "../src/public-client"; +import type { PublicServiceManifest } from "../src/public-service"; + +const endpoint = "https://api.example.com"; +const method = "com.example.getFeed"; +const digest = `sha256:${"a".repeat(64)}`; + +function token() { + const encode = (value: unknown) => + btoa(JSON.stringify(value)) + .replaceAll("+", "-") + .replaceAll("/", "_") + .replaceAll("=", ""); + return `${encode({ alg: "none" })}.${encode({ exp: Math.floor(Date.now() / 1000) + 60 })}.signature`; +} + +function manifest(): PublicServiceManifest { + return { + format: "contrail.service", + version: 1, + endpoint, + namespace: "com.example", + contract: { digest }, + lexicons: { url: `${endpoint}/lexicons/${digest}`, digest }, + status: { url: `${endpoint}/status` }, + collections: [], + methods: ["com.example.getCursor"], + serviceAuth: { + type: "atproto-service-auth", + audience: "did:web:api.example.com", + methods: [{ id: method, type: "query" }], + }, + }; +} + +function authenticatedPds(jwt: string) { + const handler = vi.fn(async (pathname: string) => { + const url = new URL(pathname, "https://pds.example.com"); + expect(url.pathname).toBe("/xrpc/com.atproto.server.getServiceAuth"); + expect(url.searchParams.get("aud")).toBe("did:web:api.example.com"); + expect(url.searchParams.get("lxm")).toBe(method); + return Response.json({ token: jwt }); + }); + return { client: new Client({ handler }), handler }; +} + +describe("public service client", () => { + it("keeps anonymous requests anonymous", async () => { + const fetcher = vi.fn(async () => Response.json({ records: [] })); + const handler = publicServiceFetchHandler({ endpoint, fetch: fetcher }); + + const response = await handler("/xrpc/com.example.listRecords", { + method: "get", + }); + + expect(response.status).toBe(200); + expect(fetcher).toHaveBeenCalledTimes(1); + }); + + it("discovers, mints, caches, and attaches method-bound tokens", async () => { + const jwt = token(); + const pds = authenticatedPds(jwt); + const requests: Array<{ url: string; authorization: string | null }> = []; + const fetcher = vi.fn(async (input: RequestInfo | URL, init?: RequestInit) => { + const url = String(input); + const authorization = new Headers(init?.headers).get("authorization"); + requests.push({ url, authorization }); + if (url.endsWith("/.well-known/contrail")) { + return Response.json(manifest()); + } + return authorization === `Bearer ${jwt}` + ? Response.json({ records: [] }) + : Response.json( + { error: "AuthenticationRequired" }, + { status: 401, headers: { "www-authenticate": "Bearer" } }, + ); + }); + const client = createPublicServiceClient({ + endpoint, + authenticatedPds: pds.client, + contractDigest: digest, + fetch: fetcher, + }); + + const first = await (client as any).get(method, { + params: { actor: "did:plc:test", feed: "network" }, + }); + const second = await (client as any).get(method, { + params: { actor: "did:plc:test", feed: "network" }, + }); + + expect(first.ok).toBe(true); + expect(second.ok).toBe(true); + expect(pds.handler).toHaveBeenCalledTimes(1); + expect( + requests.filter((request) => + request.url.includes(`/xrpc/${method}`), + ).map((request) => request.authorization), + ).toEqual([null, `Bearer ${jwt}`, `Bearer ${jwt}`]); + }); + + it("refuses runtime discovery that differs from an optional lock pin", async () => { + const pds = authenticatedPds(token()); + const fetcher = vi.fn(async (input: RequestInfo | URL) => + String(input).endsWith("/.well-known/contrail") + ? Response.json(manifest()) + : new Response(null, { status: 401 }), + ); + const handler = publicServiceFetchHandler({ + endpoint, + authenticatedPds: pds.client, + contractDigest: `sha256:${"b".repeat(64)}`, + fetch: fetcher, + }); + + await expect( + handler(`/xrpc/${method}`, { method: "get" }), + ).rejects.toThrow("contract digest mismatch"); + expect(pds.handler).not.toHaveBeenCalled(); + }); +}); diff --git a/packages/contrail/tests/public-service-e2e.test.ts b/packages/contrail/tests/public-service-e2e.test.ts index 84200f2..881be2d 100644 --- a/packages/contrail/tests/public-service-e2e.test.ts +++ b/packages/contrail/tests/public-service-e2e.test.ts @@ -2,7 +2,10 @@ import { spawnSync } from "node:child_process"; import { mkdirSync, mkdtempSync, rmSync, writeFileSync } from "node:fs"; import { join } from "node:path"; import { afterEach, describe, expect, it } from "vitest"; -import { connectPublicService } from "../src/cli/commands/connect"; +import { + connectPublicService, + ensureConsumerLexiconConfig, +} from "../src/cli/commands/connect"; import { generateLexiconTypesWithAtcute } from "../src/cli/atcute"; import { createSqliteDatabase } from "../src/adapters/sqlite"; import { createApp } from "../src/core/router"; @@ -101,17 +104,6 @@ describe("public service consumer integration", () => { const fetcher: typeof fetch = (input, init) => app.fetch(new Request(input, init)); - writeFileSync( - join(consumerRoot, "lex.config.js"), - `import { defineLexiconConfig } from "@atcute/lex-cli"; -export default defineLexiconConfig({ - generate: { - files: ["lexicons/pulled/**/*.json"], - outdir: "src/lexicons/", - }, -}); -`, - ); await connectPublicService({ endpoint: "https://api.example.com", root: consumerRoot, @@ -119,6 +111,11 @@ export default defineLexiconConfig({ lock: "contrail.lock.json", fetcher, }); + const generatedConfig = await ensureConsumerLexiconConfig({ + root: consumerRoot, + out: "lexicons/pulled", + }); + expect(generatedConfig.created).toBe(true); generateLexiconTypesWithAtcute(consumerRoot); mkdirSync(join(consumerRoot, "src"), { recursive: true }); diff --git a/packages/contrail/tsup.config.ts b/packages/contrail/tsup.config.ts index 813e522..f6632ba 100644 --- a/packages/contrail/tsup.config.ts +++ b/packages/contrail/tsup.config.ts @@ -4,6 +4,7 @@ export default defineConfig({ entry: [ "src/index.ts", "src/server.ts", + "src/public-client.ts", "src/adapters/sqlite.ts", "src/adapters/postgres.ts", "src/workers/backfill.ts", -- 2.51.2