diff --git a/Dockerfile.dockerignore b/Dockerfile.dockerignore new file mode 100644 index 000000000..a60e23cb8 --- /dev/null +++ b/Dockerfile.dockerignore @@ -0,0 +1,79 @@ +# Pi-artifacts build context (this file shadows `.dockerignore` only for the +# pi-root `Dockerfile`). Robomp builds with `dockerfile: python/robomp/Dockerfile` +# still fall back to the shared `.dockerignore` next door because they don't +# have their own ignore file. +# +# Keep this file in sync with `.dockerignore` for the shared rules; everything +# below the divider is the artifacts-only addendum. + +# ─── Shared with .dockerignore ──────────────────────────────────────────────── + +# Heavy build outputs — must never reach the build context. `target/` alone is +# >100 GB on a dev machine. +target/ +node_modules/ +dist/ +runs/ + +# Per-host scratch the pi codebase uses for parallel agents / worktrees. +.fallow/ +.worktrees/ +.wt/ +.opencode/ +.pi_config/ +.omp/plugins/ + +# VCS, editors, IDEs — irrelevant to the build, churn on every IDE keystroke. +.git/ +.npm/ +.vscode/ +.zed/ +.idea/ + +# OS + transient noise. +.DS_Store +*.swp +*.swo +*~ +*.tmp + +# Logs + profiling artifacts. +*.log +*.cpuprofile +*.heapprofile +*.heapsnapshot +CPU.* + +# Build / test side outputs. +*.tsbuildinfo +coverage/ +.nyc_output/ +__pycache__/ +compaction-results/ +changes/ + +# Generated files (the in-image build regenerates them). +packages/coding-agent/src/internal-urls/docs-index.generated.ts +packages/natives/native/.build/ +packages/natives/native/pi_natives.darwin-*.node +packages/natives/native/pi_natives.dev.node +packages/ai/test/.temp-images/ +python/omp-rpc/src/omp_rpc.egg-info/ + +# Scratch files the repo creates ad-hoc. +syntax.jsonl +out.jsonl +out.html +pi-*.html + +# Secrets. Should never be in the image regardless. +.env + +# ─── Pi-artifacts only ──────────────────────────────────────────────────────── +# Robomp's source tree is unused by the artifacts image — pi-natives + omp-rpc +# are the only outputs, and `python/omp-rpc/` is reached explicitly by the +# python-builder stage (`COPY python/omp-rpc /src`). Everything under +# `python/robomp/` (orchestrator source, web bundle, tests, container scripts) +# would otherwise be transferred as part of the `COPY . /pi/` layer and bake +# uselessly into the natives-builder cache. +python/robomp/ diff --git a/packages/ai/CHANGELOG.md b/packages/ai/CHANGELOG.md index 37c833289..5ea89e4f3 100644 --- a/packages/ai/CHANGELOG.md +++ b/packages/ai/CHANGELOG.md @@ -1,6 +1,7 @@ # Changelog ## [Unreleased] + ### Breaking Changes - Renamed public schema utilities in `@oh-my-pi/pi-ai/utils/schema` by replacing `sanitizeSchemaForGoogle`, `sanitizeSchemaForCCA`, `prepareSchemaForCCA`, and `sanitizeSchemaForMCP` with `normalizeSchemaForGoogle`, `normalizeSchemaForCCA`, and `normalizeSchemaForMCP` @@ -10,18 +11,33 @@ ### Added +- Added `AuthStorage.setConfigApiKey` / `removeConfigApiKey` / `clearConfigApiKeys` for config-sourced per-provider bearers (e.g. `models.yml` `providers..apiKey`). The new tier sits between runtime `--api-key` and stored credentials in `getApiKey`/`peekApiKey` resolution, so a bearer pinned in config now beats the broker's OAuth access token. Also suppresses OAuth `account_uuid` attribution when active, since outbound auth is the explicit config bearer, not OAuth. `describeCredentialSource` reports `"config override (models.yml)"` for visibility. +- Added per-model `additional_rate_limits` parsing to `openaiCodexUsageProvider`. The Codex `wham/usage` endpoint surfaces a separate `GPT-5.3-Codex-Spark` rate limit (`metered_feature: codex_bengalfox`) on Pro accounts; these now emit dedicated `openai-codex:spark:{primary,secondary}` `UsageLimit` entries with `scope.tier = "spark"`, mirroring how Anthropic exposes `anthropic:7d:sonnet` separately from the umbrella `anthropic:7d` bucket. The osx-widgets client already keyed spark detection off `limit.id.includes("spark")`; this populates that contract end-to-end. +- Added `GET /v1/usage` to the auth-broker API to expose aggregated usage reports from `AuthStorage.fetchUsageReports` +- Added auth-broker usage polling response handling that returns normalized usage reports plus generation timestamp for clients - Added the auth-broker subsystem (`@oh-my-pi/pi-ai/auth-broker`) for sharing OAuth credentials across machines without leaking refresh tokens. - - `startAuthBroker(...)` boots a `Bun.serve` HTTP server exposing `GET /v1/healthz`, `GET /v1/snapshot`, `POST /v1/credential` (upsert), `POST /v1/credential/:id/refresh`, and `POST /v1/credential/:id/disable`. - - `AuthBrokerClient` is the matching HTTP client used by remote clients. - - `RemoteAuthCredentialStore` is a client-side `AuthCredentialStore` that mirrors a broker snapshot in memory; mutating methods (`replace*`, `upsert*`, `delete*ForProvider`) throw because writes are server-side only. - - `AuthBrokerRefresher` is the background refresh loop that pre-refreshes credentials within `refreshSkewMs` and disables on definitive failure (`invalid_grant` / non-network 401-403). +- `startAuthBroker(...)` boots a `Bun.serve` HTTP server exposing `GET /v1/healthz`, `GET /v1/snapshot`, `POST /v1/credential` (upsert), `POST /v1/credential/:id/refresh`, and `POST /v1/credential/:id/disable`. +- `AuthBrokerClient` is the matching HTTP client used by remote clients. +- `RemoteAuthCredentialStore` is a client-side `AuthCredentialStore` that mirrors a broker snapshot in memory; mutating methods (`replace*`, `upsert*`, `delete*ForProvider`) throw because writes are server-side only. +- `AuthBrokerRefresher` is the background refresh loop that pre-refreshes credentials within `refreshSkewMs` and disables on definitive failure (`invalid_grant` / non-network 401-403). - Added `AuthStorage.exportSnapshot()`, `AuthStorage.upsertCredential(provider, credential)`, `AuthStorage.forceRefreshCredentialById(id)`, and `AuthStorage.disableCredentialById(id, cause)` public methods consumed by the auth-broker server. - Added `AuthStorageOptions.refreshOAuthCredential` override so a remote-store client can route every OAuth refresh through the broker instead of the local OAuth endpoint. - Added `REMOTE_REFRESH_SENTINEL` (`"__remote__"`) — the wire placeholder substituted for OAuth refresh tokens in broker snapshots; clients never see the real refresh token. - Exposed the OAuth provider catalog (`getOAuthProviders`, `OAuthProvider`, `OAuthProviderInfo`) and `refreshOAuthToken` through the package barrel so the coding-agent CLI can target them without reaching into `utils/oauth`. +- Added the auth-gateway subsystem (`@oh-my-pi/pi-ai/auth-gateway`) — a forward-proxy that sits between unauthenticated clients (the macOS usage widget, llm-git, robomp containers, …) and the broker. Clients send standard provider-format requests; the gateway parses them into omp's canonical `Context`, dispatches through pi-ai's `streamSimple()`, and translates the canonical event stream back to the matching wire format. `Authorization` is injected server-side so access tokens never leave the gateway host. Wire surface: +- `GET /healthz` — unauth liveness. +- `GET /v1/usage` — aggregated provider usage; 30s cache via `AuthStorage.fetchUsageReports`. +- `GET /v1/models` — model catalog (scoped to providers with credentials). +- `POST /v1/chat/completions` — OpenAI chat-completions in/out. +- `POST /v1/messages` — Anthropic messages in/out (text + thinking + tool_use blocks, SSE event taxonomy preserved). +- `POST /v1/responses` — OpenAI Responses in/out (reasoning items + function_call output items, SSE pass-through). +- Added exports from `@oh-my-pi/pi-ai/auth-gateway`: `startAuthGateway`, `AuthGatewayServerOptions`, `AuthGatewayBootOptions`, `AuthGatewayServerHandle`, `ModelResolver`, `DEFAULT_AUTH_GATEWAY_BIND`, plus per-format `parseRequest` / `encodeResponse` / `encodeStream` triples under `auth-gateway/formats/{openai-chat,anthropic-messages,openai-responses}`. +- Added `listProvidersWithEnvKey()` to enumerate every provider with an env-var fallback (used by the new migrate command in coding-agent). ### Changed +- Changed usage report caching to use a 5-minute per-credential TTL with jittered refresh timing to reduce usage endpoint rate-limit collisions +- Changed usage polling failure handling so transient errors continue serving the last known report instead of returning null and dropping the credential from usage aggregates after cache expiry - Changed `sanitizeSchemaForGoogle` to normalize snake_case schema keys (such as `any_of` and `additional_properties`) to camelCase and auto-generate `propertyOrdering` for multi-property objects - Changed strict-mode sanitization to resolve `$ref` nodes with sibling keys by inlining and merging referenced local definitions - Changed strict-mode sanitization to flatten single-entry `allOf` nodes and remove the `allOf` wrapper @@ -30,6 +46,9 @@ ### Fixed +- Fixed Claude usage fetching to retry transient `429` and `5xx` responses with exponential backoff, respecting `Retry-After` before returning failure +- Fixed auth-gateway request translation to preserve OpenAI Responses string/system message content, reasoning replay payloads, completed item text in stream item-done events, Anthropic tool-result ordering, and OpenAI Chat/Responses cached-token usage totals +- Fixed auth-gateway failure handling so unsupported request controls, upstream terminal errors, non-streaming aborts, and already-aborted client requests fail explicitly instead of being accepted, ignored, or encoded as successful HTTP 200 responses - Fixed Gemini CLI / Antigravity tool schema normalization to run the full Cloud Code Assist pipeline, matching shared Google schema handling for union/object merging and nullable extraction - Fixed stripped validation hints to be preserved as description spill text (`{key: value}` blocks) when `normalizeSchemaForGoogle` and `normalizeSchemaForCCA` drop unsupported schema keywords - Fixed `sanitizeSchemaForGoogle` to collapse nullability forms (`type:'null'` and null-bearing `anyOf` variants) into `nullable` while preserving remaining variants @@ -37,6 +56,7 @@ - Fixed `normalizeAnthropicToolSchema` to handle self-referential schemas without infinite recursion - Fixed object schema normalization so explicit open-map declarations (`additionalProperties: true` and schema-valued `additionalProperties`) are preserved instead of being converted to closed objects - Fixed unsupported schema constraints on arrays and strings (`maxItems`, `uniqueItems`, `pattern`, `minLength`, `maxLength`, and `minItems` when greater than 1) by demoting them into `description` rather than dropping them + ## [15.1.2] - 2026-05-15 ### Breaking Changes diff --git a/packages/ai/package.json b/packages/ai/package.json index d8a178b5b..6905600b1 100644 --- a/packages/ai/package.json +++ b/packages/ai/package.json @@ -73,6 +73,22 @@ "types": "./src/*.ts", "import": "./src/*.ts" }, + "./auth-broker": { + "types": "./src/auth-broker/index.ts", + "import": "./src/auth-broker/index.ts" + }, + "./auth-broker/*": { + "types": "./src/auth-broker/*.ts", + "import": "./src/auth-broker/*.ts" + }, + "./auth-gateway": { + "types": "./src/auth-gateway/index.ts", + "import": "./src/auth-gateway/index.ts" + }, + "./auth-gateway/*": { + "types": "./src/auth-gateway/*.ts", + "import": "./src/auth-gateway/*.ts" + }, "./models.json": { "types": "./src/models.json.d.ts", "import": "./src/models.json" diff --git a/packages/ai/src/auth-broker/client.ts b/packages/ai/src/auth-broker/client.ts index 5ea6eba60..ed4e325ce 100644 --- a/packages/ai/src/auth-broker/client.ts +++ b/packages/ai/src/auth-broker/client.ts @@ -5,6 +5,7 @@ * `omp auth-broker status` (liveness checks). All endpoints except * `/v1/healthz` require a bearer token. */ +import type { ZodType, infer as zInfer } from "zod/v4"; import type { AuthCredential } from "../auth-storage"; import type { CredentialDisableRequest, @@ -14,7 +15,16 @@ import type { CredentialUploadResponse, HealthzResponse, SnapshotResponse, + UsageResponse, } from "./types"; +import { + credentialDisableResponseSchema, + credentialRefreshResponseSchema, + credentialUploadResponseSchema, + healthzResponseSchema, + snapshotResponseSchema, + usageResponseSchema, +} from "./wire-schemas"; export interface AuthBrokerClientOptions { /** Base URL (e.g. `https://broker.tailnet:8765`). Trailing slashes are trimmed. */ @@ -59,30 +69,48 @@ export class AuthBrokerClient { } healthz(): Promise { - return this.#request("GET", "/v1/healthz", { auth: false }); + return this.#request("GET", "/v1/healthz", { schema: healthzResponseSchema, auth: false }); } fetchSnapshot(): Promise { - return this.#request("GET", "/v1/snapshot"); + // `snapshotResponseSchema` narrows `refresh` to the sentinel literal where + // the public type uses plain `string`; the wire shape is identical. + return this.#request("GET", "/v1/snapshot", { schema: snapshotResponseSchema }) as Promise; + } + + fetchUsage(): Promise { + // `usageResponseSchema` keeps the report array as `unknown[]` — per-provider + // usage modules own the inner shape; the broker doesn't re-validate it. + return this.#request("GET", "/v1/usage", { schema: usageResponseSchema }) as Promise; } async refreshCredential(id: number): Promise { - return this.#request("POST", `/v1/credential/${id}/refresh`); + return this.#request("POST", `/v1/credential/${id}/refresh`, { + schema: credentialRefreshResponseSchema, + }) as Promise; } async disableCredential(id: number, cause: string): Promise { const body: CredentialDisableRequest = { cause }; - return this.#request("POST", `/v1/credential/${id}/disable`, { + return this.#request("POST", `/v1/credential/${id}/disable`, { body, + schema: credentialDisableResponseSchema, }); } async uploadCredential(provider: string, credential: AuthCredential): Promise { const body: CredentialUploadRequest = { provider, credential }; - return this.#request("POST", "/v1/credential", { body }); + return this.#request("POST", "/v1/credential", { + body, + schema: credentialUploadResponseSchema, + }) as Promise; } - async #request(method: "GET" | "POST", path: string, opts: { auth?: boolean; body?: unknown } = {}): Promise { + async #request( + method: "GET" | "POST", + path: string, + opts: { schema: TSchema; auth?: boolean; body?: unknown }, + ): Promise> { const auth = opts.auth ?? true; const url = `${this.#baseUrl}${path}`; const headers: Record = { Accept: "application/json" }; @@ -109,9 +137,9 @@ export class AuthBrokerClient { body: text, }); } - if (!text) return undefined as T; + let raw: unknown; try { - return JSON.parse(text) as T; + raw = text.length === 0 ? null : JSON.parse(text); } catch (parseError) { throw new AuthBrokerError("Auth broker returned malformed JSON", { status: response.status, @@ -119,6 +147,14 @@ export class AuthBrokerClient { cause: parseError, }); } + const validated = opts.schema.safeParse(raw); + if (!validated.success) { + throw new AuthBrokerError("Auth broker response failed schema validation", { + status: response.status, + body: validated.error.message, + }); + } + return validated.data; } catch (error) { lastError = error; if (error instanceof AuthBrokerError && error.status !== undefined) { diff --git a/packages/ai/src/auth-broker/remote-store.ts b/packages/ai/src/auth-broker/remote-store.ts index f4b3b46e6..6f7687928 100644 --- a/packages/ai/src/auth-broker/remote-store.ts +++ b/packages/ai/src/auth-broker/remote-store.ts @@ -7,12 +7,17 @@ * usage reports cache TTL is ~30s, so durability across runs isn't required. */ import { logger } from "@oh-my-pi/pi-utils"; -import type { - AuthCredential, - AuthCredentialSnapshot, - AuthCredentialStore, - StoredAuthCredential, +import { + type AuthCredential, + type AuthCredentialSnapshot, + type AuthCredentialStore, + type OAuthCredential, + REMOTE_REFRESH_SENTINEL, + type StoredAuthCredential, } from "../auth-storage"; +import type { Provider } from "../types"; +import type { UsageReport } from "../usage"; +import type { OAuthCredentials } from "../utils/oauth/types"; import type { AuthBrokerClient } from "./client"; interface CacheEntry { @@ -136,6 +141,44 @@ export class RemoteAuthCredentialStore implements AuthCredentialStore { } } + /** + * Store-level hook consumed by `AuthStorage` — routes refresh through the + * broker so the actual refresh token never leaves the broker host. Returns + * the broker-redacted credential with {@link REMOTE_REFRESH_SENTINEL} in + * the `refresh` slot. + */ + async refreshOAuthCredential( + _provider: Provider, + credentialId: number, + _credential: OAuthCredential, + ): Promise { + const { entry } = await this.#client.refreshCredential(credentialId); + if (entry.credential.type !== "oauth") { + throw new Error(`Broker returned non-OAuth credential for id=${credentialId}`); + } + const refreshed = entry.credential; + return { + access: refreshed.access, + refresh: REMOTE_REFRESH_SENTINEL, + expires: refreshed.expires, + accountId: refreshed.accountId, + email: refreshed.email, + projectId: refreshed.projectId, + enterpriseUrl: refreshed.enterpriseUrl, + }; + } + + /** + * Store-level hook consumed by `AuthStorage.fetchUsageReports()` — proxies + * to the broker's `/v1/usage` endpoint. The broker's egress IP isn't + * rate-limited by Anthropic's per-IP `/usage` cap the way a heavy + * residential laptop is, so all credentials surface every cycle. + */ + async fetchUsageReports(): Promise { + const body = await this.#client.fetchUsage(); + return body.reports; + } + close(): void { if (this.#closed) return; this.#closed = true; diff --git a/packages/ai/src/auth-broker/server.ts b/packages/ai/src/auth-broker/server.ts index cd2054438..379a9d20f 100644 --- a/packages/ai/src/auth-broker/server.ts +++ b/packages/ai/src/auth-broker/server.ts @@ -13,15 +13,14 @@ import { logger } from "@oh-my-pi/pi-utils"; import type { AuthStorage } from "../auth-storage"; import { AuthBrokerRefresher } from "./refresher"; import type { - CredentialDisableRequest, CredentialDisableResponse, CredentialRefreshResponse, - CredentialUploadRequest, CredentialUploadResponse, HealthzResponse, SnapshotResponse, } from "./types"; import { DEFAULT_AUTH_BROKER_BIND, DEFAULT_REFRESH_INTERVAL_MS, DEFAULT_REFRESH_SKEW_MS } from "./types"; +import { credentialDisableRequestSchema, credentialUploadRequestSchema } from "./wire-schemas"; export interface AuthBrokerServerOptions { /** Underlying credential storage (wraps the local SQLite store on the broker). */ @@ -53,10 +52,24 @@ interface ParsedBind { port: number; } +function parsePort(raw: string, bind: string): number { + if (!/^\d+$/.test(raw)) { + throw new Error(`Invalid bind '${bind}'; port must be an integer.`); + } + const port = Number.parseInt(raw, 10); + if (!Number.isFinite(port) || port < 0 || port > 65535) { + throw new Error(`Invalid bind '${bind}'; port out of range.`); + } + return port; +} + function parseBind(raw: string): ParsedBind { const trimmed = raw.trim(); + if (trimmed.length === 0) { + throw new Error("Invalid bind; expected 'host:port' or 'port'."); + } if (/^\d+$/.test(trimmed)) { - return { hostname: "127.0.0.1", port: Number.parseInt(trimmed, 10) }; + return { hostname: "127.0.0.1", port: parsePort(trimmed, raw) }; } const lastColon = trimmed.lastIndexOf(":"); if (lastColon < 0) { @@ -64,11 +77,10 @@ function parseBind(raw: string): ParsedBind { } const hostPart = trimmed.slice(0, lastColon); const portPart = trimmed.slice(lastColon + 1); - const port = Number.parseInt(portPart, 10); - if (!Number.isFinite(port) || port < 0 || port > 65535) { - throw new Error(`Invalid bind '${raw}'; port out of range.`); + if (hostPart.length === 0) { + throw new Error(`Invalid bind '${raw}'; host must not be empty.`); } - return { hostname: hostPart, port }; + return { hostname: hostPart, port: parsePort(portPart, raw) }; } function json(status: number, body: unknown): Response { @@ -87,6 +99,38 @@ function isAuthorized(req: Request, tokens: ReadonlySet): boolean { return tokens.has(match[1].trim()); } +/** + * Parse + validate a JSON request body against a Zod schema. Returns a + * `Response` (400) on parse/validation failure so handlers can early-return. + * When `allowEmpty` is set, an empty request body is validated against `{}`. + */ +async function parseBody( + req: Request, + schema: { safeParse(input: unknown): { success: true; data: T } | { success: false; error: { message: string } } }, + options: { allowEmpty?: boolean } = {}, +): Promise<{ ok: true; data: T } | { ok: false; response: Response }> { + let raw: string; + try { + raw = await req.text(); + } catch (error) { + return { ok: false, response: json(400, { error: `Invalid request body: ${String(error)}` }) }; + } + if (raw.length === 0 && !options.allowEmpty) { + return { ok: false, response: json(400, { error: "Request body required" }) }; + } + let parsed: unknown; + try { + parsed = raw.length === 0 ? {} : JSON.parse(raw); + } catch (error) { + return { ok: false, response: json(400, { error: `Invalid JSON body: ${String(error)}` }) }; + } + const result = schema.safeParse(parsed); + if (!result.success) { + return { ok: false, response: json(400, { error: result.error.message }) }; + } + return { ok: true, data: result.data }; +} + const REFRESH_ROUTE = /^\/v1\/credential\/(\d+)\/refresh$/; const DISABLE_ROUTE = /^\/v1\/credential\/(\d+)\/disable$/; @@ -128,6 +172,24 @@ export function startAuthBroker(opts: AuthBrokerServerOptions): AuthBrokerServer logger.info("auth-broker snapshot served", { peer, credentials: body.credentials.length }); return json(200, body); } + if (req.method === "GET" && pathname === "/v1/usage") { + try { + // AuthStorage caches usage reports internally with a 30s TTL + // (USAGE_REPORT_TTL_MS) so back-to-back widget polls re-use the + // last fetch instead of hitting provider endpoints repeatedly. + const reports = (await opts.storage.fetchUsageReports?.()) ?? []; + // Drop the `raw` field — it's the provider-specific upstream body, + // large and unstable. Everything UI-relevant lives in `limits` and + // `metadata`. + const trimmed = reports.map(({ raw: _raw, ...rest }) => rest); + logger.info("auth-broker usage served", { peer, reports: trimmed.length }); + return json(200, { generatedAt: Date.now(), reports: trimmed }); + } catch (error) { + const message = error instanceof Error ? error.message : String(error); + logger.warn("auth-broker usage fetch failed", { peer, error: message }); + return json(502, { error: message }); + } + } const refreshMatch = req.method === "POST" ? pathname.match(REFRESH_ROUTE) : null; if (refreshMatch) { const id = Number.parseInt(refreshMatch[1], 10); @@ -151,13 +213,10 @@ export function startAuthBroker(opts: AuthBrokerServerOptions): AuthBrokerServer const disableMatch = req.method === "POST" ? pathname.match(DISABLE_ROUTE) : null; if (disableMatch) { const id = Number.parseInt(disableMatch[1], 10); - let cause = "disabled via auth-broker"; - try { - const body = (await req.json()) as Partial; - if (typeof body?.cause === "string" && body.cause.length > 0) cause = body.cause; - } catch { - // Empty / malformed body — default cause already set. - } + const parsed = await parseBody(req, credentialDisableRequestSchema, { allowEmpty: true }); + if (!parsed.ok) return parsed.response; + const cause = + parsed.data.cause && parsed.data.cause.length > 0 ? parsed.data.cause : "disabled via auth-broker"; const ok = opts.storage.disableCredentialById(id, cause); if (!ok) { logger.info("auth-broker disable miss", { id, peer, cause }); @@ -168,32 +227,17 @@ export function startAuthBroker(opts: AuthBrokerServerOptions): AuthBrokerServer return json(200, response); } if (req.method === "POST" && pathname === "/v1/credential") { - let body: Partial; + const parsed = await parseBody(req, credentialUploadRequestSchema); + if (!parsed.ok) return parsed.response; + const { provider, credential } = parsed.data; try { - body = (await req.json()) as Partial; - } catch (error) { - return json(400, { error: `Invalid JSON body: ${String(error)}` }); - } - if (!body || typeof body.provider !== "string" || body.provider.length === 0) { - return json(400, { error: "Missing `provider` field" }); - } - if (!body.credential || typeof body.credential !== "object") { - return json(400, { error: "Missing `credential` field" }); - } - const credential = body.credential; - if (credential.type !== "oauth" && credential.type !== "api_key") { - return json(400, { - error: `Invalid credential.type: ${String((credential as { type?: unknown }).type)}`, - }); - } - try { - const entries = opts.storage.upsertCredential(body.provider, credential); + const entries = opts.storage.upsertCredential(provider, credential); const identity = credential.type === "oauth" ? (credential.email ?? credential.accountId ?? credential.projectId ?? "(no identity)") : "(api key)"; logger.info("auth-broker credential upserted", { - provider: body.provider, + provider, type: credential.type, identity, peer, @@ -203,7 +247,7 @@ export function startAuthBroker(opts: AuthBrokerServerOptions): AuthBrokerServer return json(200, response); } catch (error) { const message = error instanceof Error ? error.message : String(error); - logger.warn("auth-broker upload failed", { provider: body.provider, peer, error: message }); + logger.warn("auth-broker upload failed", { provider, peer, error: message }); return json(500, { error: message }); } } diff --git a/packages/ai/src/auth-broker/types.ts b/packages/ai/src/auth-broker/types.ts index 59a04efea..bcabd4c7c 100644 --- a/packages/ai/src/auth-broker/types.ts +++ b/packages/ai/src/auth-broker/types.ts @@ -7,6 +7,7 @@ */ import type { AuthCredential, AuthCredentialSnapshot, AuthCredentialSnapshotEntry } from "../auth-storage"; +import type { UsageReport } from "../usage"; /** GET /v1/healthz response body. */ export interface HealthzResponse { @@ -17,6 +18,12 @@ export interface HealthzResponse { /** GET /v1/snapshot response body. */ export type SnapshotResponse = AuthCredentialSnapshot; +/** GET /v1/usage response body — matches the local `AuthStorage.fetchUsageReports` shape. */ +export interface UsageResponse { + generatedAt: number; + reports: UsageReport[]; +} + /** POST /v1/credential/:id/refresh response body. */ export interface CredentialRefreshResponse { entry: AuthCredentialSnapshotEntry; diff --git a/packages/ai/src/auth-broker/wire-schemas.ts b/packages/ai/src/auth-broker/wire-schemas.ts new file mode 100644 index 000000000..8267304de --- /dev/null +++ b/packages/ai/src/auth-broker/wire-schemas.ts @@ -0,0 +1,134 @@ +/** + * Zod schemas for the auth-broker wire protocol. + * + * Shared between the server (validates inbound request bodies) and the client + * (validates responses from the broker). Schemas mirror the TypeScript types + * in `./types.ts` 1:1; the types remain the source of truth for static typing, + * and `z.infer` is asserted-compatible with them where possible. + * + * Schemas use `.strict()` on objects with a closed set of fields so unknown + * keys are rejected — the previous implementation used a hand-rolled + * `hasOnlyFields` allowlist for the same effect. + */ +import * as z from "zod/v4"; +import { REMOTE_REFRESH_SENTINEL } from "../auth-storage"; +import { usageReportSchema } from "../usage"; + +// ─── Credential payloads ─────────────────────────────────────────────────── + +/** Real OAuth credential (broker-side) — refresh token is the actual upstream value. */ +export const oauthCredentialSchema = z + .object({ + type: z.literal("oauth"), + refresh: z.string().min(1), + access: z.string().min(1), + expires: z.number(), + enterpriseUrl: z.string().optional(), + projectId: z.string().optional(), + email: z.string().optional(), + accountId: z.string().optional(), + }) + .strict(); + +/** OAuth credential as it appears in broker snapshots — refresh replaced with sentinel. */ +export const remoteOauthCredentialSchema = oauthCredentialSchema.extend({ + refresh: z.literal(REMOTE_REFRESH_SENTINEL), +}); + +export const apiKeyCredentialSchema = z + .object({ + type: z.literal("api_key"), + key: z.string().min(1), + }) + .strict(); + +/** Discriminated union accepted on POST /v1/credential (writes). */ +export const writableAuthCredentialSchema = z.discriminatedUnion("type", [ + oauthCredentialSchema, + apiKeyCredentialSchema, +]); + +/** Discriminated union returned in snapshots (refresh is sentinel for OAuth). */ +export const snapshotCredentialSchema = z.discriminatedUnion("type", [ + remoteOauthCredentialSchema, + apiKeyCredentialSchema, +]); + +// ─── Snapshot ────────────────────────────────────────────────────────────── + +export const snapshotEntrySchema = z + .object({ + id: z.number().int(), + provider: z.string().min(1), + credential: snapshotCredentialSchema, + identityKey: z.string().nullable(), + }) + .strict(); + +export const snapshotResponseSchema = z + .object({ + generatedAt: z.number(), + credentials: z.array(snapshotEntrySchema), + }) + .strict(); + +// ─── Healthz ──────────────────────────────────────────────────────────────── + +export const healthzResponseSchema = z + .object({ + ok: z.boolean(), + version: z.string().optional(), + }) + .strict(); + +// ─── Usage ───────────────────────────────────────────────────────────────── + +/** + * Broker `/v1/usage` response. Reports are full {@link UsageReport}s minus the + * heavy provider-specific `raw` field (the server strips it before send) — we + * keep `raw` optional in the underlying schema so a misconfigured broker that + * forgot to strip still validates. + */ +export const usageResponseSchema = z + .object({ + generatedAt: z.number(), + reports: z.array(usageReportSchema), + }) + .strict(); + +// ─── Refresh ─────────────────────────────────────────────────────────────── + +export const credentialRefreshResponseSchema = z + .object({ + entry: snapshotEntrySchema, + }) + .strict(); + +// ─── Disable ─────────────────────────────────────────────────────────────── + +export const credentialDisableRequestSchema = z + .object({ + cause: z.string().optional(), + }) + .strict(); + +export const credentialDisableResponseSchema = z + .object({ + ok: z.boolean(), + }) + .strict(); + +// ─── Upload ──────────────────────────────────────────────────────────────── + +export const credentialUploadRequestSchema = z + .object({ + provider: z.string().min(1), + credential: writableAuthCredentialSchema, + }) + .strict(); + +export const credentialUploadResponseSchema = z + .object({ + entries: z.array(snapshotEntrySchema), + }) + .strict(); diff --git a/packages/ai/src/auth-gateway/http.ts b/packages/ai/src/auth-gateway/http.ts new file mode 100644 index 000000000..01c6fe237 --- /dev/null +++ b/packages/ai/src/auth-gateway/http.ts @@ -0,0 +1,32 @@ +/** + * Shared HTTP helpers for the auth-gateway routes. + * + * Centralized so we share the same JSON shape, auth check, + * and peer-resolution logic. + */ + +const JSON_HEADERS = { + "Content-Type": "application/json", + "X-Content-Type-Options": "nosniff", +} as const; +export function json(status: number, body: unknown): Response { + return new Response(JSON.stringify(body) ?? "null", { + status, + headers: JSON_HEADERS, + }); +} + +export function resolvePeer(req: Request): string { + const fwd = req.headers.get("x-forwarded-for"); + if (fwd) return fwd.split(",")[0].trim(); + return req.headers.get("x-real-ip") ?? "unknown"; +} + +export function isAuthorized(req: Request, tokens: ReadonlySet): boolean { + if (tokens.size === 0) return true; + const header = req.headers.get("authorization"); + if (!header) return false; + const match = header.match(/^Bearer\s+(.+)$/i); + if (!match) return false; + return tokens.has(match[1].trim()); +} diff --git a/packages/ai/src/auth-gateway/index.ts b/packages/ai/src/auth-gateway/index.ts new file mode 100644 index 000000000..e16648ed0 --- /dev/null +++ b/packages/ai/src/auth-gateway/index.ts @@ -0,0 +1,3 @@ +export * from "./http"; +export * from "./server"; +export * from "./types"; diff --git a/packages/ai/src/auth-gateway/server.ts b/packages/ai/src/auth-gateway/server.ts new file mode 100644 index 000000000..341d6254b --- /dev/null +++ b/packages/ai/src/auth-gateway/server.ts @@ -0,0 +1,503 @@ +/** + * omp auth-gateway HTTP server. + * + * Accepts any provider-format request (OpenAI chat-completions, Anthropic + * messages, OpenAI Responses) and dispatches through pi-ai's `streamSimple()` + * — which handles credential injection, anthropic-beta headers, codex + * websocket transport, and all the per-provider intricacies. The gateway is + * pure protocol translation: foreign wire → omp Context → pi-ai stream() → + * omp events → foreign wire. + * + * Endpoints: + * GET /healthz → unauth; ok + version + * GET /v1/usage → aggregated provider usage (30s cache via AuthStorage) + * GET /v1/models → list known models from the registry + * POST /v1/chat/completions → OpenAI chat-completions in/out + * POST /v1/messages → Anthropic messages in/out + * POST /v1/responses → OpenAI Responses in/out + */ +import { logger } from "@oh-my-pi/pi-utils"; +import type { AuthStorage } from "../auth-storage"; +import { Effort } from "../model-thinking"; +import * as anthropicMessages from "../providers/anthropic-messages-server"; +import * as openaiChat from "../providers/openai-chat-server"; +import * as openaiResponses from "../providers/openai-responses-server"; +import { streamSimple } from "../stream"; +import type { Api, AssistantMessageEventStream, Model, SimpleStreamOptions } from "../types"; +import { isAuthorized, json, resolvePeer } from "./http"; +import type { + AuthGatewayServerHandle, + AuthGatewayServerOptions, + AuthGatewayFormatModule as FormatModule, + AuthGatewayParsedRequest as ParsedFormatRequest, +} from "./types"; +import { DEFAULT_AUTH_GATEWAY_BIND } from "./types"; + +// ParsedFormatRequest / ParsedFormatOptions / FormatModule come from ./types. + +export type ModelResolver = (modelId: string) => Model | undefined; + +export interface AuthGatewayBootOptions extends AuthGatewayServerOptions { + /** Source of credentials. Caller wires this to a broker-backed AuthStorage. */ + storage: AuthStorage; + /** + * Resolve a client-requested model id to a pi-ai Model. Caller supplies + * this from a ModelRegistry (lives in `coding-agent` to avoid an inverse + * dependency in `pi-ai`). + */ + resolveModel: ModelResolver; + /** Optional supplier for `/v1/models` listing. Returns the full model array. */ + listModels?: () => Iterable>; +} + +interface ParsedBind { + hostname: string; + port: number; +} + +function parseBind(raw: string): ParsedBind { + const trimmed = raw.trim(); + if (/^\d+$/.test(trimmed)) { + return { hostname: "127.0.0.1", port: Number.parseInt(trimmed, 10) }; + } + const lastColon = trimmed.lastIndexOf(":"); + if (lastColon < 0) throw new Error(`Invalid bind '${raw}'; expected 'host:port' or 'port'.`); + const port = Number.parseInt(trimmed.slice(lastColon + 1), 10); + if (!Number.isFinite(port) || port < 0 || port > 65535) { + throw new Error(`Invalid bind '${raw}'; port out of range.`); + } + return { hostname: trimmed.slice(0, lastColon), port }; +} + +const FORMAT_ROUTES: Record = { + "/v1/chat/completions": { module: openaiChat, label: "openai-chat" }, + "/v1/messages": { module: anthropicMessages, label: "anthropic-messages" }, + "/v1/responses": { module: openaiResponses, label: "openai-responses" }, +}; + +/** + * Wire path on the upstream provider that each inbound format maps to when + * passthrough is taken. Same path as the gateway's inbound route in every + * case — that's what makes the fast-path "passthrough" rather than + * "rewrite": we forward the bytes to the same logical endpoint on the real + * provider, with `Authorization` swapped. + */ +const FORMAT_TO_UPSTREAM_PATH: Record = { + "openai-chat": "/v1/chat/completions", + "anthropic-messages": "/v1/messages", + "openai-responses": "/v1/responses", +}; + +/** + * Inbound format → set of model.api values where a 1:1 byte passthrough is + * legal. When the inbound format matches the model's native API, we skip the + * translate/rebuild round-trip and forward the request body unchanged with + * `Authorization` rewritten. Two big wins: + * - prompt caching hints (`cache_control`, etc.) flow through to upstream + * intact; the gateway no longer breaks anthropic prompt-caching; + * - provider-specific options (`metadata`, `service_tier`, `tool_choice` + * extensions, …) work without per-field allowlist maintenance here. + * + * `openai-codex-responses` is deliberately absent — codex runs over a + * websocket transport that has no equivalent inbound shape, so it always + * takes the translate path. + */ +const FORMAT_TO_PASSTHROUGH_API: Record> = { + "openai-chat": new Set(["openai-completions"]), + "anthropic-messages": new Set(["anthropic-messages"]), + "openai-responses": new Set(["openai-responses"]), +}; + +/** + * Hop-by-hop headers per RFC 7230. Stripped from both the inbound (so we don't + * forward the client's `Authorization` containing only the gateway bearer) and + * the upstream response (so we don't pass `Transfer-Encoding: chunked` back + * after we've already buffered). + */ +const HOP_BY_HOP_HEADERS = new Set([ + "connection", + "keep-alive", + "proxy-authenticate", + "proxy-authorization", + "te", + "trailers", + "transfer-encoding", + "upgrade", +]); + +// Options the caller's wire format may carry but the resolved provider can't +// honour are dropped silently in `buildStreamOptions`. We used to 400 here +// (`Unsupported option: temperature for openai-codex-responses`), but every +// realistic client (llm-git, openai SDK, anthropic SDK) bakes some of these +// defaults in without knowing which model they'll resolve to. Failing loudly +// just turned that into per-call config hell. Silent strip is what the +// upstream provider would do anyway when it ignores extra fields. + +function buildStreamOptions(parsed: ParsedFormatRequest, api: Api, signal: AbortSignal): SimpleStreamOptions { + const opts: SimpleStreamOptions = { signal }; + const { options } = parsed; + // Codex backend rejects `temperature` / `top_p` (per-model defaults only), + // so we drop them silently for that one provider. Every other unsupported + // option is just ignored by `streamSimple` if the underlying provider + // doesn't honour it. + const isCodex = api === "openai-codex-responses"; + if (options.maxOutputTokens !== undefined) opts.maxTokens = options.maxOutputTokens; + if (options.temperature !== undefined && !isCodex) opts.temperature = options.temperature; + if (options.topP !== undefined && !isCodex) opts.topP = options.topP; + if (options.topK !== undefined) opts.topK = options.topK; + if (options.toolChoice !== undefined) { + opts.toolChoice = + typeof options.toolChoice === "object" ? { type: "tool", name: options.toolChoice.name } : options.toolChoice; + } + if (options.reasoning !== undefined) opts.reasoning = options.reasoning; + if (options.hideThinkingSummary !== undefined) opts.hideThinkingSummary = options.hideThinkingSummary; + if (options.serviceTier !== undefined) opts.serviceTier = options.serviceTier; + if (options.presencePenalty !== undefined) opts.presencePenalty = options.presencePenalty; + if (options.disableReasoning !== undefined) opts.disableReasoning = options.disableReasoning; + if (options.cacheRetention !== undefined) opts.cacheRetention = options.cacheRetention; + if (options.thinkingBudget !== undefined) { + // Anthropic gives a single budget number with no effort label; bridge it + // to pi-ai's per-level map and default the level to "high" so providers + // that key off `reasoning` actually surface the budget. + opts.thinkingBudgets = { ...(opts.thinkingBudgets ?? {}), [Effort.High]: options.thinkingBudget }; + opts.reasoning ??= Effort.High; + } + return opts; +} + +function mirrorRequestAbort(req: Request): AbortController { + const controller = new AbortController(); + if (req.signal.aborted) { + controller.abort(req.signal.reason); + } else { + req.signal.addEventListener("abort", () => controller.abort(req.signal.reason), { once: true }); + } + return controller; +} + +function clientClosedResponse(): Response { + return json(499, { error: "client closed request" }); +} + +/** + * 1:1 byte passthrough fast-path. When the inbound format matches the model's + * native API (per {@link FORMAT_TO_PASSTHROUGH_API}), we skip parse + translate + * + re-emit and forward the request body as-is to the upstream provider with + * `Authorization` rewritten to the real access token. Provider-specific + * features (anthropic prompt caching, openai `service_tier`, tool-choice + * extensions, …) pass through unchanged. + * + * `body` is the already-parsed JSON; we re-serialize it to bytes for the + * upstream request. Re-serialization is intentional — Bun's `Request#json()` + * consumes the underlying stream, so the original bytes aren't available + * anyway, and any client-side whitespace/key-order difference is irrelevant + * to every provider this gateway targets. + */ +async function handlePassthrough( + route: { module: FormatModule; label: string }, + model: Model, + body: unknown, + apiKey: string, + req: Request, + peer: string, + signal: AbortSignal, +): Promise { + const wirePath = FORMAT_TO_UPSTREAM_PATH[route.label]; + if (!wirePath) { + // Shouldn't happen — caller already confirmed FORMAT_TO_PASSTHROUGH_API + // has an entry for this label, which implies a wire path exists. + return json(500, { error: `No upstream wire path for format ${route.label}` }); + } + const baseUrl = model.baseUrl.replace(/\/+$/, ""); + const upstreamUrl = `${baseUrl}${wirePath}`; + + const upstreamHeaders = new Headers(); + req.headers.forEach((value, key) => { + const lower = key.toLowerCase(); + // Strip every header the client uses to identify itself to the gateway. + // The gateway is the only thing that should be telling the upstream + // provider who's calling; the client's bearer (or anthropic's `x-api-key`, + // which omp's own anthropic provider always sends alongside `Authorization`) + // is just access control INTO the gateway and would otherwise leak to + // upstream as a 401-inducing bogus credential. + if (lower === "authorization" || lower === "x-api-key") return; + if (lower === "host" || lower === "content-length") return; + if (HOP_BY_HOP_HEADERS.has(lower)) return; + upstreamHeaders.set(key, value); + }); + upstreamHeaders.set("Authorization", `Bearer ${apiKey}`); + + let upstream: Response; + try { + upstream = await fetch(upstreamUrl, { + method: req.method, + headers: upstreamHeaders, + body: JSON.stringify(body), + signal, + }); + } catch (error) { + if (signal.aborted) return clientClosedResponse(); + const message = error instanceof Error ? error.message : String(error); + logger.warn("auth-gateway passthrough upstream failed", { + format: route.label, + provider: model.provider, + model: model.id, + upstream: upstreamUrl, + peer, + error: message, + }); + return json(502, { error: message }); + } + + logger.info("auth-gateway passthrough", { + format: route.label, + provider: model.provider, + model: model.id, + upstream: upstreamUrl, + status: upstream.status, + peer, + }); + + // Pass body straight through without buffering. Strip hop-by-hop headers + // from upstream, plus `content-encoding` and `content-length`: Bun's + // `fetch` transparently decodes gzip/br/deflate bodies but leaves the + // `Content-Encoding` header intact — forwarding it makes the client try to + // re-decode plain bytes and crash with `ZlibError`. `content-length` is + // stale too once Bun re-frames the response. + const outboundHeaders = new Headers(); + upstream.headers.forEach((value, key) => { + const lower = key.toLowerCase(); + if (HOP_BY_HOP_HEADERS.has(lower)) return; + if (lower === "content-encoding" || lower === "content-length") return; + outboundHeaders.set(key, value); + }); + return new Response(upstream.body, { + status: upstream.status, + statusText: upstream.statusText, + headers: outboundHeaders, + }); +} + +async function handleFormatEndpoint( + route: { module: FormatModule; label: string }, + bootOpts: AuthGatewayBootOptions, + req: Request, + peer: string, +): Promise { + const controller = mirrorRequestAbort(req); + if (controller.signal.aborted) return clientClosedResponse(); + + let body: unknown; + try { + body = await req.json(); + } catch (error) { + if (controller.signal.aborted) return clientClosedResponse(); + return json(400, { error: `Invalid JSON body: ${String(error)}` }); + } + if (controller.signal.aborted) return clientClosedResponse(); + + // All three supported wire formats put the model id on a top-level `model` + // field. Read it without running the full strict schema so the passthrough + // fast-path doesn't block on provider-specific fields the schema would + // otherwise reject (anthropic `metadata`, openai `service_tier`, …). + const modelId = + typeof body === "object" && body !== null && typeof (body as { model?: unknown }).model === "string" + ? (body as { model: string }).model + : undefined; + if (!modelId) { + return json(400, { error: "Missing top-level `model` field" }); + } + + const model = bootOpts.resolveModel(modelId); + if (!model) { + return json(404, { error: `Unknown model: ${modelId}` }); + } + + // pi-ai's stream() does NOT consult AuthStorage — the caller (us) is + // expected to resolve the credential and pass it as `options.apiKey`. + // For OAuth providers this returns the access token (refreshed via the + // broker override on AuthStorage when needed). + let apiKey: string | undefined; + try { + apiKey = await bootOpts.storage.getApiKey(model.provider, undefined, { modelId: model.id }); + } catch (error) { + if (controller.signal.aborted) return clientClosedResponse(); + const message = error instanceof Error ? error.message : String(error); + logger.warn("auth-gateway getApiKey threw", { provider: model.provider, peer, error: message }); + return json(502, { error: message }); + } + if (controller.signal.aborted) return clientClosedResponse(); + if (!apiKey) { + return json(401, { error: `No credential available for provider ${model.provider}` }); + } + + // Fast path: 1:1 byte passthrough when the inbound format matches the + // model's native API. Skips schema validation entirely — provider-specific + // fields (prompt caching, service tier, …) flow through unchanged. + const passthroughApis = FORMAT_TO_PASSTHROUGH_API[route.label]; + if (passthroughApis?.has(model.api)) { + return handlePassthrough(route, model, body, apiKey, req, peer, controller.signal); + } + + // Translate path: parse + validate against the strict format schema, + // rebuild as omp's canonical Context, dispatch through pi-ai's + // streamSimple, encode the canonical event stream back to the inbound + // format. Used when the inbound wire format and the selected model's + // native API differ (e.g. /v1/chat/completions targeting an Anthropic + // model, or /v1/responses targeting openai-codex over websocket). + let parsed: ParsedFormatRequest; + try { + parsed = route.module.parseRequest(body); + } catch (error) { + const message = error instanceof Error ? error.message : String(error); + return json(400, { error: message }); + } + if (controller.signal.aborted) return clientClosedResponse(); + + const streamOpts = buildStreamOptions(parsed, model.api, controller.signal); + streamOpts.apiKey = apiKey; + + logger.info("auth-gateway request", { + format: route.label, + model: parsed.modelId, + resolvedProvider: model.provider, + resolvedModel: model.id, + stream: parsed.stream, + peer, + }); + + let events: AssistantMessageEventStream; + try { + if (controller.signal.aborted) return clientClosedResponse(); + events = streamSimple(model, parsed.context, streamOpts); + } catch (error) { + const message = error instanceof Error ? error.message : String(error); + logger.warn("auth-gateway streamSimple threw", { format: route.label, error: message, peer }); + return json(502, { error: message }); + } + + if (!parsed.stream) { + try { + if (controller.signal.aborted) return clientClosedResponse(); + const message = await events.result(); + if (message.stopReason === "aborted" || message.stopReason === "error") { + const errorMessage = + message.errorMessage ?? + (message.stopReason === "aborted" ? "Request was aborted" : "Upstream request failed"); + logger.warn("auth-gateway non-streaming failed", { + format: route.label, + reason: message.stopReason, + error: errorMessage, + peer, + }); + return json(message.stopReason === "aborted" ? 499 : 502, { error: errorMessage }); + } + return json(200, route.module.encodeResponse(message, parsed.modelId)); + } catch (error) { + if (controller.signal.aborted) return clientClosedResponse(); + const errMsg = error instanceof Error ? error.message : String(error); + logger.warn("auth-gateway non-streaming aborted", { format: route.label, error: errMsg, peer }); + return json(502, { error: errMsg }); + } + } + if (controller.signal.aborted) return clientClosedResponse(); + + const sseStream = route.module.encodeStream(events, parsed.modelId, parsed.options); + return new Response(sseStream, { + status: 200, + headers: { + "Content-Type": "text/event-stream; charset=utf-8", + "Cache-Control": "no-cache", + Connection: "keep-alive", + }, + }); +} + +/** + * Snapshot of `GET /v1/usage` — fetchUsageReports already caches reports at 30s TTL + * inside AuthStorage, so this handler is a thin wrapper that surfaces the same + * data to HTTP callers (notably the macOS usage widget). + */ +async function handleUsage(storage: AuthStorage): Promise { + const reports = (await storage.fetchUsageReports?.()) ?? []; + // Drop the heavy provider-specific `raw` payload — UI consumers only need + // `limits` + `metadata`. Match the broker's `/v1/usage` shape so a single + // client struct (Swift widget, llm-git, ...) works against either endpoint. + const trimmed = reports.map(({ raw: _raw, ...rest }) => rest); + return json(200, { generatedAt: Date.now(), reports: trimmed }); +} + +function handleModelsList(opts: AuthGatewayBootOptions): Response { + const list = opts.listModels ? Array.from(opts.listModels()) : []; + const data = list.map(model => ({ + id: model.id, + object: "model" as const, + owned_by: model.provider, + api: model.api, + })); + return json(200, { object: "list", data }); +} + +export function startAuthGateway(opts: AuthGatewayBootOptions): AuthGatewayServerHandle { + const bind = parseBind(opts.bind ?? DEFAULT_AUTH_GATEWAY_BIND); + const tokens = new Set(opts.bearerTokens); + const version = opts.version; + + const server = Bun.serve({ + hostname: bind.hostname, + port: bind.port, + fetch: async (req): Promise => { + const url = new URL(req.url); + const pathname = url.pathname; + const peer = resolvePeer(req); + try { + if (req.method === "GET" && pathname === "/healthz") { + return json(200, { ok: true, version }); + } + if (!isAuthorized(req, tokens)) { + logger.info("auth-gateway request unauthorized", { method: req.method, path: pathname, peer }); + return json(401, { error: "unauthorized" }); + } + + // Aggregated usage — backed by AuthStorage's 30s cache. Same shape as + // the broker's `/v1/usage`, so widget/llm-git speak to either with the + // same client struct. + if (req.method === "GET" && pathname === "/v1/usage") { + return await handleUsage(opts.storage); + } + + // Provider-format dispatch. + const formatRoute = FORMAT_ROUTES[pathname]; + if (formatRoute && req.method === "POST") { + return await handleFormatEndpoint(formatRoute, opts, req, peer); + } + + // Model catalog. + if (req.method === "GET" && pathname === "/v1/models") { + return handleModelsList(opts); + } + + return json(404, { error: `No route: ${req.method} ${pathname}` }); + } catch (error) { + logger.error("auth-gateway handler crashed", { + method: req.method, + path: pathname, + peer, + error: String(error), + }); + return json(500, { error: "internal error" }); + } + }, + }); + + const boundHost = server.hostname ?? bind.hostname; + const boundPort = server.port ?? bind.port; + return { + url: `http://${boundHost}:${boundPort}`, + port: boundPort, + hostname: boundHost, + close: async () => { + server.stop(true); + }, + }; +} diff --git a/packages/ai/src/auth-gateway/types.ts b/packages/ai/src/auth-gateway/types.ts new file mode 100644 index 000000000..584f6e970 --- /dev/null +++ b/packages/ai/src/auth-gateway/types.ts @@ -0,0 +1,83 @@ +import type { Effort } from "../model-thinking"; +import type { AssistantMessage, AssistantMessageEventStream, CacheRetention, Context, ServiceTier } from "../types"; + +/** + * Wire types for the omp auth-gateway. + * + * The gateway sits between unauthenticated clients (containerized omp, + * llm-git, …) and the broker. It accepts provider-format HTTP requests + * (OpenAI chat-completions / Anthropic messages / OpenAI Responses), + * dispatches them through pi-ai's `streamSimple()`, and translates the + * canonical event stream back to the matching wire format. The gateway + * injects `Authorization` server-side so clients never see access tokens. + */ + +/** Default bind. Loopback-only — front with reverse proxy for remote access. */ +export const DEFAULT_AUTH_GATEWAY_BIND = "127.0.0.1:4000"; + +export type AuthGatewayToolChoice = "auto" | "none" | "required" | { name: string }; + +export interface AuthGatewayParsedRequestOptions { + maxOutputTokens?: number; + temperature?: number; + topP?: number; + topK?: number; + stopSequences?: string[]; + toolChoice?: AuthGatewayToolChoice; + /** Effort-level reasoning request (OpenAI Responses / Chat `reasoning_effort`). */ + reasoning?: Effort; + /** Force-disable reasoning (Anthropic `thinking: { type: "disabled" }`). */ + disableReasoning?: boolean; + /** + * Token budget for thinking (Anthropic `thinking.budget_tokens`). Bridged to + * pi-ai via `thinkingBudgets[high]` when the wire format only carries a + * single budget number with no effort label. + */ + thinkingBudget?: number; + /** Suppress the provider's reasoning summary stream. */ + hideThinkingSummary?: boolean; + /** OpenAI service tier (auto|default|flex|scale|priority). */ + serviceTier?: ServiceTier; + /** Presence penalty (OpenAI). */ + presencePenalty?: number; + /** Cache retention hint derived from inbound `cache_control` markers. */ + cacheRetention?: CacheRetention; + /** + * Provider-specific request controls that need server-side routing support + * but aren't yet first-class on this interface. + */ + extra?: Record; +} + +export interface AuthGatewayParsedRequest { + modelId: string; + context: Context; + stream: boolean; + options: AuthGatewayParsedRequestOptions; +} + +export interface AuthGatewayFormatModule { + parseRequest(body: unknown): AuthGatewayParsedRequest; + encodeResponse(message: AssistantMessage, requestedModelId: string): Record; + encodeStream( + events: AssistantMessageEventStream, + requestedModelId: string, + options?: AuthGatewayParsedRequestOptions, + ): ReadableStream; +} + +export interface AuthGatewayServerOptions { + /** Listen address. Default `127.0.0.1:4000`. */ + bind?: string; + /** Accept any of these bearer tokens. Empty allows unauthenticated calls. */ + bearerTokens: string[]; + /** Version surfaced on `/healthz`. */ + version?: string; +} + +export interface AuthGatewayServerHandle { + url: string; + port: number; + hostname: string; + close(): Promise; +} diff --git a/packages/ai/src/auth-storage.ts b/packages/ai/src/auth-storage.ts index c5e69cb6e..db6e4768d 100644 --- a/packages/ai/src/auth-storage.ts +++ b/packages/ai/src/auth-storage.ts @@ -136,9 +136,30 @@ export interface AuthCredentialStore { replaceAuthCredentialsForProvider(provider: string, credentials: AuthCredential[]): StoredAuthCredential[]; upsertAuthCredentialForProvider(provider: string, credential: AuthCredential): StoredAuthCredential[]; deleteAuthCredentialsForProvider(provider: string, disabledCause: string): void; - getCache(key: string): string | null; + getCache(key: string, options?: { includeExpired?: boolean }): string | null; setCache(key: string, value: string, expiresAtSec: number): void; cleanExpiredCache(): void; + /** + * Optional store-supplied OAuth refresh. When present, `AuthStorage` uses + * it before the per-provider local refresh path. `RemoteAuthCredentialStore` + * implements this against the broker; SQLite stores leave it undefined. + * + * Precedence: `AuthStorageOptions.refreshOAuthCredential` > this hook > local. + */ + refreshOAuthCredential?( + provider: Provider, + credentialId: number, + credential: OAuthCredential, + ): Promise; + /** + * Optional store-supplied aggregate usage fetch. When present, `AuthStorage` + * routes `fetchUsageReports()` here instead of fanning out per-credential. + * `RemoteAuthCredentialStore` proxies to the broker (whose datacenter IP + * isn't rate-limited like a heavy residential client). + * + * Precedence: `AuthStorageOptions.fetchUsageReports` > this hook > local fan-out. + */ + fetchUsageReports?(): Promise; } // ───────────────────────────────────────────────────────────────────────────── @@ -204,6 +225,17 @@ export type AuthStorageOptions = { * - `"broker http://can.internal:8765"` */ sourceLabel?: string; + /** + * Override `fetchUsageReports`. When set, `AuthStorage.fetchUsageReports` + * calls this instead of fanning out per-credential. The primary use case is + * routing through a broker that egresses from a less-throttled IP — e.g. a + * residential laptop trips Anthropic's per-IP rate limit on the usage + * endpoint and drops 2-of-5 credentials, while the VPS broker gets all 5. + * + * Implementations may return null when no usage data is available; the + * AuthStorage caller surfaces that to its own consumer unchanged. + */ + fetchUsageReports?: () => Promise; }; // ───────────────────────────────────────────────────────────────────────────── @@ -238,8 +270,22 @@ const DEFAULT_USAGE_PROVIDER_MAP = new Map( ); const USAGE_CACHE_PREFIX = "usage_cache:"; -const USAGE_REPORT_TTL_MS = 30_000; -const DEFAULT_USAGE_REQUEST_TIMEOUT_MS = 3_000; +// 5 min stale tolerance. Anthropic / OpenAI rate-limit /usage hard at the IP +// level so we can't fetch all N credentials every cycle; with a long cache +// each credential's last-known value sticks visible while peers retry. UI +// data (5h / 7d / monthly limits) is fine being a few minutes stale. +const USAGE_REPORT_TTL_MS = 5 * 60_000; +const USAGE_LAST_GOOD_RETENTION_MS = 24 * 60 * 60_000; +/** + * Per-credential cool-down after a usage fetch fails. While this window is + * active we serve the last successful value to avoid dropping the credential + * from the report; without a previous value we just return null and retry + * on the next poll. + */ +const USAGE_FAILURE_BACKOFF_MS = 10_000; +// Bumped from 3s — Claude usage retries up to 3 times with exponential backoff +// (~3.5s total worst case); a tight per-request budget aborts retries mid-cycle. +const DEFAULT_USAGE_REQUEST_TIMEOUT_MS = 10_000; const DEFAULT_OAUTH_REFRESH_TIMEOUT_MS = 10_000; /** * Cap on the buffered credential_disabled backlog held while no handler is attached. @@ -255,6 +301,7 @@ type UsageCacheEntry = { interface UsageCache { get(key: string): UsageCacheEntry | undefined; + getStale(key: string): UsageCacheEntry | undefined; set(key: string, entry: UsageCacheEntry): void; cleanup?(): void; } @@ -328,9 +375,17 @@ class AuthStorageUsageCache implements UsageCache { return parseUsageCacheEntry(raw); } + getStale(key: string): UsageCacheEntry | undefined { + const raw = this.store.getCache(`${USAGE_CACHE_PREFIX}${key}`, { includeExpired: true }); + if (!raw) return undefined; + return parseUsageCacheEntry(raw); + } + set(key: string, entry: UsageCacheEntry): void { const payload = JSON.stringify({ value: entry.value, expiresAt: entry.expiresAt }); - this.store.setCache(`${USAGE_CACHE_PREFIX}${key}`, payload, Math.floor(entry.expiresAt / 1000)); + const durableExpiresAt = + entry.value === null ? entry.expiresAt : Math.max(entry.expiresAt, Date.now() + USAGE_LAST_GOOD_RETENTION_MS); + this.store.setCache(`${USAGE_CACHE_PREFIX}${key}`, payload, Math.floor(durableExpiresAt / 1000)); } cleanup(): void { @@ -359,6 +414,7 @@ export class AuthStorage { /** Provider -> credentials cache, populated from store on reload(). */ #data: Map = new Map(); #runtimeOverrides: Map = new Map(); + #configOverrides: Map = new Map(); /** Tracks next credential index per provider:type key for round-robin distribution (non-session use). */ #providerRoundRobinIndex: Map = new Map(); /** Tracks the last used credential per provider for a session (used for rate-limit switching). */ @@ -377,6 +433,7 @@ export class AuthStorage { #store: AuthCredentialStore; #configValueResolver: (config: string) => Promise; #refreshOAuthCredentialOverride?: AuthStorageOptions["refreshOAuthCredential"]; + #fetchUsageReportsOverride?: AuthStorageOptions["fetchUsageReports"]; #sourceLabel?: string; #credentialDisabledListeners: Set<(event: CredentialDisabledEvent) => void | Promise> = new Set(); /** @@ -399,6 +456,7 @@ export class AuthStorage { this.#usageFetch = options.usageFetch ?? fetch; this.#usageRequestTimeoutMs = options.usageRequestTimeoutMs ?? DEFAULT_USAGE_REQUEST_TIMEOUT_MS; this.#refreshOAuthCredentialOverride = options.refreshOAuthCredential; + this.#fetchUsageReportsOverride = options.fetchUsageReports; this.#sourceLabel = options.sourceLabel; if (options.onCredentialDisabled) { // Constructor-registered subscribers are permanent for this AuthStorage's lifetime; @@ -482,6 +540,35 @@ export class AuthStorage { this.#runtimeOverrides.delete(provider); } + /** + * Register a per-provider API key sourced from user configuration + * (e.g. `models.yml` `providers..apiKey`). Higher priority than + * stored credentials and OAuth tokens — when the user pins a key in + * config, that key is what authenticates outbound requests, regardless + * of whatever the broker happens to have loaded for that provider. + * + * Lower priority than {@link setRuntimeApiKey} so a CLI `--api-key` + * still wins for the duration of a single invocation. + */ + setConfigApiKey(provider: string, apiKey: string): void { + this.#configOverrides.set(provider, apiKey); + } + + /** + * Remove a single config-sourced API key override. + */ + removeConfigApiKey(provider: string): void { + this.#configOverrides.delete(provider); + } + + /** + * Drop every config-sourced API key. Called by `ModelRegistry` before + * re-parsing `models.yml` so removed entries actually disappear. + */ + clearConfigApiKeys(): void { + this.#configOverrides.clear(); + } + /** * Set a fallback resolver for API keys not found in storage or env vars. * Used for custom provider keys from models.json. @@ -879,6 +966,7 @@ export class AuthStorage { */ hasAuth(provider: string): boolean { if (this.#runtimeOverrides.has(provider)) return true; + if (this.#configOverrides.has(provider)) return true; if (this.#getCredentialsForProvider(provider).length > 0) return true; if (getEnvApiKey(provider)) return true; if (this.#fallbackResolver?.(provider)) return true; @@ -913,8 +1001,9 @@ export class AuthStorage { const oauthCredentials = allCredentials.filter((c): c is OAuthCredential => c.type === "oauth"); if (oauthCredentials.length === 0) return undefined; - // Runtime override always returns before recording a session credential. - if (this.#runtimeOverrides.has(provider)) return undefined; + // Runtime / config overrides bypass OAuth account_uuid attribution — the + // caller is authenticating with an explicit key, not the broker's OAuth. + if (this.#runtimeOverrides.has(provider) || this.#configOverrides.has(provider)) return undefined; // Prefer the session-sticky credential when available. const sessionPref = this.#getSessionCredential(provider, sessionId); @@ -1467,6 +1556,7 @@ export class AuthStorage { const cacheKey = this.#buildUsageReportCacheKey(request); const now = Date.now(); const cached = this.#usageCache.get(cacheKey); + // Fresh cache hit: return whatever's there (success or null fallback). if (cached && cached.expiresAt > now) { return cached.value; } @@ -1476,11 +1566,27 @@ export class AuthStorage { const promise = (async () => { const report = await this.#fetchUsageUncached(request, timeoutMs); + const ttlJitter = USAGE_REPORT_TTL_MS * (Math.random() * 0.5 - 0.25); if (report !== null) { - this.#usageCache.set(cacheKey, { value: report, expiresAt: Date.now() + USAGE_REPORT_TTL_MS }); + // Success: stagger per-credential cache expiry so all accounts don't + // refresh in the same window — Anthropic / OpenAI rate-limit `/usage` + // per source IP regardless of account, and synchronized 5-credential + // fan-out trips 429s every cycle. With ±25% jitter on TTL the refresh + // times decorrelate within a few cycles. + this.#usageCache.set(cacheKey, { value: report, expiresAt: Date.now() + USAGE_REPORT_TTL_MS + ttlJitter }); return report; } - return cached?.value ?? null; + // Failure: cache the LAST GOOD value (if any) with a short jittered TTL + // so the credential cools down briefly without dropping out of the + // report. If we never had a good value, return null this cycle and + // don't write — let the next poll retry. + const lastGood = this.#usageCache.getStale(cacheKey)?.value ?? null; + if (lastGood !== null) { + const backoffJitter = USAGE_FAILURE_BACKOFF_MS * (Math.random() * 0.5 - 0.25); + const coolDown = Date.now() + USAGE_FAILURE_BACKOFF_MS + backoffJitter; + this.#usageCache.set(cacheKey, { value: lastGood, expiresAt: coolDown }); + } + return lastGood; })().finally(() => { this.#usageRequestInFlight.delete(cacheKey); }); @@ -1693,6 +1799,14 @@ export class AuthStorage { async fetchUsageReports(options?: { baseUrlResolver?: (provider: Provider) => string | undefined; }): Promise { + // Caller override > store-level hook > local per-credential fan-out. + // `RemoteAuthCredentialStore` implements the store hook so a gateway + // backed by a broker automatically routes usage to the broker without + // needing the caller to wire it explicitly. + const override = this.#fetchUsageReportsOverride ?? this.#store.fetchUsageReports?.bind(this.#store); + if (override) { + return override(); + } if (!this.#usageProviderResolver) return null; const requests = this.#collectUsageRequests(options); @@ -1702,12 +1816,12 @@ export class AuthStorage { providers: [...new Set(requests.map(request => request.provider))].sort(), }); + // Per-credential caching with jitter lives in #fetchUsageCached, so we + // don't store the aggregated result here — doing so locks the widget to + // a single decorrelation snapshot for 30s, defeating the jitter (some + // accounts can be missing from one fetch and present in the next; the + // aggregate cache freezes whichever set landed first). const cacheKey = this.#buildUsageReportsCacheKey(requests); - const now = Date.now(); - const cached = this.#usageCache.get(cacheKey); - if (cached && cached.expiresAt > now) { - return cached.value; - } const inFlight = this.#usageReportsInFlight.get(cacheKey); if (inFlight) return inFlight; @@ -1728,10 +1842,8 @@ export class AuthStorage { ); const reports = results.filter((report): report is UsageReport => report !== null); const deduped = this.#dedupeUsageReports(reports); - if (deduped.length > 0) { - this.#usageCache.set(cacheKey, { value: deduped, expiresAt: Date.now() + USAGE_REPORT_TTL_MS }); - } - const resolved = deduped.length > 0 ? deduped : (cached?.value ?? []); + // no outer cache write — see comment above. + const resolved = deduped; this.#usageLogger?.debug("Usage fetch resolved", { reports: resolved.map(report => { const accountLabel = @@ -1868,8 +1980,12 @@ export class AuthStorage { primaryDrainRate: number; orderPos: number; }> = []; - // Pre-fetch usage reports in parallel for non-blocked credentials - const usageResults = await Promise.all( + // Pre-fetch usage reports in parallel for non-blocked credentials. + // Wrap with a timeout so slow/429'd fetches don't indefinitely block + // credential selection — better to pick a credential without usage data + // than to hang the agent waiting for rate-limited usage endpoints. + const usageTimeout = Math.max(5000, this.#usageRequestTimeoutMs * 1.5); + const usagePromise = Promise.all( args.order.map(async idx => { const selection = args.credentials[idx]; if (!selection) return null; @@ -1882,6 +1998,14 @@ export class AuthStorage { return { selection, usage, usageChecked: true, blockedUntil: undefined as number | undefined }; }), ); + const usageResults = await Promise.race([usagePromise, Bun.sleep(usageTimeout).then(() => null)]).then( + result => + result ?? + args.order.map(idx => { + const selection = args.credentials[idx]; + return selection ? { selection, usage: null, usageChecked: false, blockedUntil: undefined } : null; + }), + ); for (let orderPos = 0; orderPos < usageResults.length; orderPos += 1) { const result = usageResults[orderPos]; @@ -2055,8 +2179,13 @@ export class AuthStorage { ): Promise { if (Date.now() < credential.expires) return credential; let refreshPromise: Promise; - if (this.#refreshOAuthCredentialOverride && credentialId !== undefined) { - refreshPromise = this.#refreshOAuthCredentialOverride(provider, credentialId, credential); + // Caller override > store-level hook > local per-provider refresh. + // `RemoteAuthCredentialStore` exposes the hook so a broker-backed gateway + // routes refresh through the broker without explicit wiring. + const storeRefresh = this.#store.refreshOAuthCredential?.bind(this.#store); + const overrideRefresh = this.#refreshOAuthCredentialOverride ?? storeRefresh; + if (overrideRefresh && credentialId !== undefined) { + refreshPromise = overrideRefresh(provider, credentialId, credential); } else { const customProvider = getOAuthProvider(provider); if (customProvider) { @@ -2273,6 +2402,11 @@ export class AuthStorage { return runtimeKey; } + const configKey = this.#configOverrides.get(provider); + if (configKey) { + return configKey; + } + const apiKeySelection = this.#selectCredentialByType(provider, "api_key"); if (apiKeySelection) { return this.#configValueResolver(apiKeySelection.credential.key); @@ -2303,10 +2437,11 @@ export class AuthStorage { * Get API key for a provider. * Priority: * 1. Runtime override (CLI --api-key) - * 2. API key from storage - * 3. OAuth token from storage (auto-refreshed) - * 4. Environment variable - * 5. Fallback resolver (models.json custom providers) + * 2. Config override (models.yml `providers..apiKey`) + * 3. API key from storage + * 4. OAuth token from storage (auto-refreshed) + * 5. Environment variable + * 6. Fallback resolver (models.yml custom providers, last-resort) */ async getApiKey(provider: string, sessionId?: string, options?: AuthApiKeyOptions): Promise { // Runtime override takes highest priority @@ -2315,6 +2450,16 @@ export class AuthStorage { return runtimeKey; } + // Config override: explicit apiKey pinned in models.yml beats the broker's + // OAuth credentials. The user redirected a provider at a custom baseUrl + // (e.g. an auth-gateway) and supplied the bearer for that endpoint — + // honor it instead of forwarding an upstream OAuth token that the proxy + // won't accept. + const configKey = this.#configOverrides.get(provider); + if (configKey) { + return configKey; + } + const apiKeySelection = this.#selectCredentialByType(provider, "api_key", sessionId); if (apiKeySelection) { this.#recordSessionCredential(provider, sessionId, "api_key", apiKeySelection.index); @@ -2458,11 +2603,12 @@ export class AuthStorage { /** * Describe where the active credential for a provider came from. * - * Surfaces three layers, highest precedence first: + * Surfaces four layers, highest precedence first: * 1. Runtime override (`--api-key`). - * 2. Stored credential (the one this session is currently sticky to, or the + * 2. Config override (`models.yml` `providers..apiKey`). + * 3. Stored credential (the one this session is currently sticky to, or the * one round-robin would pick next when no session id is supplied). - * 3. Env var / fallback resolver — when no stored credential exists. + * 4. Env var / fallback resolver — when no stored credential exists. * * The string is purely informational; consumers must not parse it. */ @@ -2470,6 +2616,9 @@ export class AuthStorage { if (this.#runtimeOverrides.has(provider)) { return "runtime override (--api-key)"; } + if (this.#configOverrides.has(provider)) { + return "config override (models.yml)"; + } const baseLabel = this.#sourceLabel ?? "local store"; const stored = this.#getStoredCredentials(provider); @@ -2702,6 +2851,7 @@ export class SqliteAuthCredentialStore implements AuthCredentialStore { #deleteByProviderStmt: Statement; #hardDeleteStmt: Statement; #getCacheStmt: Statement; + #getCacheIncludingExpiredStmt: Statement; #upsertCacheStmt: Statement; #deleteExpiredCacheStmt: Statement; #closed = false; @@ -2738,6 +2888,7 @@ export class SqliteAuthCredentialStore implements AuthCredentialStore { this.#getCacheStmt = this.#db.prepare( `SELECT value FROM cache WHERE key = ? AND expires_at > ${SQLITE_NOW_EPOCH}`, ); + this.#getCacheIncludingExpiredStmt = this.#db.prepare("SELECT value FROM cache WHERE key = ?"); this.#upsertCacheStmt = this.#db.prepare( "INSERT INTO cache (key, value, expires_at) VALUES (?, ?, ?) ON CONFLICT(key) DO UPDATE SET value = excluded.value, expires_at = excluded.expires_at", ); @@ -3154,9 +3305,10 @@ export class SqliteAuthCredentialStore implements AuthCredentialStore { } } - getCache(key: string): string | null { + getCache(key: string, options?: { includeExpired?: boolean }): string | null { try { - const row = this.#getCacheStmt.get(key) as { value?: string } | undefined; + const stmt = options?.includeExpired === true ? this.#getCacheIncludingExpiredStmt : this.#getCacheStmt; + const row = stmt.get(key) as { value?: string } | undefined; return row?.value ?? null; } catch { return null; @@ -3259,6 +3411,7 @@ export class SqliteAuthCredentialStore implements AuthCredentialStore { this.#deleteByProviderStmt.finalize(); this.#hardDeleteStmt.finalize(); this.#getCacheStmt.finalize(); + this.#getCacheIncludingExpiredStmt.finalize(); this.#upsertCacheStmt.finalize(); this.#deleteExpiredCacheStmt.finalize(); this.#db.close(); diff --git a/packages/ai/src/index.ts b/packages/ai/src/index.ts index e007493cd..7d71d9e7e 100644 --- a/packages/ai/src/index.ts +++ b/packages/ai/src/index.ts @@ -1,6 +1,8 @@ export { type ZodType, z } from "zod/v4"; export * from "./api-registry"; export * from "./auth-broker"; +export { type AuthGatewayBootOptions, type ModelResolver, startAuthGateway } from "./auth-gateway/server"; +export * from "./auth-gateway/types"; export * from "./auth-storage"; export * from "./model-cache"; export * from "./model-manager"; diff --git a/packages/ai/src/providers/anthropic-messages-server-schema.ts b/packages/ai/src/providers/anthropic-messages-server-schema.ts new file mode 100644 index 000000000..767aacaca --- /dev/null +++ b/packages/ai/src/providers/anthropic-messages-server-schema.ts @@ -0,0 +1,201 @@ +/** + * Zod schemas for the Anthropic Messages API request shape we accept on the + * gateway. Mirrors https://docs.anthropic.com/en/api/messages — only the + * shapes the gateway actually understands; unsupported fields are caught with + * `.refine(...)` so the error mentions them explicitly. + * + * Used by `anthropic-messages.ts:parseRequest` to validate the inbound JSON + * before walking it into pi-ai's canonical `Context`. + */ +import type { + ContentBlockParam, + ImageBlockParam, + MessageCreateParams, + MessageParam, + TextBlockParam, + Tool, + ToolChoice, +} from "@anthropic-ai/sdk/resources/messages"; +import * as z from "zod/v4"; + +// `cache_control` is accepted and translated to pi-ai's per-request +// `cacheRetention` (any `ttl: "1h"` marker upgrades the request to "long"; +// any other ephemeral marker maps to "short"). The walker doesn't try to +// preserve per-block breakpoints — pi-ai's anthropic provider re-applies them +// against the rebuilt outbound request anyway. +export const cacheControlSchema = z + .object({ + type: z.literal("ephemeral"), + ttl: z.union([z.literal("1h"), z.literal("5m")]).optional(), + }) + .loose(); + +// ─── Sources / inner shapes ───────────────────────────────────────────────── + +export const base64ImageSourceSchema = z.object({ + type: z.literal("base64"), + data: z.string().min(1), + media_type: z.string().min(1), +}); + +const textBlockSchema = z.object({ + type: z.literal("text"), + text: z.string(), + cache_control: cacheControlSchema.optional(), +}); + +const imageBlockSchema = z.object({ + type: z.literal("image"), + source: base64ImageSourceSchema, + cache_control: cacheControlSchema.optional(), +}); + +const thinkingBlockSchema = z.object({ + type: z.literal("thinking"), + thinking: z.string(), + signature: z.string().optional(), +}); + +const redactedThinkingBlockSchema = z.object({ + type: z.literal("redacted_thinking"), + data: z.string(), +}); + +const toolUseBlockSchema = z.object({ + type: z.literal("tool_use"), + id: z.string().min(1), + name: z.string().min(1), + input: z.record(z.string(), z.unknown()).optional(), +}); + +const toolResultContentBlockSchema = z.discriminatedUnion("type", [textBlockSchema, imageBlockSchema]); + +const toolResultBlockSchema = z.object({ + type: z.literal("tool_result"), + tool_use_id: z.string().min(1), + content: z.union([z.string(), z.array(toolResultContentBlockSchema)]).optional(), + is_error: z.boolean().optional(), + cache_control: cacheControlSchema.optional(), +}); + +// ─── System ──────────────────────────────────────────────────────────────── + +const systemBlockSchema = z.object({ + type: z.literal("text"), + text: z.string(), + cache_control: cacheControlSchema.optional(), +}); + +export const systemSchema = z.union([z.string(), z.array(systemBlockSchema)]).optional(); + +// ─── Messages ────────────────────────────────────────────────────────────── + +const userContentBlockSchema = z.discriminatedUnion("type", [textBlockSchema, imageBlockSchema, toolResultBlockSchema]); + +const assistantContentBlockSchema = z.discriminatedUnion("type", [ + textBlockSchema, + thinkingBlockSchema, + redactedThinkingBlockSchema, + toolUseBlockSchema, +]); + +export const userMessageSchema = z.object({ + role: z.literal("user"), + content: z.union([z.string(), z.array(userContentBlockSchema)]), +}); + +export const assistantMessageSchema = z.object({ + role: z.literal("assistant"), + content: z.union([z.string(), z.array(assistantContentBlockSchema)]), +}); + +export const messageSchema = z.discriminatedUnion("role", [userMessageSchema, assistantMessageSchema]); + +// ─── Tools ───────────────────────────────────────────────────────────────── + +export const toolSchema = z.object({ + name: z.string().min(1), + description: z.string().optional(), + input_schema: z.record(z.string(), z.unknown()), + cache_control: cacheControlSchema.optional(), +}); + +// ─── Tool choice ─────────────────────────────────────────────────────────── + +export const toolChoiceSchema = z + .discriminatedUnion("type", [ + z.object({ type: z.literal("auto"), disable_parallel_tool_use: z.unknown().optional() }), + z.object({ type: z.literal("any"), disable_parallel_tool_use: z.unknown().optional() }), + z.object({ type: z.literal("none"), disable_parallel_tool_use: z.unknown().optional() }), + z.object({ + type: z.literal("tool"), + name: z.string().min(1), + disable_parallel_tool_use: z.unknown().optional(), + }), + ]) + .refine(value => value.disable_parallel_tool_use === undefined, { + message: "tool_choice.disable_parallel_tool_use is not supported by this gateway", + }); + +// ─── Thinking ────────────────────────────────────────────────────────────── + +// Anthropic's three thinking shapes. `enabled` requires a budget; `disabled` +// suppresses reasoning even on models that default it on; `adaptive` lets the +// provider pick the budget on the fly. Extra hints (`display: "omitted"`, …) +// are accepted but ignored on the translate path. +export const thinkingConfigSchema = z.discriminatedUnion("type", [ + z.object({ + type: z.literal("enabled"), + budget_tokens: z.number(), + display: z.unknown().optional(), + }), + z.object({ + type: z.literal("disabled"), + display: z.unknown().optional(), + }), + z.object({ + type: z.literal("adaptive"), + budget_tokens: z.number().optional(), + display: z.unknown().optional(), + }), +]); + +// ─── Top-level request ───────────────────────────────────────────────────── + +export const anthropicMessagesRequestSchema = z.object({ + model: z.string().min(1), + messages: z.array(messageSchema), + max_tokens: z.number(), + system: systemSchema, + tools: z.array(toolSchema).optional(), + tool_choice: toolChoiceSchema.optional(), + temperature: z.number().optional(), + top_p: z.number().optional(), + top_k: z.number().optional(), + stop_sequences: z.array(z.string()).optional(), + stream: z.boolean().optional(), + thinking: thinkingConfigSchema.optional(), + // Spec fields that the gateway tolerates but doesn't translate. Anthropic + // clients commonly send `metadata: { user_id }` — failing the request just + // because we can't route it is hostile. They're accepted permissively and + // silently dropped on the translate path. + metadata: z.unknown().optional(), + container: z.unknown().optional(), + context_management: z.unknown().optional(), + mcp_servers: z.unknown().optional(), + service_tier: z.unknown().optional(), +}); + +/** + * Public types are sourced from the upstream Anthropic SDK so the gateway + * stays in lock-step with the canonical API surface; the schemas above are + * runtime validators for the subset we actually accept. + */ +export type AnthropicMessagesRequest = MessageCreateParams; +export type AnthropicSystem = MessageCreateParams["system"]; +export type AnthropicMessage = MessageParam; +export type AnthropicUserContentBlock = ContentBlockParam; +export type AnthropicAssistantContentBlock = ContentBlockParam; +export type AnthropicTool = Tool; +export type AnthropicToolChoice = ToolChoice; +export type AnthropicToolResultContent = TextBlockParam | ImageBlockParam; diff --git a/packages/ai/src/providers/anthropic-messages-server.ts b/packages/ai/src/providers/anthropic-messages-server.ts new file mode 100644 index 000000000..63bf05153 --- /dev/null +++ b/packages/ai/src/providers/anthropic-messages-server.ts @@ -0,0 +1,558 @@ +import type { + AssistantMessage, + AssistantMessageEventStream, + Message, + RedactedThinkingContent, + StopReason, + TextContent, + ThinkingContent, + Tool, + ToolCall, + ToolResultMessage, + UserMessage, +} from "../types"; +import { + type AnthropicAssistantContentBlock, + type AnthropicMessage, + type AnthropicSystem, + type AnthropicTool, + type AnthropicToolChoice, + type AnthropicToolResultContent, + type AnthropicUserContentBlock, + anthropicMessagesRequestSchema, +} from "./anthropic-messages-server-schema"; + +/** + * Anthropic Messages API (https://docs.anthropic.com/en/api/messages) ↔ pi-ai + * gateway translation. Inbound: foreign HTTP body → omp Context. Outbound: + * omp AssistantMessage[Stream] → Anthropic-shaped JSON / SSE. + */ + +import type { AuthGatewayParsedRequest as ParsedRequest } from "../auth-gateway/types"; + +export type { ParsedRequest }; + +// --------------------------------------------------------------------------- +// Inbound parsing +// --------------------------------------------------------------------------- + +type ImageContentPart = { type: "image"; data: string; mimeType: string }; + +function buildSystemPrompt(raw: AnthropicSystem): string[] | undefined { + if (raw === undefined) return undefined; + if (typeof raw === "string") return raw.length > 0 ? [raw] : undefined; + const parts = raw.map(block => block.text).filter(text => text.length > 0); + return parts.length > 0 ? [parts.join("\n\n")] : undefined; +} + +function makeUserMessage(parts: (TextContent | ImageContentPart)[], timestamp: number): UserMessage { + return { + role: "user", + content: parts.length === 1 && parts[0].type === "text" ? parts[0].text : parts, + timestamp, + }; +} + +function toolResultPartsFromBlocks( + content: AnthropicToolResultContent[] | string | undefined, +): (TextContent | ImageContentPart)[] { + if (content === undefined) return []; + if (typeof content === "string") return [{ type: "text", text: content }]; + const out: (TextContent | ImageContentPart)[] = []; + for (const block of content) { + if (block.type === "text") { + out.push({ type: "text", text: block.text }); + continue; + } + // block.type === "image" — schema only accepts base64 sources. + if (block.source.type === "base64") { + out.push({ type: "image", data: block.source.data, mimeType: block.source.media_type }); + } + } + return out; +} + +function walkUserContent( + blocks: string | AnthropicUserContentBlock[], + timestamp: number, +): (UserMessage | ToolResultMessage)[] { + const messages: (UserMessage | ToolResultMessage)[] = []; + const userParts: (TextContent | ImageContentPart)[] = []; + const flush = () => { + if (userParts.length === 0) return; + messages.push(makeUserMessage(userParts.splice(0), timestamp)); + }; + if (typeof blocks === "string") { + if (blocks.length > 0) userParts.push({ type: "text", text: blocks }); + flush(); + return messages; + } + for (const block of blocks) { + if (block.type === "text") { + userParts.push({ type: "text", text: block.text }); + } else if (block.type === "image") { + if (block.source.type !== "base64") continue; + userParts.push({ type: "image", data: block.source.data, mimeType: block.source.media_type }); + } else if (block.type === "tool_result") { + // tool_result blocks must follow any plain text/image siblings. + if (userParts.length > 0) { + throw new Error("anthropic-messages: user text/image blocks before tool_result are not supported"); + } + messages.push({ + role: "toolResult", + toolCallId: block.tool_use_id, + // Anthropic tool_results don't carry the tool name; downstream can rehydrate. + toolName: "", + content: toolResultPartsFromBlocks(block.content as AnthropicToolResultContent[] | string | undefined), + isError: block.is_error === true, + timestamp, + }); + } + } + flush(); + return messages; +} + +function walkAssistantContent( + blocks: string | AnthropicAssistantContentBlock[], +): (TextContent | ThinkingContent | RedactedThinkingContent | ToolCall)[] { + const out: (TextContent | ThinkingContent | RedactedThinkingContent | ToolCall)[] = []; + if (typeof blocks === "string") { + if (blocks.length > 0) out.push({ type: "text", text: blocks }); + return out; + } + for (const block of blocks) { + switch (block.type) { + case "text": + out.push({ type: "text", text: block.text }); + break; + case "thinking": { + const tc: ThinkingContent = { type: "thinking", thinking: block.thinking }; + if (block.signature !== undefined) tc.thinkingSignature = block.signature; + out.push(tc); + break; + } + case "redacted_thinking": + out.push({ type: "redactedThinking", data: block.data }); + break; + case "tool_use": + out.push({ + type: "toolCall", + id: block.id, + name: block.name, + arguments: block.input ?? {}, + }); + break; + } + } + return out; +} + +function walkTools(tools: AnthropicTool[] | undefined): Tool[] | undefined { + if (!tools) return undefined; + return tools.map(tool => ({ + name: tool.name, + description: tool.description ?? "", + parameters: tool.input_schema as Record, + })); +} + +function mapToolChoice(choice: AnthropicToolChoice | undefined): ParsedRequest["options"]["toolChoice"] { + if (!choice) return undefined; + switch (choice.type) { + case "auto": + return "auto"; + case "any": + return "required"; + case "none": + return "none"; + case "tool": + return { name: choice.name }; + } +} + +type AnthropicCacheControl = { type: "ephemeral"; ttl?: "1h" | "5m" }; +type HasCacheControl = { cache_control?: AnthropicCacheControl }; + +function readCacheControl(value: unknown): AnthropicCacheControl | undefined { + if (value === null || typeof value !== "object") return undefined; + const cc = (value as HasCacheControl).cache_control; + if (!cc || typeof cc !== "object" || cc.type !== "ephemeral") return undefined; + return cc; +} + +/** + * Anthropic clients annotate caching breakpoints per block via + * `cache_control: { type: "ephemeral", ttl?: "1h"|"5m" }`. pi-ai's + * `cacheRetention` is per-request, not per-block, and its anthropic provider + * re-applies breakpoints itself on the rebuilt outbound wire. Scan every + * block once and return the strongest retention requested: any `ttl: "1h"` + * promotes the request to "long", anything else ephemeral maps to "short". + */ +function deriveCacheRetention(data: { + system?: unknown; + messages: readonly unknown[]; + tools?: readonly unknown[]; +}): "short" | "long" | undefined { + let strongest: "short" | "long" | undefined; + const visit = (cc: AnthropicCacheControl | undefined): void => { + if (!cc) return; + if (cc.ttl === "1h") strongest = "long"; + else strongest ??= "short"; + }; + if (Array.isArray(data.system)) { + for (const block of data.system) visit(readCacheControl(block)); + } + for (const message of data.messages) { + if (message === null || typeof message !== "object") continue; + const content = (message as { content?: unknown }).content; + if (!Array.isArray(content)) continue; + for (const block of content) visit(readCacheControl(block)); + } + if (data.tools) { + for (const tool of data.tools) visit(readCacheControl(tool)); + } + return strongest; +} + +export function parseRequest(body: unknown): ParsedRequest { + const parsed = anthropicMessagesRequestSchema.safeParse(body); + if (!parsed.success) { + throw new Error(`anthropic-messages: ${parsed.error.message}`); + } + const data = parsed.data; + + const now = Date.now(); + const messages: Message[] = []; + for (const message of data.messages as AnthropicMessage[]) { + if (message.role === "user") { + for (const m of walkUserContent(message.content, now)) messages.push(m); + } else { + const assistant: AssistantMessage = { + role: "assistant", + content: walkAssistantContent(message.content), + api: "anthropic-messages", + provider: "anthropic", + model: data.model, + usage: emptyUsage(), + stopReason: "stop", + timestamp: now, + }; + messages.push(assistant); + } + } + + const options: ParsedRequest["options"] = { + maxOutputTokens: data.max_tokens, + }; + if (data.temperature !== undefined) options.temperature = data.temperature; + if (data.top_p !== undefined) options.topP = data.top_p; + if (data.top_k !== undefined) options.topK = data.top_k; + if (data.stop_sequences) options.stopSequences = data.stop_sequences; + const toolChoice = mapToolChoice(data.tool_choice as AnthropicToolChoice | undefined); + if (toolChoice !== undefined) options.toolChoice = toolChoice; + if (data.thinking) { + switch (data.thinking.type) { + case "enabled": + options.thinkingBudget = data.thinking.budget_tokens; + break; + case "disabled": + options.disableReasoning = true; + break; + case "adaptive": + if (data.thinking.budget_tokens !== undefined) { + options.thinkingBudget = data.thinking.budget_tokens; + } + break; + } + } + const cacheRetention = deriveCacheRetention(data); + if (cacheRetention !== undefined) options.cacheRetention = cacheRetention; + + return { + modelId: data.model, + context: { + systemPrompt: buildSystemPrompt(data.system as AnthropicSystem), + messages, + tools: walkTools(data.tools as AnthropicTool[] | undefined), + }, + stream: data.stream === true, + options, + }; +} + +function emptyUsage(): AssistantMessage["usage"] { + return { + input: 0, + output: 0, + cacheRead: 0, + cacheWrite: 0, + totalTokens: 0, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, + }; +} + +// --------------------------------------------------------------------------- +// Outbound encoding +// --------------------------------------------------------------------------- + +function newMessageId(): string { + const hex = (globalThis.crypto?.randomUUID?.() ?? randomFallback()).replace(/-/g, "").slice(0, 24); + return `msg_${hex}`; +} + +function randomFallback(): string { + // Sufficient for tests / environments without crypto.randomUUID + const buf = new Uint8Array(16); + for (let i = 0; i < 16; i++) buf[i] = Math.floor(Math.random() * 256); + const hex = Array.from(buf, b => b.toString(16).padStart(2, "0")).join(""); + return `${hex.slice(0, 8)}-${hex.slice(8, 12)}-${hex.slice(12, 16)}-${hex.slice(16, 20)}-${hex.slice(20)}`; +} + +function mapStopReasonOut(reason: StopReason): "end_turn" | "max_tokens" | "tool_use" { + switch (reason) { + case "length": + return "max_tokens"; + case "toolUse": + return "tool_use"; + default: + return "end_turn"; + } +} + +function encodeContentBlocks(message: AssistantMessage): Record[] { + const blocks: Record[] = []; + for (const c of message.content) { + switch (c.type) { + case "text": + blocks.push({ type: "text", text: c.text }); + break; + case "thinking": { + const b: Record = { type: "thinking", thinking: c.thinking }; + if (c.thinkingSignature) b.signature = c.thinkingSignature; + blocks.push(b); + break; + } + case "redactedThinking": + blocks.push({ type: "redacted_thinking", data: c.data }); + break; + case "toolCall": + blocks.push({ type: "tool_use", id: c.id, name: c.name, input: c.arguments ?? {} }); + break; + } + } + return blocks; +} + +function encodeUsage(message: AssistantMessage): Record { + return { + input_tokens: message.usage.input, + output_tokens: message.usage.output, + cache_read_input_tokens: message.usage.cacheRead, + cache_creation_input_tokens: message.usage.cacheWrite, + }; +} + +export function encodeResponse(message: AssistantMessage, requestedModelId: string): Record { + if (message.stopReason === "error" || message.stopReason === "aborted") { + throw new Error(message.errorMessage ?? `anthropic-messages: upstream ${message.stopReason}`); + } + return { + id: message.responseId ?? newMessageId(), + type: "message", + role: "assistant", + model: requestedModelId, + content: encodeContentBlocks(message), + stop_reason: mapStopReasonOut(message.stopReason), + stop_sequence: null, + usage: encodeUsage(message), + }; +} + +// --------------------------------------------------------------------------- +// Streaming encoder +// --------------------------------------------------------------------------- + +const ENCODER = new TextEncoder(); + +function sseFrame(event: string, data: Record): Uint8Array { + return ENCODER.encode(`event: ${event}\ndata: ${JSON.stringify(data)}\n\n`); +} + +type BlockKind = "text" | "thinking" | "tool_use"; + +interface OpenBlock { + index: number; + kind: BlockKind; +} + +export function encodeStream( + events: AssistantMessageEventStream, + requestedModelId: string, +): ReadableStream { + return new ReadableStream({ + async start(controller) { + const messageId = newMessageId(); + let started = false; + const open = new Map(); + + const ensureStart = (partial: AssistantMessage) => { + if (started) return; + started = true; + controller.enqueue( + sseFrame("message_start", { + type: "message_start", + message: { + id: messageId, + type: "message", + role: "assistant", + model: requestedModelId, + content: [], + stop_reason: null, + stop_sequence: null, + usage: encodeUsage(partial), + }, + }), + ); + }; + + const closeBlock = (index: number) => { + if (!open.has(index)) return; + controller.enqueue(sseFrame("content_block_stop", { type: "content_block_stop", index })); + open.delete(index); + }; + + try { + for await (const ev of events) { + switch (ev.type) { + case "start": + ensureStart(ev.partial); + break; + case "text_start": { + ensureStart(ev.partial); + open.set(ev.contentIndex, { index: ev.contentIndex, kind: "text" }); + controller.enqueue( + sseFrame("content_block_start", { + type: "content_block_start", + index: ev.contentIndex, + content_block: { type: "text", text: "" }, + }), + ); + break; + } + case "text_delta": + controller.enqueue( + sseFrame("content_block_delta", { + type: "content_block_delta", + index: ev.contentIndex, + delta: { type: "text_delta", text: ev.delta }, + }), + ); + break; + case "text_end": + closeBlock(ev.contentIndex); + break; + case "thinking_start": { + ensureStart(ev.partial); + open.set(ev.contentIndex, { index: ev.contentIndex, kind: "thinking" }); + controller.enqueue( + sseFrame("content_block_start", { + type: "content_block_start", + index: ev.contentIndex, + content_block: { type: "thinking", thinking: "" }, + }), + ); + break; + } + case "thinking_delta": + controller.enqueue( + sseFrame("content_block_delta", { + type: "content_block_delta", + index: ev.contentIndex, + delta: { type: "thinking_delta", thinking: ev.delta }, + }), + ); + break; + case "thinking_end": { + const c = ev.partial.content[ev.contentIndex]; + if (c?.type === "thinking" && c.thinkingSignature) { + controller.enqueue( + sseFrame("content_block_delta", { + type: "content_block_delta", + index: ev.contentIndex, + delta: { type: "signature_delta", signature: c.thinkingSignature }, + }), + ); + } + closeBlock(ev.contentIndex); + break; + } + case "toolcall_start": { + ensureStart(ev.partial); + const tc = ev.partial.content[ev.contentIndex] as ToolCall | undefined; + open.set(ev.contentIndex, { index: ev.contentIndex, kind: "tool_use" }); + controller.enqueue( + sseFrame("content_block_start", { + type: "content_block_start", + index: ev.contentIndex, + content_block: { + type: "tool_use", + id: tc?.id ?? "", + name: tc?.name ?? "", + input: {}, + }, + }), + ); + break; + } + case "toolcall_delta": + controller.enqueue( + sseFrame("content_block_delta", { + type: "content_block_delta", + index: ev.contentIndex, + delta: { type: "input_json_delta", partial_json: ev.delta }, + }), + ); + break; + case "toolcall_end": + closeBlock(ev.contentIndex); + break; + case "done": { + for (const idx of [...open.keys()]) closeBlock(idx); + controller.enqueue( + sseFrame("message_delta", { + type: "message_delta", + delta: { stop_reason: mapStopReasonOut(ev.reason), stop_sequence: null }, + usage: encodeUsage(ev.message), + }), + ); + controller.enqueue(sseFrame("message_stop", { type: "message_stop" })); + controller.close(); + return; + } + case "error": { + const msg = ev.error.errorMessage ?? "stream error"; + controller.enqueue( + sseFrame("error", { type: "error", error: { type: "api_error", message: msg } }), + ); + controller.close(); + return; + } + } + } + // stream ended without explicit done; close gracefully + for (const idx of [...open.keys()]) closeBlock(idx); + controller.enqueue(sseFrame("message_stop", { type: "message_stop" })); + controller.close(); + } catch (err) { + controller.enqueue( + sseFrame("error", { + type: "error", + error: { type: "api_error", message: err instanceof Error ? err.message : String(err) }, + }), + ); + controller.close(); + } + }, + }); +} diff --git a/packages/ai/src/providers/openai-chat-server-schema.ts b/packages/ai/src/providers/openai-chat-server-schema.ts new file mode 100644 index 000000000..0e79f4bcb --- /dev/null +++ b/packages/ai/src/providers/openai-chat-server-schema.ts @@ -0,0 +1,152 @@ +/** + * Zod schemas for the OpenAI chat-completions request shape we accept on the + * gateway. Mirrors https://platform.openai.com/docs/api-reference/chat — only + * the shapes the gateway translation layer understands. Unsupported fields + * inside `stream_options` are rejected explicitly. + */ +import type { + ChatCompletionContentPart, + ChatCompletionCreateParams, + ChatCompletionMessageParam, + ChatCompletionMessageToolCall, + ChatCompletionTool, + ChatCompletionToolChoiceOption, +} from "openai/resources/chat/completions"; +import * as z from "zod/v4"; + +// ─── User-message content parts ───────────────────────────────────────────── + +export const textPartSchema = z.object({ + type: z.literal("text"), + text: z.string(), +}); + +/** + * OpenAI documents `image_url` as either `{ url: string }` or — older clients — + * a bare string. Accept both shapes; downstream we extract a URL. + */ +export const imagePartSchema = z.object({ + type: z.literal("image_url"), + image_url: z.union([z.string(), z.object({ url: z.string() })]), +}); + +export const userContentPartSchema = z.union([textPartSchema, imagePartSchema]); + +// ─── Tool calls / tools ───────────────────────────────────────────────────── + +export const toolCallSchema = z.object({ + id: z.string(), + type: z.literal("function").optional(), + function: z.object({ + name: z.string(), + arguments: z.string(), + }), +}); + +export const toolSchema = z.object({ + type: z.literal("function"), + function: z.object({ + name: z.string().min(1), + description: z.string().optional(), + parameters: z.record(z.string(), z.unknown()).optional(), + }), +}); + +// ─── Tool choice ──────────────────────────────────────────────────────────── + +export const toolChoiceSchema = z.union([ + z.literal("auto"), + z.literal("none"), + z.literal("required"), + z.object({ + type: z.literal("function"), + function: z.object({ name: z.string().min(1) }), + }), +]); + +// ─── Messages ─────────────────────────────────────────────────────────────── + +const baseContent = z.union([z.string(), z.array(userContentPartSchema)]); + +export const systemMessageSchema = z.object({ + role: z.literal("system"), + content: baseContent, +}); + +export const developerMessageSchema = z.object({ + role: z.literal("developer"), + content: baseContent, +}); + +export const userMessageSchema = z.object({ + role: z.literal("user"), + content: baseContent, +}); + +export const assistantMessageSchema = z.object({ + role: z.literal("assistant"), + content: baseContent.optional(), + tool_calls: z.array(toolCallSchema).optional(), +}); + +export const toolMessageSchema = z.object({ + role: z.literal("tool"), + content: baseContent.optional(), + tool_call_id: z.string().optional(), +}); + +export const messageSchema = z.discriminatedUnion("role", [ + systemMessageSchema, + developerMessageSchema, + userMessageSchema, + assistantMessageSchema, + toolMessageSchema, +]); + +// ─── Stream options ───────────────────────────────────────────────────────── + +export const streamOptionsSchema = z + .object({ + include_usage: z.boolean().optional(), + }) + .strict(); + +// ─── Stop sequences ───────────────────────────────────────────────────────── + +export const stopSchema = z.union([z.string(), z.array(z.string())]); + +// ─── Top-level request ────────────────────────────────────────────────────── + +export const openaiChatRequestSchema = z.object({ + model: z.string().min(1), + messages: z.array(messageSchema), + tools: z.array(toolSchema).optional(), + tool_choice: toolChoiceSchema.optional(), + max_tokens: z.number().optional(), + max_completion_tokens: z.number().optional(), + temperature: z.number().optional(), + top_p: z.number().optional(), + stop: stopSchema.optional(), + stream: z.boolean().optional(), + stream_options: streamOptionsSchema.optional(), + // Passthroughs surfaced to providers via options.extra. We accept any JSON + // for them — the provider validates further if it cares. + response_format: z.unknown().optional(), + seed: z.number().optional(), + presence_penalty: z.number().optional(), + frequency_penalty: z.number().optional(), + logit_bias: z.record(z.string(), z.number()).optional(), + user: z.string().optional(), +}); + +/** + * Public types are sourced from the OpenAI SDK so the gateway stays in + * lock-step with the canonical API surface; the schemas above are runtime + * validators for the subset we actually accept. + */ +export type OpenAIChatRequest = ChatCompletionCreateParams; +export type OpenAIChatMessage = ChatCompletionMessageParam; +export type OpenAIChatToolCall = ChatCompletionMessageToolCall; +export type OpenAIChatTool = ChatCompletionTool; +export type OpenAIChatToolChoice = ChatCompletionToolChoiceOption; +export type OpenAIChatContentPart = ChatCompletionContentPart; diff --git a/packages/ai/src/providers/openai-chat-server.ts b/packages/ai/src/providers/openai-chat-server.ts new file mode 100644 index 000000000..f2a66f9df --- /dev/null +++ b/packages/ai/src/providers/openai-chat-server.ts @@ -0,0 +1,484 @@ +import { randomUUID } from "node:crypto"; +/** + * Parsed inbound OpenAI chat-completions request, ready to feed into pi-ai + * `stream(model, context, options)`. + */ +import type { AuthGatewayParsedRequest as ParsedRequest } from "../auth-gateway/types"; +import type { + AssistantMessage, + AssistantMessageEventStream, + Context, + ImageContent, + Message, + StopReason, + TextContent, + Tool, + ToolCall, + ToolResultMessage, + TSchema, +} from "../types"; +import { + type OpenAIChatContentPart, + type OpenAIChatMessage, + type OpenAIChatTool, + type OpenAIChatToolCall, + type OpenAIChatToolChoice, + openaiChatRequestSchema, +} from "./openai-chat-server-schema"; + +export type { ParsedRequest }; + +// --------------------------------------------------------------------------- +// parseRequest +// --------------------------------------------------------------------------- + +export function parseRequest(body: unknown): ParsedRequest { + const parsed = openaiChatRequestSchema.safeParse(body); + if (!parsed.success) { + throw new Error(`openai-chat: ${parsed.error.message}`); + } + const data = parsed.data; + + const now = Date.now(); + const systemParts: string[] = []; + const messages: Message[] = []; + + for (const m of data.messages as OpenAIChatMessage[]) { + switch (m.role) { + case "system": { + const text = stringifyContent(m.content); + if (text.length > 0) systemParts.push(text); + break; + } + case "developer": + messages.push({ role: "developer", content: parseUserLikeContent(m.content), timestamp: now }); + break; + case "user": + messages.push({ role: "user", content: parseUserLikeContent(m.content), timestamp: now }); + break; + case "assistant": + messages.push( + buildAssistantMessage( + (m.content ?? undefined) as string | OpenAIChatContentPart[] | undefined, + m.tool_calls, + data.model, + now, + ), + ); + break; + case "tool": + messages.push(buildToolMessage(m.content, m.tool_call_id, now)); + break; + } + } + + const tools = data.tools ? buildTools(data.tools as OpenAIChatTool[]) : undefined; + + const context: Context = { + messages, + ...(systemParts.length > 0 ? { systemPrompt: [systemParts.join("\n\n")] } : {}), + ...(tools ? { tools } : {}), + }; + + // Prefer max_completion_tokens (newer) over max_tokens. + const maxOutputTokens = data.max_completion_tokens ?? data.max_tokens; + const stopSequences = normalizeStop(data.stop); + const toolChoice = normalizeToolChoice(data.tool_choice); + const includeStreamingUsage = data.stream_options?.include_usage === true; + + const extra: Record = {}; + let hasExtra = false; + const carry = (key: K, value: unknown) => { + if (value === undefined) return; + extra[key] = value; + hasExtra = true; + }; + carry("response_format", data.response_format); + carry("seed", data.seed); + carry("presence_penalty", data.presence_penalty); + carry("frequency_penalty", data.frequency_penalty); + carry("logit_bias", data.logit_bias); + carry("user", data.user); + if (includeStreamingUsage) { + extra.includeStreamingUsage = true; + hasExtra = true; + } + + return { + modelId: data.model, + context, + stream: data.stream === true, + options: { + ...(maxOutputTokens !== undefined ? { maxOutputTokens } : {}), + ...(data.temperature !== undefined ? { temperature: data.temperature } : {}), + ...(data.top_p !== undefined ? { topP: data.top_p } : {}), + ...(stopSequences ? { stopSequences } : {}), + ...(toolChoice !== undefined ? { toolChoice } : {}), + ...(hasExtra ? { extra } : {}), + }, + }; +} + +function stringifyContent(content: string | OpenAIChatContentPart[] | undefined): string { + if (content === undefined) return ""; + if (typeof content === "string") return content; + const out: string[] = []; + for (const part of content) { + if (part.type === "text") out.push(part.text); + } + return out.join(""); +} + +function parseUserLikeContent( + content: string | OpenAIChatContentPart[] | undefined, +): string | (TextContent | ImageContent)[] { + if (content === undefined) return ""; + if (typeof content === "string") return content; + const parts: (TextContent | ImageContent)[] = []; + for (const part of content) { + if (part.type === "text") { + parts.push({ type: "text", text: part.text }); + continue; + } + if (part.type !== "image_url") continue; + const url = typeof part.image_url === "string" ? part.image_url : part.image_url.url; + const decoded = decodeDataUri(url); + if (decoded) { + parts.push({ type: "image", data: decoded.data, mimeType: decoded.mimeType }); + } else { + // No image fetcher available in the gateway; surface as a text placeholder so + // downstream providers still receive a coherent message. + parts.push({ type: "text", text: `[image: ${url}]` }); + } + } + return parts; +} + +function decodeDataUri(url: string): { data: string; mimeType: string } | undefined { + if (!url.startsWith("data:")) return undefined; + const comma = url.indexOf(","); + if (comma < 0) return undefined; + const header = url.slice(5, comma); + const payload = url.slice(comma + 1); + const isBase64 = header.endsWith(";base64"); + const mimeType = (isBase64 ? header.slice(0, -";base64".length) : header) || "application/octet-stream"; + const data = isBase64 ? payload : Buffer.from(decodeURIComponent(payload), "utf8").toString("base64"); + return { data, mimeType }; +} + +function buildAssistantMessage( + content: string | OpenAIChatContentPart[] | undefined, + toolCalls: OpenAIChatToolCall[] | undefined, + modelId: string, + now: number, +): AssistantMessage { + const parts: AssistantMessage["content"] = []; + const text = stringifyContent(content); + if (text.length > 0) parts.push({ type: "text", text }); + if (toolCalls) { + for (const raw of toolCalls) { + // Schema only accepts type:"function" (or omitted); narrow the SDK + // union here so the custom-tool variant doesn't trip TS. + if (raw.type !== undefined && raw.type !== "function") continue; + const fn = (raw as { function: { name: string; arguments: string } }).function; + const argsStr = fn.arguments; + let args: Record = {}; + if (argsStr.length > 0) { + try { + const v: unknown = JSON.parse(argsStr); + args = + v && typeof v === "object" && !Array.isArray(v) ? (v as Record) : { __raw: argsStr }; + } catch { + args = { __raw: argsStr }; + } + } + const call: ToolCall = { type: "toolCall", id: raw.id, name: fn.name, arguments: args }; + parts.push(call); + } + } + return { + role: "assistant", + content: parts, + api: "openai-completions", + provider: "openai", + model: modelId, + usage: { + input: 0, + output: 0, + cacheRead: 0, + cacheWrite: 0, + totalTokens: 0, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, + }, + stopReason: "stop", + timestamp: now, + }; +} + +function buildToolMessage( + content: string | OpenAIChatContentPart[] | undefined, + toolCallId: string | undefined, + now: number, +): ToolResultMessage { + return { + role: "toolResult", + toolCallId: toolCallId ?? "", + // OpenAI chat-completions doesn't carry the tool name on tool-role messages; + // downstream providers that need it tolerate an empty string. + toolName: "", + content: [{ type: "text", text: stringifyContent(content) }], + isError: false, + timestamp: now, + }; +} + +function buildTools(tools: OpenAIChatTool[]): Tool[] | undefined { + if (tools.length === 0) return undefined; + const out: Tool[] = []; + for (const t of tools) { + if (t.type !== "function") continue; + out.push({ + name: t.function.name, + description: t.function.description ?? "", + parameters: (t.function.parameters ?? {}) as Record as TSchema, + }); + } + return out; +} + +function normalizeStop(value: string | string[] | undefined): string[] | undefined { + if (value === undefined) return undefined; + if (typeof value === "string") return [value]; + return value.length > 0 ? value : undefined; +} + +function normalizeToolChoice(value: OpenAIChatToolChoice | undefined): ParsedRequest["options"]["toolChoice"] { + if (value === undefined) return undefined; + if (value === "auto" || value === "none" || value === "required") return value; + if ("function" in value) return { name: value.function.name }; + return undefined; +} + +// --------------------------------------------------------------------------- +// encodeResponse (non-streaming) +// --------------------------------------------------------------------------- + +export function encodeResponse(message: AssistantMessage, requestedModelId: string): Record { + const { text, toolCalls } = flattenAssistant(message); + + const responseMessage: Record = { + role: "assistant", + content: text.length > 0 ? text : null, + }; + if (toolCalls.length > 0) { + responseMessage.tool_calls = toolCalls.map(tc => ({ + id: tc.id, + type: "function", + function: { name: tc.name, arguments: stringifyArgs(tc.arguments) }, + })); + } + + return { + id: makeId(), + object: "chat.completion", + created: Math.floor(Date.now() / 1000), + model: requestedModelId, + choices: [ + { + index: 0, + message: responseMessage, + finish_reason: mapFinishReason(message.stopReason, toolCalls.length > 0), + }, + ], + usage: buildUsage(message), + }; +} + +function buildUsage(message: AssistantMessage): Record { + const promptTokens = message.usage.input + message.usage.cacheRead + message.usage.cacheWrite; + return { + prompt_tokens: promptTokens, + completion_tokens: message.usage.output, + total_tokens: promptTokens + message.usage.output, + prompt_tokens_details: { cached_tokens: message.usage.cacheRead }, + }; +} + +function flattenAssistant(message: AssistantMessage): { text: string; toolCalls: ToolCall[] } { + let text = ""; + const toolCalls: ToolCall[] = []; + for (const part of message.content) { + switch (part.type) { + case "text": + text += part.text; + break; + case "toolCall": + toolCalls.push(part); + break; + // thinking / redactedThinking: dropped — openai chat-completions has no reasoning channel. + } + } + return { text, toolCalls }; +} + +function isOnlyRaw(args: Record): boolean { + for (const k in args) { + if (k !== "__raw") return false; + } + return true; +} + +function stringifyArgs(args: Record): string { + // `__raw` is our fallback marker for un-parseable inbound args; preserve it verbatim on the way out. + if (typeof args.__raw === "string" && isOnlyRaw(args)) return args.__raw; + try { + return JSON.stringify(args); + } catch { + return "{}"; + } +} + +function mapFinishReason(reason: StopReason, hasToolCalls: boolean): string { + if (reason === "toolUse" || (hasToolCalls && reason === "stop")) return "tool_calls"; + if (reason === "length") return "length"; + return "stop"; +} + +function makeId(): string { + return `chatcmpl-${randomUUID()}`; +} + +// --------------------------------------------------------------------------- +// encodeStream (SSE) +// --------------------------------------------------------------------------- + +export function encodeStream( + events: AssistantMessageEventStream, + requestedModelId: string, + options?: ParsedRequest["options"], +): ReadableStream { + const encoder = new TextEncoder(); + const id = makeId(); + const created = Math.floor(Date.now() / 1000); + const includeUsage = options?.extra?.includeStreamingUsage === true; + + const baseChunk = (delta: Record, finishReason: string | null) => ({ + id, + object: "chat.completion.chunk", + created, + model: requestedModelId, + choices: [{ index: 0, delta, finish_reason: finishReason }], + ...(includeUsage ? { usage: null } : {}), + }); + + const writeSse = (controller: ReadableStreamDefaultController, payload: unknown): void => { + controller.enqueue(encoder.encode(`data: ${JSON.stringify(payload)}\n\n`)); + }; + + const writeUsage = (controller: ReadableStreamDefaultController, message: AssistantMessage): void => { + writeSse(controller, { + id, + object: "chat.completion.chunk", + created, + model: requestedModelId, + choices: [], + usage: buildUsage(message), + }); + }; + + return new ReadableStream({ + async start(controller) { + // contentIndex (from pi-ai events) -> tool_calls index on the wire. + const toolIndexByContentIndex = new Map(); + let nextToolIndex = 0; + let hasToolCalls = false; + let finishReason: string = "stop"; + + try { + // Initial role chunk. + writeSse(controller, baseChunk({ role: "assistant" }, null)); + + for await (const event of events) { + switch (event.type) { + case "text_delta": + if (event.delta.length > 0) { + writeSse(controller, baseChunk({ content: event.delta }, null)); + } + break; + + case "toolcall_start": { + hasToolCalls = true; + const idx = nextToolIndex++; + toolIndexByContentIndex.set(event.contentIndex, idx); + const partial = event.partial.content[event.contentIndex]; + const call = partial && partial.type === "toolCall" ? partial : undefined; + writeSse( + controller, + baseChunk( + { + tool_calls: [ + { + index: idx, + id: call?.id ?? "", + type: "function", + function: { name: call?.name ?? "", arguments: "" }, + }, + ], + }, + null, + ), + ); + break; + } + + case "toolcall_delta": { + const idx = toolIndexByContentIndex.get(event.contentIndex); + if (idx === undefined) break; + writeSse( + controller, + baseChunk({ tool_calls: [{ index: idx, function: { arguments: event.delta } }] }, null), + ); + break; + } + + case "done": + finishReason = + event.reason === "toolUse" + ? "tool_calls" + : event.reason === "length" + ? "length" + : hasToolCalls + ? "tool_calls" + : "stop"; + writeSse(controller, baseChunk({}, finishReason)); + if (includeUsage) writeUsage(controller, event.message); + controller.enqueue(encoder.encode("data: [DONE]\n\n")); + controller.close(); + return; + + case "error": { + const msg = event.error.errorMessage ?? "stream error"; + writeSse(controller, { error: { message: msg, type: "upstream_error" } }); + controller.close(); + return; + } + + // Drop start / *_start / *_end / thinking_* — chat-completions wire only + // surfaces deltas and the terminal finish_reason. + default: + break; + } + } + + // Stream ended without a terminal `done` (defensive). Close gracefully. + writeSse(controller, baseChunk({}, hasToolCalls ? "tool_calls" : "stop")); + controller.enqueue(encoder.encode("data: [DONE]\n\n")); + controller.close(); + } catch (err) { + const msg = err instanceof Error ? err.message : String(err); + writeSse(controller, { error: { message: msg, type: "upstream_error" } }); + controller.close(); + } + }, + }); +} diff --git a/packages/ai/src/providers/openai-responses-server-schema.ts b/packages/ai/src/providers/openai-responses-server-schema.ts new file mode 100644 index 000000000..b2d870d3e --- /dev/null +++ b/packages/ai/src/providers/openai-responses-server-schema.ts @@ -0,0 +1,201 @@ +/** + * Zod schemas for the OpenAI Responses API request shape we accept on the + * gateway. Mirrors https://platform.openai.com/docs/api-reference/responses — + * only the item types the gateway translation layer understands. Unsupported + * controls (background/include/metadata/prompt/…) are caught explicitly with + * `.refine(...)` so the error message names them. + */ +import type { + EasyInputMessage, + ResponseCreateParams, + ResponseFunctionToolCall, + ResponseInputContent, + ResponseInputItem, + ResponseOutputMessage, + ResponseReasoningItem, + Tool as ResponsesTool, +} from "openai/resources/responses/responses"; +import * as z from "zod/v4"; + +// ─── Input items ──────────────────────────────────────────────────────────── + +const inputTextSchema = z.object({ + type: z.literal("input_text"), + text: z.string(), +}); + +const outputTextSchema = z.object({ + type: z.literal("output_text"), + text: z.string(), +}); + +const summaryTextSchema = z.object({ + type: z.literal("summary_text"), + text: z.string(), +}); + +const reasoningTextSchema = z.object({ + type: z.literal("reasoning_text"), + text: z.string(), +}); + +const plainTextSchema = z.object({ + type: z.literal("text"), + text: z.string(), +}); + +const inputContentBlockSchema = z.union([inputTextSchema, plainTextSchema]); +const outputContentBlockSchema = z.union([outputTextSchema, plainTextSchema]); + +const userMessageItemSchema = z.object({ + type: z.literal("message").optional(), + role: z.union([z.literal("user"), z.literal("developer")]), + content: z.union([z.string(), z.array(inputContentBlockSchema)]).optional(), +}); + +const systemMessageItemSchema = z.object({ + type: z.literal("message").optional(), + role: z.literal("system"), + content: z.union([z.string(), z.array(inputContentBlockSchema)]).optional(), +}); + +const assistantMessageItemSchema = z.object({ + type: z.literal("message").optional(), + role: z.literal("assistant"), + content: z.union([z.string(), z.array(outputContentBlockSchema)]).optional(), +}); + +const reasoningItemSchema = z.object({ + type: z.literal("reasoning"), + id: z.string().optional(), + summary: z.array(summaryTextSchema).optional(), + content: z.array(reasoningTextSchema).optional(), +}); + +const functionCallItemSchema = z.object({ + type: z.literal("function_call"), + id: z.string().optional(), + call_id: z.string().min(1), + name: z.string().min(1), + arguments: z.string().optional(), +}); + +const functionCallOutputItemSchema = z.object({ + type: z.literal("function_call_output"), + call_id: z.string().min(1), + output: z.string().optional(), +}); + +/** + * An input item is one of the union members below. The convenience shape + * `{role, content}` (no `type`) is mapped to "message" before validation in + * the walker — schemas here only handle the canonical {type, ...} forms. + */ +export const inputItemSchema = z.union([ + userMessageItemSchema, + systemMessageItemSchema, + assistantMessageItemSchema, + reasoningItemSchema, + functionCallItemSchema, + functionCallOutputItemSchema, + // Tolerated but not bridged (file_search_call, web_search_call, …). + z.object({ type: z.string() }), +]); + +// Variant types alias the canonical SDK union members so the walker can +// narrow them cleanly. The convenience "message" shape (no `type` field) maps +// to EasyInputMessage; the explicit form maps to ResponseInputItem.Message. +export type OpenAIResponsesUserItem = EasyInputMessage | ResponseInputItem.Message; +export type OpenAIResponsesSystemItem = EasyInputMessage | ResponseInputItem.Message; +export type OpenAIResponsesAssistantItem = EasyInputMessage | ResponseOutputMessage; +export type OpenAIResponsesReasoningItem = ResponseReasoningItem; +export type OpenAIResponsesFunctionCallItem = ResponseFunctionToolCall; +export type OpenAIResponsesFunctionCallOutputItem = ResponseInputItem.FunctionCallOutput; + +// ─── Tools ────────────────────────────────────────────────────────────────── + +export const toolSchema = z.object({ + type: z.literal("function"), + name: z.string().min(1), + description: z.string().optional(), + parameters: z.record(z.string(), z.unknown()).optional(), + strict: z.boolean().optional(), +}); + +// Built-in tool entries (web_search, file_search, …) — accepted but skipped +// by the walker. +const builtinToolSchema = z.object({ + type: z.string(), +}); + +// ─── Tool choice ──────────────────────────────────────────────────────────── + +export const toolChoiceSchema = z.union([ + z.literal("auto"), + z.literal("none"), + z.literal("required"), + z.object({ + type: z.literal("function"), + name: z.string().min(1), + }), +]); + +// ─── Reasoning config ─────────────────────────────────────────────────────── + +export const reasoningConfigSchema = z.object({ + effort: z.string().optional(), + summary: z.string().optional(), +}); + +// ─── Stop ─────────────────────────────────────────────────────────────────── + +export const stopSchema = z.union([z.string(), z.array(z.string()), z.null()]); + +// ─── Top-level request ────────────────────────────────────────────────────── + +const refuse = (field: string) => + z + .unknown() + .refine(v => v === undefined, { message: `openai-responses: unsupported option \`${field}\`` }) + .optional(); + +export const openaiResponsesRequestSchema = z.object({ + model: z.string().min(1), + input: z.union([z.string(), z.array(inputItemSchema)]).optional(), + instructions: z.union([z.string(), z.null()]).optional(), + tools: z.array(z.union([toolSchema, builtinToolSchema])).optional(), + tool_choice: toolChoiceSchema.optional(), + max_output_tokens: z.number().optional(), + temperature: z.number().optional(), + top_p: z.number().optional(), + stop: stopSchema.optional(), + stream: z.boolean().optional(), + reasoning: reasoningConfigSchema.optional(), + store: z.boolean().optional(), + previous_response_id: z.string().optional(), + parallel_tool_calls: z.boolean().optional(), + service_tier: z.string().optional(), + presence_penalty: z.number().optional(), + // Explicitly rejected. + background: refuse("background"), + include: refuse("include"), + metadata: refuse("metadata"), + prompt: refuse("prompt"), + safety_identifier: refuse("safety_identifier"), + text: refuse("text"), + top_logprobs: refuse("top_logprobs"), + truncation: refuse("truncation"), + user: refuse("user"), +}); + +/** + * Public types are sourced from the OpenAI SDK so the gateway stays in + * lock-step with the canonical API surface; the schemas above are runtime + * validators for the subset we actually accept. + */ +export type OpenAIResponsesRequest = ResponseCreateParams; +export type OpenAIResponsesInputItem = ResponseInputItem; +export type OpenAIResponsesTool = ResponsesTool; +export type OpenAIResponsesToolChoice = NonNullable; +export type OpenAIResponsesInputContent = ResponseInputContent; +export type OpenAIResponsesOutputContent = ResponseOutputMessage["content"][number]; diff --git a/packages/ai/src/providers/openai-responses-server.ts b/packages/ai/src/providers/openai-responses-server.ts new file mode 100644 index 000000000..bfc4c980c --- /dev/null +++ b/packages/ai/src/providers/openai-responses-server.ts @@ -0,0 +1,889 @@ +/** + * OpenAI Responses HTTP wire-format ↔ omp Context bridge for the auth-gateway. + * + * Inbound: parses `POST /v1/responses` request bodies into a {@link ParsedRequest}. + * Outbound: encodes omp's {@link AssistantMessage} (and event stream) back into + * the documented `response.*` SSE taxonomy or the non-streaming JSON shape. + * + * Spec: https://platform.openai.com/docs/api-reference/responses + * Inverse direction (source-of-truth for item shapes): ../../providers/openai-responses.ts + * + * Note: images and other non-text input/output parts are not emitted by this + * encoder (omp's TextContent/ImageContent split is preserved on input but the + * Responses format documents far more part types than we exercise here). + */ + +import type { AuthGatewayParsedRequest as ParsedRequest } from "../auth-gateway/types"; +import type { + AssistantMessage, + AssistantMessageEventStream, + Context, + Message, + TextContent, + ThinkingContent, + Tool, + ToolCall, +} from "../types"; + +export type { ParsedRequest }; + +function isReasoningEffort(value: unknown): value is NonNullable { + return value === "minimal" || value === "low" || value === "medium" || value === "high" || value === "xhigh"; +} + +function isServiceTier(value: unknown): value is NonNullable { + return value === "auto" || value === "default" || value === "flex" || value === "scale" || value === "priority"; +} + +// ─── helpers ──────────────────────────────────────────────────────────────── + +function uuidNoDashes(): string { + return crypto.randomUUID().replace(/-/g, ""); +} + +function makeRespId(): string { + return `resp_${uuidNoDashes()}`; +} + +function makeMsgId(): string { + return `msg_${uuidNoDashes()}`; +} + +function makeReasoningId(): string { + return `rs_${uuidNoDashes()}`; +} + +function makeFuncCallId(): string { + return `fc_${uuidNoDashes()}`; +} + +import { + type OpenAIResponsesFunctionCallItem, + type OpenAIResponsesFunctionCallOutputItem, + type OpenAIResponsesInputContent, + type OpenAIResponsesOutputContent, + type OpenAIResponsesReasoningItem, + type OpenAIResponsesTool, + type OpenAIResponsesToolChoice, + openaiResponsesRequestSchema, +} from "./openai-responses-server-schema"; + +function isObj(v: unknown): v is Record { + return typeof v === "object" && v !== null && !Array.isArray(v); +} + +function asString(v: unknown): string | undefined { + return typeof v === "string" ? v : undefined; +} + +// ─── inbound parser ───────────────────────────────────────────────────────── + +function extractReasoningTextFromItem(item: OpenAIResponsesReasoningItem): string { + const fromContent = (item.content ?? []).map(c => c.text).join(""); + if (fromContent) return fromContent; + return (item.summary ?? []).map(c => c.text).join(""); +} + +function inputTextOf(blocks: OpenAIResponsesInputContent[] | string | undefined): string | TextContent[] { + if (typeof blocks === "string") return blocks; + if (!blocks) return []; + const parts: TextContent[] = []; + for (const block of blocks) { + if (block.type === "input_text") parts.push({ type: "text", text: block.text }); + } + return parts.length === 1 ? parts[0].text : parts; +} + +function outputTextOf(blocks: OpenAIResponsesOutputContent[] | string | undefined): TextContent[] { + if (typeof blocks === "string") return blocks.length > 0 ? [{ type: "text", text: blocks }] : []; + if (!blocks) return []; + const out: TextContent[] = []; + for (const block of blocks) { + if (block.type === "output_text") out.push({ type: "text", text: block.text }); + } + return out; +} + +function mapToolChoice(value: OpenAIResponsesToolChoice | undefined): ParsedRequest["options"]["toolChoice"] { + if (value === undefined) return undefined; + if (value === "auto" || value === "none" || value === "required") return value; + // Schema only validates ToolChoiceFunction; narrow defensively against the + // wider SDK union (allowed/types/mcp/custom/apply_patch/shell variants). + if ("type" in value && value.type === "function" && "name" in value) return { name: value.name }; + return undefined; +} + +function buildTools(tools: Array | undefined): Tool[] | undefined { + if (!tools) return undefined; + const out: Tool[] = []; + for (const t of tools) { + // Skip non-function tools (web_search_call, file_search_call, …). + if (t.type !== "function") continue; + const fn = t as Extract; + const tool: Tool = { + name: fn.name, + description: fn.description ?? "", + parameters: (fn.parameters ?? {}) as Tool["parameters"], + }; + if (fn.strict !== undefined && fn.strict !== null) tool.strict = fn.strict; + out.push(tool); + } + return out.length > 0 ? out : undefined; +} + +function ensureAssistantPlaceholder(messages: Message[], modelId: string, now: number): AssistantMessage { + const last = messages[messages.length - 1]; + if (last && last.role === "assistant") return last; + const placeholder: AssistantMessage = { + role: "assistant", + content: [], + api: "openai-responses", + provider: "openai", + model: modelId, + usage: { + input: 0, + output: 0, + cacheRead: 0, + cacheWrite: 0, + totalTokens: 0, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, + }, + stopReason: "stop", + timestamp: now, + }; + messages.push(placeholder); + return placeholder; +} + +export function parseRequest(body: unknown): ParsedRequest { + const parsed = openaiResponsesRequestSchema.safeParse(body); + if (!parsed.success) { + throw new Error(`openai-responses: ${parsed.error.message}`); + } + const data = parsed.data; + + const now = Date.now(); + const messages: Message[] = []; + const systemPrompt: string[] = []; + + if (typeof data.instructions === "string" && data.instructions.length > 0) { + systemPrompt.push(data.instructions); + } + + if (typeof data.input === "string") { + messages.push({ role: "user", content: data.input, timestamp: now }); + } else if (data.input) { + for (const item of data.input) { + // Items may omit `type` and rely on `role` (the convenience shape). + const effectiveType = item.type ?? ("role" in item ? "message" : undefined); + if (effectiveType === "message") { + const msg = item as { + role?: string; + content?: OpenAIResponsesInputContent[] | OpenAIResponsesOutputContent[] | string; + }; + switch (msg.role) { + case "system": { + const text = inputTextOf(msg.content as OpenAIResponsesInputContent[] | string | undefined); + const flat = typeof text === "string" ? text : text.map(p => p.text).join(""); + if (flat.length > 0) systemPrompt.push(flat); + break; + } + case "user": + case "developer": { + const content = inputTextOf(msg.content as OpenAIResponsesInputContent[] | string | undefined); + messages.push({ role: msg.role, content, timestamp: now }); + break; + } + case "assistant": { + const parts = outputTextOf(msg.content as OpenAIResponsesOutputContent[] | string | undefined); + messages.push({ + role: "assistant", + content: parts, + api: "openai-responses", + provider: "openai", + model: data.model, + usage: { + input: 0, + output: 0, + cacheRead: 0, + cacheWrite: 0, + totalTokens: 0, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, + }, + stopReason: "stop", + timestamp: now, + }); + break; + } + } + continue; + } + if (effectiveType === "reasoning") { + const reasoning = item as OpenAIResponsesReasoningItem; + const text = extractReasoningTextFromItem(reasoning); + const thinking: ThinkingContent = { + type: "thinking", + thinking: text, + thinkingSignature: JSON.stringify(reasoning), + ...(reasoning.id ? { itemId: reasoning.id } : {}), + }; + ensureAssistantPlaceholder(messages, data.model, now).content.push(thinking); + continue; + } + if (effectiveType === "function_call") { + const call = item as OpenAIResponsesFunctionCallItem; + const argsRaw = call.arguments ?? "{}"; + let args: Record; + try { + const parsed: unknown = JSON.parse(argsRaw); + args = isObj(parsed) ? parsed : {}; + } catch { + throw new Error(`openai-responses: function_call ${call.call_id} has invalid JSON arguments`); + } + const toolCall: ToolCall = { + type: "toolCall", + id: call.call_id, + name: call.name, + arguments: args, + ...(call.id ? { thoughtSignature: call.id } : {}), + }; + ensureAssistantPlaceholder(messages, data.model, now).content.push(toolCall); + continue; + } + if (effectiveType === "function_call_output") { + const output = item as OpenAIResponsesFunctionCallOutputItem; + // Find the matching tool call name from earlier assistant content. + let toolName = ""; + for (let i = messages.length - 1; i >= 0; i--) { + const m = messages[i]; + if (m.role !== "assistant") continue; + for (const c of m.content) { + if (c.type === "toolCall" && c.id === output.call_id) { + toolName = c.name; + break; + } + } + if (toolName) break; + } + messages.push({ + role: "toolResult", + toolCallId: output.call_id, + toolName, + content: [{ type: "text", text: typeof output.output === "string" ? output.output : "" }], + isError: false, + timestamp: now, + }); + } + // Other item types are tolerated but not bridged. + } + } + + const tools = buildTools(data.tools); + const context: Context = { + ...(systemPrompt.length > 0 ? { systemPrompt } : {}), + messages, + ...(tools ? { tools } : {}), + }; + + const options: ParsedRequest["options"] = {}; + if (data.max_output_tokens !== undefined) options.maxOutputTokens = data.max_output_tokens; + if (data.temperature !== undefined) options.temperature = data.temperature; + if (data.top_p !== undefined) options.topP = data.top_p; + if (data.stop !== undefined && data.stop !== null) { + options.stopSequences = typeof data.stop === "string" ? [data.stop] : data.stop; + } + const toolChoice = mapToolChoice(data.tool_choice); + if (toolChoice !== undefined) options.toolChoice = toolChoice; + if (data.reasoning?.effort && isReasoningEffort(data.reasoning.effort)) { + options.reasoning = data.reasoning.effort; + } + // OpenAI summary "auto"|"concise"|"detailed" → request a visible summary; + // absent → leave pi-ai's default. The "none" / absence inverse maps to + // `hideThinkingSummary: true`. + if (data.reasoning?.summary === undefined) { + // no-op; provider decides + } else if (data.reasoning.summary === "none") { + options.hideThinkingSummary = true; + } + if (data.service_tier !== undefined && isServiceTier(data.service_tier)) { + options.serviceTier = data.service_tier; + } + if (data.presence_penalty !== undefined) options.presencePenalty = data.presence_penalty; + // `store`, `previous_response_id`, `parallel_tool_calls` are accepted by the + // schema for forward-compatibility but not yet plumbed through pi-ai. + + return { + modelId: data.model, + context, + stream: data.stream === true, + options, + }; +} + +// ─── output item builders (shared by streaming + non-streaming encoders) ──── + +type ReasoningOutputItem = { + type: "reasoning"; + id: string; + summary: Array<{ type?: string; text?: string }>; + content?: Array<{ type: "reasoning_text"; text: string }>; +} & Record; + +type OutputItem = + | ReasoningOutputItem + | { + type: "message"; + id: string; + role: "assistant"; + status: "completed"; + content: Array<{ type: "output_text"; text: string; annotations: never[] }>; + } + | { type: "function_call"; id: string; call_id: string; name: string; arguments: string; status: "completed" }; + +type ResponseStatus = "completed" | "in_progress" | "failed" | "incomplete"; + +function responseStatusForStopReason(message: AssistantMessage): ResponseStatus { + if (message.stopReason === "length") return "incomplete"; + if (message.stopReason === "error" || message.stopReason === "aborted") return "failed"; + return "completed"; +} + +function buildReasoningItem(part: ThinkingContent): ReasoningOutputItem { + if (part.thinkingSignature) { + try { + const parsed: unknown = JSON.parse(part.thinkingSignature); + if (isObj(parsed) && parsed.type === "reasoning") { + const id = part.itemId ?? asString(parsed.id) ?? makeReasoningId(); + return { ...parsed, type: "reasoning", id } as ReasoningOutputItem; + } + } catch { + // Not a serialized Responses reasoning item; fall back to raw thinking text. + } + } + return { + type: "reasoning", + id: part.itemId ?? makeReasoningId(), + summary: [], + content: [{ type: "reasoning_text", text: part.thinking }], + }; +} + +function reasoningItemId(part: ThinkingContent): string { + if (part.itemId) return part.itemId; + if (part.thinkingSignature) { + try { + const parsed: unknown = JSON.parse(part.thinkingSignature); + if (isObj(parsed)) { + const id = asString(parsed.id); + if (id) return id; + } + } catch { + // Not a serialized Responses reasoning item. + } + } + return makeReasoningId(); +} + +/** + * Walk the assistant content array and group consecutive TextContent into a + * single message item; each ThinkingContent / ToolCall is its own item. + */ +function buildOutputItems(message: AssistantMessage): OutputItem[] { + const out: OutputItem[] = []; + let pendingMessage: Extract | null = null; + const flushMessage = () => { + if (pendingMessage) { + out.push(pendingMessage); + pendingMessage = null; + } + }; + + for (const part of message.content) { + if (part.type === "text") { + if (!pendingMessage) { + pendingMessage = { + type: "message", + id: makeMsgId(), + role: "assistant", + status: "completed", + content: [], + }; + } + pendingMessage.content.push({ type: "output_text", text: part.text, annotations: [] }); + } else if (part.type === "thinking") { + flushMessage(); + out.push(buildReasoningItem(part)); + } else if (part.type === "toolCall") { + flushMessage(); + const id = part.thoughtSignature ?? makeFuncCallId(); + out.push({ + type: "function_call", + id, + call_id: part.id, + name: part.name, + arguments: JSON.stringify(part.arguments ?? {}), + status: "completed", + }); + } + // RedactedThinking is silently dropped — no direct Responses wire representation. + } + flushMessage(); + return out; +} + +function buildUsage(message: AssistantMessage): Record { + const u = message.usage; + const inputTokens = u.input + u.cacheRead + u.cacheWrite; + return { + input_tokens: inputTokens, + input_tokens_details: { cached_tokens: u.cacheRead }, + output_tokens: u.output, + output_tokens_details: { reasoning_tokens: u.reasoningTokens ?? 0 }, + total_tokens: inputTokens + u.output, + }; +} + +function buildResponseEnvelope( + message: AssistantMessage, + requestedModelId: string, + id: string, + status: ResponseStatus, + items: OutputItem[] | [], + usage: Record | null, +): Record { + return { + id, + object: "response", + created_at: Math.floor(message.timestamp / 1000), + status, + model: requestedModelId, + output: items, + usage, + ...(status === "incomplete" ? { incomplete_details: { reason: "max_output_tokens" } } : {}), + ...(status === "failed" ? { error: { message: message.errorMessage ?? "response failed" } } : {}), + }; +} + +// ─── encodeResponse (non-streaming) ───────────────────────────────────────── + +export function encodeResponse(message: AssistantMessage, requestedModelId: string): Record { + const items = buildOutputItems(message); + return buildResponseEnvelope( + message, + requestedModelId, + makeRespId(), + responseStatusForStopReason(message), + items, + buildUsage(message), + ); +} + +// ─── encodeStream ─────────────────────────────────────────────────────────── + +interface OpenMessage { + kind: "message"; + itemId: string; + outputIndex: number; + contentIndex: number; + currentPartText: string; + content: Array<{ type: "output_text"; text: string; annotations: never[] }>; +} +interface OpenReasoning { + kind: "reasoning"; + itemId: string; + outputIndex: number; + reasoningText: string; +} +interface OpenFunctionCall { + kind: "function_call"; + itemId: string; + outputIndex: number; + callId: string; + name: string; + argsText: string; +} +type OpenItem = OpenMessage | OpenReasoning | OpenFunctionCall; + +function sseEvent(name: string, data: unknown): string { + return `event: ${name}\ndata: ${JSON.stringify(data)}\n\n`; +} + +export function encodeStream( + events: AssistantMessageEventStream, + requestedModelId: string, +): ReadableStream { + const encoder = new TextEncoder(); + const responseId = makeRespId(); + let sequenceNumber = 0; + const seq = () => sequenceNumber++; + + return new ReadableStream({ + async start(controller) { + const emit = (name: string, data: Record) => { + controller.enqueue(encoder.encode(sseEvent(name, { type: name, sequence_number: seq(), ...data }))); + }; + const emitDone = () => controller.enqueue(encoder.encode("data: [DONE]\n\n")); + + let createdAt = Math.floor(Date.now() / 1000); + let outputIndex = 0; + const state: { open: OpenItem | null } = { open: null }; + const finishedItems: OutputItem[] = []; + + const openMessage = (): OpenMessage => { + const itemId = makeMsgId(); + const item = { + type: "message" as const, + id: itemId, + status: "in_progress", + role: "assistant" as const, + content: [] as Array<{ type: "output_text"; text: string; annotations: never[] }>, + }; + emit("response.output_item.added", { output_index: outputIndex, item }); + const next: OpenMessage = { + kind: "message", + itemId, + outputIndex, + contentIndex: 0, + currentPartText: "", + content: [], + }; + state.open = next; + return next; + }; + + const openReasoning = (partial: AssistantMessage, contentIndex: number): OpenReasoning => { + const part = partial.content[contentIndex]; + const itemId = part && part.type === "thinking" ? reasoningItemId(part) : makeReasoningId(); + const item = { + type: "reasoning" as const, + id: itemId, + summary: [] as never[], + content: [] as Array<{ type: "reasoning_text"; text: string }>, + }; + emit("response.output_item.added", { output_index: outputIndex, item }); + const next: OpenReasoning = { kind: "reasoning", itemId, outputIndex, reasoningText: "" }; + state.open = next; + return next; + }; + + const openToolCall = (partial: AssistantMessage, contentIndex: number): OpenFunctionCall => { + const part = partial.content[contentIndex]; + const tc = part && part.type === "toolCall" ? part : undefined; + const itemId = tc?.thoughtSignature ?? makeFuncCallId(); + const callId = tc?.id ?? ""; + const name = tc?.name ?? ""; + const item = { + type: "function_call" as const, + id: itemId, + call_id: callId, + name, + arguments: "", + status: "in_progress", + }; + emit("response.output_item.added", { output_index: outputIndex, item }); + const next: OpenFunctionCall = { kind: "function_call", itemId, outputIndex, callId, name, argsText: "" }; + state.open = next; + return next; + }; + + const closeOpen = () => { + if (!state.open) return; + if (state.open.kind === "message") { + // (No defensive part-close needed; text_end always flushes the part before + // the next non-text event triggers closeOpen.) + const item = { + type: "message", + id: state.open.itemId, + status: "completed", + role: "assistant", + content: state.open.content, + }; + emit("response.output_item.done", { output_index: state.open.outputIndex, item }); + finishedItems.push({ + type: "message", + id: state.open.itemId, + role: "assistant", + status: "completed", + content: state.open.content, + }); + } else if (state.open.kind === "reasoning") { + const item = { + type: "reasoning", + id: state.open.itemId, + summary: [], + content: [{ type: "reasoning_text", text: state.open.reasoningText ?? "" }], + }; + emit("response.output_item.done", { output_index: state.open.outputIndex, item }); + finishedItems.push({ + type: "reasoning", + id: state.open.itemId, + summary: [], + content: [{ type: "reasoning_text", text: state.open.reasoningText ?? "" }], + }); + } else { + const args = state.open.argsText ?? ""; + const item = { + type: "function_call", + id: state.open.itemId, + call_id: state.open.callId ?? "", + name: state.open.name ?? "", + arguments: args, + status: "completed", + }; + emit("response.output_item.done", { output_index: state.open.outputIndex, item }); + finishedItems.push({ + type: "function_call", + id: state.open.itemId, + call_id: state.open.callId ?? "", + name: state.open.name ?? "", + arguments: args, + status: "completed", + }); + } + outputIndex++; + state.open = null; + }; + + try { + let finalMessage: AssistantMessage | null = null; + let failureMessage: AssistantMessage | null = null; + + for await (const ev of events) { + switch (ev.type) { + case "start": { + createdAt = Math.floor((ev.partial.timestamp || Date.now()) / 1000); + controller.enqueue( + encoder.encode( + sseEvent("response.created", { + type: "response.created", + sequence_number: seq(), + response: { + id: responseId, + object: "response", + created_at: createdAt, + status: "in_progress", + model: requestedModelId, + output: [], + usage: null, + }, + }), + ), + ); + break; + } + case "text_start": { + let cur: OpenMessage; + if (state.open && state.open.kind === "message") { + // continue same message item, new content part + cur = state.open; + cur.currentPartText = ""; + } else { + if (state.open) closeOpen(); + cur = openMessage(); + } + const part = { type: "output_text", text: "", annotations: [] as never[] }; + emit("response.content_part.added", { + item_id: cur.itemId, + output_index: cur.outputIndex, + content_index: cur.contentIndex, + part, + }); + break; + } + case "text_delta": { + if (!state.open || state.open.kind !== "message") break; + const cur: OpenMessage = state.open; + cur.currentPartText += ev.delta; + emit("response.output_text.delta", { + item_id: cur.itemId, + output_index: cur.outputIndex, + content_index: cur.contentIndex, + delta: ev.delta, + logprobs: [], + }); + break; + } + case "text_end": { + if (!state.open || state.open.kind !== "message") break; + const cur: OpenMessage = state.open; + const text = ev.content ?? cur.currentPartText; + emit("response.output_text.done", { + item_id: cur.itemId, + output_index: cur.outputIndex, + content_index: cur.contentIndex, + text, + logprobs: [], + }); + cur.content.push({ type: "output_text", text, annotations: [] }); + emit("response.content_part.done", { + item_id: cur.itemId, + output_index: cur.outputIndex, + content_index: cur.contentIndex, + part: { type: "output_text", text, annotations: [] }, + }); + cur.contentIndex += 1; + cur.currentPartText = ""; + break; + } + case "thinking_start": { + if (state.open) closeOpen(); + openReasoning(ev.partial, ev.contentIndex); + break; + } + case "thinking_delta": { + if (!state.open || state.open.kind !== "reasoning") break; + const cur: OpenReasoning = state.open; + cur.reasoningText += ev.delta; + emit("response.reasoning_text.delta", { + item_id: cur.itemId, + output_index: cur.outputIndex, + content_index: 0, + delta: ev.delta, + }); + break; + } + case "thinking_end": { + if (!state.open || state.open.kind !== "reasoning") break; + const cur: OpenReasoning = state.open; + const text = ev.content ?? cur.reasoningText; + cur.reasoningText = text; + emit("response.reasoning_text.done", { + item_id: cur.itemId, + output_index: cur.outputIndex, + content_index: 0, + text, + }); + closeOpen(); + break; + } + case "toolcall_start": { + if (state.open) closeOpen(); + openToolCall(ev.partial, ev.contentIndex); + break; + } + case "toolcall_delta": { + if (!state.open || state.open.kind !== "function_call") break; + const cur: OpenFunctionCall = state.open; + cur.argsText += ev.delta; + emit("response.function_call_arguments.delta", { + item_id: cur.itemId, + output_index: cur.outputIndex, + delta: ev.delta, + }); + break; + } + case "toolcall_end": { + if (!state.open || state.open.kind !== "function_call") break; + const cur: OpenFunctionCall = state.open; + // Finalize from the canonical ToolCall. arguments live as an object on the omp side; + // the wire wants the JSON string the model emitted, which streamed deltas accumulated. + const argsJson = cur.argsText || JSON.stringify(ev.toolCall.arguments ?? {}); + cur.argsText = argsJson; + cur.callId = ev.toolCall.id; + cur.name = ev.toolCall.name; + if (ev.toolCall.thoughtSignature) cur.itemId = ev.toolCall.thoughtSignature; + emit("response.function_call_arguments.done", { + item_id: cur.itemId, + output_index: cur.outputIndex, + arguments: argsJson, + name: cur.name, + }); + closeOpen(); + break; + } + case "done": { + finalMessage = ev.message; + break; + } + case "error": { + failureMessage = ev.error; + break; + } + } + } + + if (failureMessage) { + if (state.open) closeOpen(); + controller.enqueue( + encoder.encode( + sseEvent("response.failed", { + type: "response.failed", + sequence_number: seq(), + response: { + id: responseId, + object: "response", + created_at: createdAt, + status: "failed", + model: requestedModelId, + output: finishedItems, + error: { message: failureMessage.errorMessage ?? "stream failed" }, + }, + }), + ), + ); + emitDone(); + controller.close(); + return; + } + + if (state.open) closeOpen(); + const message = finalMessage ?? ((await events.result().catch(() => null)) as AssistantMessage | null); + + // Build the canonical output from the final message so non-streaming + // readers see the exact same shape they'd get from encodeResponse(). + const items = message ? buildOutputItems(message) : finishedItems; + const usage = message ? buildUsage(message) : null; + const status = message ? responseStatusForStopReason(message) : "completed"; + const terminalEvent = + status === "incomplete" + ? "response.incomplete" + : status === "failed" + ? "response.failed" + : "response.completed"; + controller.enqueue( + encoder.encode( + sseEvent(terminalEvent, { + type: terminalEvent, + sequence_number: seq(), + response: { + id: responseId, + object: "response", + created_at: createdAt, + status, + model: requestedModelId, + output: items, + usage, + ...(status === "incomplete" ? { incomplete_details: { reason: "max_output_tokens" } } : {}), + ...(status === "failed" + ? { error: { message: message?.errorMessage ?? "response failed" } } + : {}), + }, + }), + ), + ); + emitDone(); + controller.close(); + } catch (err) { + controller.enqueue( + encoder.encode( + sseEvent("response.failed", { + type: "response.failed", + sequence_number: seq(), + response: { + id: responseId, + object: "response", + created_at: Math.floor(Date.now() / 1000), + status: "failed", + model: requestedModelId, + output: [], + error: { message: err instanceof Error ? err.message : String(err) }, + }, + }), + ), + ); + emitDone(); + controller.close(); + } + }, + }); +} diff --git a/packages/ai/src/stream.ts b/packages/ai/src/stream.ts index b2f0d80e2..ced7ee158 100644 --- a/packages/ai/src/stream.ts +++ b/packages/ai/src/stream.ts @@ -176,6 +176,15 @@ export function getEnvApiKey(provider: string): string | undefined { return resolver?.(); } +/** + * Enumerate every provider that has an env-var fallback for `getEnvApiKey`. + * Used by `omp auth-broker migrate --include-env` to discover env-sourced keys + * that should be uploaded to the broker. + */ +export function listProvidersWithEnvKey(): string[] { + return Object.keys(serviceProviderMap); +} + export function stream( model: Model, context: Context, diff --git a/packages/ai/src/usage.ts b/packages/ai/src/usage.ts index 376e172e6..ab4c9e91d 100644 --- a/packages/ai/src/usage.ts +++ b/packages/ai/src/usage.ts @@ -4,8 +4,8 @@ * Provides a normalized schema to represent multiple limit windows, model tiers, * and shared quotas across providers. */ +import * as z from "zod/v4"; import type { Provider } from "./types"; - export type UsageUnit = "percent" | "tokens" | "requests" | "usd" | "minutes" | "bytes" | "unknown"; export type UsageStatus = "ok" | "warning" | "exhausted" | "unknown"; @@ -72,6 +72,58 @@ export interface UsageReport { raw?: unknown; } +// ─── Zod schemas (wire-shape validation for the broker `/v1/usage` endpoint) ─ + +export const usageUnitSchema = z.enum(["percent", "tokens", "requests", "usd", "minutes", "bytes", "unknown"]); +export const usageStatusSchema = z.enum(["ok", "warning", "exhausted", "unknown"]); + +export const usageWindowSchema = z.object({ + id: z.string(), + label: z.string(), + durationMs: z.number().optional(), + resetsAt: z.number().optional(), +}); + +export const usageAmountSchema = z.object({ + used: z.number().optional(), + limit: z.number().optional(), + remaining: z.number().optional(), + usedFraction: z.number().optional(), + remainingFraction: z.number().optional(), + unit: usageUnitSchema, +}); + +export const usageScopeSchema = z.object({ + provider: z.string(), + accountId: z.string().optional(), + projectId: z.string().optional(), + orgId: z.string().optional(), + modelId: z.string().optional(), + tier: z.string().optional(), + windowId: z.string().optional(), + shared: z.boolean().optional(), +}); + +export const usageLimitSchema = z.object({ + id: z.string(), + label: z.string(), + scope: usageScopeSchema, + window: usageWindowSchema.optional(), + amount: usageAmountSchema, + status: usageStatusSchema.optional(), + notes: z.array(z.string()).optional(), +}); + +export const usageReportSchema = z.object({ + provider: z.string(), + fetchedAt: z.number(), + limits: z.array(usageLimitSchema), + metadata: z.record(z.string(), z.unknown()).optional(), + // `raw` is provider-specific and may be anything; the broker strips it before + // sending the report over the wire, so accept-but-ignore here. + raw: z.unknown().optional(), +}); + /** Optional logger for usage fetchers. */ export interface UsageLogger { debug(message: string, meta?: Record): void; diff --git a/packages/ai/src/usage/claude.ts b/packages/ai/src/usage/claude.ts index c4d0fe5e8..e4331c31e 100644 --- a/packages/ai/src/usage/claude.ts +++ b/packages/ai/src/usage/claude.ts @@ -1,3 +1,4 @@ +import { scheduler } from "node:timers/promises"; import type { CredentialRankingStrategy, UsageAmount, @@ -14,7 +15,7 @@ import { isRecord, toNumber } from "../utils"; const DEFAULT_ENDPOINT = "https://api.anthropic.com/api/oauth"; const FIVE_HOURS_MS = 5 * 60 * 60 * 1000; const SEVEN_DAYS_MS = 7 * 24 * 60 * 60 * 1000; -const MAX_RETRIES = 3; +const MAX_ATTEMPTS = 3; const BASE_RETRY_DELAY_MS = 500; const CLAUDE_HEADERS = { @@ -90,6 +91,11 @@ function getPayloadString(payload: Record, key: string): string return typeof value === "string" && value.trim() ? value.trim() : undefined; } +function getNestedPayloadString(payload: Record, key: string, nestedKey: string): string | undefined { + const nested = payload[key]; + return isRecord(nested) ? getPayloadString(nested, nestedKey) : undefined; +} + function extractUsageIdentity(payload: ClaudeUsageResponse, orgId?: string): { accountId?: string; email?: string } { if (!isRecord(payload)) return { accountId: orgId }; const accountId = @@ -99,16 +105,60 @@ function extractUsageIdentity(payload: ClaudeUsageResponse, orgId?: string): { a getPayloadString(payload, "userId") ?? getPayloadString(payload, "org_id") ?? getPayloadString(payload, "orgId") ?? + getNestedPayloadString(payload, "account", "uuid") ?? + getNestedPayloadString(payload, "account", "id") ?? + getNestedPayloadString(payload, "organization", "uuid") ?? + getNestedPayloadString(payload, "organization", "id") ?? + getNestedPayloadString(payload, "user", "uuid") ?? + getNestedPayloadString(payload, "user", "id") ?? orgId; const email = getPayloadString(payload, "email") ?? getPayloadString(payload, "user_email") ?? - getPayloadString(payload, "userEmail"); + getPayloadString(payload, "userEmail") ?? + getNestedPayloadString(payload, "account", "email") ?? + getNestedPayloadString(payload, "user", "email"); return { accountId, email }; } function hasUsageData(payload: ClaudeUsageResponse): boolean { - return Boolean(payload.five_hour || payload.seven_day || payload.seven_day_opus || payload.seven_day_sonnet); + return ( + parseBucket(payload.five_hour)?.utilization !== undefined || + parseBucket(payload.seven_day)?.utilization !== undefined || + parseBucket(payload.seven_day_opus)?.utilization !== undefined || + parseBucket(payload.seven_day_sonnet)?.utilization !== undefined + ); +} + +function isRetryableStatus(status: number): boolean { + return status === 429 || (status >= 500 && status < 600); +} + +function isAbortError(error: unknown, signal?: AbortSignal): boolean { + if (signal?.aborted) return true; + if (!isRecord(error)) return false; + return error.name === "AbortError" || error.name === "TimeoutError"; +} + +function retryDelayMs(attempt: number, retryAfter: string | null): number { + const baseline = BASE_RETRY_DELAY_MS * 2 ** attempt; + if (!retryAfter?.trim()) return baseline; + const seconds = Number.parseFloat(retryAfter); + if (Number.isFinite(seconds)) return Math.max(baseline, Math.max(0, seconds * 1000)); + const dateDelay = Date.parse(retryAfter) - Date.now(); + return Number.isFinite(dateDelay) ? Math.max(baseline, Math.max(0, dateDelay)) : baseline; +} + +async function waitBeforeRetry(attempt: number, retryAfter: string | null, signal?: AbortSignal): Promise { + if (signal?.aborted) return false; + if (attempt >= MAX_ATTEMPTS - 1) return false; + try { + await scheduler.wait(retryDelayMs(attempt, retryAfter), { signal }); + return !signal?.aborted; + } catch (error) { + if (isAbortError(error, signal)) return false; + throw error; + } } async function fetchUsagePayload( @@ -117,29 +167,49 @@ async function fetchUsagePayload( ctx: UsageFetchContext, signal?: AbortSignal, ): Promise { + if (signal?.aborted) return null; + let lastPayload: ClaudeUsageResponse | null = null; let lastOrgId: string | undefined; - for (let attempt = 0; attempt < MAX_RETRIES; attempt++) { + for (let attempt = 0; attempt < MAX_ATTEMPTS; attempt++) { try { const response = await ctx.fetch(url, { headers, signal }); - if (!response.ok) { - ctx.logger?.warn("Claude usage fetch failed", { status: response.status, statusText: response.statusText }); - return null; - } - const payload = (await response.json()) as ClaudeUsageResponse; - lastPayload = payload; const orgId = response.headers.get("anthropic-organization-id")?.trim() || undefined; lastOrgId = orgId ?? lastOrgId; - if (payload && isRecord(payload) && hasUsageData(payload)) { - return { payload, orgId }; - } - } catch (error) { - ctx.logger?.warn("Claude usage fetch error", { error: String(error) }); - return null; - } - if (attempt < MAX_RETRIES - 1) { - await Bun.sleep(BASE_RETRY_DELAY_MS * 2 ** attempt); + if (!response.ok) { + const retryable = isRetryableStatus(response.status); + ctx.logger?.warn("Claude usage fetch failed", { + status: response.status, + statusText: response.statusText, + attempt, + willRetry: retryable && attempt < MAX_ATTEMPTS - 1, + }); + if (!retryable) return null; + if (!(await waitBeforeRetry(attempt, response.headers.get("retry-after"), signal))) break; + continue; + } + + const parsed = (await response.json()) as unknown; + if (isRecord(parsed)) { + const payload = parsed as ClaudeUsageResponse; + lastPayload = payload; + if (hasUsageData(payload)) return { payload, orgId }; + } + + ctx.logger?.warn("Claude usage response missing usage data", { + attempt, + willRetry: attempt < MAX_ATTEMPTS - 1, + }); + if (!(await waitBeforeRetry(attempt, null, signal))) break; + } catch (error) { + if (isAbortError(error, signal)) return null; + ctx.logger?.warn("Claude usage fetch error", { + error: String(error), + attempt, + willRetry: attempt < MAX_ATTEMPTS - 1, + }); + if (!(await waitBeforeRetry(attempt, null, signal))) break; } } @@ -147,40 +217,47 @@ async function fetchUsagePayload( } interface ClaudeProfile { + uuid?: string; + email?: string; account?: { uuid?: string; email?: string; }; } +function extractProfileIdentity(profile: ClaudeProfile | null): { accountId?: string; email?: string } { + if (!profile || !isRecord(profile)) return {}; + const account = isRecord(profile.account) ? profile.account : undefined; + return { + accountId: + (typeof profile.uuid === "string" && profile.uuid.trim() ? profile.uuid.trim() : undefined) ?? + (typeof account?.uuid === "string" && account.uuid.trim() ? account.uuid.trim() : undefined), + email: + (typeof profile.email === "string" && profile.email.trim() ? profile.email.trim() : undefined) ?? + (typeof account?.email === "string" && account.email.trim() ? account.email.trim() : undefined), + }; +} + async function fetchProfile( baseUrl: string, headers: Record, ctx: UsageFetchContext, signal?: AbortSignal, ): Promise { + if (signal?.aborted) return null; const url = `${baseUrl}/profile`; try { const response = await ctx.fetch(url, { headers, signal }); if (!response.ok) return null; - return (await response.json()) as ClaudeProfile; - } catch { + const payload = (await response.json()) as unknown; + return isRecord(payload) ? (payload as ClaudeProfile) : null; + } catch (error) { + if (isAbortError(error, signal)) return null; + ctx.logger?.debug("Claude profile fetch error", { error: String(error) }); return null; } } -async function resolveEmail( - params: UsageFetchParams, - ctx: UsageFetchContext, - baseUrl: string, - headers: Record, -): Promise { - if (params.credential.email) return params.credential.email; - - const profile = await fetchProfile(baseUrl, headers, ctx, params.signal); - return profile?.account?.email; -} - function buildUsageAmount(utilization: number | undefined): UsageAmount | undefined { if (utilization === undefined) return undefined; const clamped = Math.min(Math.max(utilization, 0), 100); @@ -303,17 +380,23 @@ async function fetchClaudeUsage(params: UsageFetchParams, ctx: UsageFetchContext if (limits.length === 0) return null; const identity = extractUsageIdentity(payload, orgId); - const accountId = identity.accountId ?? credential.accountId; - const email = identity.email ?? (await resolveEmail(params, ctx, baseUrl, headers)); + let accountId = identity.accountId ?? credential.accountId; + let email = identity.email ?? credential.email; + if ((!accountId || !email) && !params.signal?.aborted) { + const profileIdentity = extractProfileIdentity(await fetchProfile(baseUrl, headers, ctx, params.signal)); + accountId = accountId ?? profileIdentity.accountId; + email = email ?? profileIdentity.email; + } const report: UsageReport = { provider: params.provider, fetchedAt: Date.now(), limits, metadata: { - accountId, - email, endpoint: url, + ...(accountId ? { accountId } : {}), + ...(email ? { email } : {}), + ...(orgId ? { orgId } : {}), }, raw: payload, }; diff --git a/packages/ai/src/usage/openai-codex.ts b/packages/ai/src/usage/openai-codex.ts index a42cdc90e..b4ff37f3b 100644 --- a/packages/ai/src/usage/openai-codex.ts +++ b/packages/ai/src/usage/openai-codex.ts @@ -31,9 +31,16 @@ interface CodexUsageRateLimitPayload { secondary_window?: CodexUsageWindowPayload | null; } +interface CodexUsageAdditionalRateLimitPayload { + limit_name?: string; + metered_feature?: string; + rate_limit?: CodexUsageRateLimitPayload | null; +} + interface CodexUsagePayload { plan_type?: string; rate_limit?: CodexUsageRateLimitPayload | null; + additional_rate_limits?: CodexUsageAdditionalRateLimitPayload[] | null; } interface ParsedUsageWindow { @@ -43,12 +50,22 @@ interface ParsedUsageWindow { resetAt?: number; } +interface ParsedAdditionalUsage { + limitName?: string; + meteredFeature?: string; + allowed?: boolean; + limitReached?: boolean; + primary?: ParsedUsageWindow; + secondary?: ParsedUsageWindow; +} + interface ParsedUsage { planType?: string; allowed?: boolean; limitReached?: boolean; primary?: ParsedUsageWindow; secondary?: ParsedUsageWindow; + additional: ParsedAdditionalUsage[]; raw: CodexUsagePayload; } @@ -124,20 +141,45 @@ function parseUsageWindow(payload: unknown): ParsedUsageWindow | undefined { }; } +function parseAdditionalRateLimit(payload: unknown): ParsedAdditionalUsage | null { + if (!isRecord(payload)) return null; + const limitName = typeof payload.limit_name === "string" ? payload.limit_name : undefined; + const meteredFeature = typeof payload.metered_feature === "string" ? payload.metered_feature : undefined; + const rateLimit = isRecord(payload.rate_limit) ? payload.rate_limit : undefined; + if (!rateLimit) return null; + const primary = parseUsageWindow(rateLimit.primary_window); + const secondary = parseUsageWindow(rateLimit.secondary_window); + const allowed = toBoolean(rateLimit.allowed); + const limitReached = toBoolean(rateLimit.limit_reached); + if (!primary && !secondary && allowed === undefined && limitReached === undefined) return null; + return { limitName, meteredFeature, allowed, limitReached, primary, secondary }; +} + function parseUsagePayload(payload: unknown): ParsedUsage | null { if (!isRecord(payload)) return null; const planType = typeof payload.plan_type === "string" ? payload.plan_type : undefined; const rateLimit = isRecord(payload.rate_limit) ? payload.rate_limit : undefined; - if (!rateLimit) return null; + const additionalRaw = Array.isArray(payload.additional_rate_limits) ? payload.additional_rate_limits : []; + const additional = additionalRaw + .map(parseAdditionalRateLimit) + .filter((value): value is ParsedAdditionalUsage => value !== null); + if (!rateLimit && additional.length === 0) return null; const parsed: ParsedUsage = { planType, - allowed: toBoolean(rateLimit.allowed), - limitReached: toBoolean(rateLimit.limit_reached), - primary: parseUsageWindow(rateLimit.primary_window), - secondary: parseUsageWindow(rateLimit.secondary_window), + allowed: rateLimit ? toBoolean(rateLimit.allowed) : undefined, + limitReached: rateLimit ? toBoolean(rateLimit.limit_reached) : undefined, + primary: rateLimit ? parseUsageWindow(rateLimit.primary_window) : undefined, + secondary: rateLimit ? parseUsageWindow(rateLimit.secondary_window) : undefined, + additional, raw: payload as CodexUsagePayload, }; - if (!parsed.primary && !parsed.secondary && parsed.allowed === undefined && parsed.limitReached === undefined) { + if ( + !parsed.primary && + !parsed.secondary && + parsed.allowed === undefined && + parsed.limitReached === undefined && + parsed.additional.length === 0 + ) { return null; } return parsed; @@ -251,6 +293,56 @@ function buildUsageLimit(args: { status: buildUsageStatus(amount.usedFraction, args.limitReached), }; } +function additionalLimitSlug(args: { limitName?: string; meteredFeature?: string }): string { + const probe = `${args.limitName ?? ""} ${args.meteredFeature ?? ""}`.toLowerCase(); + if (probe.includes("spark") || probe.includes("bengalfox")) return "spark"; + const source = (args.meteredFeature ?? args.limitName ?? "extra").toLowerCase(); + return ( + source + .replace(/^codex[-_]/, "") + .replace(/[^a-z0-9]+/g, "-") + .replace(/^-+|-+$/g, "") || "extra" + ); +} + +function additionalDisplayName(slug: string, limitName?: string): string { + if (slug === "spark") return "Spark"; + if (limitName) return limitName; + return slug.replace( + /(^|-)([a-z])/g, + (_match, sep: string, ch: string) => `${sep === "-" ? " " : ""}${ch.toUpperCase()}`, + ); +} + +function buildAdditionalUsageLimit(args: { + key: "primary" | "secondary"; + slug: string; + displayName: string; + window: ParsedUsageWindow; + accountId?: string; + limitReached?: boolean; + limitName?: string; + meteredFeature?: string; + nowMs: number; +}): UsageLimit { + const usageWindow = buildUsageWindow(args.window, args.key, args.nowMs); + const amount = buildUsageAmount(args.window); + return { + id: `openai-codex:${args.slug}:${args.key}`, + label: `${usageWindow.label} (${args.displayName})`, + scope: { + provider: "openai-codex", + accountId: args.accountId, + tier: args.slug, + modelId: args.limitName, + windowId: usageWindow.id, + shared: true, + }, + window: usageWindow, + amount, + status: buildUsageStatus(amount.usedFraction, args.limitReached), + }; +} export const openaiCodexUsageProvider: UsageProvider = { id: "openai-codex", @@ -327,6 +419,40 @@ export const openaiCodexUsageProvider: UsageProvider = { }), ); } + for (const extra of parsed?.additional ?? []) { + const slug = additionalLimitSlug({ limitName: extra.limitName, meteredFeature: extra.meteredFeature }); + const displayName = additionalDisplayName(slug, extra.limitName); + if (extra.primary) { + limits.push( + buildAdditionalUsageLimit({ + key: "primary", + slug, + displayName, + window: extra.primary, + accountId, + limitReached: extra.limitReached, + limitName: extra.limitName, + meteredFeature: extra.meteredFeature, + nowMs, + }), + ); + } + if (extra.secondary) { + limits.push( + buildAdditionalUsageLimit({ + key: "secondary", + slug, + displayName, + window: extra.secondary, + accountId, + limitReached: extra.limitReached, + limitName: extra.limitName, + meteredFeature: extra.meteredFeature, + nowMs, + }), + ); + } + } const report: UsageReport = { provider: "openai-codex", diff --git a/packages/ai/test/auth-gateway-anthropic-messages.test.ts b/packages/ai/test/auth-gateway-anthropic-messages.test.ts new file mode 100644 index 000000000..1a9834d55 --- /dev/null +++ b/packages/ai/test/auth-gateway-anthropic-messages.test.ts @@ -0,0 +1,465 @@ +import { describe, expect, it } from "bun:test"; +import { encodeResponse, encodeStream, parseRequest } from "../src/providers/anthropic-messages-server"; +import type { AssistantMessage, AssistantMessageEvent, ToolResultMessage } from "../src/types"; +import { AssistantMessageEventStream } from "../src/utils/event-stream"; + +function emptyUsage(): AssistantMessage["usage"] { + return { + input: 0, + output: 0, + cacheRead: 0, + cacheWrite: 0, + totalTokens: 0, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, + }; +} + +function makeStream(events: AssistantMessageEvent[]): AssistantMessageEventStream { + const s = new AssistantMessageEventStream(); + queueMicrotask(() => { + for (const ev of events) s.push(ev); + s.end(); + }); + return s; +} + +interface SseEvent { + event: string; + data: Record; +} + +async function collectSse(stream: ReadableStream): Promise { + const reader = stream.getReader(); + const decoder = new TextDecoder(); + let buf = ""; + const out: SseEvent[] = []; + while (true) { + const { value, done } = await reader.read(); + if (done) break; + buf += decoder.decode(value, { stream: true }); + } + buf += decoder.decode(); + for (const chunk of buf.split("\n\n")) { + if (!chunk.trim()) continue; + let event = ""; + let dataLine = ""; + for (const line of chunk.split("\n")) { + if (line.startsWith("event: ")) event = line.slice(7); + else if (line.startsWith("data: ")) dataLine = line.slice(6); + } + out.push({ event, data: JSON.parse(dataLine) as Record }); + } + return out; +} + +describe("anthropic-messages parseRequest", () => { + it("parses system + user + assistant(thinking,text,tool_use) + tool_result", () => { + const parsed = parseRequest({ + model: "claude-opus-4-7", + max_tokens: 1024, + temperature: 0.2, + top_p: 0.9, + stop_sequences: ["\n\n"], + tool_choice: { type: "any" }, + thinking: { type: "enabled", budget_tokens: 2048 }, + system: [ + { type: "text", text: "You are X" }, + { type: "text", text: "Be brief." }, + ], + tools: [ + { + name: "lookup", + description: "find a thing", + input_schema: { type: "object", properties: { q: { type: "string" } }, required: ["q"] }, + }, + ], + messages: [ + { role: "user", content: "hi" }, + { + role: "assistant", + content: [ + { type: "thinking", thinking: "hmm", signature: "sig-1" }, + { type: "redacted_thinking", data: "REDACTED" }, + { type: "text", text: "calling tool" }, + { type: "tool_use", id: "toolu_abc", name: "lookup", input: { q: "x" } }, + ], + }, + { + role: "user", + content: [ + { + type: "tool_result", + tool_use_id: "toolu_abc", + content: [{ type: "text", text: "result text" }], + is_error: false, + }, + ], + }, + { + role: "user", + content: [ + { + type: "tool_result", + tool_use_id: "toolu_def", + content: "string body", + is_error: true, + }, + { type: "text", text: "and another result coming" }, + ], + }, + ], + }); + + expect(parsed.modelId).toBe("claude-opus-4-7"); + expect(parsed.stream).toBe(false); + expect(parsed.context.systemPrompt).toEqual(["You are X\n\nBe brief."]); + expect(parsed.options.maxOutputTokens).toBe(1024); + expect(parsed.options.temperature).toBe(0.2); + expect(parsed.options.topP).toBe(0.9); + expect(parsed.options.stopSequences).toEqual(["\n\n"]); + expect(parsed.options.toolChoice).toBe("required"); + expect(parsed.options.thinkingBudget).toBe(2048); + expect(parsed.options.extra).toBeUndefined(); + + expect(parsed.context.tools).toHaveLength(1); + const tool = parsed.context.tools![0]!; + expect(tool.name).toBe("lookup"); + expect(tool.description).toBe("find a thing"); + expect(tool.parameters).toEqual({ + type: "object", + properties: { q: { type: "string" } }, + required: ["q"], + }); + + // messages: user("hi"), assistant(4 blocks), toolResult(toolu_abc), + // toolResult(toolu_def), user("and another result coming") + const msgs = parsed.context.messages; + expect(msgs).toHaveLength(5); + + expect(msgs[0]).toMatchObject({ role: "user", content: "hi" }); + + const asst = msgs[1]; + expect(asst.role).toBe("assistant"); + if (asst.role !== "assistant") throw new Error(); + expect(asst.content).toEqual([ + { type: "thinking", thinking: "hmm", thinkingSignature: "sig-1" }, + { type: "redactedThinking", data: "REDACTED" }, + { type: "text", text: "calling tool" }, + { type: "toolCall", id: "toolu_abc", name: "lookup", arguments: { q: "x" } }, + ]); + expect(asst.api).toBe("anthropic-messages"); + expect(asst.provider).toBe("anthropic"); + expect(asst.model).toBe("claude-opus-4-7"); + + const tr1 = msgs[2] as ToolResultMessage; + expect(tr1.role).toBe("toolResult"); + expect(tr1.toolCallId).toBe("toolu_abc"); + expect(tr1.isError).toBe(false); + expect(tr1.content).toEqual([{ type: "text", text: "result text" }]); + + const tr2 = msgs[3] as ToolResultMessage; + expect(tr2.role).toBe("toolResult"); + expect(tr2.toolCallId).toBe("toolu_def"); + expect(tr2.isError).toBe(true); + expect(tr2.content).toEqual([{ type: "text", text: "string body" }]); + + expect(msgs[4]).toMatchObject({ role: "user", content: "and another result coming" }); + }); + + it("maps tool_choice variants and suppresses user wrappers that hold only tool_result", () => { + const auto = parseRequest({ + model: "m", + max_tokens: 8, + tool_choice: { type: "auto" }, + messages: [{ role: "user", content: "hi" }], + }); + expect(auto.options.toolChoice).toBe("auto"); + + const named = parseRequest({ + model: "m", + max_tokens: 8, + tool_choice: { type: "tool", name: "lookup" }, + messages: [{ role: "user", content: "hi" }], + }); + expect(named.options.toolChoice).toEqual({ name: "lookup" }); + + const onlyResult = parseRequest({ + model: "m", + max_tokens: 8, + messages: [ + { + role: "user", + content: [{ type: "tool_result", tool_use_id: "t1", content: [{ type: "text", text: "ok" }] }], + }, + ], + }); + // no user wrapper, just the toolResult + expect(onlyResult.context.messages).toHaveLength(1); + expect(onlyResult.context.messages[0]!.role).toBe("toolResult"); + }); + + it("rejects ambiguous user text before tool_result blocks", () => { + expect(() => + parseRequest({ + model: "m", + max_tokens: 8, + messages: [ + { + role: "user", + content: [ + { type: "text", text: "this would replay before the tool result" }, + { type: "tool_result", tool_use_id: "t1", content: "ok" }, + ], + }, + ], + }), + ).toThrow(/tool_result/i); + }); + + it("rejects missing required fields and unsupported request controls", () => { + expect(() => parseRequest({})).toThrow(/model/); + expect(() => parseRequest({ model: "m", messages: [] })).toThrow(/max_tokens/); + expect(() => parseRequest({ model: "m", max_tokens: 1 })).toThrow(/messages/); + const topK = parseRequest({ model: "m", max_tokens: 1, messages: [{ role: "user", content: "hi" }], top_k: 50 }); + expect(topK.options.topK).toBe(50); + // `metadata` is tolerated permissively now (Anthropic clients ship it + // by default with `user_id`); it should parse without throwing and + // surface nothing on the parsed options. + const withMetadata = parseRequest({ + model: "m", + max_tokens: 1, + messages: [{ role: "user", content: "hi" }], + metadata: { user_id: "u_1" }, + }); + expect(withMetadata.options.extra).toBeUndefined(); + }); +}); + +describe("anthropic-messages encodeResponse", () => { + it("encodes text + thinking + tool_use with correct ordering and stop_reason mapping", () => { + const message: AssistantMessage = { + role: "assistant", + content: [ + { type: "thinking", thinking: "let me think", thinkingSignature: "sig-xyz" }, + { type: "text", text: "calling tool now" }, + { type: "toolCall", id: "toolu_999", name: "lookup", arguments: { q: "hello" } }, + ], + api: "anthropic-messages", + provider: "anthropic", + model: "claude-opus-4-7", + usage: { ...emptyUsage(), input: 12, output: 34, cacheRead: 5, cacheWrite: 7, totalTokens: 58 }, + stopReason: "toolUse", + timestamp: 1000, + }; + const encoded = encodeResponse(message, "claude-opus-4-7"); + expect(encoded.type).toBe("message"); + expect(encoded.role).toBe("assistant"); + expect(encoded.model).toBe("claude-opus-4-7"); + expect(encoded.stop_reason).toBe("tool_use"); + expect(encoded.stop_sequence).toBeNull(); + expect(encoded.usage).toEqual({ + input_tokens: 12, + output_tokens: 34, + cache_read_input_tokens: 5, + cache_creation_input_tokens: 7, + }); + expect(encoded.content).toEqual([ + { type: "thinking", thinking: "let me think", signature: "sig-xyz" }, + { type: "text", text: "calling tool now" }, + { type: "tool_use", id: "toolu_999", name: "lookup", input: { q: "hello" } }, + ]); + expect(typeof encoded.id).toBe("string"); + expect((encoded.id as string).startsWith("msg_")).toBe(true); + }); + + it("maps stop reasons and rejects upstream terminal errors", () => { + const base: AssistantMessage = { + role: "assistant", + content: [], + api: "anthropic-messages", + provider: "anthropic", + model: "m", + usage: emptyUsage(), + stopReason: "stop", + timestamp: 0, + }; + expect(encodeResponse({ ...base, stopReason: "stop" }, "m").stop_reason).toBe("end_turn"); + expect(encodeResponse({ ...base, stopReason: "length" }, "m").stop_reason).toBe("max_tokens"); + expect(encodeResponse({ ...base, stopReason: "toolUse" }, "m").stop_reason).toBe("tool_use"); + expect(() => encodeResponse({ ...base, stopReason: "error", errorMessage: "upstream failed" }, "m")).toThrow( + /upstream failed/, + ); + expect(() => encodeResponse({ ...base, stopReason: "aborted", errorMessage: "request aborted" }, "m")).toThrow( + /request aborted/, + ); + }); +}); + +describe("anthropic-messages encodeStream", () => { + it("emits thinking_delta + signature_delta + text_delta + tool_use input_json_delta + message_stop", async () => { + const finalMessage: AssistantMessage = { + role: "assistant", + content: [ + { type: "thinking", thinking: "thoughts", thinkingSignature: "SIG" }, + { type: "text", text: "hi there" }, + { type: "toolCall", id: "toolu_1", name: "go", arguments: { x: 1 } }, + ], + api: "anthropic-messages", + provider: "anthropic", + model: "claude-opus-4-7", + usage: { ...emptyUsage(), input: 11, output: 42, cacheRead: 3, cacheWrite: 5 }, + stopReason: "toolUse", + timestamp: 0, + }; + + const partialAfterThinkingEnd: AssistantMessage = { + ...finalMessage, + content: [{ type: "thinking", thinking: "thoughts", thinkingSignature: "SIG" }], + }; + const partialAtToolStart: AssistantMessage = { + ...finalMessage, + content: [ + { type: "thinking", thinking: "thoughts", thinkingSignature: "SIG" }, + { type: "text", text: "hi there" }, + { type: "toolCall", id: "toolu_1", name: "go", arguments: {} }, + ], + }; + + const events: AssistantMessageEvent[] = [ + { type: "start", partial: finalMessage }, + { type: "thinking_start", contentIndex: 0, partial: finalMessage }, + { type: "thinking_delta", contentIndex: 0, delta: "thoughts", partial: finalMessage }, + { type: "thinking_end", contentIndex: 0, content: "thoughts", partial: partialAfterThinkingEnd }, + { type: "text_start", contentIndex: 1, partial: finalMessage }, + { type: "text_delta", contentIndex: 1, delta: "hi ", partial: finalMessage }, + { type: "text_delta", contentIndex: 1, delta: "there", partial: finalMessage }, + { type: "text_end", contentIndex: 1, content: "hi there", partial: finalMessage }, + { type: "toolcall_start", contentIndex: 2, partial: partialAtToolStart }, + { type: "toolcall_delta", contentIndex: 2, delta: '{"x":', partial: partialAtToolStart }, + { type: "toolcall_delta", contentIndex: 2, delta: "1}", partial: partialAtToolStart }, + { + type: "toolcall_end", + contentIndex: 2, + toolCall: { type: "toolCall", id: "toolu_1", name: "go", arguments: { x: 1 } }, + partial: finalMessage, + }, + { type: "done", reason: "toolUse", message: finalMessage }, + ]; + + const sse = await collectSse(encodeStream(makeStream(events), "claude-opus-4-7")); + + // Sequence check + const types = sse.map(e => e.event); + expect(types).toEqual([ + "message_start", + "content_block_start", + "content_block_delta", + "content_block_delta", // signature_delta + "content_block_stop", + "content_block_start", + "content_block_delta", + "content_block_delta", + "content_block_stop", + "content_block_start", + "content_block_delta", + "content_block_delta", + "content_block_stop", + "message_delta", + "message_stop", + ]); + + // message_start payload + const start = sse[0]!.data as { + type: string; + message: { id: string; model: string; role: string; usage: Record }; + }; + expect(start.type).toBe("message_start"); + expect(start.message.model).toBe("claude-opus-4-7"); + expect(start.message.role).toBe("assistant"); + expect(start.message.id.startsWith("msg_")).toBe(true); + expect(start.message.usage).toEqual({ + input_tokens: 11, + output_tokens: 42, + cache_read_input_tokens: 3, + cache_creation_input_tokens: 5, + }); + + // thinking block_start + expect(sse[1]!.data).toEqual({ + type: "content_block_start", + index: 0, + content_block: { type: "thinking", thinking: "" }, + }); + expect(sse[2]!.data).toEqual({ + type: "content_block_delta", + index: 0, + delta: { type: "thinking_delta", thinking: "thoughts" }, + }); + expect(sse[3]!.data).toEqual({ + type: "content_block_delta", + index: 0, + delta: { type: "signature_delta", signature: "SIG" }, + }); + expect(sse[4]!.data).toEqual({ type: "content_block_stop", index: 0 }); + + // text block + expect(sse[5]!.data).toEqual({ + type: "content_block_start", + index: 1, + content_block: { type: "text", text: "" }, + }); + expect(sse[6]!.data).toEqual({ + type: "content_block_delta", + index: 1, + delta: { type: "text_delta", text: "hi " }, + }); + + // tool_use block + expect(sse[9]!.data).toEqual({ + type: "content_block_start", + index: 2, + content_block: { type: "tool_use", id: "toolu_1", name: "go", input: {} }, + }); + expect(sse[10]!.data).toEqual({ + type: "content_block_delta", + index: 2, + delta: { type: "input_json_delta", partial_json: '{"x":' }, + }); + + // message_delta with mapped stop_reason + expect(sse[13]!.data).toEqual({ + type: "message_delta", + delta: { stop_reason: "tool_use", stop_sequence: null }, + usage: { + input_tokens: 11, + output_tokens: 42, + cache_read_input_tokens: 3, + cache_creation_input_tokens: 5, + }, + }); + + expect(sse[14]!.data).toEqual({ type: "message_stop" }); + }); + + it("emits an error event when the upstream stream errors", async () => { + const errMessage: AssistantMessage = { + role: "assistant", + content: [], + api: "anthropic-messages", + provider: "anthropic", + model: "m", + usage: emptyUsage(), + stopReason: "error", + errorMessage: "boom", + timestamp: 0, + }; + const events: AssistantMessageEvent[] = [ + { type: "start", partial: errMessage }, + { type: "error", reason: "error", error: errMessage }, + ]; + const sse = await collectSse(encodeStream(makeStream(events), "m")); + const last = sse.at(-1)!; + expect(last.event).toBe("error"); + expect(last.data).toEqual({ type: "error", error: { type: "api_error", message: "boom" } }); + }); +}); diff --git a/packages/ai/test/auth-gateway-openai-chat.test.ts b/packages/ai/test/auth-gateway-openai-chat.test.ts new file mode 100644 index 000000000..ef138343f --- /dev/null +++ b/packages/ai/test/auth-gateway-openai-chat.test.ts @@ -0,0 +1,297 @@ +import { describe, expect, it } from "bun:test"; +import { encodeResponse, encodeStream, parseRequest } from "../src/providers/openai-chat-server"; +import type { AssistantMessage, AssistantMessageEvent, AssistantMessageEventStream } from "../src/types"; + +function makeEventStream(events: AssistantMessageEvent[], final: AssistantMessage): AssistantMessageEventStream { + async function* iter() { + for (const e of events) yield e; + } + const stream = iter() as unknown as AssistantMessageEventStream; + (stream as { result(): Promise }).result = async () => final; + return stream; +} + +async function collectStream(stream: ReadableStream): Promise { + const reader = stream.getReader(); + const decoder = new TextDecoder(); + let buf = ""; + for (;;) { + const { value, done } = await reader.read(); + if (done) break; + buf += decoder.decode(value, { stream: true }); + } + buf += decoder.decode(); + return buf.split("\n\n").filter(s => s.length > 0); +} + +function parseSseLine(line: string): unknown { + const stripped = line.replace(/^data: /, ""); + if (stripped === "[DONE]") return "[DONE]"; + return JSON.parse(stripped); +} + +const baseUsage = { + input: 0, + output: 0, + cacheRead: 0, + cacheWrite: 0, + totalTokens: 0, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, +}; + +function emptyAssistant(): AssistantMessage { + return { + role: "assistant", + content: [], + api: "openai-completions", + provider: "openai", + model: "gpt-test", + usage: baseUsage, + stopReason: "stop", + timestamp: 0, + }; +} + +describe("auth-gateway openai-chat: parseRequest", () => { + it("converts a full request into a Context", () => { + const parsed = parseRequest({ + model: "gpt-5.2", + messages: [ + { role: "system", content: "you are X" }, + { role: "system", content: "also Y" }, + { role: "user", content: "hi" }, + { + role: "assistant", + content: "hello", + tool_calls: [ + { + id: "call_1", + type: "function", + function: { name: "lookup", arguments: '{"q":"a"}' }, + }, + { + id: "call_2", + type: "function", + function: { name: "broken", arguments: "not-json" }, + }, + ], + }, + { role: "tool", tool_call_id: "call_1", content: "result-text" }, + ], + tools: [ + { + type: "function", + function: { + name: "lookup", + description: "look something up", + parameters: { type: "object", properties: { q: { type: "string" } } }, + }, + }, + ], + stream: true, + max_tokens: 512, + max_completion_tokens: 1024, + temperature: 0.2, + top_p: 0.9, + stop: ["\n\n"], + tool_choice: { type: "function", function: { name: "lookup" } }, + response_format: { type: "json_object" }, + stream_options: { include_usage: true }, + }); + + expect(parsed.modelId).toBe("gpt-5.2"); + expect(parsed.stream).toBe(true); + expect(parsed.context.systemPrompt).toEqual(["you are X\n\nalso Y"]); + expect(parsed.context.messages).toHaveLength(3); + + const [user, assistant, tool] = parsed.context.messages; + expect(user.role).toBe("user"); + expect(assistant.role).toBe("assistant"); + if (assistant.role !== "assistant") throw new Error("unreachable"); + expect(assistant.api).toBe("openai-completions"); + expect(assistant.provider).toBe("openai"); + expect(assistant.model).toBe("gpt-5.2"); + expect(assistant.content[0]).toEqual({ type: "text", text: "hello" }); + const call1 = assistant.content[1]; + const call2 = assistant.content[2]; + if (call1.type !== "toolCall" || call2.type !== "toolCall") throw new Error("unreachable"); + expect(call1.id).toBe("call_1"); + expect(call1.name).toBe("lookup"); + expect(call1.arguments).toEqual({ q: "a" }); + // Un-parseable args fall back to __raw passthrough. + expect(call2.arguments).toEqual({ __raw: "not-json" }); + + expect(tool.role).toBe("toolResult"); + if (tool.role !== "toolResult") throw new Error("unreachable"); + expect(tool.toolCallId).toBe("call_1"); + expect(tool.toolName).toBe(""); + expect(tool.content).toEqual([{ type: "text", text: "result-text" }]); + + expect(parsed.context.tools).toHaveLength(1); + expect(parsed.context.tools?.[0].name).toBe("lookup"); + + // max_completion_tokens wins over max_tokens. + expect(parsed.options.maxOutputTokens).toBe(1024); + expect(parsed.options.temperature).toBe(0.2); + expect(parsed.options.topP).toBe(0.9); + expect(parsed.options.stopSequences).toEqual(["\n\n"]); + expect(parsed.options.toolChoice).toEqual({ name: "lookup" }); + expect(parsed.options.extra).toEqual({ + response_format: { type: "json_object" }, + includeStreamingUsage: true, + }); + }); + + it("rejects missing required fields", () => { + expect(() => parseRequest({ messages: [] })).toThrow(/model/); + expect(() => parseRequest({ model: "x" })).toThrow(/messages/); + }); + + it("falls back to max_tokens when max_completion_tokens is absent", () => { + const parsed = parseRequest({ model: "m", messages: [], max_tokens: 256 }); + expect(parsed.options.maxOutputTokens).toBe(256); + expect(parsed.stream).toBe(false); + }); +}); + +describe("auth-gateway openai-chat: encodeResponse", () => { + it("serializes text + tool calls with finish_reason=tool_calls", () => { + const message: AssistantMessage = { + ...emptyAssistant(), + content: [ + { type: "text", text: "the answer is " }, + { type: "thinking", thinking: "private reasoning" }, // dropped + { type: "toolCall", id: "call_42", name: "compute", arguments: { x: 1 } }, + ], + usage: { ...baseUsage, input: 10, output: 20, cacheRead: 4, cacheWrite: 6, totalTokens: 40 }, + stopReason: "toolUse", + }; + + const out = encodeResponse(message, "gpt-5.2"); + expect(out.object).toBe("chat.completion"); + expect(out.model).toBe("gpt-5.2"); + expect(typeof out.id).toBe("string"); + expect(String(out.id).startsWith("chatcmpl-")).toBe(true); + + const choices = out.choices as Array<{ + index: number; + message: { role: string; content: string | null; tool_calls?: unknown }; + finish_reason: string; + }>; + expect(choices).toHaveLength(1); + expect(choices[0].finish_reason).toBe("tool_calls"); + expect(choices[0].message.role).toBe("assistant"); + expect(choices[0].message.content).toBe("the answer is "); + expect(choices[0].message.tool_calls).toEqual([ + { id: "call_42", type: "function", function: { name: "compute", arguments: '{"x":1}' } }, + ]); + + expect(out.usage).toEqual({ + prompt_tokens: 20, + prompt_tokens_details: { cached_tokens: 4 }, + completion_tokens: 20, + total_tokens: 40, + }); + }); + + it("maps length stop reason and emits null content when text is empty", () => { + const message: AssistantMessage = { ...emptyAssistant(), stopReason: "length" }; + const out = encodeResponse(message, "gpt-test"); + const choices = out.choices as Array<{ finish_reason: string; message: { content: string | null } }>; + expect(choices[0].finish_reason).toBe("length"); + expect(choices[0].message.content).toBeNull(); + }); +}); + +describe("auth-gateway openai-chat: encodeStream", () => { + it("emits role chunk, text deltas, tool_call deltas with sequential indexes, then [DONE]", async () => { + const partial = emptyAssistant(); + // Pre-populate partial.content so toolcall_start can look up id/name by contentIndex. + partial.content = [ + { type: "text", text: "" }, + { type: "toolCall", id: "call_A", name: "tool_a", arguments: {} }, + { type: "toolCall", id: "call_B", name: "tool_b", arguments: {} }, + ]; + const events: AssistantMessageEvent[] = [ + { type: "text_start", contentIndex: 0, partial }, + { type: "text_delta", contentIndex: 0, delta: "Hi ", partial }, + { type: "text_delta", contentIndex: 0, delta: "there", partial }, + { type: "text_end", contentIndex: 0, content: "Hi there", partial }, + { type: "toolcall_start", contentIndex: 1, partial }, + { type: "toolcall_delta", contentIndex: 1, delta: '{"a":', partial }, + { type: "toolcall_delta", contentIndex: 1, delta: "1}", partial }, + { type: "toolcall_start", contentIndex: 2, partial }, + { type: "toolcall_delta", contentIndex: 2, delta: "{}", partial }, + { + type: "done", + reason: "toolUse", + message: { ...partial, stopReason: "toolUse" }, + }, + ]; + + const stream = encodeStream(makeEventStream(events, partial), "gpt-5.2"); + const lines = await collectStream(stream); + const payloads = lines.map(parseSseLine); + + expect(payloads[payloads.length - 1]).toBe("[DONE]"); + + const chunks = payloads.slice(0, -1) as Array<{ + id: string; + object: string; + model: string; + choices: Array<{ delta: Record; finish_reason: string | null }>; + }>; + + // First chunk is the role announcement. + expect(chunks[0].object).toBe("chat.completion.chunk"); + expect(chunks[0].model).toBe("gpt-5.2"); + expect(chunks[0].choices[0].delta).toEqual({ role: "assistant" }); + expect(chunks[0].choices[0].finish_reason).toBeNull(); + + // All chunks share the same id. + const id = chunks[0].id; + for (const c of chunks) expect(c.id).toBe(id); + + // Collect text deltas. + const textDeltas = chunks.map(c => c.choices[0].delta.content).filter((v): v is string => typeof v === "string"); + expect(textDeltas.join("")).toBe("Hi there"); + + // Collect tool_call deltas; verify index sequence. + const toolDeltas: Array<{ index: number; id?: string; function?: { name?: string; arguments?: string } }> = []; + for (const c of chunks) { + const tc = c.choices[0].delta.tool_calls; + if (Array.isArray(tc)) toolDeltas.push(...(tc as typeof toolDeltas)); + } + // Two starts (index 0 and 1, NOT contentIndex 1 and 2) plus three arg deltas. + const starts = toolDeltas.filter(t => typeof t.id === "string" && t.id.length > 0); + expect(starts.map(s => s.index)).toEqual([0, 1]); + expect(starts[0].id).toBe("call_A"); + expect(starts[0].function?.name).toBe("tool_a"); + expect(starts[1].id).toBe("call_B"); + expect(starts[1].function?.name).toBe("tool_b"); + + // Argument deltas use the wire index, not the contentIndex. + const argDeltas = toolDeltas.filter(t => typeof t.function?.arguments === "string" && !t.id); + expect(argDeltas.map(d => [d.index, d.function?.arguments])).toEqual([ + [0, '{"a":'], + [0, "1}"], + [1, "{}"], + ]); + + // Penultimate chunk carries finish_reason. + const finishChunk = chunks[chunks.length - 1]; + expect(finishChunk.choices[0].delta).toEqual({}); + expect(finishChunk.choices[0].finish_reason).toBe("tool_calls"); + }); + + it("emits an error envelope when the stream errors", async () => { + const partial = emptyAssistant(); + const errorMessage: AssistantMessage = { ...partial, errorMessage: "upstream went away" }; + const events: AssistantMessageEvent[] = [{ type: "error", reason: "error", error: errorMessage }]; + const stream = encodeStream(makeEventStream(events, partial), "gpt-test"); + const lines = await collectStream(stream); + expect(lines).toHaveLength(2); // role chunk + error envelope + const payloads = lines.map(parseSseLine) as Array>; + expect(payloads[1]).toEqual({ error: { message: "upstream went away", type: "upstream_error" } }); + }); +}); diff --git a/packages/ai/test/auth-gateway-openai-responses.test.ts b/packages/ai/test/auth-gateway-openai-responses.test.ts new file mode 100644 index 000000000..d08bc038a --- /dev/null +++ b/packages/ai/test/auth-gateway-openai-responses.test.ts @@ -0,0 +1,533 @@ +import { describe, expect, it } from "bun:test"; +import { Effort } from "../src/model-thinking"; +import { encodeResponse, encodeStream, parseRequest } from "../src/providers/openai-responses-server"; +import type { AssistantMessage } from "../src/types"; +import { AssistantMessageEventStream } from "../src/utils/event-stream"; + +function zeroUsage(): AssistantMessage["usage"] { + return { + input: 0, + output: 0, + cacheRead: 0, + cacheWrite: 0, + totalTokens: 0, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, + }; +} + +async function collectStream(stream: ReadableStream): Promise { + const reader = stream.getReader(); + const decoder = new TextDecoder(); + let out = ""; + while (true) { + const { value, done } = await reader.read(); + if (done) break; + out += decoder.decode(value); + } + return out; +} + +interface SseFrame { + event: string; + data: Record | string; +} + +function parseSse(raw: string): SseFrame[] { + const frames: SseFrame[] = []; + for (const chunk of raw.split("\n\n")) { + if (!chunk.trim()) continue; + let event = ""; + let dataLine = ""; + for (const line of chunk.split("\n")) { + if (line.startsWith("event: ")) event = line.slice("event: ".length); + else if (line.startsWith("data: ")) dataLine = line.slice("data: ".length); + } + if (dataLine === "[DONE]") { + frames.push({ event: event || "done_sentinel", data: "[DONE]" }); + } else if (dataLine) { + const parsed: unknown = JSON.parse(dataLine); + if (parsed && typeof parsed === "object") { + frames.push({ event, data: parsed as Record }); + } + } + } + return frames; +} + +describe("openai-responses parseRequest", () => { + it("parses an input array with mixed message + reasoning + function_call + function_call_output", () => { + const reasoningItem = { + type: "reasoning", + id: "rs_abc", + summary: [], + content: [{ type: "reasoning_text", text: "The user wants arithmetic." }], + }; + const parsed = parseRequest({ + model: "gpt-5.3-codex-spark", + instructions: "You are X", + input: [ + { type: "message", role: "user", content: [{ type: "input_text", text: "what's 2+2?" }] }, + { + type: "message", + role: "assistant", + content: [{ type: "output_text", text: "Let me think." }], + }, + reasoningItem, + { + type: "function_call", + id: "fc_item_999", + call_id: "call_42", + name: "math", + arguments: '{"a":2,"b":2}', + }, + { type: "function_call_output", call_id: "call_42", output: "4" }, + ], + tools: [ + { + type: "function", + name: "math", + description: "Do arithmetic", + parameters: { type: "object", properties: { a: { type: "number" }, b: { type: "number" } } }, + strict: true, + }, + ], + tool_choice: { type: "function", name: "math" }, + max_output_tokens: 1024, + temperature: 0.1, + top_p: 0.9, + reasoning: { effort: "high", summary: "detailed" }, + store: true, + previous_response_id: "resp_prev", + stream: true, + }); + + expect(parsed.modelId).toBe("gpt-5.3-codex-spark"); + expect(parsed.stream).toBe(true); + expect(parsed.context.systemPrompt).toEqual(["You are X"]); + + const msgs = parsed.context.messages; + expect(msgs).toHaveLength(3); + + // 1. user + expect(msgs[0]!.role).toBe("user"); + const u = msgs[0]!; + if (u.role !== "user") throw new Error("expected user"); + expect(u.content).toBe("what's 2+2?"); + + // 2. assistant with text + reasoning + toolCall + const a = msgs[1]!; + if (a.role !== "assistant") throw new Error("expected assistant"); + expect(a.api).toBe("openai-responses"); + expect(a.provider).toBe("openai"); + expect(a.model).toBe("gpt-5.3-codex-spark"); + expect(a.content).toHaveLength(3); + expect(a.content[0]).toMatchObject({ type: "text", text: "Let me think." }); + expect(a.content[1]).toMatchObject({ + type: "thinking", + thinking: "The user wants arithmetic.", + thinkingSignature: JSON.stringify(reasoningItem), + itemId: "rs_abc", + }); + // Critical: call_id and item id are distinct. + expect(a.content[2]).toMatchObject({ + type: "toolCall", + id: "call_42", + name: "math", + arguments: { a: 2, b: 2 }, + thoughtSignature: "fc_item_999", + }); + + // 3. toolResult + const tr = msgs[2]!; + if (tr.role !== "toolResult") throw new Error("expected toolResult"); + expect(tr.toolCallId).toBe("call_42"); + expect(tr.toolName).toBe("math"); + expect(tr.content).toEqual([{ type: "text", text: "4" }]); + expect(tr.isError).toBe(false); + + expect(parsed.context.tools).toHaveLength(1); + expect(parsed.context.tools![0]).toMatchObject({ name: "math", strict: true }); + + expect(parsed.options.maxOutputTokens).toBe(1024); + expect(parsed.options.temperature).toBe(0.1); + expect(parsed.options.topP).toBe(0.9); + expect(parsed.options.toolChoice).toEqual({ name: "math" }); + expect(parsed.options.reasoning).toBe(Effort.High); + // `reasoning.summary: "detailed"` is treated as the default visible-summary + // case (only "none" toggles hideThinkingSummary). + expect(parsed.options.hideThinkingSummary).toBeUndefined(); + // `store` and `previous_response_id` are accepted by the schema but not + // plumbed through pi-ai — they no longer leak into options.extra. + expect(parsed.options.extra).toBeUndefined(); + }); + + it("accepts a bare string input and rejects a missing model", () => { + const parsed = parseRequest({ model: "m", input: "hi" }); + expect(parsed.context.messages).toHaveLength(1); + const m = parsed.context.messages[0]!; + if (m.role !== "user") throw new Error("expected user"); + expect(m.content).toBe("hi"); + + expect(() => parseRequest({ input: "hi" })).toThrow(/model/); + }); + + it("preserves string message content and system input items", () => { + const parsed = parseRequest({ + model: "m", + instructions: "top-level instructions", + input: [ + { role: "system", content: "system from easy input" }, + { role: "user", content: "hello" }, + { role: "assistant", content: "hi there" }, + { + type: "message", + role: "system", + content: [{ type: "input_text", text: "structured system" }], + }, + ], + }); + + expect(parsed.context.systemPrompt).toEqual([ + "top-level instructions", + "system from easy input", + "structured system", + ]); + expect(parsed.context.messages).toHaveLength(2); + const user = parsed.context.messages[0]!; + const assistant = parsed.context.messages[1]!; + if (user.role !== "user") throw new Error("expected user"); + if (assistant.role !== "assistant") throw new Error("expected assistant"); + expect(user.content).toBe("hello"); + expect(assistant.content).toEqual([{ type: "text", text: "hi there" }]); + }); + + it("creates a synthetic assistant when reasoning comes before any assistant message", () => { + const reasoningItem = { + type: "reasoning", + id: "rs_x", + content: [{ type: "reasoning_text", text: "hmm" }], + }; + const parsed = parseRequest({ + model: "m", + input: [reasoningItem], + }); + expect(parsed.context.messages).toHaveLength(1); + const a = parsed.context.messages[0]!; + if (a.role !== "assistant") throw new Error("expected synthetic assistant"); + expect(a.content).toHaveLength(1); + expect(a.content[0]).toMatchObject({ + type: "thinking", + thinking: "hmm", + thinkingSignature: JSON.stringify(reasoningItem), + itemId: "rs_x", + }); + }); +}); + +describe("openai-responses encodeResponse", () => { + it("encodes reasoning + message + function_call output items", () => { + const reasoningItem = { + type: "reasoning", + id: "rs_signed", + summary: [], + content: [{ type: "reasoning_text", text: "thinking aloud" }], + }; + const message: AssistantMessage = { + role: "assistant", + api: "openai-responses", + provider: "openai", + model: "gpt-5", + content: [ + { + type: "thinking", + thinking: "thinking aloud", + thinkingSignature: JSON.stringify(reasoningItem), + itemId: "rs_signed", + }, + { type: "text", text: "Hello " }, + { type: "text", text: "world" }, + { + type: "toolCall", + id: "call_t1", + name: "math", + arguments: { a: 1, b: 2 }, + thoughtSignature: "fc_item_t1", + }, + ], + usage: { + ...zeroUsage(), + input: 10, + output: 20, + cacheRead: 4, + cacheWrite: 6, + reasoningTokens: 5, + }, + stopReason: "toolUse", + timestamp: 1_700_000_000_000, + }; + + const body = encodeResponse(message, "gpt-5-requested"); + + expect(body.object).toBe("response"); + expect(body.status).toBe("completed"); + expect(body.model).toBe("gpt-5-requested"); + expect(body.created_at).toBe(1_700_000_000); + expect(typeof body.id).toBe("string"); + expect((body.id as string).startsWith("resp_")).toBe(true); + + const output = body.output as Array>; + expect(output).toHaveLength(3); + + expect(output[0]).toEqual(reasoningItem); + + // Consecutive text collapses into one message item with two parts. + expect(output[1]!.type).toBe("message"); + expect(output[1]!.role).toBe("assistant"); + const parts = output[1]!.content as Array<{ type: string; text: string; annotations: never[] }>; + expect(parts).toEqual([ + { type: "output_text", text: "Hello ", annotations: [] }, + { type: "output_text", text: "world", annotations: [] }, + ]); + + // function_call: wire id (thoughtSignature) and call_id are distinct. + expect(output[2]).toMatchObject({ + type: "function_call", + id: "fc_item_t1", + call_id: "call_t1", + name: "math", + arguments: '{"a":1,"b":2}', + status: "completed", + }); + + expect(body.usage).toEqual({ + input_tokens: 20, + input_tokens_details: { cached_tokens: 4 }, + output_tokens: 20, + output_tokens_details: { reasoning_tokens: 5 }, + total_tokens: 40, + }); + }); + + it("marks length-limited responses incomplete", () => { + const message: AssistantMessage = { + role: "assistant", + api: "openai-responses", + provider: "openai", + model: "gpt-5", + content: [{ type: "text", text: "partial" }], + usage: zeroUsage(), + stopReason: "length", + timestamp: 1_700_000_000_000, + }; + + const body = encodeResponse(message, "gpt-5-requested"); + + expect(body.status).toBe("incomplete"); + expect(body.incomplete_details).toEqual({ reason: "max_output_tokens" }); + }); +}); + +describe("openai-responses encodeStream", () => { + it("emits response.created, reasoning_text.delta, output_text.delta, function_call_arguments.delta, response.completed, [DONE]", async () => { + const stream = new AssistantMessageEventStream(); + + const partial: AssistantMessage = { + role: "assistant", + api: "openai-responses", + provider: "openai", + model: "gpt-5", + content: [], + usage: zeroUsage(), + stopReason: "stop", + timestamp: 1_700_000_000_000, + }; + + const finalMessage: AssistantMessage = { + role: "assistant", + api: "openai-responses", + provider: "openai", + model: "gpt-5", + content: [ + { type: "thinking", thinking: "step 1", thinkingSignature: "rs_s1", itemId: "rs_s1" }, + { type: "text", text: "Hi!" }, + { + type: "toolCall", + id: "call_x", + name: "math", + arguments: { a: 1 }, + thoughtSignature: "fc_x", + }, + ], + usage: { ...zeroUsage(), input: 1, output: 2 }, + stopReason: "toolUse", + timestamp: 1_700_000_000_000, + }; + + // Push events asynchronously while consumer reads. + const partialWithThinking: AssistantMessage = { + ...partial, + content: [{ type: "thinking", thinking: "", thinkingSignature: "rs_s1", itemId: "rs_s1" }], + }; + const partialWithToolCall: AssistantMessage = { + ...partial, + content: [ + { type: "thinking", thinking: "step 1", thinkingSignature: "rs_s1", itemId: "rs_s1" }, + { type: "text", text: "Hi!" }, + { type: "toolCall", id: "call_x", name: "math", arguments: {}, thoughtSignature: "fc_x" }, + ], + }; + + queueMicrotask(() => { + stream.push({ type: "start", partial }); + stream.push({ type: "thinking_start", contentIndex: 0, partial: partialWithThinking }); + stream.push({ type: "thinking_delta", contentIndex: 0, delta: "step ", partial: partialWithThinking }); + stream.push({ type: "thinking_delta", contentIndex: 0, delta: "1", partial: partialWithThinking }); + stream.push({ type: "thinking_end", contentIndex: 0, content: "step 1", partial: partialWithThinking }); + stream.push({ type: "text_start", contentIndex: 1, partial }); + stream.push({ type: "text_delta", contentIndex: 1, delta: "Hi", partial }); + stream.push({ type: "text_delta", contentIndex: 1, delta: "!", partial }); + stream.push({ type: "text_end", contentIndex: 1, content: "Hi!", partial }); + stream.push({ type: "toolcall_start", contentIndex: 2, partial: partialWithToolCall }); + stream.push({ type: "toolcall_delta", contentIndex: 2, delta: '{"a":', partial: partialWithToolCall }); + stream.push({ type: "toolcall_delta", contentIndex: 2, delta: "1}", partial: partialWithToolCall }); + stream.push({ + type: "toolcall_end", + contentIndex: 2, + toolCall: { + type: "toolCall", + id: "call_x", + name: "math", + arguments: { a: 1 }, + thoughtSignature: "fc_x", + }, + partial: partialWithToolCall, + }); + stream.push({ type: "done", reason: "toolUse", message: finalMessage }); + }); + + const raw = await collectStream(encodeStream(stream, "gpt-5-requested")); + const frames = parseSse(raw); + const names = frames.map(f => f.event); + + // Ordering: created → thinking flow → message flow → tool-call flow → completed → [DONE] + expect(names[0]).toBe("response.created"); + expect(names[names.length - 1]).toBe("done_sentinel"); + expect(frames[frames.length - 1]!.data).toBe("[DONE]"); + + // Spot-check critical events appear in the expected order. + const idxCreated = names.indexOf("response.created"); + const idxReasoningDelta = names.indexOf("response.reasoning_text.delta"); + const idxReasoningDone = names.indexOf("response.reasoning_text.done"); + const idxTextDelta = names.indexOf("response.output_text.delta"); + const idxTextDone = names.indexOf("response.output_text.done"); + const idxArgsDelta = names.indexOf("response.function_call_arguments.delta"); + const idxMessageDone = frames.findIndex( + f => + f.event === "response.output_item.done" && + (f.data as Record).item && + ((f.data as Record).item as Record).type === "message", + ); + const idxArgsDone = names.indexOf("response.function_call_arguments.done"); + const idxCompleted = names.indexOf("response.completed"); + + expect(idxCreated).toBeGreaterThanOrEqual(0); + expect(idxReasoningDelta).toBeGreaterThan(idxCreated); + expect(idxReasoningDone).toBeGreaterThan(idxReasoningDelta); + expect(idxTextDelta).toBeGreaterThan(idxReasoningDone); + expect(idxTextDone).toBeGreaterThan(idxTextDelta); + expect(idxArgsDelta).toBeGreaterThan(idxTextDone); + expect(idxArgsDone).toBeGreaterThan(idxArgsDelta); + expect(idxCompleted).toBeGreaterThan(idxArgsDone); + + // reasoning_text.delta must carry item_id matching the signature, and output_index 0. + const reasoningDelta = frames[idxReasoningDelta]!.data as Record; + expect(reasoningDelta.item_id).toBe("rs_s1"); + expect(reasoningDelta.output_index).toBe(0); + expect(reasoningDelta.delta).toBe("step "); + + // output_text.delta's item_id is a new msg_*, output_index moved on past the reasoning item. + const textDelta = frames[idxTextDelta]!.data as Record; + expect(typeof textDelta.item_id).toBe("string"); + expect((textDelta.item_id as string).startsWith("msg_")).toBe(true); + expect(textDelta.output_index).toBe(1); + expect(textDelta.delta).toBe("Hi"); + expect(textDelta.logprobs).toEqual([]); + + const textDone = frames[idxTextDone]!.data as Record; + expect(textDone.text).toBe("Hi!"); + expect(textDone.logprobs).toEqual([]); + + const messageDone = frames[idxMessageDone]!.data as Record; + expect(messageDone.output_index).toBe(1); + expect(messageDone.item).toMatchObject({ + type: "message", + status: "completed", + content: [{ type: "output_text", text: "Hi!", annotations: [] }], + }); + + // function_call_arguments.delta uses the fc_* wire id, NOT call_x. + const argsDelta = frames[idxArgsDelta]!.data as Record; + expect(argsDelta.item_id).toBe("fc_x"); + expect(argsDelta.output_index).toBe(2); + expect(argsDelta.delta).toBe('{"a":'); + + const argsDone = frames[idxArgsDone]!.data as Record; + expect(argsDone.item_id).toBe("fc_x"); + expect(argsDone.arguments).toBe('{"a":1}'); + expect(argsDone.name).toBe("math"); + + // response.completed: assert the final response object carries the full output items + // and that call_id ≠ id for the function_call item. + const completed = frames[idxCompleted]!.data as Record; + const response = completed.response as Record; + expect(response.status).toBe("completed"); + expect(response.model).toBe("gpt-5-requested"); + const output = response.output as Array>; + expect(output).toHaveLength(3); + expect(output[0]!.type).toBe("reasoning"); + expect(output[1]!.type).toBe("message"); + expect(output[2]).toMatchObject({ + type: "function_call", + id: "fc_x", + call_id: "call_x", + name: "math", + arguments: '{"a":1}', + }); + // Critical gotcha: id and call_id are distinct. + expect(output[2]!.id).not.toBe(output[2]!.call_id); + }); + + it("emits response.incomplete for length-limited streams", async () => { + const stream = new AssistantMessageEventStream(); + const message: AssistantMessage = { + role: "assistant", + api: "openai-responses", + provider: "openai", + model: "gpt-5", + content: [{ type: "text", text: "partial" }], + usage: { ...zeroUsage(), output: 1 }, + stopReason: "length", + timestamp: 1_700_000_000_000, + }; + + queueMicrotask(() => { + stream.push({ type: "start", partial: { ...message, content: [] } }); + stream.push({ type: "text_start", contentIndex: 0, partial: message }); + stream.push({ type: "text_delta", contentIndex: 0, delta: "partial", partial: message }); + stream.push({ type: "text_end", contentIndex: 0, content: "partial", partial: message }); + stream.push({ type: "done", reason: "length", message }); + }); + + const raw = await collectStream(encodeStream(stream, "gpt-5-requested")); + const frames = parseSse(raw); + const names = frames.map(f => f.event); + const idxIncomplete = names.indexOf("response.incomplete"); + + expect(idxIncomplete).toBeGreaterThan(-1); + expect(names).not.toContain("response.completed"); + const incomplete = frames[idxIncomplete]!.data as Record; + const response = incomplete.response as Record; + expect(response.status).toBe("incomplete"); + expect(response.incomplete_details).toEqual({ reason: "max_output_tokens" }); + }); +}); diff --git a/packages/ai/test/auth-storage-config-override.test.ts b/packages/ai/test/auth-storage-config-override.test.ts new file mode 100644 index 000000000..16e140255 --- /dev/null +++ b/packages/ai/test/auth-storage-config-override.test.ts @@ -0,0 +1,125 @@ +import { afterEach, beforeEach, describe, expect, test } from "bun:test"; +import * as fs from "node:fs/promises"; +import * as os from "node:os"; +import * as path from "node:path"; +import { type AuthCredentialStore, AuthStorage, SqliteAuthCredentialStore } from "../src/auth-storage"; +import { withEnv } from "./helpers"; + +const SUPPRESS_ANTHROPIC_ENV = { + ANTHROPIC_API_KEY: undefined, + ANTHROPIC_OAUTH_TOKEN: undefined, +} as const; + +describe("AuthStorage config-override apiKey", () => { + let tempDir = ""; + let store: AuthCredentialStore | null = null; + let authStorage: AuthStorage | null = null; + + beforeEach(async () => { + tempDir = await fs.mkdtemp(path.join(os.tmpdir(), "pi-ai-auth-config-override-")); + store = await SqliteAuthCredentialStore.open(path.join(tempDir, "agent.db")); + authStorage = new AuthStorage(store); + }); + + afterEach(async () => { + store?.close(); + store = null; + authStorage = null; + if (tempDir) { + await fs.rm(tempDir, { recursive: true, force: true }); + tempDir = ""; + } + }); + + async function seedOAuth(provider: string, access: string): Promise { + if (!authStorage) throw new Error("test setup failed"); + await authStorage.set(provider, [ + { + type: "oauth", + access, + refresh: `${access}-refresh`, + expires: Date.now() + 60 * 60_000, + }, + ]); + } + + test("setConfigApiKey beats OAuth access token for getApiKey", async () => { + await withEnv(SUPPRESS_ANTHROPIC_ENV, async () => { + if (!authStorage) throw new Error("test setup failed"); + await seedOAuth("anthropic", "oauth-from-broker"); + authStorage.setConfigApiKey("anthropic", "gateway-bearer"); + + expect(await authStorage.getApiKey("anthropic")).toBe("gateway-bearer"); + expect(await authStorage.peekApiKey("anthropic")).toBe("gateway-bearer"); + }); + }); + + test("runtime override (--api-key) still beats setConfigApiKey", async () => { + await withEnv(SUPPRESS_ANTHROPIC_ENV, async () => { + if (!authStorage) throw new Error("test setup failed"); + await seedOAuth("anthropic", "oauth-from-broker"); + authStorage.setConfigApiKey("anthropic", "gateway-bearer"); + authStorage.setRuntimeApiKey("anthropic", "cli-flag-bearer"); + + expect(await authStorage.getApiKey("anthropic")).toBe("cli-flag-bearer"); + }); + }); + + test("removeConfigApiKey restores OAuth resolution", async () => { + await withEnv(SUPPRESS_ANTHROPIC_ENV, async () => { + if (!authStorage) throw new Error("test setup failed"); + await seedOAuth("anthropic", "oauth-from-broker"); + authStorage.setConfigApiKey("anthropic", "gateway-bearer"); + expect(await authStorage.getApiKey("anthropic")).toBe("gateway-bearer"); + + authStorage.removeConfigApiKey("anthropic"); + expect(await authStorage.getApiKey("anthropic")).toBe("oauth-from-broker"); + }); + }); + + test("clearConfigApiKeys drops every config override at once", async () => { + await withEnv(SUPPRESS_ANTHROPIC_ENV, async () => { + if (!authStorage) throw new Error("test setup failed"); + await seedOAuth("anthropic", "oauth-anthropic"); + await seedOAuth("openai-codex", "oauth-codex"); + authStorage.setConfigApiKey("anthropic", "gateway-bearer-A"); + authStorage.setConfigApiKey("openai-codex", "gateway-bearer-B"); + + authStorage.clearConfigApiKeys(); + + expect(await authStorage.getApiKey("anthropic")).toBe("oauth-anthropic"); + expect(await authStorage.getApiKey("openai-codex")).toBe("oauth-codex"); + }); + }); + + test("setConfigApiKey suppresses OAuth account_uuid attribution", async () => { + await withEnv(SUPPRESS_ANTHROPIC_ENV, async () => { + if (!authStorage) throw new Error("test setup failed"); + await authStorage.set("anthropic", [ + { + type: "oauth", + access: "oauth-with-account", + refresh: "r", + expires: Date.now() + 60 * 60_000, + accountId: "acc-123", + }, + ]); + // Sanity: without override, accountId is exposed. + expect(authStorage.getOAuthAccountId("anthropic")).toBe("acc-123"); + + authStorage.setConfigApiKey("anthropic", "gateway-bearer"); + // With an explicit config bearer in play, OAuth account attribution + // must NOT leak — outbound auth is the gateway bearer, not OAuth. + expect(authStorage.getOAuthAccountId("anthropic")).toBeUndefined(); + }); + }); + + test("describeCredentialSource reports config override", async () => { + await withEnv(SUPPRESS_ANTHROPIC_ENV, async () => { + if (!authStorage) throw new Error("test setup failed"); + await seedOAuth("anthropic", "oauth-from-broker"); + authStorage.setConfigApiKey("anthropic", "gateway-bearer"); + expect(authStorage.describeCredentialSource("anthropic")).toBe("config override (models.yml)"); + }); + }); +}); diff --git a/packages/ai/test/auth-storage-usage-cache.test.ts b/packages/ai/test/auth-storage-usage-cache.test.ts new file mode 100644 index 000000000..7869af845 --- /dev/null +++ b/packages/ai/test/auth-storage-usage-cache.test.ts @@ -0,0 +1,265 @@ +/** + * Tests for the new usage-cache contracts introduced after the broker + * migration surfaced Anthropic per-IP rate limits: + * + * 1. Per-credential cache stores the last successful report; failures + * DON'T overwrite a stale-but-good entry with null. + * 2. With a stale-but-good entry, a failure serves the previous value + * (cached for a short cool-down) instead of dropping the credential + * from the report. + * 3. Without a previous value, a failure returns null and DOES NOT cache — + * the next poll retries on the next request. + */ +import { afterEach, beforeEach, describe, expect, it, vi } from "bun:test"; +import { + type AuthCredential, + type AuthCredentialStore, + AuthStorage, + type StoredAuthCredential, +} from "../src/auth-storage"; +import type { UsageReport } from "../src/usage"; +import * as claudeUsage from "../src/usage/claude"; + +function anthropicReports(reports: UsageReport[] | null): UsageReport[] { + return (reports ?? []).filter(r => r.provider === "anthropic"); +} + +/** + * Force every cache entry to look stale to AuthStorage WITHOUT dropping the + * value. The cache layer is two-tier: the store-level `expiresAtSec` controls + * whether `getCache` returns anything at all, and the JSON payload's own + * `expiresAt` is what AuthStorage compares against `Date.now()` to decide if + * the entry is fresh. Mutating only the inner expiresAt simulates time + * passing while keeping the last-good value reachable for the failure path. + */ +function expireCachePayloads(store: ObservableStore): void { + for (const [key, entry] of store.cache) { + try { + const parsed = JSON.parse(entry.value); + parsed.expiresAt = 1; // positive but already in the past (epoch ms) + store.cache.set(key, { value: JSON.stringify(parsed), expiresAtSec: entry.expiresAtSec }); + } catch { + // Non-JSON entries — leave alone. + } + } +} + +interface CacheEntry { + value: string; + expiresAtSec: number; +} + +interface ObservableStore extends AuthCredentialStore { + cache: Map; +} + +/** + * Minimal in-memory `AuthCredentialStore` exposing the cache so we can + * assert what AuthStorage writes to it during usage fetches. + */ +function makeStore(rows: StoredAuthCredential[]): ObservableStore { + const cache = new Map(); + return { + cache, + close() {}, + listAuthCredentials() { + return rows; + }, + updateAuthCredential() {}, + deleteAuthCredential() {}, + tryDisableAuthCredentialIfMatches() { + return false; + }, + replaceAuthCredentialsForProvider() { + return rows; + }, + upsertAuthCredentialForProvider() { + return rows; + }, + deleteAuthCredentialsForProvider() {}, + getCache(key) { + const entry = cache.get(key); + if (!entry) return null; + if (entry.expiresAtSec * 1000 <= Date.now()) return null; + return entry.value; + }, + setCache(key, value, expiresAtSec) { + cache.set(key, { value, expiresAtSec }); + }, + cleanExpiredCache() {}, + }; +} + +function oauthRow(id: number, email: string): StoredAuthCredential { + const credential: AuthCredential = { + type: "oauth", + access: `oat-${id}`, + refresh: `refresh-${id}`, + expires: Date.now() + 3_600_000, + accountId: `account-${id}`, + email, + }; + return { id, provider: "anthropic", credential, disabledCause: null }; +} + +function makeReport(account: string): UsageReport { + return { + provider: "anthropic", + fetchedAt: Date.now(), + limits: [ + { + id: "anthropic:5h", + label: "5 Hour", + scope: { provider: "anthropic", windowId: "5h" }, + window: { id: "5h", label: "5 Hour" }, + amount: { used: 42, limit: 100, unit: "percent" }, + status: "ok", + }, + ], + metadata: { email: account, accountId: `account-${account}` }, + }; +} + +describe("AuthStorage usage cache: last-good failure fallback", () => { + let store: ObservableStore; + let storage: AuthStorage; + + beforeEach(async () => { + store = makeStore([oauthRow(1, "a@example.com")]); + storage = new AuthStorage(store); + await storage.reload(); + }); + + afterEach(() => { + storage.close(); + vi.restoreAllMocks(); + }); + + it("caches a successful report and replays it on a second poll", async () => { + let calls = 0; + const goldReport = makeReport("a@example.com"); + vi.spyOn(claudeUsage.claudeUsageProvider, "fetchUsage").mockImplementation(async () => { + calls += 1; + return goldReport; + }); + + const first = anthropicReports(await storage.fetchUsageReports()); + expect(first).toHaveLength(1); + expect(calls).toBe(1); + + const second = anthropicReports(await storage.fetchUsageReports()); + expect(second).toHaveLength(1); + // Cache hit — provider was NOT called a second time. + expect(calls).toBe(1); + }); + + it("does NOT cache a failure when no previous good value exists — retries next poll", async () => { + let calls = 0; + vi.spyOn(claudeUsage.claudeUsageProvider, "fetchUsage").mockImplementation(async () => { + calls += 1; + return null; + }); + + const first = anthropicReports(await storage.fetchUsageReports()); + expect(first).toHaveLength(0); + expect(calls).toBe(1); + + const second = anthropicReports(await storage.fetchUsageReports()); + // No previous value → no cache write → retry on next poll. + expect(calls).toBe(2); + expect(second).toHaveLength(0); + }); + + it("serves last-good value through a failure cycle", async () => { + let calls = 0; + const goldReport = makeReport("a@example.com"); + vi.spyOn(claudeUsage.claudeUsageProvider, "fetchUsage").mockImplementation(async () => { + calls += 1; + if (calls === 1) return goldReport; + return null; + }); + + // First poll: real fetch → cached. + const first = anthropicReports(await storage.fetchUsageReports()); + expect(first).toHaveLength(1); + expect(calls).toBe(1); + + // Force every cached entry to expire so the next poll refetches. + // Bun's `bun:test` doesn't ship setSystemTime, so we manipulate the + // observable store cache directly — equivalent to advancing time past + // the success TTL. + expireCachePayloads(store); + + // Second poll: cache expired → refetch → provider returns null → + // AuthStorage falls back to last-good and the report stays populated. + const second = anthropicReports(await storage.fetchUsageReports()); + expect(calls).toBe(2); + expect(second).toHaveLength(1); + // The fallback value must be the SAME report (not a synthetic empty one). + expect(second?.[0]?.limits[0]?.amount.used).toBe(42); + }); + + it("re-attempts the failing credential after the cool-down expires", async () => { + let calls = 0; + const goldReport = makeReport("a@example.com"); + vi.spyOn(claudeUsage.claudeUsageProvider, "fetchUsage").mockImplementation(async () => { + calls += 1; + // Succeed on attempt 1, fail on 2, succeed on 3. + if (calls === 2) return null; + return goldReport; + }); + + const first = anthropicReports(await storage.fetchUsageReports()); + expect(first).toHaveLength(1); + expect(calls).toBe(1); + + // Expire success cache → poll 2 fetches and 429s → cool-down written. + expireCachePayloads(store); + const second = anthropicReports(await storage.fetchUsageReports()); + expect(second).toHaveLength(1); // last-good fallback + expect(calls).toBe(2); + + // Expire the cool-down → poll 3 refetches → success. + expireCachePayloads(store); + const third = anthropicReports(await storage.fetchUsageReports()); + expect(third).toHaveLength(1); + expect(calls).toBe(3); + }); +}); + +describe("AuthStorage usage cache: jitter", () => { + it("writes per-credential cache TTLs with ±25% jitter so refreshes decorrelate", async () => { + const store = makeStore([oauthRow(1, "a@example.com"), oauthRow(2, "b@example.com")]); + const storage = new AuthStorage(store); + await storage.reload(); + try { + const goldA = makeReport("a@example.com"); + const goldB = makeReport("b@example.com"); + vi.spyOn(claudeUsage.claudeUsageProvider, "fetchUsage").mockImplementation(async params => { + return params.credential.email === "a@example.com" ? goldA : goldB; + }); + + await storage.fetchUsageReports(); + + // The store-level TTL is bumped to the 24h durable-retention floor so + // `getStale` can recover last-good values; the freshness TTL we actually + // jitter lives in the JSON payload. Read that, not the store TTL. + const freshExpiries: number[] = []; + for (const entry of store.cache.values()) { + if (entry.value.length === 0) continue; + const parsed = JSON.parse(entry.value); + if (typeof parsed?.expiresAt === "number") freshExpiries.push(parsed.expiresAt); + } + expect(freshExpiries.length).toBeGreaterThanOrEqual(2); + const now = Date.now(); + for (const expiry of freshExpiries) { + const delta = expiry - now; + expect(delta).toBeGreaterThan(3.5 * 60_000); + expect(delta).toBeLessThan(6.5 * 60_000); + } + } finally { + storage.close(); + vi.restoreAllMocks(); + } + }); +}); diff --git a/packages/ai/test/claude-usage-retry.test.ts b/packages/ai/test/claude-usage-retry.test.ts new file mode 100644 index 000000000..73a1cd8d4 --- /dev/null +++ b/packages/ai/test/claude-usage-retry.test.ts @@ -0,0 +1,178 @@ +import { afterEach, describe, expect, it, vi } from "bun:test"; +import { setTimeout as setTimeoutCb } from "node:timers"; +import type { UsageFetchContext } from "../src/usage"; +import { claudeUsageProvider } from "../src/usage/claude"; + +const VALID_PAYLOAD = { + five_hour: { utilization: 42, resets_at: new Date(Date.now() + 5 * 60_000).toISOString() }, +}; + +function jsonResponse(status: number, body: unknown, headers: Record = {}): Response { + return new Response(JSON.stringify(body), { + status, + headers: { "Content-Type": "application/json", ...headers }, + }); +} + +function makeContext(fetchImpl: typeof fetch): UsageFetchContext { + return { fetch: fetchImpl }; +} + +function baseParams() { + return { + provider: "anthropic" as const, + credential: { + type: "oauth" as const, + accessToken: "oat-test", + accountId: "org_test", + email: "user@example.com", + expiresAt: Date.now() + 60_000, + }, + }; +} + +describe("claudeUsageProvider retry contract", () => { + afterEach(() => { + vi.restoreAllMocks(); + }); + + it("retries on 429 and succeeds on a later attempt", async () => { + let attempt = 0; + const fetchMock = (async () => { + attempt += 1; + if (attempt < 3) return jsonResponse(429, { error: "rate_limited" }); + return jsonResponse(200, VALID_PAYLOAD); + }) as unknown as typeof fetch; + + const report = await claudeUsageProvider.fetchUsage(baseParams(), makeContext(fetchMock)); + expect(report).not.toBeNull(); + expect(attempt).toBe(3); + expect(report?.limits[0]?.amount.used).toBe(42); + }); + + it("retries on 503 then succeeds", async () => { + let attempt = 0; + const fetchMock = (async () => { + attempt += 1; + if (attempt === 1) return jsonResponse(503, { error: "unavailable" }); + return jsonResponse(200, VALID_PAYLOAD); + }) as unknown as typeof fetch; + + const report = await claudeUsageProvider.fetchUsage(baseParams(), makeContext(fetchMock)); + expect(report).not.toBeNull(); + expect(attempt).toBe(2); + }); + + it("does NOT retry on 401 — permanent for this credential", async () => { + let attempt = 0; + const fetchMock = (async () => { + attempt += 1; + return jsonResponse(401, { error: "unauthorized" }); + }) as unknown as typeof fetch; + + const report = await claudeUsageProvider.fetchUsage(baseParams(), makeContext(fetchMock)); + expect(report).toBeNull(); + expect(attempt).toBe(1); + }); + + it("does NOT retry on 404 — permanent for this credential", async () => { + let attempt = 0; + const fetchMock = (async () => { + attempt += 1; + return jsonResponse(404, { error: "not_found" }); + }) as unknown as typeof fetch; + + const report = await claudeUsageProvider.fetchUsage(baseParams(), makeContext(fetchMock)); + expect(report).toBeNull(); + expect(attempt).toBe(1); + }); + + it("returns null after MAX_RETRIES of consecutive 429s", async () => { + let attempt = 0; + const fetchMock = (async () => { + attempt += 1; + return jsonResponse(429, { error: "rate_limited" }); + }) as unknown as typeof fetch; + + // Provider's MAX_RETRIES is 3; provider sleeps BASE_RETRY_DELAY_MS * 2^attempt + // between attempts — total worst-case ~1.5s, well within our test budget. + const report = await claudeUsageProvider.fetchUsage(baseParams(), makeContext(fetchMock)); + expect(report).toBeNull(); + expect(attempt).toBe(3); + }); + + it("honours Retry-After when retrying a 429", async () => { + let attempt = 0; + const callTimes: number[] = []; + const fetchMock = (async () => { + attempt += 1; + callTimes.push(Date.now()); + if (attempt === 1) { + // Retry-After: 1 second. Provider must wait ~1s before re-attempting. + return jsonResponse(429, { error: "rate_limited" }, { "retry-after": "1" }); + } + return jsonResponse(200, VALID_PAYLOAD); + }) as unknown as typeof fetch; + + const t0 = Date.now(); + const report = await claudeUsageProvider.fetchUsage(baseParams(), makeContext(fetchMock)); + const elapsed = Date.now() - t0; + expect(report).not.toBeNull(); + expect(attempt).toBe(2); + // Allow generous slop (Bun scheduling jitter) but ensure we actually waited + // closer to the Retry-After than to the default 500ms backoff. + expect(elapsed).toBeGreaterThanOrEqual(800); + expect(callTimes[1] - callTimes[0]).toBeGreaterThanOrEqual(800); + }); + + it("aborts the retry sleep when the signal fires mid-backoff", async () => { + let attempt = 0; + const fetchMock = (async (_url: string | URL, init?: RequestInit) => { + attempt += 1; + if (init?.signal?.aborted) throw new Error("AbortError"); + if (attempt === 1) { + // Pretend Anthropic wants us to back off for 60s. Without + // `scheduler.wait({ signal })` the provider would stall through + // the timeout; with it, the abort rejects the sleep promptly. + return jsonResponse(429, { error: "rate_limited" }, { "retry-after": "60" }); + } + return jsonResponse(200, VALID_PAYLOAD); + }) as unknown as typeof fetch; + + const controller = new AbortController(); + setTimeoutCb(() => controller.abort(), 150); + + const t0 = Date.now(); + const report = await claudeUsageProvider.fetchUsage( + { ...baseParams(), signal: controller.signal }, + makeContext(fetchMock), + ); + const elapsed = Date.now() - t0; + expect(report).toBeNull(); + expect(elapsed).toBeLessThan(3_000); + expect(attempt).toBe(1); + }); + + it("falls back to lastPayload when retries exhausted with stale-but-valid data", async () => { + // Provider keeps lastPayload across attempts — if the upstream returns + // a 200 with a recognized shape but no usage data, we keep iterating. + // If we then 429 forever, we return what we have (null in this case). + let attempt = 0; + const fetchMock = (async () => { + attempt += 1; + if (attempt === 1) { + // 200 OK but no usage payload — provider continues to next attempt + // (waiting for fresh data) rather than returning immediately. + return jsonResponse(200, {}); + } + return jsonResponse(429, { error: "rate_limited" }); + }) as unknown as typeof fetch; + + const report = await claudeUsageProvider.fetchUsage(baseParams(), makeContext(fetchMock)); + // The 200 set lastPayload but had no usage data; 429s mean no further + // successes. lastPayload survives but has no usage data → no limits. + // Specifically: report is null (since lastPayload has nothing to expose). + expect(report).toBeNull(); + expect(attempt).toBe(3); + }); +}); diff --git a/packages/ai/test/openai-codex-usage.test.ts b/packages/ai/test/openai-codex-usage.test.ts new file mode 100644 index 000000000..626646402 --- /dev/null +++ b/packages/ai/test/openai-codex-usage.test.ts @@ -0,0 +1,128 @@ +/** + * Codex usage parser regressions. The widget client (osx-widgets) keys spark + * detection off `limit.id.includes("spark")`, so the parser MUST surface + * `additional_rate_limits[].metered_feature == "codex_bengalfox"` (the upstream + * codename for GPT-5.3-Codex-Spark) as separate `UsageLimit` entries with + * `spark` in the id. If this contract breaks, both the TUI and the macOS + * widget lose per-model visibility. + */ +import { describe, expect, it } from "bun:test"; +import { openaiCodexUsageProvider } from "../src/usage/openai-codex"; + +const accessTokenFixture = (() => { + const header = Buffer.from(JSON.stringify({ alg: "none", typ: "JWT" })).toString("base64url"); + const body = Buffer.from( + JSON.stringify({ + "https://api.openai.com/auth": { chatgpt_account_id: "acct-fixture" }, + "https://api.openai.com/profile": { email: "fixture@example.com" }, + }), + ).toString("base64url"); + return `${header}.${body}.sig`; +})(); + +function makePayload() { + return { + plan_type: "pro", + rate_limit: { + allowed: true, + limit_reached: false, + primary_window: { used_percent: 4, limit_window_seconds: 17940, reset_at: 2_000_000_000 }, + secondary_window: { used_percent: 1, limit_window_seconds: 604740, reset_at: 2_000_500_000 }, + }, + additional_rate_limits: [ + { + limit_name: "GPT-5.3-Codex-Spark", + metered_feature: "codex_bengalfox", + rate_limit: { + allowed: true, + limit_reached: false, + primary_window: { used_percent: 17, limit_window_seconds: 18000, reset_at: 2_000_001_000 }, + secondary_window: { used_percent: 61, limit_window_seconds: 604800, reset_at: 2_000_600_000 }, + }, + }, + ], + }; +} + +function fakeFetch(payload: unknown): typeof fetch { + const fn = async () => + new Response(JSON.stringify(payload), { status: 200, headers: { "content-type": "application/json" } }); + return fn as unknown as typeof fetch; +} + +describe("openai-codex usage parser", () => { + it("emits primary + secondary limits from the main rate_limit block", async () => { + const report = await openaiCodexUsageProvider.fetchUsage( + { + provider: "openai-codex", + credential: { type: "oauth", accessToken: accessTokenFixture, accountId: "acct-1", email: "u@example.com" }, + }, + { fetch: fakeFetch(makePayload()) }, + ); + expect(report).not.toBeNull(); + const main = report?.limits.filter(l => l.id === "openai-codex:primary" || l.id === "openai-codex:secondary"); + expect(main?.map(l => l.id)).toEqual(["openai-codex:primary", "openai-codex:secondary"]); + expect(main?.[0].scope.tier).toBe("pro"); + expect(main?.[0].amount.usedFraction).toBeCloseTo(0.04, 5); + }); + + it("surfaces additional_rate_limits as spark UsageLimit entries the widget can detect", async () => { + const report = await openaiCodexUsageProvider.fetchUsage( + { + provider: "openai-codex", + credential: { type: "oauth", accessToken: accessTokenFixture, accountId: "acct-1", email: "u@example.com" }, + }, + { fetch: fakeFetch(makePayload()) }, + ); + const spark = report?.limits.filter(l => l.id.includes("spark")); + expect(spark?.map(l => l.id)).toEqual(["openai-codex:spark:primary", "openai-codex:spark:secondary"]); + expect(spark?.[0].label).toBe("5 hours (Spark)"); + expect(spark?.[1].label).toBe("7 days (Spark)"); + expect(spark?.[0].scope.tier).toBe("spark"); + expect(spark?.[0].scope.modelId).toBe("GPT-5.3-Codex-Spark"); + expect(spark?.[0].amount.usedFraction).toBeCloseTo(0.17, 5); + expect(spark?.[1].amount.usedFraction).toBeCloseTo(0.61, 5); + }); + + it("treats bengalfox codename as spark even without explicit limit_name", async () => { + const payload = makePayload(); + payload.additional_rate_limits[0].limit_name = undefined as unknown as string; + const report = await openaiCodexUsageProvider.fetchUsage( + { + provider: "openai-codex", + credential: { type: "oauth", accessToken: accessTokenFixture, accountId: "acct-1", email: "u@example.com" }, + }, + { fetch: fakeFetch(payload) }, + ); + const spark = report?.limits.find(l => l.id === "openai-codex:spark:primary"); + expect(spark).toBeTruthy(); + expect(spark?.scope.tier).toBe("spark"); + }); + + it("returns a report even when only additional_rate_limits are present (no main rate_limit)", async () => { + const payload = { + plan_type: "pro", + rate_limit: null, + additional_rate_limits: [ + { + limit_name: "GPT-5.3-Codex-Spark", + metered_feature: "codex_bengalfox", + rate_limit: { + allowed: true, + limit_reached: false, + primary_window: { used_percent: 5, limit_window_seconds: 18000, reset_at: 2_000_000_000 }, + }, + }, + ], + }; + const report = await openaiCodexUsageProvider.fetchUsage( + { + provider: "openai-codex", + credential: { type: "oauth", accessToken: accessTokenFixture, accountId: "acct-1", email: "u@example.com" }, + }, + { fetch: fakeFetch(payload) }, + ); + expect(report).not.toBeNull(); + expect(report?.limits.map(l => l.id)).toEqual(["openai-codex:spark:primary"]); + }); +}); diff --git a/packages/coding-agent/CHANGELOG.md b/packages/coding-agent/CHANGELOG.md index 589720019..133cdf667 100644 --- a/packages/coding-agent/CHANGELOG.md +++ b/packages/coding-agent/CHANGELOG.md @@ -10,6 +10,7 @@ ### Added +- `ModelRegistry` now promotes `models.yml` `providers..apiKey` entries to `AuthStorage`'s new config-override tier (above OAuth, below `--api-key`). Pinning a bearer in `models.yml` was previously a no-op when the broker had an OAuth credential for the same provider — the OAuth access token won and got sent unmodified to whatever `baseUrl` you redirected to, which an auth-gateway in front of that endpoint rightly rejected with 401. The override is now honored, and is cleared/repopulated atomically on `models.yml` reload (`#reloadStaticModels` calls `clearConfigApiKeys` before re-parsing). Use case: route `anthropic` / `openai-codex` to `http://llm-gateway.internal:4000` with the gateway's own bearer. - Added `omp auth-broker` subcommand for running and consuming a hosted credential vault. - `serve [--bind=host:port]` — boots a local broker against the SQLite store at `$AGENT_DB_PATH`. - `token [--regenerate]` — prints (and rotates) the bearer token stored at `~/.omp/auth-broker.token`. @@ -18,11 +19,24 @@ - `import [--provider=] [--include-disabled] [--dry-run]` — imports CLIProxyAPI-style JSON credential dumps (`~/.cliproxy/auth/*.json`). When `OMP_AUTH_BROKER_URL` is configured, credentials are uploaded to the remote broker via `POST /v1/credential`; otherwise they go into the local SQLite store. JSON `type` is mapped to omp providers (`claude` → `anthropic`, `codex` → `openai-codex`, `gemini[-cli]` → `google-gemini-cli`, `antigravity` → `google-antigravity`); `--provider` overrides the mapping for unrecognized types. - `status` — pings the configured remote broker (`OMP_AUTH_BROKER_URL`). - Added remote credential vault support to `discoverAuthStorage`. Configure via env (`OMP_AUTH_BROKER_URL` / `OMP_AUTH_BROKER_TOKEN`) or by setting `auth.broker.url` and `auth.broker.token` in `~/.omp/agent/config.yml` (hidden from the settings UI; supports `!command` resolution). Falls back to `~/.omp/auth-broker.token` when no token is provided inline. Otherwise behavior is unchanged. +- Added `omp auth-broker migrate --from-local [--include-env] [--include-oauth] [--dry-run]` — uploads local SQLite credentials (and optionally env-var API keys) to the configured broker. Skips anything already on the broker via identity-key matching. OAuth is skipped by default (handled via `cliproxy` import). Idempotent on re-runs. +- Added `omp auth-gateway` subcommand for running a forward-proxy that hides access tokens from less-trusted clients: + - `serve [--bind=…]` — boots the gateway against the configured broker. Listens on `127.0.0.1:4000` by default. + - `token [--regenerate]` — manages the gateway bearer token at `~/.omp/auth-gateway.token` (separate from the broker bearer). + - `status` — verifies gateway config and authenticated broker readiness. + - One wire surface: `POST /v1/chat/completions` (OpenAI chat-completions), `POST /v1/messages` (Anthropic messages), `POST /v1/responses` (OpenAI Responses), `GET /v1/usage` (aggregated, 30s-cached), `GET /v1/models` (catalog). Model id in the request body selects which omp provider/model services it; the gateway translates wire format ↔ omp canonical `Context` and dispatches through `pi-ai` `streamSimple()`. Container deployments (robomp, etc.) get inference auth without ever holding access tokens or the broker bearer. ### Changed - Changed TTSR `interruptMode` semantics so a non-interrupting decision on a tool-source match now folds the rule reminder into that specific tool's `toolResult` content instead of queuing a loop-wide deferred follow-up turn. Text/thinking matches keep the previous deferred-injection behavior. +### Fixed + +- Fixed `omp auth-gateway` request handling to reject unsupported OpenAI/Anthropic protocol controls with 400 instead of accepting and ignoring them, propagate upstream error/abort terminal states as failures, preserve Responses reasoning and completed text items, accept string/system Responses messages, and keep Anthropic tool-result ordering valid. +- Fixed gateway usage reporting to include cached-token totals for OpenAI Chat/Responses and to serve the last good cached report during transient upstream usage fetch failures. +- Fixed auth-gateway request cancellation for requests that are already aborted before dispatch. +- Fixed `/login` and `/logout` provider selector overflowing tall provider lists off-screen on small terminals. The selector now scrolls a 10-item window centered on the highlighted entry, shows a `(n/total)` indicator when windowed, and accepts PageUp/PageDown for faster navigation. + ## [15.1.2] - 2026-05-15 ### Fixed diff --git a/packages/coding-agent/src/cli.ts b/packages/coding-agent/src/cli.ts index fe077137f..d9dd04c88 100755 --- a/packages/coding-agent/src/cli.ts +++ b/packages/coding-agent/src/cli.ts @@ -18,7 +18,7 @@ procmgr.scrubProcessEnv(); * CLI entry point — registers all commands explicitly and delegates to the * lightweight CLI runner from pi-utils. */ -import { type CommandEntry, run } from "@oh-my-pi/pi-utils/cli"; +import { type CliConfig, type CommandEntry, run } from "@oh-my-pi/pi-utils/cli"; if (Bun.semver.order(Bun.version, MIN_BUN_VERSION) < 0) { process.stderr.write( @@ -33,6 +33,7 @@ const commands: CommandEntry[] = [ { name: "launch", load: () => import("./commands/launch").then(m => m.default) }, { name: "acp", load: () => import("./commands/acp").then(m => m.default) }, { name: "auth-broker", load: () => import("./commands/auth-broker").then(m => m.default) }, + { name: "auth-gateway", load: () => import("./commands/auth-gateway").then(m => m.default) }, { name: "agents", load: () => import("./commands/agents").then(m => m.default) }, { name: "commit", load: () => import("./commands/commit").then(m => m.default) }, { name: "config", load: () => import("./commands/config").then(m => m.default) }, @@ -48,7 +49,7 @@ const commands: CommandEntry[] = [ { name: "search", load: () => import("./commands/web-search").then(m => m.default), aliases: ["q"] }, ]; -async function showHelp(config: import("@oh-my-pi/pi-utils/cli").CliConfig): Promise { +async function showHelp(config: CliConfig): Promise { const { renderRootHelp } = await import("@oh-my-pi/pi-utils/cli"); const { getExtraHelpText } = await import("./cli/args"); renderRootHelp(config); diff --git a/packages/coding-agent/src/cli/auth-broker-cli.ts b/packages/coding-agent/src/cli/auth-broker-cli.ts index a354d64e8..c4973c1bf 100644 --- a/packages/coding-agent/src/cli/auth-broker-cli.ts +++ b/packages/coding-agent/src/cli/auth-broker-cli.ts @@ -8,6 +8,9 @@ * via SSH tunnel into a remote broker host. * - `import ` — imports CLIProxyAPI-style JSON credentials into * the local SQLite store (typical use: `import ~/.cliproxy/auth`). + * - `migrate --from-local [--include-env] [--include-oauth] [--dry-run]` — + * uploads local SQLite + env API keys to the broker, skipping anything + * the broker already has. * - `status` — health-pings the configured remote broker. */ import * as crypto from "node:crypto"; @@ -16,10 +19,13 @@ import * as os from "node:os"; import * as path from "node:path"; import { AuthBrokerClient, + type AuthCredential, AuthStorage, type CredentialDisabledEvent, DEFAULT_AUTH_BROKER_BIND, + getEnvApiKey, getOAuthProviders, + listProvidersWithEnvKey, type OAuthCredential, type OAuthProvider, SqliteAuthCredentialStore, @@ -30,7 +36,7 @@ import { $ } from "bun"; import chalk from "chalk"; import { resolveAuthBrokerConfig } from "../session/auth-broker-config"; -export type AuthBrokerAction = "serve" | "token" | "login" | "logout" | "status" | "import"; +export type AuthBrokerAction = "serve" | "token" | "login" | "logout" | "status" | "import" | "migrate"; export interface AuthBrokerCommandArgs { action: AuthBrokerAction; @@ -45,10 +51,16 @@ export interface AuthBrokerCommandArgs { source?: string; /** `import`: keep credentials whose JSON had `disabled: true`. */ includeDisabled?: boolean; + /** `migrate`: also upload local OAuth (default: api_key only, since OAuth is via cliproxy import). */ + includeOauth?: boolean; + /** `migrate`: also capture env-var API keys for providers not yet on broker. */ + includeEnv?: boolean; + /** `migrate`: required `--from-local` source. Reserved for future sources. */ + fromLocal?: boolean; }; } -const ACTIONS: readonly AuthBrokerAction[] = ["serve", "token", "login", "logout", "import", "status"]; +const ACTIONS: readonly AuthBrokerAction[] = ["serve", "token", "login", "logout", "import", "migrate", "status"]; /** Callback ports baked from the per-provider OAuth flow modules. */ const CALLBACK_PORTS: Record = { @@ -453,6 +465,204 @@ async function runImport(flags: AuthBrokerCommandArgs["flags"]): Promise { } } +// ─── Migrate: local SQLite + env → broker ────────────────────────────── + +interface MigratePlanEntry { + source: "local-sqlite" | "env"; + provider: string; + credential: AuthCredential; + identity: string; +} + +interface MigrateSkip { + source: "local-sqlite" | "env"; + provider: string; + identity: string; + reason: string; +} + +function credentialIdentity(provider: string, credential: AuthCredential): string { + if (credential.type === "api_key") return "(api key)"; + return credential.email ?? credential.accountId ?? credential.projectId ?? `<${provider} oauth>`; +} + +/** + * Build the set of "identities already on the broker" so re-runs are idempotent. + * For OAuth, identity = email|accountId|projectId. For api_key, we collapse + * to a single marker per provider (broker has no concept of "multiple api keys + * per provider with different identities"; upsert would coalesce them). + */ +function indexBrokerSnapshot(snapshot: { + credentials: Array<{ + provider: string; + credential: { type: string; email?: string; accountId?: string; projectId?: string }; + }>; +}): Map> { + const out = new Map>(); + for (const entry of snapshot.credentials) { + const ids = out.get(entry.provider) ?? new Set(); + if (entry.credential.type === "api_key") { + ids.add("@api_key"); + } else { + if (entry.credential.email) ids.add(`email:${entry.credential.email}`); + if (entry.credential.accountId) ids.add(`accountId:${entry.credential.accountId}`); + if (entry.credential.projectId) ids.add(`projectId:${entry.credential.projectId}`); + } + out.set(entry.provider, ids); + } + return out; +} + +function brokerAlreadyHas(existing: Map>, provider: string, credential: AuthCredential): boolean { + const ids = existing.get(provider); + if (!ids) return false; + if (credential.type === "api_key") return ids.has("@api_key"); + if (credential.email && ids.has(`email:${credential.email}`)) return true; + if (credential.accountId && ids.has(`accountId:${credential.accountId}`)) return true; + if (credential.projectId && ids.has(`projectId:${credential.projectId}`)) return true; + return false; +} + +async function runMigrate(flags: AuthBrokerCommandArgs["flags"]): Promise { + const brokerConfig = await resolveAuthBrokerConfig(); + if (!brokerConfig) { + throw new Error( + "OMP_AUTH_BROKER_URL must be set (or `auth.broker.url` in config.yml). `migrate` uploads local credentials to a configured broker.", + ); + } + if (flags.fromLocal !== true) { + throw new Error( + "`omp auth-broker migrate` requires an explicit source. Pass `--from-local` to migrate from the local SQLite store and env vars.", + ); + } + + const client = new AuthBrokerClient({ url: brokerConfig.url, token: brokerConfig.token }); + const snapshot = await client.fetchSnapshot(); + const existing = indexBrokerSnapshot(snapshot); + + const plan: MigratePlanEntry[] = []; + const skipped: MigrateSkip[] = []; + + // 1. Local SQLite rows. + const localDbPath = getAgentDbPath(); + const localStore = await SqliteAuthCredentialStore.open(localDbPath); + const plannedApiKeyProviders = new Set(); + try { + for (const row of localStore.listAuthCredentials()) { + const identity = credentialIdentity(row.provider, row.credential); + if (row.credential.type === "oauth" && flags.includeOauth !== true) { + skipped.push({ + source: "local-sqlite", + provider: row.provider, + identity, + reason: "OAuth from local SQLite skipped by default (use --include-oauth)", + }); + continue; + } + if (brokerAlreadyHas(existing, row.provider, row.credential)) { + skipped.push({ + source: "local-sqlite", + provider: row.provider, + identity, + reason: "already on broker", + }); + continue; + } + if (row.credential.type === "api_key" && plannedApiKeyProviders.has(row.provider)) { + skipped.push({ + source: "local-sqlite", + provider: row.provider, + identity, + reason: "another local api_key for this provider already planned", + }); + continue; + } + if (row.credential.type === "api_key") plannedApiKeyProviders.add(row.provider); + plan.push({ source: "local-sqlite", provider: row.provider, credential: row.credential, identity }); + } + } finally { + localStore.close(); + } + + // 2. Env-var API keys (opt-in). + if (flags.includeEnv === true) { + for (const provider of listProvidersWithEnvKey()) { + const envValue = getEnvApiKey(provider); + if (!envValue) continue; + if (envValue === "") continue; // Bedrock/Vertex sentinels — not literal keys. + const credential: AuthCredential = { type: "api_key", key: envValue }; + if (brokerAlreadyHas(existing, provider, credential)) { + skipped.push({ + source: "env", + provider, + identity: "(api key)", + reason: "already on broker (provider has an api_key)", + }); + continue; + } + // Also skip if local SQLite already produced an entry for this provider in this batch. + if (plan.some(p => p.provider === provider && p.credential.type === "api_key")) { + skipped.push({ + source: "env", + provider, + identity: "(api key)", + reason: "local SQLite already supplied an api_key for this provider", + }); + continue; + } + plan.push({ source: "env", provider, credential, identity: "(api key)" }); + } + } + + if (flags.json) { + process.stdout.write( + `${JSON.stringify({ + dryRun: flags.dryRun === true, + plan: plan.map(p => ({ source: p.source, provider: p.provider, identity: p.identity })), + skipped, + })}\n`, + ); + } else { + for (const skip of skipped) { + process.stdout.write( + `${chalk.yellow("skip")} [${skip.source}] ${skip.provider} ${skip.identity}: ${skip.reason}\n`, + ); + } + } + + if (plan.length === 0) { + if (!flags.json) process.stdout.write("Nothing to migrate.\n"); + return; + } + + if (flags.dryRun === true) { + if (!flags.json) { + process.stdout.write(`Dry run — would upload ${plan.length} credential(s):\n`); + for (const entry of plan) { + process.stdout.write(` [${entry.source}] ${entry.provider} ${entry.identity}\n`); + } + } + return; + } + + for (const entry of plan) { + try { + await client.uploadCredential(entry.provider, entry.credential); + if (!flags.json) { + process.stdout.write(`${chalk.green("uploaded")} [${entry.source}] ${entry.provider} ${entry.identity}\n`); + } + } catch (error) { + const message = error instanceof Error ? error.message : String(error); + if (flags.json) { + process.stdout.write(`${JSON.stringify({ error: message, provider: entry.provider })}\n`); + } else { + process.stdout.write(`${chalk.red("failed")} [${entry.source}] ${entry.provider}: ${message}\n`); + } + process.exitCode = 1; + } + } +} + async function runStatus(flags: AuthBrokerCommandArgs["flags"]): Promise { const cfg = await resolveAuthBrokerConfig(); if (!cfg) { @@ -497,6 +707,9 @@ export async function runAuthBrokerCommand(cmd: AuthBrokerCommandArgs): Promise< case "import": await runImport(cmd.flags); return; + case "migrate": + await runMigrate(cmd.flags); + return; case "status": await runStatus(cmd.flags); return; diff --git a/packages/coding-agent/src/cli/auth-gateway-cli.ts b/packages/coding-agent/src/cli/auth-gateway-cli.ts new file mode 100644 index 000000000..28b369d78 --- /dev/null +++ b/packages/coding-agent/src/cli/auth-gateway-cli.ts @@ -0,0 +1,311 @@ +/** + * `omp auth-gateway` command handlers. + * + * Boots a forward-proxy server that lets less-trusted clients (the macOS + * usage widget, robomp containers, …) make provider API calls without ever + * seeing the access token. The gateway is itself a broker client and + * resolves credentials through the configured broker (via the same + * `OMP_AUTH_BROKER_URL` / `auth.broker.url` precedence used elsewhere). + * + * Sub-verbs: + * - `serve [--bind=…]` — boots the gateway against the configured broker. + * - `token` / `token --regenerate` — manages the gateway bearer token file. + * - `status` — prints the locally-stored gateway token and bind hint. + */ +import * as crypto from "node:crypto"; +import * as fs from "node:fs/promises"; +import * as path from "node:path"; +import { + type Api, + AuthBrokerClient, + AuthStorage, + DEFAULT_AUTH_GATEWAY_BIND, + type GeneratedProvider, + getBundledModels, + getBundledProviders, + type Model, + RemoteAuthCredentialStore, + type SnapshotResponse, + startAuthGateway, +} from "@oh-my-pi/pi-ai"; +import { getConfigRootDir, isEnoent, VERSION } from "@oh-my-pi/pi-utils"; +import chalk from "chalk"; +import { type AuthBrokerClientConfig, resolveAuthBrokerConfig } from "../session/auth-broker-config"; + +export type AuthGatewayAction = "serve" | "token" | "status"; + +export interface AuthGatewayCommandArgs { + action: AuthGatewayAction; + flags: { + json?: boolean; + bind?: string; + regenerate?: boolean; + /** + * Disable bearer-token auth on inbound requests. Useful when the gateway + * is bound to loopback (the default `127.0.0.1:4000`) and you don't want + * to wire token-paste plumbing into every local client. + */ + noAuth?: boolean; + }; +} + +const ACTIONS: readonly AuthGatewayAction[] = ["serve", "token", "status"]; + +function getTokenFilePath(): string { + return path.join(getConfigRootDir(), "auth-gateway.token"); +} + +async function readToken(): Promise { + try { + const raw = await Bun.file(getTokenFilePath()).text(); + const trimmed = raw.trim(); + return trimmed.length > 0 ? trimmed : null; + } catch (err) { + if (isEnoent(err)) return null; + throw err; + } +} + +async function writeToken(token: string): Promise { + const file = getTokenFilePath(); + await fs.mkdir(path.dirname(file), { recursive: true, mode: 0o700 }); + await fs.writeFile(file, token, { mode: 0o600 }); + try { + await fs.chmod(file, 0o600); + } catch { + // Best-effort (e.g. Windows). + } +} + +function generateToken(): string { + return crypto.randomBytes(32).toString("base64url"); +} + +async function ensureToken(): Promise { + const existing = await readToken(); + if (existing) return existing; + const token = generateToken(); + await writeToken(token); + return token; +} + +function createBrokerClient(brokerConfig: AuthBrokerClientConfig): AuthBrokerClient { + return new AuthBrokerClient({ url: brokerConfig.url, token: brokerConfig.token }); +} + +async function fetchBrokerSnapshot(client: AuthBrokerClient): Promise { + return client.fetchSnapshot(); +} + +async function runServe(flags: AuthGatewayCommandArgs["flags"]): Promise { + const brokerConfig = await resolveAuthBrokerConfig(); + if (!brokerConfig) { + throw new Error( + "`omp auth-gateway serve` requires OMP_AUTH_BROKER_URL (or `auth.broker.url`/`auth.broker.token` in config.yml). The gateway is itself a broker client.", + ); + } + const bind = flags.bind ?? DEFAULT_AUTH_GATEWAY_BIND; + const gatewayToken = flags.noAuth ? null : await ensureToken(); + + // Build a broker-backed AuthStorage — same pattern as discoverAuthStorage() + // in sdk.ts. The gateway never touches local SQLite. + const client = createBrokerClient(brokerConfig); + const initialSnapshot = await fetchBrokerSnapshot(client); + const store = new RemoteAuthCredentialStore({ client, initialSnapshot }); + // Refresh + usage both flow through the store's broker hooks automatically — + // `RemoteAuthCredentialStore.refreshOAuthCredential` and `.fetchUsageReports`. + // AuthStorage discovers them when no explicit option overrides them, so the + // gateway only needs to construct the store and pass it in. + const storage = new AuthStorage(store, { + sourceLabel: `broker ${brokerConfig.url}`, + }); + await storage.reload(); + + // Build the model resolver + catalog from pi-ai's bundled metadata, scoped + // to providers we hold credentials for. Format handlers ask `resolveModel` + // to translate a client-requested `model` field into a pi-ai `Model` + // before dispatch; `listModels` powers `/v1/models`. + const snapshot = storage.exportSnapshot(); + const providersWithCreds = new Set(); + for (const entry of snapshot.credentials) providersWithCreds.add(entry.provider); + const modelById = new Map>(); + for (const provider of getBundledProviders()) { + if (!providersWithCreds.has(provider)) continue; + for (const model of getBundledModels(provider as GeneratedProvider)) { + // First-write-wins so a canonical model id collisions across providers + // stick to the provider listed first by getBundledProviders. + if (!modelById.has(model.id)) modelById.set(model.id, model); + } + } + + const handle = startAuthGateway({ + storage, + bind, + bearerTokens: gatewayToken ? [gatewayToken] : [], + version: VERSION, + resolveModel: (id: string) => modelById.get(id), + listModels: () => modelById.values(), + }); + process.stdout.write(`auth-gateway listening on ${handle.url}\n`); + if (gatewayToken) { + process.stdout.write(`bearer token: ${getTokenFilePath()} (chmod 0600)\n`); + } else { + process.stdout.write(`auth: disabled (--no-auth) — any client can call this gateway\n`); + } + process.stdout.write(`upstream broker: ${brokerConfig.url}\n`); + + const stopped = Promise.withResolvers(); + let shutdownStarted = false; + const stop = async (signal: NodeJS.Signals): Promise => { + if (shutdownStarted) return; + shutdownStarted = true; + process.stdout.write(`\nReceived ${signal}, shutting down...\n`); + let closeError: unknown; + try { + await handle.close(); + } catch (error) { + closeError = error; + } finally { + storage.close(); + } + if (closeError) { + stopped.reject(closeError); + } else { + stopped.resolve(); + } + }; + const onSigint = (): void => { + void stop("SIGINT"); + }; + const onSigterm = (): void => { + void stop("SIGTERM"); + }; + process.once("SIGINT", onSigint); + process.once("SIGTERM", onSigterm); + + try { + await stopped.promise; + } finally { + process.off("SIGINT", onSigint); + process.off("SIGTERM", onSigterm); + } +} + +async function runToken(flags: AuthGatewayCommandArgs["flags"]): Promise { + if (flags.regenerate) { + const next = generateToken(); + await writeToken(next); + if (flags.json) { + process.stdout.write(`${JSON.stringify({ token: next, path: getTokenFilePath() })}\n`); + } else { + process.stdout.write(`${next}\n`); + } + return; + } + const token = await ensureToken(); + if (flags.json) { + process.stdout.write(`${JSON.stringify({ token, path: getTokenFilePath() })}\n`); + } else { + process.stdout.write(`${token}\n`); + } +} + +async function runStatus(flags: AuthGatewayCommandArgs["flags"]): Promise { + const token = await readToken(); + const brokerConfig = await resolveAuthBrokerConfig(); + const tokenFile = getTokenFilePath(); + if (!brokerConfig) { + const status = { + ready: false, + reason: "not_configured", + tokenFile, + tokenPresent: token !== null, + broker: null, + brokerConfigured: false, + brokerAuthenticated: false, + }; + if (flags.json) { + process.stdout.write(`${JSON.stringify(status)}\n`); + } else { + process.stdout.write(`${chalk.yellow("No broker configured.")} Set OMP_AUTH_BROKER_URL.\n`); + process.stdout.write( + `token: ${status.tokenPresent ? chalk.green("present") : chalk.red("missing")} at ${status.tokenFile}\n`, + ); + } + process.exitCode = 1; + return; + } + + try { + const snapshot = await fetchBrokerSnapshot(createBrokerClient(brokerConfig)); + const tokenPresent = token !== null; + const status = { + ready: tokenPresent, + reason: tokenPresent ? null : "token_missing", + tokenFile, + tokenPresent, + broker: brokerConfig.url, + brokerConfigured: true, + brokerAuthenticated: true, + credentialCount: snapshot.credentials.length, + }; + if (flags.json) { + process.stdout.write(`${JSON.stringify(status)}\n`); + } else { + const brokerLine = `upstream broker: ${brokerConfig.url} (${snapshot.credentials.length} credential${ + snapshot.credentials.length === 1 ? "" : "s" + })`; + process.stdout.write(`${tokenPresent ? chalk.green("ready") : chalk.yellow("not ready")} ${brokerLine}\n`); + process.stdout.write( + `token: ${tokenPresent ? chalk.green("present") : chalk.red("missing")} at ${status.tokenFile}\n`, + ); + if (!tokenPresent) { + process.stdout.write( + "Run `omp auth-gateway token` or `omp auth-gateway serve` to create a bearer token.\n", + ); + } + } + if (!tokenPresent) process.exitCode = 1; + } catch (error) { + const message = error instanceof Error ? error.message : String(error); + const status = { + ready: false, + reason: "broker_unavailable", + tokenFile, + tokenPresent: token !== null, + broker: brokerConfig.url, + brokerConfigured: true, + brokerAuthenticated: false, + error: message, + }; + if (flags.json) { + process.stdout.write(`${JSON.stringify(status)}\n`); + } else { + process.stdout.write(`${chalk.red("FAILED")} upstream broker: ${brokerConfig.url}: ${message}\n`); + process.stdout.write( + `token: ${status.tokenPresent ? chalk.green("present") : chalk.red("missing")} at ${status.tokenFile}\n`, + ); + } + process.exitCode = 1; + } +} + +export async function runAuthGatewayCommand(cmd: AuthGatewayCommandArgs): Promise { + switch (cmd.action) { + case "serve": + await runServe(cmd.flags); + return; + case "token": + await runToken(cmd.flags); + return; + case "status": + await runStatus(cmd.flags); + return; + default: { + const _exhaustive: never = cmd.action; + throw new Error(`Unknown auth-gateway action: ${String(_exhaustive)}`); + } + } +} + +export { ACTIONS as AUTH_GATEWAY_ACTIONS }; diff --git a/packages/coding-agent/src/commands/auth-broker.ts b/packages/coding-agent/src/commands/auth-broker.ts index ddb250836..c6beb1f14 100644 --- a/packages/coding-agent/src/commands/auth-broker.ts +++ b/packages/coding-agent/src/commands/auth-broker.ts @@ -39,7 +39,16 @@ export default class AuthBroker extends Command { "include-disabled": Flags.boolean({ description: "Import credentials whose JSON has `disabled: true` (import)", }), - "dry-run": Flags.boolean({ description: "Print actions without executing (import / login --via)" }), + "from-local": Flags.boolean({ + description: "migrate source: local SQLite + env vars (required for `migrate`)", + }), + "include-env": Flags.boolean({ + description: "Capture env-var API keys for providers not yet on broker (migrate)", + }), + "include-oauth": Flags.boolean({ + description: "Also upload OAuth from local SQLite during migrate (default skips them)", + }), + "dry-run": Flags.boolean({ description: "Print actions without executing (import / login --via / migrate)" }), }; static examples = [ @@ -51,6 +60,8 @@ export default class AuthBroker extends Command { "# Remote login over SSH tunnel\n omp auth-broker login anthropic --via=user@broker", "# Import a CLIProxyAPI auth dump\n omp auth-broker import ~/.cliproxy/auth", "# Import a single CLIProxyAPI JSON, overriding the provider mapping\n omp auth-broker import ~/.cliproxy/auth/claude-foo.json --provider anthropic", + "# Preview a migration from local store + env vars to the configured broker\n omp auth-broker migrate --from-local --include-env --dry-run", + "# Apply the migration\n omp auth-broker migrate --from-local --include-env", "# Health-check the configured remote broker\n omp auth-broker status", ]; @@ -73,6 +84,9 @@ export default class AuthBroker extends Command { provider: action === "import" ? flags.provider : (args.source ?? flags.provider), source: args.source, includeDisabled: flags["include-disabled"], + fromLocal: flags["from-local"], + includeEnv: flags["include-env"], + includeOauth: flags["include-oauth"], dryRun: flags["dry-run"], }, }; diff --git a/packages/coding-agent/src/commands/auth-gateway.ts b/packages/coding-agent/src/commands/auth-gateway.ts new file mode 100644 index 000000000..6b91c52ee --- /dev/null +++ b/packages/coding-agent/src/commands/auth-gateway.ts @@ -0,0 +1,61 @@ +/** + * `omp auth-gateway` — run a forward proxy that injects auth from the broker. + */ +import { Args, Command, Flags, renderCommandHelp } from "@oh-my-pi/pi-utils/cli"; +import { + AUTH_GATEWAY_ACTIONS, + type AuthGatewayAction, + type AuthGatewayCommandArgs, + runAuthGatewayCommand, +} from "../cli/auth-gateway-cli"; +import { initTheme } from "../modes/theme/theme"; + +export default class AuthGateway extends Command { + static description = "Run an auth-gateway forward proxy backed by the configured broker"; + + static args = { + action: Args.string({ + description: "Sub-command", + required: false, + options: [...AUTH_GATEWAY_ACTIONS], + }), + }; + + static flags = { + json: Flags.boolean({ description: "Output JSON (token/status)" }), + bind: Flags.string({ description: "Bind address for `serve` (host:port)", char: "b" }), + regenerate: Flags.boolean({ description: "Regenerate the gateway bearer token (token)" }), + "no-auth": Flags.boolean({ + description: + "Disable inbound bearer-token auth (serve). Useful when bound to loopback — any caller is allowed.", + }), + }; + + static examples = [ + "# Boot the gateway against the configured broker\n omp auth-gateway serve", + "# Boot on a non-default port\n omp auth-gateway serve --bind=127.0.0.1:4000", + "# Print the gateway bearer token (creates one on first run)\n omp auth-gateway token", + "# Rotate the gateway bearer token\n omp auth-gateway token --regenerate", + "# Run on loopback without any bearer (anyone on this host can call)\n omp auth-gateway serve --no-auth", + "# Show local gateway + broker config status\n omp auth-gateway status", + ]; + + async run(): Promise { + const { args, flags } = await this.parse(AuthGateway); + if (!args.action) { + renderCommandHelp("omp", "auth-gateway", AuthGateway); + return; + } + const cmd: AuthGatewayCommandArgs = { + action: args.action as AuthGatewayAction, + flags: { + json: flags.json, + bind: flags.bind, + regenerate: flags.regenerate, + noAuth: flags["no-auth"], + }, + }; + await initTheme(); + await runAuthGatewayCommand(cmd); + } +} diff --git a/packages/coding-agent/src/commands/launch.ts b/packages/coding-agent/src/commands/launch.ts index c74392592..8c513ad77 100644 --- a/packages/coding-agent/src/commands/launch.ts +++ b/packages/coding-agent/src/commands/launch.ts @@ -23,7 +23,7 @@ export default class Index extends Command { static flags = { model: Flags.string({ - description: 'Model to use (fuzzy match: "opus", "gpt-5.2", or "p-openai/gpt-5.2")', + description: 'Model to use (fuzzy match: "opus", "gpt-5.2", or "openai/gpt-5.2")', }), smol: Flags.string({ description: "Smol/fast model for lightweight tasks (or PI_SMOL_MODEL env)", diff --git a/packages/coding-agent/src/config/model-registry.ts b/packages/coding-agent/src/config/model-registry.ts index 1a8dd41ee..361ec281b 100644 --- a/packages/coding-agent/src/config/model-registry.ts +++ b/packages/coding-agent/src/config/model-registry.ts @@ -792,6 +792,10 @@ export class ModelRegistry { this.#customProviderApiKeys.clear(); this.#keylessProviders.clear(); this.#discoverableProviders = []; + // Drop config-sourced apiKeys from AuthStorage before reload; entries + // removed from models.yml must actually disappear from the resolver, not + // linger from the previous parse. The post-load setters below repopulate. + this.authStorage.clearConfigApiKeys(); // Restore runtime API keys before #loadModels — survives because // #loadModels only calls .set() on #customProviderApiKeys, never reassigns it. for (const [k, v] of this.#runtimeProviderApiKeys) { @@ -1117,9 +1121,14 @@ export class ModelRegistry { }); } - // Always store API key for fallback resolver + // Store API key for fallback resolver AND register as config override + // so it wins over OAuth tokens from the broker — when the user pins a + // bearer in models.yml (e.g. for an auth-gateway baseUrl), that bearer + // must authenticate the outbound request. if (providerConfig.apiKey) { this.#customProviderApiKeys.set(providerName, providerConfig.apiKey); + const resolved = resolveApiKeyConfig(providerConfig.apiKey); + if (resolved) this.authStorage.setConfigApiKey(providerName, resolved); } // Parse per-model overrides @@ -1766,6 +1775,8 @@ export class ModelRegistry { if (modelDefs.length === 0) continue; // Override-only, no custom models if (providerConfig.apiKey) { this.#customProviderApiKeys.set(providerName, providerConfig.apiKey); + const resolved = resolveApiKeyConfig(providerConfig.apiKey); + if (resolved) this.authStorage.setConfigApiKey(providerName, resolved); } for (const modelDef of modelDefs) { const providerCompat = providerConfig.disableStrictTools @@ -2008,6 +2019,7 @@ export class ModelRegistry { this.#runtimeProviderApiKeys.delete(providerName); this.#runtimeProviderOverrides.delete(providerName); this.#runtimeModelOverlays = this.#runtimeModelOverlays.filter(overlay => overlay.provider !== providerName); + this.authStorage.removeConfigApiKey(providerName); } /** @@ -2115,6 +2127,8 @@ export class ModelRegistry { this.#customProviderApiKeys.set(providerName, config.apiKey); // Persist runtime API keys so they survive #reloadStaticModels() cycles this.#runtimeProviderApiKeys.set(providerName, config.apiKey); + const resolved = resolveApiKeyConfig(config.apiKey); + if (resolved) this.authStorage.setConfigApiKey(providerName, resolved); } if (config.models && config.models.length > 0) { diff --git a/packages/coding-agent/src/modes/components/oauth-selector.ts b/packages/coding-agent/src/modes/components/oauth-selector.ts index dbdf3d1f4..b0cd6ca5f 100644 --- a/packages/coding-agent/src/modes/components/oauth-selector.ts +++ b/packages/coding-agent/src/modes/components/oauth-selector.ts @@ -5,6 +5,8 @@ import { theme } from "../../modes/theme/theme"; import { matchesSelectCancel } from "../../modes/utils/keybinding-matchers"; import type { AuthStorage } from "../../session/auth-storage"; import { DynamicBorder } from "./dynamic-border"; + +const OAUTH_SELECTOR_MAX_VISIBLE = 10; /** * Component that renders an OAuth provider selector. */ @@ -144,7 +146,16 @@ export class OAuthSelectorComponent extends Container { } #updateList(): void { this.#listContainer.clear(); - for (let i = 0; i < this.#allProviders.length; i++) { + + const total = this.#allProviders.length; + const maxVisible = OAUTH_SELECTOR_MAX_VISIBLE; + const startIndex = + total <= maxVisible + ? 0 + : Math.max(0, Math.min(this.#selectedIndex - Math.floor(maxVisible / 2), total - maxVisible)); + const endIndex = Math.min(startIndex + maxVisible, total); + + for (let i = startIndex; i < endIndex; i++) { const provider = this.#allProviders[i]; if (!provider) continue; const isSelected = i === this.#selectedIndex; @@ -163,8 +174,14 @@ export class OAuthSelectorComponent extends Container { this.#listContainer.addChild(new TruncatedText(line, 0, 0)); } + // Scroll indicator when list is windowed + if (startIndex > 0 || endIndex < total) { + const scrollInfo = theme.fg("muted", ` (${this.#selectedIndex + 1}/${total})`); + this.#listContainer.addChild(new TruncatedText(scrollInfo, 0, 0)); + } + // Show "no providers" if empty - if (this.#allProviders.length === 0) { + if (total === 0) { const message = this.#mode === "login" ? "No OAuth providers available" : "No OAuth providers logged in. Use /login first."; this.#listContainer.addChild(new TruncatedText(theme.fg("muted", ` ${message}`), 0, 0)); @@ -191,6 +208,25 @@ export class OAuthSelectorComponent extends Container { this.#statusMessage = undefined; this.#updateList(); } + // Page up - jump up by one visible page + else if (matchesKey(keyData, "pageUp")) { + if (this.#allProviders.length > 0) { + this.#selectedIndex = Math.max(0, this.#selectedIndex - OAUTH_SELECTOR_MAX_VISIBLE); + } + this.#statusMessage = undefined; + this.#updateList(); + } + // Page down - jump down by one visible page + else if (matchesKey(keyData, "pageDown")) { + if (this.#allProviders.length > 0) { + this.#selectedIndex = Math.min( + this.#allProviders.length - 1, + this.#selectedIndex + OAUTH_SELECTOR_MAX_VISIBLE, + ); + } + this.#statusMessage = undefined; + this.#updateList(); + } // Enter else if (matchesKey(keyData, "enter") || matchesKey(keyData, "return") || keyData === "\n") { const selectedProvider = this.#allProviders[this.#selectedIndex]; diff --git a/packages/coding-agent/src/modes/controllers/command-controller.ts b/packages/coding-agent/src/modes/controllers/command-controller.ts index 9f426b675..d6b329dd0 100644 --- a/packages/coding-agent/src/modes/controllers/command-controller.ts +++ b/packages/coding-agent/src/modes/controllers/command-controller.ts @@ -508,7 +508,8 @@ export class CommandController { return; } - const output = renderUsageReports(usageReports, theme, Date.now()); + const availableWidth = Math.max(40, (this.ctx.ui.terminal.columns ?? 100) - 2); + const output = renderUsageReports(usageReports, theme, Date.now(), availableWidth); this.ctx.chatContainer.addChild(new Spacer(1)); this.ctx.chatContainer.addChild(new Text(output, 1, 0)); this.ctx.ui.requestRender(); @@ -1242,8 +1243,8 @@ export class CommandController { } } -const BAR_WIDTH = 24; -const COLUMN_WIDTH = BAR_WIDTH + 2; +const BAR_WIDTH_MAX = 24; +const BAR_WIDTH_MIN = 4; function renderJobLine(job: AsyncJobSnapshotItem, now: number): string { const duration = formatDuration(Math.max(0, now - job.startTime)); @@ -1449,20 +1450,42 @@ function resolveStatusColor(status: UsageLimit["status"]): "success" | "warning" return "dim"; } -function renderUsageBar(limit: UsageLimit, uiTheme: typeof theme): string { +function renderUsageBar(limit: UsageLimit, uiTheme: typeof theme, barWidth: number): string { const fraction = resolveFraction(limit); if (fraction === undefined) { - return uiTheme.fg("dim", `[${"·".repeat(BAR_WIDTH)}]`); + return uiTheme.fg("dim", `[${"·".repeat(barWidth)}]`); } const clamped = Math.min(Math.max(fraction, 0), 1); - const filled = Math.round(clamped * BAR_WIDTH); + const filled = Math.round(clamped * barWidth); const filledBar = "█".repeat(filled); - const emptyBar = "░".repeat(Math.max(0, BAR_WIDTH - filled)); + const emptyBar = "░".repeat(Math.max(0, barWidth - filled)); const color = resolveStatusColor(limit.status); return `${uiTheme.fg("dim", "[")}${uiTheme.fg(color, filledBar)}${uiTheme.fg("dim", emptyBar)}${uiTheme.fg("dim", "]")}`; } -function renderUsageReports(reports: UsageReport[], uiTheme: typeof theme, nowMs: number): string { +/** + * Pick a per-column width so n bars + a trailing amount string fit in `available` columns. + * Falls back to the minimum when the terminal is too narrow rather than wrapping. + */ +function resolveColumnWidth(count: number, available: number, trailing: number): number { + if (count <= 0) return BAR_WIDTH_MAX + 2; + const indent = 2; + const gaps = count - 1; + const spaceForBars = available - indent - gaps - (trailing > 0 ? trailing + 1 : 0); + const ideal = Math.floor(spaceForBars / count); + const min = BAR_WIDTH_MIN + 2; + const max = BAR_WIDTH_MAX + 2; + if (ideal < min) return min; + if (ideal > max) return max; + return ideal; +} + +function renderUsageReports( + reports: UsageReport[], + uiTheme: typeof theme, + nowMs: number, + availableWidth: number, +): string { const lines: string[] = []; const latestFetchedAt = Math.max(...reports.map(report => report.fetchedAt ?? 0)); const headerSuffix = latestFetchedAt ? ` (${formatDuration(nowMs - latestFetchedAt)} ago)` : ""; @@ -1532,12 +1555,18 @@ function renderUsageReports(reports: UsageReport[], uiTheme: typeof theme, nowMs const windowSuffix = formatWindowSuffix(group.label, group.windowLabel, uiTheme); lines.push(`${statusIcon} ${uiTheme.bold(group.label)} ${windowSuffix}`.trim()); + const amountText = formatAggregateAmount(sortedLimits); + const columnWidth = resolveColumnWidth(sortedLimits.length, availableWidth, visibleWidth(amountText)); + const barWidth = columnWidth - 2; const accountLabels = sortedLimits.map((limit, index) => - padColumn(formatAccountHeader(limit, sortedReports[index], index, nowMs), COLUMN_WIDTH), + padColumn( + truncateJobLabel(formatAccountHeader(limit, sortedReports[index], index, nowMs), columnWidth), + columnWidth, + ), ); lines.push(` ${accountLabels.join(" ")}`.trimEnd()); - const bars = sortedLimits.map(limit => padColumn(renderUsageBar(limit, uiTheme), COLUMN_WIDTH)); - lines.push(` ${bars.join(" ")} ${formatAggregateAmount(sortedLimits)}`.trimEnd()); + const bars = sortedLimits.map(limit => padColumn(renderUsageBar(limit, uiTheme, barWidth), columnWidth)); + lines.push(` ${bars.join(" ")} ${amountText}`.trimEnd()); const resetText = sortedLimits.length <= 1 ? resolveResetRange(sortedLimits, nowMs) : null; if (resetText) { lines.push(` ${uiTheme.fg("dim", resetText)}`.trimEnd()); diff --git a/packages/coding-agent/src/sdk.ts b/packages/coding-agent/src/sdk.ts index e5ca3f646..6f3332006 100644 --- a/packages/coding-agent/src/sdk.ts +++ b/packages/coding-agent/src/sdk.ts @@ -94,12 +94,7 @@ import { } from "./secrets"; import { AgentSession } from "./session/agent-session"; import { resolveAuthBrokerConfig } from "./session/auth-broker-config"; -import { - AuthBrokerClient, - AuthStorage, - REMOTE_REFRESH_SENTINEL, - RemoteAuthCredentialStore, -} from "./session/auth-storage"; +import { AuthBrokerClient, AuthStorage, RemoteAuthCredentialStore } from "./session/auth-storage"; import { convertToLlm } from "./session/messages"; import { SessionManager } from "./session/session-manager"; import { closeAllConnections } from "./ssh/connection-manager"; @@ -339,27 +334,11 @@ export async function discoverAuthStorage(agentDir: string = getDefaultAgentDir( const client = new AuthBrokerClient({ url: brokerConfig.url, token: brokerConfig.token }); const initialSnapshot = await client.fetchSnapshot(); const store = new RemoteAuthCredentialStore({ client, initialSnapshot }); + // Refresh + usage hooks live on RemoteAuthCredentialStore; AuthStorage + // discovers them automatically when no explicit option overrides them. const storage = new AuthStorage(store, { configValueResolver: resolveConfigValue, sourceLabel: `broker ${brokerConfig.url}`, - refreshOAuthCredential: async (_provider, credentialId, _credential) => { - const { entry } = await client.refreshCredential(credentialId); - if (entry.credential.type !== "oauth") { - throw new Error(`Broker returned non-OAuth credential for id=${credentialId}`); - } - const refreshed = entry.credential; - return { - access: refreshed.access, - // Sentinel — AuthStorage stores it back into the in-memory snapshot, - // but a refresh through the broker is the only legal way to mint tokens. - refresh: REMOTE_REFRESH_SENTINEL, - expires: refreshed.expires, - accountId: refreshed.accountId, - email: refreshed.email, - projectId: refreshed.projectId, - enterpriseUrl: refreshed.enterpriseUrl, - }; - }, }); await storage.reload(); return storage; diff --git a/packages/coding-agent/test/model-registry.test.ts b/packages/coding-agent/test/model-registry.test.ts index 911d96666..64d4017cd 100644 --- a/packages/coding-agent/test/model-registry.test.ts +++ b/packages/coding-agent/test/model-registry.test.ts @@ -341,11 +341,11 @@ describe("ModelRegistry", () => { test("applies explicit equivalence overrides from config", () => { writeRawModelsConfig({ providers: { - "p-anthropic": providerConfig("https://demo.example.com/v1", [{ id: "corp-sonnet" }]), + "proxy-anthropic": providerConfig("https://demo.example.com/v1", [{ id: "corp-sonnet" }]), }, equivalence: { overrides: { - "p-anthropic/corp-sonnet": "claude-sonnet-4-5", + "proxy-anthropic/corp-sonnet": "claude-sonnet-4-5", }, }, }); @@ -353,7 +353,7 @@ describe("ModelRegistry", () => { const registry = new ModelRegistry(authStorage, modelsJsonPath); const variants = registry.getCanonicalVariants("claude-sonnet-4-5"); - expect(variants.some(variant => variant.selector === "p-anthropic/corp-sonnet")).toBe(true); + expect(variants.some(variant => variant.selector === "proxy-anthropic/corp-sonnet")).toBe(true); }); test("exclusions keep variants out of canonical grouping", () => { @@ -2032,7 +2032,7 @@ describe("ModelRegistry", () => { describe("provider auth: oauth", () => { test("models from a provider with auth: oauth are marked isOAuth=true", async () => { writeRawModelsJson({ - "p-anthropic": { + "proxy-anthropic": { baseUrl: "https://proxy.example.com", apiKey: "literal-key", api: "anthropic-messages", @@ -2050,19 +2050,19 @@ describe("ModelRegistry", () => { ], }, }); - await authStorage.setRuntimeApiKey("p-anthropic", "literal-key"); + await authStorage.setRuntimeApiKey("proxy-anthropic", "literal-key"); const registry = new ModelRegistry(authStorage, modelsJsonPath); await registry.refresh("offline"); - const model = registry.find("p-anthropic", "claude-sonnet-4-5"); + const model = registry.find("proxy-anthropic", "claude-sonnet-4-5"); expect(model).toBeDefined(); expect(model?.isOAuth).toBe(true); }); test("anthropic-messages providers default to isOAuth=true even without explicit auth", async () => { writeRawModelsJson({ - "p-anthropic": { + "proxy-anthropic": { baseUrl: "https://proxy.example.com", apiKey: "literal-key", api: "anthropic-messages", @@ -2079,19 +2079,19 @@ describe("ModelRegistry", () => { ], }, }); - await authStorage.setRuntimeApiKey("p-anthropic", "literal-key"); + await authStorage.setRuntimeApiKey("proxy-anthropic", "literal-key"); const registry = new ModelRegistry(authStorage, modelsJsonPath); await registry.refresh("offline"); - const model = registry.find("p-anthropic", "claude-sonnet-4-5"); + const model = registry.find("proxy-anthropic", "claude-sonnet-4-5"); expect(model).toBeDefined(); expect(model?.isOAuth).toBe(true); }); test("auth: apiKey opts out of the anthropic-messages default", async () => { writeRawModelsJson({ - "p-anthropic": { + "proxy-anthropic": { baseUrl: "https://proxy.example.com", apiKey: "literal-key", api: "anthropic-messages", @@ -2109,19 +2109,19 @@ describe("ModelRegistry", () => { ], }, }); - await authStorage.setRuntimeApiKey("p-anthropic", "literal-key"); + await authStorage.setRuntimeApiKey("proxy-anthropic", "literal-key"); const registry = new ModelRegistry(authStorage, modelsJsonPath); await registry.refresh("offline"); - const model = registry.find("p-anthropic", "claude-sonnet-4-5"); + const model = registry.find("proxy-anthropic", "claude-sonnet-4-5"); expect(model).toBeDefined(); expect(model?.isOAuth).toBeUndefined(); }); test("non-anthropic apis do not get the OAuth default", async () => { writeRawModelsJson({ - "p-openai": { + "proxy-openai": { baseUrl: "https://proxy.example.com/v1", apiKey: "literal-key", api: "openai-completions", @@ -2138,12 +2138,12 @@ describe("ModelRegistry", () => { ], }, }); - await authStorage.setRuntimeApiKey("p-openai", "literal-key"); + await authStorage.setRuntimeApiKey("proxy-openai", "literal-key"); const registry = new ModelRegistry(authStorage, modelsJsonPath); await registry.refresh("offline"); - const model = registry.find("p-openai", "gpt-5"); + const model = registry.find("proxy-openai", "gpt-5"); expect(model).toBeDefined(); expect(model?.isOAuth).toBeUndefined(); }); diff --git a/python/robomp/.env.example b/python/robomp/.env.example index ccf350b23..4b21c2107 100644 --- a/python/robomp/.env.example +++ b/python/robomp/.env.example @@ -93,8 +93,8 @@ GITHUB_TOKEN= # ============================================================================= # Either a single model id or a comma-separated pool — roboomp picks one # uniformly at random per task. Use the `/` form that matches -# your ~/.omp/agent/models.yml (which is mounted into the container). -ROBOMP_MODEL=p-anthropic/claude-sonnet-4-6 +# your ~/.omp/agent/models.container.yml (which is mounted into the container as models.yml). +ROBOMP_MODEL=anthropic/claude-sonnet-4-6 # off|low|medium|high ROBOMP_THINKING=high # Optional provider override (passed to `omp --provider`). diff --git a/python/robomp/AGENTS.md b/python/robomp/AGENTS.md index da301d174..a42ad3d56 100644 --- a/python/robomp/AGENTS.md +++ b/python/robomp/AGENTS.md @@ -96,7 +96,7 @@ Lint + format: TypeScript via Biome (config in `biome.json`), Python via Ruff (c - `src/robomp/dashboard.py` — single-page HTML dashboard served from `/`. - `pyproject.toml` — packaging + pytest config (`asyncio_mode = "auto"`, `testpaths = ["tests"]`). - `Dockerfile` — slim runtime; consumes `oh-my-pi/artifacts:dev` (built from `/work/pi/Dockerfile`) for `pi_natives.linux-*.node` + `omp_rpc-*.whl`. Tini entrypoint, exposes `8080`, `VOLUME /data`. -- `docker-compose.yml` — `build.args.PI_ARTIFACTS_IMAGE`, mounts `$PI_ROOT:/work/pi:ro`, `./data:/data`, `~/.omp/agent/models.yml:ro`, `extra_hosts: llm-gateway.internal:host-gateway`. +- `docker-compose.yml` — `build.args.PI_ARTIFACTS_IMAGE`, mounts `$PI_ROOT:/work/pi:ro`, `./data:/data`, `~/.omp/agent/models.container.yml:ro` (mapped to `models.yml` inside the container — kept separate from the host's `~/.omp/agent/models.yml` so the host omp doesn't pick up gateway routing intended only for the container), `extra_hosts: llm-gateway.internal:host-gateway`. - `entrypoint.sh` — validates `PI_ROOT`, creates `/data/{workspaces,logs}` + build caches. - `.env.example` — authoritative list of required runtime env vars. - `README.md` — full architecture + operational reference. Authoritative for end-to-end flow, host-tool spec, security posture, and configuration reference. diff --git a/python/robomp/Dockerfile b/python/robomp/Dockerfile index a689460be..e3edf530c 100644 --- a/python/robomp/Dockerfile +++ b/python/robomp/Dockerfile @@ -1,4 +1,4 @@ -# syntax=docker/dockerfile:1.7 +# syntax=docker/dockerfile:1.7-labs ############################################################################### # roboomp — orchestrator image # @@ -30,14 +30,14 @@ FROM ${PI_ARTIFACTS_IMAGE} AS pi-artifacts ############################ FROM oven/bun:1.3.14-slim AS web-builder WORKDIR /work -# The repo is a Bun workspace (`workspaces: ["web"]` at the root). Install -# from the root lockfile so the web subpackage resolves against the same -# pinned dependency graph used locally. +# Build context is the pi monorepo root, so the web-builder stage installs +# from pi's bun.lock — that's how `web/package.json` resolves its `catalog:` +# references against the workspace-wide catalog declared at pi root. COPY package.json bun.lock ./ -COPY web/package.json ./web/package.json -RUN bun install --frozen-lockfile -COPY web/ ./web/ -RUN bun --cwd=web run build +COPY python/robomp/web/package.json ./python/robomp/web/package.json +RUN bun install --filter robomp-web +COPY --exclude=node_modules --exclude=dist python/robomp/web/ ./python/robomp/web/ +RUN bun --cwd=python/robomp/web run build ############################ # 3) runtime — slim image with everything roboomp needs at boot. @@ -107,9 +107,9 @@ RUN printf '%s\n' \ # roboomp itself. Drop the Vite-built dashboard into the package tree before # `pip install` so it lands in the installed wheel (`static/**/*` is declared # as package-data in pyproject.toml). -COPY pyproject.toml ./ -COPY src/ ./src/ -COPY --from=web-builder /work/web/dist/ ./src/robomp/static/ +COPY python/robomp/pyproject.toml ./ +COPY python/robomp/src/ ./src/ +COPY --from=web-builder /work/python/robomp/web/dist/ ./src/robomp/static/ RUN pip install --upgrade pip \ && pip install \ "fastapi>=0.112" "uvicorn[standard]>=0.30" "httpx>=0.27" \ @@ -121,7 +121,7 @@ RUN mkdir -p /srv/agent-home/.agent /srv/agent-home/.omp/agent \ && mkdir -p /srv/agent-home-stage/.agent /srv/agent-home-stage/.omp/agent \ && printf '[install]\nbackend = "copyfile"\n' > /srv/agent-home/.bunfig.toml -COPY entrypoint.sh /usr/local/bin/robomp-entrypoint +COPY python/robomp/entrypoint.sh /usr/local/bin/robomp-entrypoint RUN chmod +x /usr/local/bin/robomp-entrypoint VOLUME ["/data"] diff --git a/python/robomp/README.md b/python/robomp/README.md index 6afd4652e..2f0a33378 100644 --- a/python/robomp/README.md +++ b/python/robomp/README.md @@ -47,7 +47,7 @@ into the `tool_calls` table with credential-redacted args and results. ## Setup Requires Docker Compose v2 and a LiteLLM-style proxy on the host that your -`~/.omp/agent/models.yml` points at. roboomp lives inside the oh-my-pi +`~/.omp/agent/models.container.yml` points at (mounted into the container as `models.yml`; kept under a separate filename on the host so the host omp doesn't route through the gateway). roboomp lives inside the oh-my-pi monorepo at `python/robomp/`; both the docker build context and the `/work/pi` bind mount default to the parent monorepo (`../..`). Override `PI_ROOT` only if you want a different oh-my-pi checkout backing the build @@ -187,7 +187,7 @@ The integration test spawns a real `omp --mode rpc` against an | `refusing to push: working tree is dirty` | Uncommitted agent edits. Or just call `gh_open_pr`, which auto-commits `bun run fix` output. | | `bun check failed before PR creation` | Fix the reported failure and retry `gh_open_pr`. | | `Failed to load pi_natives` | Wrong arch / missing native. `bun run pi-artifacts` then `bun run build`. | -| `No API key found for ` | `~/.omp/agent/models.yml` mount missing or provider id mismatch with `ROBOMP_MODEL`. | +| `No API key found for ` | `~/.omp/agent/models.container.yml` mount missing or provider id mismatch with `ROBOMP_MODEL`. | ## Layout diff --git a/python/robomp/docker-compose.yml b/python/robomp/docker-compose.yml index 7b658d679..afef75aae 100644 --- a/python/robomp/docker-compose.yml +++ b/python/robomp/docker-compose.yml @@ -11,11 +11,14 @@ services: # ─────────────────────────────────────────────────────────────────────────── robomp: build: - context: . - dockerfile: Dockerfile + # pi root: gives the web-builder stage access to the workspace + # bun.lock + catalog (web/package.json refs `catalog:` versions). + # python/robomp/data is excluded via pi's .dockerignore. + context: ../.. + dockerfile: python/robomp/Dockerfile args: # Tag of the pre-built artifacts image produced by `bun run pi-artifacts` - # (sources: /work/pi/Dockerfile). Override per-environment as needed. + # (sources: pi root /Dockerfile). Override per-environment as needed. PI_ARTIFACTS_IMAGE: oh-my-pi/artifacts:dev image: robomp:dev container_name: robomp @@ -41,7 +44,7 @@ services: ROBOMP_REVIEWER_BOTS: ${ROBOMP_REVIEWER_BOTS:-} # --- model selection --- - ROBOMP_MODEL: ${ROBOMP_MODEL:-p-anthropic/claude-sonnet-4-6} + ROBOMP_MODEL: ${ROBOMP_MODEL:-anthropic/claude-sonnet-4-6} ROBOMP_PROVIDER: ${ROBOMP_PROVIDER:-} ROBOMP_THINKING: ${ROBOMP_THINKING:-high} @@ -86,7 +89,7 @@ services: # root-owned, world-readable files under /srv/agent-home; the agent # subprocess runs with HOME=/srv/agent-home, so ~/.omp and ~/.agent # resolve there without exposing mutable host mounts. - - ${HOME}/.omp/agent/models.yml:/srv/agent-home-stage/.omp/agent/models.yml:ro + - ${HOME}/.omp/agent/models.container.yml:/srv/agent-home-stage/.omp/agent/models.yml:ro - ${HOME}/.agent/AGENT.md:/srv/agent-home-stage/.agent/AGENTS.md:ro - ${HOME}/.agent/rules:/srv/agent-home-stage/.agent/rules:ro ports: diff --git a/python/robomp/src/robomp/config.py b/python/robomp/src/robomp/config.py index f9855af18..bec0815d5 100644 --- a/python/robomp/src/robomp/config.py +++ b/python/robomp/src/robomp/config.py @@ -57,7 +57,7 @@ class Settings(BaseSettings): gh_proxy_git_timeout_seconds: float = Field(60.0, alias="ROBOMP_GH_PROXY_GIT_TIMEOUT_SECONDS") # Model selection - model: str = Field("p-anthropic/claude-sonnet-4-6", alias="ROBOMP_MODEL") + model: str = Field("anthropic/claude-sonnet-4-6", alias="ROBOMP_MODEL") provider: str | None = Field(None, alias="ROBOMP_PROVIDER") thinking_level: ThinkingLevel = Field("high", alias="ROBOMP_THINKING") diff --git a/python/robomp/tests/test_config.py b/python/robomp/tests/test_config.py index f7d62c254..5f1f3e627 100644 --- a/python/robomp/tests/test_config.py +++ b/python/robomp/tests/test_config.py @@ -102,14 +102,14 @@ def test_model_pool_single(env: dict[str, str]) -> None: def test_model_pool_csv_parses(monkeypatch: pytest.MonkeyPatch, env: dict[str, str]) -> None: monkeypatch.setenv( "ROBOMP_MODEL", - " p-codex/gpt-5.4 , p-anthropic/claude-sonnet-4-6 ,, p-anthropic/claude-opus-4-7 ", + " codex/gpt-5.4 , anthropic/claude-sonnet-4-6 ,, anthropic/claude-opus-4-7 ", ) reset_settings_cache() cfg = Settings() # type: ignore[call-arg] assert cfg.model_pool == ( - "p-codex/gpt-5.4", - "p-anthropic/claude-sonnet-4-6", - "p-anthropic/claude-opus-4-7", + "codex/gpt-5.4", + "anthropic/claude-sonnet-4-6", + "anthropic/claude-opus-4-7", ) diff --git a/python/robomp/tests/test_pragmas.py b/python/robomp/tests/test_pragmas.py index c698ae2db..71ce47abe 100644 --- a/python/robomp/tests/test_pragmas.py +++ b/python/robomp/tests/test_pragmas.py @@ -101,21 +101,21 @@ def test_pragma_value_last_wins() -> None: def test_resolve_model_alias_precedence() -> None: - pool = ("p-anthropic/claude-sonnet-4-6", "p-openai/gpt-5.5", "p-openai/gpt-5.5-mini") + pool = ("anthropic/claude-sonnet-4-6", "openai/gpt-5.5", "openai/gpt-5.5-mini") # Short-name-after-slash beats substring. - assert resolve_model_alias("gpt-5.5", pool) == "p-openai/gpt-5.5" + assert resolve_model_alias("gpt-5.5", pool) == "openai/gpt-5.5" # Substring is fallback. - assert resolve_model_alias("gpt", pool) == "p-openai/gpt-5.5" - assert resolve_model_alias("claude", pool) == "p-anthropic/claude-sonnet-4-6" + assert resolve_model_alias("gpt", pool) == "openai/gpt-5.5" + assert resolve_model_alias("claude", pool) == "anthropic/claude-sonnet-4-6" def test_resolve_model_alias_full_id() -> None: - pool = ("p-openai/gpt-5.5", "p-anthropic/claude-sonnet-4-6") - assert resolve_model_alias("p-openai/gpt-5.5", pool) == "p-openai/gpt-5.5" + pool = ("openai/gpt-5.5", "anthropic/claude-sonnet-4-6") + assert resolve_model_alias("openai/gpt-5.5", pool) == "openai/gpt-5.5" def test_resolve_model_alias_no_match() -> None: - pool = ("p-anthropic/claude-sonnet-4-6",) + pool = ("anthropic/claude-sonnet-4-6",) assert resolve_model_alias("gpt", pool) is None assert resolve_model_alias("", pool) is None diff --git a/python/robomp/tests/test_worker_pragmas.py b/python/robomp/tests/test_worker_pragmas.py index 80ed64852..3c98eb260 100644 --- a/python/robomp/tests/test_worker_pragmas.py +++ b/python/robomp/tests/test_worker_pragmas.py @@ -12,7 +12,7 @@ from robomp.worker import DirectiveInfo, _resolve_pragma_overrides def settings_with_pool(monkeypatch: pytest.MonkeyPatch, env: dict[str, str]) -> Settings: monkeypatch.setenv( "ROBOMP_MODEL", - "p-anthropic/claude-sonnet-4-6,p-openai/gpt-5.5,p-openai/gpt-5.5-mini", + "anthropic/claude-sonnet-4-6,openai/gpt-5.5,openai/gpt-5.5-mini", ) reset_settings_cache() return Settings() # type: ignore[call-arg] @@ -30,14 +30,14 @@ def test_directive_without_pragmas_means_no_override(settings_with_pool: Setting def test_model_pragma_resolves_to_pool_entry(settings_with_pool: Settings) -> None: directive = DirectiveInfo(body="run", author="can1357", pragmas=(("model", "gpt"),)) model_override, thinking_override = _resolve_pragma_overrides(directive, settings_with_pool) - assert model_override == "p-openai/gpt-5.5" + assert model_override == "openai/gpt-5.5" assert thinking_override is None def test_model_alias_exact_short_name(settings_with_pool: Settings) -> None: directive = DirectiveInfo(body="run", author="can1357", pragmas=(("model", "gpt-5.5-mini"),)) model_override, _ = _resolve_pragma_overrides(directive, settings_with_pool) - assert model_override == "p-openai/gpt-5.5-mini" + assert model_override == "openai/gpt-5.5-mini" def test_unmatched_model_alias_falls_back_to_random_pick(settings_with_pool: Settings) -> None: @@ -66,7 +66,7 @@ def test_both_pragmas_resolved_together(settings_with_pool: Settings) -> None: pragmas=(("model", "claude"), ("thinking", "medium")), ) model_override, thinking_override = _resolve_pragma_overrides(directive, settings_with_pool) - assert model_override == "p-anthropic/claude-sonnet-4-6" + assert model_override == "anthropic/claude-sonnet-4-6" assert thinking_override == "medium" @@ -77,4 +77,4 @@ def test_last_value_wins_for_duplicate_keys(settings_with_pool: Settings) -> Non pragmas=(("model", "claude"), ("model", "gpt")), ) model_override, _ = _resolve_pragma_overrides(directive, settings_with_pool) - assert model_override == "p-openai/gpt-5.5" + assert model_override == "openai/gpt-5.5" diff --git a/scripts/bench-edit-hashline-sep.ts b/scripts/bench-edit-hashline-sep.ts index 3e9779157..b89e1f69e 100644 --- a/scripts/bench-edit-hashline-sep.ts +++ b/scripts/bench-edit-hashline-sep.ts @@ -17,7 +17,7 @@ const SEPARATORS = ["~", "%", "÷", ">", ":"] as const; const MODELS = [ "openrouter/z-ai/glm-4.7:nitro", "openai/gpt-5.4-nano", - "p-anthropic/claude-sonnet-4-6", + "anthropic/claude-sonnet-4-6", ] as const; const CONCURRENCY = 3; diff --git a/scripts/eval-bench-runs.ts b/scripts/eval-bench-runs.ts index fcf29375d..5b0777fd2 100644 --- a/scripts/eval-bench-runs.ts +++ b/scripts/eval-bench-runs.ts @@ -205,8 +205,7 @@ function fmtNum(value: number): string { } function shortModel(model: string): string { - const cleaned = model.replace(/^p-anthropic\//, "anthropic/"); - const segs = cleaned.split("/"); + const segs = model.split("/"); return segs[segs.length - 1].replace(/:nitro/, ""); }