diff --git a/crates/pi-natives/src/pty.rs b/crates/pi-natives/src/pty.rs index a41d66859..391100b57 100644 --- a/crates/pi-natives/src/pty.rs +++ b/crates/pi-natives/src/pty.rs @@ -217,7 +217,8 @@ fn run_pty_sync( ct: task::CancelToken, ) -> Result { let pty_system = native_pty_system(); - ct.heartbeat().map_err(|err| Error::from_reason(format!("PTY setup cancelled before openpty: {err}")))?; + ct.heartbeat() + .map_err(|err| Error::from_reason(format!("PTY setup cancelled before openpty: {err}")))?; const PTY_STARTUP_TIMEOUT: Duration = Duration::from_secs(5); let pair = if cfg!(windows) { @@ -227,9 +228,9 @@ fn run_pty_sync( let (tx, rx) = mpsc::channel(); std::thread::spawn(move || { let result = pty_system.openpty(PtySize { - rows: config.rows, - cols: config.cols, - pixel_width: 0, + rows: config.rows, + cols: config.cols, + pixel_width: 0, pixel_height: 0, }); let _ = tx.send(result); @@ -237,16 +238,18 @@ fn run_pty_sync( match rx.recv_timeout(PTY_STARTUP_TIMEOUT) { Ok(Ok(pair)) => pair, Ok(Err(e)) => return Err(Error::from_reason(format!("Failed to open PTY: {e}"))), - Err(_) => return Err(Error::from_reason( - "PTY creation timed out (5s). ConPTY may be unavailable on this system.", - )), + Err(_) => { + return Err(Error::from_reason( + "PTY creation timed out (5s). ConPTY may be unavailable on this system.", + )); + }, } } else { pty_system .openpty(PtySize { - rows: config.rows, - cols: config.cols, - pixel_width: 0, + rows: config.rows, + cols: config.cols, + pixel_width: 0, pixel_height: 0, }) .map_err(|err| Error::from_reason(format!("Failed to open PTY: {err}")))? @@ -273,14 +276,16 @@ fn run_pty_sync( cmd.env(key, value); } } - ct.heartbeat().map_err(|err| Error::from_reason(format!("PTY setup cancelled before spawn: {err}")))?; + ct.heartbeat() + .map_err(|err| Error::from_reason(format!("PTY setup cancelled before spawn: {err}")))?; let mut child = pair .slave .spawn_command(cmd) .map_err(|err| Error::from_reason(format!("Failed to spawn PTY command: {err}")))?; drop(pair.slave); - ct.heartbeat().map_err(|err| Error::from_reason(format!("PTY setup cancelled before reader: {err}")))?; + ct.heartbeat() + .map_err(|err| Error::from_reason(format!("PTY setup cancelled before reader: {err}")))?; let master = pair.master; let mut writer = master diff --git a/packages/ai/CHANGELOG.md b/packages/ai/CHANGELOG.md index 5ea89e4f3..c3e926885 100644 --- a/packages/ai/CHANGELOG.md +++ b/packages/ai/CHANGELOG.md @@ -1,7 +1,6 @@ # 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` @@ -11,10 +10,13 @@ ### Added +- Added support for Anthropic image message parts with `type: "url"` and `type: "file"` sources +- Added `stopSequences` and `frequencyPenalty` to shared stream options and wired them through to OpenAI request translation +- Added optional request cancellation support to auth-broker interactions by propagating `AbortSignal` into health, snapshot, usage, and refresh calls - 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 auth-broker usage polling response handling that returns normalized usage reports plus generation timestamp for clients (5-min per-credential cache via `AuthStorage`) - 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. @@ -26,16 +28,20 @@ - 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/usage` — aggregated provider usage; 5-min per-credential 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 exports from `@oh-my-pi/pi-ai/auth-gateway`: `startAuthGateway`, `AuthGatewayServerOptions`, `AuthGatewayBootOptions`, `AuthGatewayServerHandle`, `ModelResolver`, `DEFAULT_AUTH_GATEWAY_BIND`. Per-format `parseRequest` / `encodeResponse` / `encodeStream` triples are reachable via the `./providers/*` subpath as `openai-chat-server`, `anthropic-messages-server`, and `openai-responses-server`. - Added `listProvidersWithEnvKey()` to enumerate every provider with an env-var fallback (used by the new migrate command in coding-agent). ### Changed +- Changed auth-gateway parsing for OpenAI chat-completions and Responses to ignore unsupported SDK-only fields instead of rejecting requests +- Changed auth-gateway protocol handling to include CORS headers on responses and support browser-origin requests +- Changed prompt-cache handling to resolve cache keys from request metadata and headers and preserve them through protocol translation +- Changed Anthropic messages parsing to forward request `metadata` through to downstream execution - 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 @@ -46,6 +52,11 @@ ### Fixed +- Fixed `RemoteAuthCredentialStore.getUsageReport` to return the matching credential-specific usage report and coalesce parallel callers into one broker `/v1/usage` fetch +- Fixed auth-broker credential upload validation to reject the remote refresh-token sentinel and prevent storing a non-refresh value +- Fixed OpenAI Responses streaming output to emit `reasoning_summary_text` events and parse/send `summary_text` reasoning payloads +- Fixed Anthropic stop-sequence handling by trimming requests to the API limit of four entries before forwarding +- Fixed prompt caching behavior across protocol translations so cached-token usage is preserved when Anthropic and OpenAI requests are routed through each other - 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 @@ -57,6 +68,10 @@ - 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 +### Security + +- Hardened auth-gateway bearer-token checks with constant-time comparison to avoid timing-side-channel leaks + ## [15.1.2] - 2026-05-15 ### Breaking Changes diff --git a/packages/ai/src/auth-broker/client.ts b/packages/ai/src/auth-broker/client.ts index ed4e325ce..4b93a3d99 100644 --- a/packages/ai/src/auth-broker/client.ts +++ b/packages/ai/src/auth-broker/client.ts @@ -68,48 +68,60 @@ export class AuthBrokerClient { this.#fetch = opts.fetchImpl ?? fetch; } - healthz(): Promise { - return this.#request("GET", "/v1/healthz", { schema: healthzResponseSchema, auth: false }); + healthz(signal?: AbortSignal): Promise { + return this.#request("GET", "/v1/healthz", { schema: healthzResponseSchema, auth: false, signal }); } - fetchSnapshot(): Promise { + fetchSnapshot(signal?: AbortSignal): Promise { // `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; + return this.#request("GET", "/v1/snapshot", { + schema: snapshotResponseSchema, + signal, + }) 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; + fetchUsage(signal?: AbortSignal): Promise { + // Validates the envelope (`generatedAt`, `reports[].provider`, `limits`, + // `metadata`) but leaves provider-specific extension fields permissive so + // the broker can ship new shapes ahead of the client. `raw` is accepted + // but normally stripped by the broker before send. + return this.#request("GET", "/v1/usage", { schema: usageResponseSchema, signal }) as Promise; } - async refreshCredential(id: number): Promise { + async refreshCredential(id: number, signal?: AbortSignal): Promise { return this.#request("POST", `/v1/credential/${id}/refresh`, { schema: credentialRefreshResponseSchema, + signal, }) as Promise; } - async disableCredential(id: number, cause: string): Promise { + async disableCredential(id: number, cause: string, signal?: AbortSignal): Promise { const body: CredentialDisableRequest = { cause }; return this.#request("POST", `/v1/credential/${id}/disable`, { body, schema: credentialDisableResponseSchema, + signal, }); } - async uploadCredential(provider: string, credential: AuthCredential): Promise { + async uploadCredential( + provider: string, + credential: AuthCredential, + signal?: AbortSignal, + ): Promise { const body: CredentialUploadRequest = { provider, credential }; return this.#request("POST", "/v1/credential", { body, schema: credentialUploadResponseSchema, + signal, }) as Promise; } async #request( method: "GET" | "POST", path: string, - opts: { schema: TSchema; auth?: boolean; body?: unknown }, + opts: { schema: TSchema; auth?: boolean; body?: unknown; signal?: AbortSignal }, ): Promise> { const auth = opts.auth ?? true; const url = `${this.#baseUrl}${path}`; @@ -121,14 +133,25 @@ export class AuthBrokerClient { headers["Content-Type"] = "application/json"; } + // Fast-fail when the caller's signal is already aborted — avoids spinning + // up a fetch + timer that the first `await` would just abort anyway. + if (opts.signal?.aborted) { + throw new AuthBrokerError("Auth broker request aborted", { cause: opts.signal.reason }); + } + let lastError: unknown; for (let attempt = 0; attempt <= this.#maxRetries; attempt += 1) { + // Compose caller's signal with the per-attempt timeout so either + // source can cancel the in-flight fetch. `AbortSignal.any` is the + // supported merge primitive in Bun ≥ 1.0 / Node ≥ 20. + const timeoutSignal = AbortSignal.timeout(this.#timeoutMs); + const signal = opts.signal ? AbortSignal.any([opts.signal, timeoutSignal]) : timeoutSignal; try { const response = await this.#fetch(url, { method, headers, body: payload, - signal: AbortSignal.timeout(this.#timeoutMs), + signal, }); const text = await response.text(); if (!response.ok) { @@ -157,6 +180,10 @@ export class AuthBrokerClient { return validated.data; } catch (error) { lastError = error; + // Caller-driven abort wins over retry — the caller said stop. + if (opts.signal?.aborted) { + throw new AuthBrokerError("Auth broker request aborted", { cause: opts.signal.reason }); + } if (error instanceof AuthBrokerError && error.status !== undefined) { // HTTP errors (4xx/5xx) don't retry — caller knows what to do. throw error; diff --git a/packages/ai/src/auth-broker/remote-store.ts b/packages/ai/src/auth-broker/remote-store.ts index 6f7687928..386749b8b 100644 --- a/packages/ai/src/auth-broker/remote-store.ts +++ b/packages/ai/src/auth-broker/remote-store.ts @@ -4,7 +4,8 @@ * `upsert*`, `delete*ForProvider`) throw because login flows are server-side. * * Cache (`getCache`/`setCache`/`cleanExpiredCache`) is in-memory and ephemeral — - * usage reports cache TTL is ~30s, so durability across runs isn't required. + * usage reports cache TTL is 5 minutes per credential, so durability across + * runs isn't required. */ import { logger } from "@oh-my-pi/pi-utils"; import { @@ -20,11 +21,24 @@ import type { UsageReport } from "../usage"; import type { OAuthCredentials } from "../utils/oauth/types"; import type { AuthBrokerClient } from "./client"; +/** + * Client-side TTL for the aggregate `/v1/usage` response. Set below the + * broker server's own 30s usage cache so we typically pick up the broker's + * cached value instead of re-walking the network — but high enough to absorb + * the parallel fan-out from `#rankOAuthSelections` into a single round-trip. + */ +const USAGE_CACHE_TTL_MS = 15_000; + interface CacheEntry { value: string; expiresAtSec: number; } +interface UsageCacheEntry { + reports: UsageReport[]; + fetchedAt: number; +} + export interface RemoteAuthCredentialStoreOptions { client: AuthBrokerClient; /** @@ -38,6 +52,8 @@ export class RemoteAuthCredentialStore implements AuthCredentialStore { readonly #client: AuthBrokerClient; #snapshot: AuthCredentialSnapshot; #cache: Map = new Map(); + #usageCache?: UsageCacheEntry; + #usageInflight?: Promise; #closed = false; constructor(opts: RemoteAuthCredentialStoreOptions) { @@ -151,8 +167,9 @@ export class RemoteAuthCredentialStore implements AuthCredentialStore { _provider: Provider, credentialId: number, _credential: OAuthCredential, + signal?: AbortSignal, ): Promise { - const { entry } = await this.#client.refreshCredential(credentialId); + const { entry } = await this.#client.refreshCredential(credentialId, signal); if (entry.credential.type !== "oauth") { throw new Error(`Broker returned non-OAuth credential for id=${credentialId}`); } @@ -174,9 +191,78 @@ export class RemoteAuthCredentialStore implements AuthCredentialStore { * 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; + async fetchUsageReports(signal?: AbortSignal): Promise { + return this.#raceWithSignal(this.#loadUsageReports(), signal); + } + + /** + * Per-credential usage hook consumed by `AuthStorage.#getUsageReport`. Pulls + * the aggregate broker `/v1/usage` once and serves all callers from the + * same response (coalesced + cached), then matches the credential to a + * report by provider + identity (accountId / email / projectId). + * + * The broker already aggregates with its own 30s TTL on the server side; our + * 15s client TTL is below that so we usually re-use the broker's cache too. + */ + async getUsageReport( + provider: Provider, + credential: OAuthCredential, + signal?: AbortSignal, + ): Promise { + const reports = await this.#raceWithSignal(this.#loadUsageReports(), signal); + if (!reports) return null; + return matchUsageReport(reports, provider, credential); + } + + /** + * Reject the awaited promise when the caller's signal aborts, without + * affecting the shared upstream fetch. Used to give each caller their + * own cancel without one caller's abort cascading into a peer's in-flight + * request through the single-flight `#usageInflight`. + */ + #raceWithSignal(promise: Promise, signal?: AbortSignal): Promise { + if (!signal) return promise; + if (signal.aborted) return Promise.reject(new Error("auth-broker request aborted")); + return new Promise((resolve, reject) => { + const onAbort = (): void => { + signal.removeEventListener("abort", onAbort); + reject(new Error("auth-broker request aborted")); + }; + signal.addEventListener("abort", onAbort, { once: true }); + promise.then( + value => { + signal.removeEventListener("abort", onAbort); + resolve(value); + }, + err => { + signal.removeEventListener("abort", onAbort); + reject(err); + }, + ); + }); + } + + #loadUsageReports(): Promise { + const cached = this.#usageCache; + if (cached && Date.now() - cached.fetchedAt < USAGE_CACHE_TTL_MS) { + return Promise.resolve(cached.reports); + } + if (this.#usageInflight) return this.#usageInflight; + const inflight = this.#client + .fetchUsage() + .then(body => { + this.#usageCache = { reports: body.reports, fetchedAt: Date.now() }; + return body.reports; + }) + .catch(error => { + logger.warn("auth-broker usage fetch failed", { error: String(error) }); + return null; + }) + .finally(() => { + this.#usageInflight = undefined; + }); + this.#usageInflight = inflight; + return inflight; } close(): void { @@ -185,3 +271,59 @@ export class RemoteAuthCredentialStore implements AuthCredentialStore { this.#cache.clear(); } } + +/** + * Match a broker-supplied usage report to a specific OAuth credential. The + * broker returns aggregate reports across all credentials it manages, so we + * pick the one whose identity (accountId / email / projectId) lines up with + * the credential the caller is asking about. + * + * Falls back to the lone candidate when only one matches the provider; falls + * through to `null` when nothing matches, which `AuthStorage` treats as "no + * usage data" (ranking proceeds without a usage signal for this credential). + */ +function matchUsageReport(reports: UsageReport[], provider: Provider, credential: OAuthCredential): UsageReport | null { + const candidates = reports.filter(report => report.provider === provider); + if (candidates.length === 0) return null; + if (candidates.length === 1) return candidates[0]; + const accountId = credential.accountId?.trim().toLowerCase(); + const email = credential.email?.trim().toLowerCase(); + const projectId = credential.projectId?.trim().toLowerCase(); + for (const report of candidates) { + if (reportMatchesIdentity(report, accountId, email, projectId)) return report; + } + return null; +} + +function reportMatchesIdentity( + report: UsageReport, + accountId: string | undefined, + email: string | undefined, + projectId: string | undefined, +): boolean { + const metadata = (report.metadata ?? {}) as Record; + if (accountId) { + const metaAccount = readMetadataString(metadata, "accountId") ?? readMetadataString(metadata, "account_id"); + if (metaAccount && metaAccount.toLowerCase() === accountId) return true; + for (const limit of report.limits) { + if (limit.scope.accountId?.toLowerCase() === accountId) return true; + } + } + if (email) { + const metaEmail = readMetadataString(metadata, "email"); + if (metaEmail && metaEmail.toLowerCase() === email) return true; + } + if (projectId) { + const metaProject = readMetadataString(metadata, "projectId") ?? readMetadataString(metadata, "project_id"); + if (metaProject && metaProject.toLowerCase() === projectId) return true; + for (const limit of report.limits) { + if (limit.scope.projectId?.toLowerCase() === projectId) return true; + } + } + return false; +} + +function readMetadataString(metadata: Record, key: string): string | undefined { + const value = metadata[key]; + return typeof value === "string" && value.trim().length > 0 ? value.trim() : undefined; +} diff --git a/packages/ai/src/auth-broker/server.ts b/packages/ai/src/auth-broker/server.ts index 379a9d20f..cc1de341d 100644 --- a/packages/ai/src/auth-broker/server.ts +++ b/packages/ai/src/auth-broker/server.ts @@ -11,6 +11,7 @@ */ import { logger } from "@oh-my-pi/pi-utils"; import type { AuthStorage } from "../auth-storage"; +import { parseBind } from "../utils/parse-bind"; import { AuthBrokerRefresher } from "./refresher"; import type { CredentialDisableResponse, @@ -47,42 +48,6 @@ export interface AuthBrokerServerHandle { close(): Promise; } -interface ParsedBind { - hostname: string; - 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: parsePort(trimmed, raw) }; - } - const lastColon = trimmed.lastIndexOf(":"); - if (lastColon < 0) { - throw new Error(`Invalid bind '${raw}'; expected 'host:port' or 'port'.`); - } - const hostPart = trimmed.slice(0, lastColon); - const portPart = trimmed.slice(lastColon + 1); - if (hostPart.length === 0) { - throw new Error(`Invalid bind '${raw}'; host must not be empty.`); - } - return { hostname: hostPart, port: parsePort(portPart, raw) }; -} - function json(status: number, body: unknown): Response { return new Response(JSON.stringify(body), { status, @@ -174,10 +139,12 @@ export function startAuthBroker(opts: AuthBrokerServerOptions): AuthBrokerServer } 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 + // AuthStorage caches usage reports internally with a 5-minute per-credential + // 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?.()) ?? []; + // `req.signal` propagates HTTP-client disconnects all the way to the + // per-caller cancel without touching the shared upstream fetch. + const reports = (await opts.storage.fetchUsageReports?.({ signal: req.signal })) ?? []; // Drop the `raw` field — it's the provider-specific upstream body, // large and unstable. Everything UI-relevant lives in `limits` and // `metadata`. @@ -194,7 +161,7 @@ export function startAuthBroker(opts: AuthBrokerServerOptions): AuthBrokerServer if (refreshMatch) { const id = Number.parseInt(refreshMatch[1], 10); try { - const entry = await opts.storage.forceRefreshCredentialById(id); + const entry = await opts.storage.forceRefreshCredentialById(id, req.signal); const body: CredentialRefreshResponse = { entry }; logger.info("auth-broker credential refreshed", { id, diff --git a/packages/ai/src/auth-broker/wire-schemas.ts b/packages/ai/src/auth-broker/wire-schemas.ts index 8267304de..34411f0fa 100644 --- a/packages/ai/src/auth-broker/wire-schemas.ts +++ b/packages/ai/src/auth-broker/wire-schemas.ts @@ -20,7 +20,17 @@ import { usageReportSchema } from "../usage"; export const oauthCredentialSchema = z .object({ type: z.literal("oauth"), - refresh: z.string().min(1), + refresh: z + .string() + .min(1) + // Reject the sentinel literal on writes: if a client somehow round-trips + // a snapshot back into POST /v1/credential, accepting the sentinel as a + // real refresh token would silently break that credential's refresh + // forever (the broker would store `"__remote__"` and try to use it as + // the upstream refresh token). + .refine(value => value !== REMOTE_REFRESH_SENTINEL, { + message: `refresh token must not equal the remote sentinel (${REMOTE_REFRESH_SENTINEL})`, + }), access: z.string().min(1), expires: z.number(), enterpriseUrl: z.string().optional(), diff --git a/packages/ai/src/auth-gateway/http.ts b/packages/ai/src/auth-gateway/http.ts index 01c6fe237..3e79e56c0 100644 --- a/packages/ai/src/auth-gateway/http.ts +++ b/packages/ai/src/auth-gateway/http.ts @@ -4,11 +4,13 @@ * Centralized so we share the same JSON shape, auth check, * and peer-resolution logic. */ +import { timingSafeEqual as nodeTimingSafeEqual } from "node:crypto"; 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, @@ -22,11 +24,171 @@ export function resolvePeer(req: Request): string { return req.headers.get("x-real-ip") ?? "unknown"; } +/** + * Constant-time byte comparison. Falls back to a manual XOR accumulator if + * `node:crypto.timingSafeEqual` isn't available. Always processes every byte + * of the longer input so length itself doesn't leak via timing. + */ +export function timingSafeEqual(a: Uint8Array, b: Uint8Array): boolean { + if (a.length === b.length && typeof nodeTimingSafeEqual === "function") { + return nodeTimingSafeEqual(a, b); + } + const len = Math.max(a.length, b.length); + let diff = a.length ^ b.length; + for (let i = 0; i < len; i++) { + // Out-of-range reads return undefined → coerce to 0 via `| 0`. + const av = (i < a.length ? a[i] : 0) | 0; + const bv = (i < b.length ? b[i] : 0) | 0; + diff |= av ^ bv; + } + return diff === 0; +} + +const TOKEN_ENCODER = new TextEncoder(); + 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()); + const presented = TOKEN_ENCODER.encode(match[1].trim()); + // Iterate every allowed token regardless of early hits so the result + // timing reflects the full set, not the position of the match. + let ok = false; + for (const tok of tokens) { + const expected = TOKEN_ENCODER.encode(tok); + if (timingSafeEqual(presented, expected)) ok = true; + } + return ok; +} + +/** + * Allow-list of inbound request headers that the gateway captures and forwards + * to the underlying parsers (which decide whether to surface them to the + * provider). Case-insensitive; `x-stainless-` is a prefix match. + */ +const PASSTHROUGH_HEADER_NAMES: Record = { + "anthropic-beta": true, + "anthropic-version": true, + "openai-organization": true, + "openai-project": true, + "openai-beta": true, + // Codex / ChatGPT-OAuth backend headers (see openai-codex/constants.ts). + // `session_id` and `conversation_id` thread the upstream session so prompt + // caching and per-conversation rate limiting work; `chatgpt-account-id` and + // `originator` identify the calling account and client surface. + "chatgpt-account-id": true, + originator: true, + session_id: true, + conversation_id: true, + // Vendor-neutral cache-identity headers. The gateway also reads these to + // populate `options.promptCacheKey` (see `resolvePromptCacheKey` below) + // so explicit client hints win over the derived fallback. + "x-prompt-cache-key": true, + "x-session-id": true, + "x-conversation-id": true, +}; + +/** + * Extract allow-listed passthrough headers from an inbound request. Keys are + * lowercased; empty values are dropped. Called once per request in + * `handleFormatEndpoint`; parsers then read `options.headers`. + */ +export function captureRequestHeaders(headers: Headers): Record { + const out: Record = {}; + headers.forEach((value, key) => { + if (!value) return; + const lower = key.toLowerCase(); + if (PASSTHROUGH_HEADER_NAMES[lower] || lower.startsWith("x-stainless-")) { + out[lower] = value; + } + }); + return out; +} + +/** + * Priority order for resolving a client-supplied prompt-cache identity. The + * first non-empty value wins. When none are present, the gateway derives a + * stable UUID from the request's stable parts. + */ +const CACHE_KEY_HEADERS: readonly string[] = [ + "x-prompt-cache-key", + "session_id", + "conversation_id", + "x-session-id", + "x-conversation-id", +]; + +function readBodyCacheKey(body: unknown): string | undefined { + if (body === null || typeof body !== "object") return undefined; + const root = body as Record; + // Explicit body fields (OpenAI Responses / Chat). + const direct = root.prompt_cache_key; + if (typeof direct === "string" && direct.length > 0) return direct; + // Nested `metadata` (Codex CLI / Anthropic clients that route a session + // identifier through the metadata bag). + const metadata = root.metadata; + if (metadata === null || typeof metadata !== "object") return undefined; + const meta = metadata as Record; + for (const field of ["prompt_cache_key", "session_id", "conversation_id"] as const) { + const v = meta[field]; + if (typeof v === "string" && v.length > 0) return v; + } + return undefined; +} + +/** + * Resolve a prompt-cache identity from inbound request body + headers. + * Order of precedence (first wins): + * 1. Body `prompt_cache_key` + * 2. Body `metadata.{prompt_cache_key,session_id,conversation_id}` + * 3. Header `x-prompt-cache-key` + * 4. Header `session_id` / `conversation_id` (Codex / ChatGPT-OAuth surface) + * 5. Header `x-session-id` / `x-conversation-id` (common informal) + * Returns undefined when none present; the gateway then derives a stable + * UUID from the request's stable parts. + */ +export function resolvePromptCacheKey(body: unknown, headers?: Headers): string | undefined { + const fromBody = readBodyCacheKey(body); + if (fromBody) return fromBody; + if (!headers) return undefined; + for (const name of CACHE_KEY_HEADERS) { + const v = headers.get(name); + if (v && v.length > 0) return v; + } + return undefined; +} + +const CORS_HEADERS: Record = { + "Access-Control-Allow-Origin": "*", + "Access-Control-Allow-Methods": "GET, POST, OPTIONS", + "Access-Control-Allow-Headers": + "authorization, content-type, anthropic-version, anthropic-beta, openai-organization, openai-project, x-stainless-*, x-api-key", + "Access-Control-Max-Age": "86400", +}; + +/** + * CORS headers for the auth-gateway. Currently echoes a wildcard origin; the + * request is accepted so future tightening can mirror `Origin` without + * threading the request through every caller. + */ +export function corsHeaders(_req: Request): Record { + return { ...CORS_HEADERS }; +} + +/** + * Re-emit `response` with CORS headers merged. The original response body is + * passed through unchanged. Used by the gateway wrapper so every outbound + * format-endpoint response carries the same CORS surface as the preflight. + */ +export function withCors(response: Response, req: Request): Response { + const headers = new Headers(response.headers); + const cors = corsHeaders(req); + for (const k in cors) headers.set(k, cors[k]); + return new Response(response.body, { + status: response.status, + statusText: response.statusText, + headers, + }); } diff --git a/packages/ai/src/auth-gateway/server.ts b/packages/ai/src/auth-gateway/server.ts index 341d6254b..79eb5a85f 100644 --- a/packages/ai/src/auth-gateway/server.ts +++ b/packages/ai/src/auth-gateway/server.ts @@ -10,7 +10,7 @@ * * Endpoints: * GET /healthz → unauth; ok + version - * GET /v1/usage → aggregated provider usage (30s cache via AuthStorage) + * GET /v1/usage → aggregated provider usage (5-min per-credential 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 @@ -24,7 +24,8 @@ 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 { parseBind } from "../utils/parse-bind"; +import { captureRequestHeaders, corsHeaders, isAuthorized, json, resolvePeer, withCors } from "./http"; import type { AuthGatewayServerHandle, AuthGatewayServerOptions, @@ -50,24 +51,8 @@ export interface AuthGatewayBootOptions extends AuthGatewayServerOptions { 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 }; -} +// `parseBind` lives in ../utils/parse-bind so the gateway and broker can't +// drift on accepted inputs (e.g. empty hostname, IPv6 brackets). const FORMAT_ROUTES: Record = { "/v1/chat/completions": { module: openaiChat, label: "openai-chat" }, @@ -75,55 +60,10 @@ const FORMAT_ROUTES: Record = { "/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", -]); +// (passthrough fast-path removed — it bypassed pi-ai provider logic, in +// particular the Anthropic Claude-Code OAuth system-prompt prefix injection. +// Every request now takes the translate path so credential-specific request +// shaping always applies.) // 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 @@ -133,6 +73,46 @@ const HOP_BY_HOP_HEADERS = new Set([ // just turned that into per-call config hell. Silent strip is what the // upstream provider would do anyway when it ignores extra fields. +/** + * Derive a stable cache identity from the parts of the request that don't + * change turn-to-turn within a logical conversation: model id, system prompt, + * tool definitions, and the first message (the conversation seed). Codex-class + * backends only cache prefixes when an explicit `prompt_cache_key` is set; + * without one, two requests with the same prefix but different trailing + * messages don't coalesce. This bridges Anthropic-style clients (which signal + * caching via `cache_control` markers rather than an opaque key) to Codex's + * keyed model so cross-protocol caching "just works". + * + * Including the first message scopes the key to one logical conversation: + * two different chats with the same system prompt no longer share a cache + * bucket and can't trample each other's prefix-tree entries. + * + * Anthropic-backed requests ignore `sessionId`; the key is harmless there. + */ +function deriveSessionId(parsed: ParsedFormatRequest): string { + const { modelId, context } = parsed; + const parts: string[] = [modelId]; + if (context.systemPrompt && context.systemPrompt.length > 0) { + parts.push(context.systemPrompt.join("\n\n")); + } + if (context.tools && context.tools.length > 0) { + parts.push(JSON.stringify(context.tools)); + } + const first = context.messages?.[0]; + if (first) { + // Strip timestamp / provider metadata so the hash is stable across turns + // of the same conversation (omp re-stamps every parsed Message). role + + // content is what's actually on the wire. + parts.push(JSON.stringify({ role: first.role, content: first.content })); + } + const seed = parts.join("\u0000"); + const hex = new Bun.CryptoHasher("sha256").update(seed).digest("hex"); + // Format the leading 128 bits as a v4-shape UUID (8-4-4-4-12). Codex's + // `normalizeOpenAIResponsesPromptCacheKey` accepts ≤64 chars verbatim, so + // the 36-char UUID flows through unchanged. + return `${hex.slice(0, 8)}-${hex.slice(8, 12)}-${hex.slice(12, 16)}-${hex.slice(16, 20)}-${hex.slice(20, 32)}`; +} + function buildStreamOptions(parsed: ParsedFormatRequest, api: Api, signal: AbortSignal): SimpleStreamOptions { const opts: SimpleStreamOptions = { signal }; const { options } = parsed; @@ -145,26 +125,116 @@ function buildStreamOptions(parsed: ParsedFormatRequest, api: Api, signal: Abort 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.minP !== undefined) opts.minP = options.minP; + if (options.stopSequences !== undefined) opts.stopSequences = options.stopSequences; + if (options.presencePenalty !== undefined) opts.presencePenalty = options.presencePenalty; + if (options.frequencyPenalty !== undefined) opts.frequencyPenalty = options.frequencyPenalty; + if (options.repetitionPenalty !== undefined) opts.repetitionPenalty = options.repetitionPenalty; + if (options.metadata !== undefined) opts.metadata = options.metadata; + if (options.headers !== undefined) opts.headers = { ...(opts.headers ?? {}), ...options.headers }; 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.disableReasoning !== undefined) opts.disableReasoning = options.disableReasoning; 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; + // Client-supplied `prompt_cache_key` wins; otherwise derive a stable + // key from the model + system + tools so prefix caching engages on + // Codex-class backends across turns of the same logical conversation. + opts.sessionId = options.promptCacheKey ?? deriveSessionId(parsed); + if (options.thinkingBudgets) { + opts.thinkingBudgets = { ...(opts.thinkingBudgets ?? {}), ...options.thinkingBudgets }; + } + if (options.explicitThinkingBudgetTokens !== undefined) { + // Mirror Rust's `resolve_thinking_budget`: explicit budget pins onto + // whichever effort the client requested (or High when unspecified) and + // ALSO sets the effort so providers that gate on `reasoning` actually + // surface the budget. + const effort = options.reasoning ?? Effort.High; + opts.thinkingBudgets = { + ...(opts.thinkingBudgets ?? {}), + [effort]: options.explicitThinkingBudgetTokens, + }; + opts.reasoning ??= effort; + } + // Fields that don't yet have a matching pi-ai `SimpleStreamOptions` slot. + // Surfaced once in debug logs so they show up when wiring a new provider, + // but NEVER widened into `options.extra` — every consumer would have to + // re-implement the typed parse to read them back out. + // TODO(pi-ai): land first-class fields and replace these blocks. + if ( + options.parallelToolCalls !== undefined || + options.previousResponseId !== undefined || + options.seed !== undefined || + options.logitBias !== undefined || + options.user !== undefined || + options.responseFormat !== undefined + ) { + logger.debug("auth-gateway dropped unsupported typed options", { + api, + parallelToolCalls: options.parallelToolCalls, + previousResponseId: options.previousResponseId, + seed: options.seed, + hasLogitBias: options.logitBias !== undefined, + user: options.user, + hasResponseFormat: options.responseFormat !== undefined, + }); } return opts; } +/** + * Classify an upstream / gateway-internal error into a status code and a + * provider-style error type tag. Used by `handleFormatEndpoint` / + * `handlePassthrough` to drive `route.module.formatError` so every wire + * format emits its native envelope shape. + */ +function classifyGatewayError(err: unknown): { status: number; type: string; message: string } { + const message = err instanceof Error ? err.message : String(err); + const lower = message.toLowerCase(); + + // Custom pi-ai errors may attach a numeric `status` property; honor it + // when present and pick the matching tag. + const statusProp = + typeof err === "object" && err !== null && typeof (err as { status?: unknown }).status === "number" + ? (err as { status: number }).status | 0 + : undefined; + if (statusProp !== undefined) { + if (statusProp === 401 || statusProp === 403) + return { status: statusProp, type: "authentication_error", message }; + if (statusProp === 429) return { status: 429, type: "rate_limit_error", message }; + if (statusProp >= 400 && statusProp < 500) return { status: statusProp, type: "invalid_request_error", message }; + if (statusProp >= 500) return { status: statusProp, type: "upstream_error", message }; + } + + if (err instanceof Error && err.name === "AbortError") return { status: 499, type: "request_aborted", message }; + if (lower.includes("aborted") || lower.includes("abortsignal")) { + return { status: 499, type: "request_aborted", message }; + } + if ( + lower.includes("401") || + lower.includes("403") || + lower.includes("unauthorized") || + lower.includes("forbidden") + ) { + return { status: 401, type: "authentication_error", message }; + } + if (lower.includes("429") || lower.includes("rate") || lower.includes("quota")) { + return { status: 429, type: "rate_limit_error", message }; + } + if (lower.includes("unsupported") || lower.includes("invalid")) { + return { status: 400, type: "invalid_request_error", message }; + } + return { status: 502, type: "upstream_error", message }; +} + +function clientClosedResponse(route: { module: FormatModule }): Response { + return route.module.formatError(499, "request_aborted", "client closed request"); +} + function mirrorRequestAbort(req: Request): AbortController { const controller = new AbortController(); if (req.signal.aborted) { @@ -175,108 +245,7 @@ function mirrorRequestAbort(req: Request): AbortController { 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, - }); -} +// (handlePassthrough removed — see note above.) async function handleFormatEndpoint( route: { module: FormatModule; label: string }, @@ -285,32 +254,31 @@ async function handleFormatEndpoint( peer: string, ): Promise { const controller = mirrorRequestAbort(req); - if (controller.signal.aborted) return clientClosedResponse(); + if (controller.signal.aborted) return clientClosedResponse(route); 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(route); + return route.module.formatError(400, "invalid_request_error", `Invalid JSON body: ${String(error)}`); } - if (controller.signal.aborted) return clientClosedResponse(); + if (controller.signal.aborted) return clientClosedResponse(route); // 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`, …). + // field. Read it without running the full strict schema so the route can + // produce a coherent error envelope when the model id is missing. 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" }); + return route.module.formatError(400, "invalid_request_error", "Missing top-level `model` field"); } const model = bootOpts.resolveModel(modelId); if (!model) { - return json(404, { error: `Unknown model: ${modelId}` }); + return route.module.formatError(404, "invalid_request_error", `Unknown model: ${modelId}`); } // pi-ai's stream() does NOT consult AuthStorage — the caller (us) is @@ -319,40 +287,48 @@ async function handleFormatEndpoint( // broker override on AuthStorage when needed). let apiKey: string | undefined; try { - apiKey = await bootOpts.storage.getApiKey(model.provider, undefined, { modelId: model.id }); + apiKey = await bootOpts.storage.getApiKey(model.provider, undefined, { + modelId: model.id, + signal: controller.signal, + }); } 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(route); + const classified = classifyGatewayError(error); + logger.warn("auth-gateway getApiKey threw", { provider: model.provider, peer, error: classified.message }); + return route.module.formatError(classified.status, classified.type, classified.message); } - if (controller.signal.aborted) return clientClosedResponse(); + if (controller.signal.aborted) return clientClosedResponse(route); if (!apiKey) { - return json(401, { error: `No credential available for provider ${model.provider}` }); + return route.module.formatError( + 401, + "authentication_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). + // 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. There is no + // passthrough fast-path — every request flows through pi-ai so that + // credential-specific request shaping (OAuth Claude-Code prefix, beta + // headers, codex websocket transport, …) always applies. let parsed: ParsedFormatRequest; try { - parsed = route.module.parseRequest(body); + parsed = route.module.parseRequest(body, req.headers); } catch (error) { + if (controller.signal.aborted) return clientClosedResponse(route); const message = error instanceof Error ? error.message : String(error); - return json(400, { error: message }); + return route.module.formatError(400, "invalid_request_error", message); } - if (controller.signal.aborted) return clientClosedResponse(); + // Merge gateway-captured passthrough headers under the parser's own + // captures. Parsers that set `options.headers` themselves win (they may + // have stripped or normalized values); the gateway's allow-list fills in + // anything they didn't touch. + { + const captured = captureRequestHeaders(req.headers); + parsed.options.headers = { ...captured, ...(parsed.options.headers ?? {}) }; + } + if (controller.signal.aborted) return clientClosedResponse(route); const streamOpts = buildStreamOptions(parsed, model.api, controller.signal); streamOpts.apiKey = apiKey; @@ -368,17 +344,17 @@ async function handleFormatEndpoint( let events: AssistantMessageEventStream; try { - if (controller.signal.aborted) return clientClosedResponse(); + if (controller.signal.aborted) return clientClosedResponse(route); 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 }); + const classified = classifyGatewayError(error); + logger.warn("auth-gateway streamSimple threw", { format: route.label, error: classified.message, peer }); + return route.module.formatError(classified.status, classified.type, classified.message); } if (!parsed.stream) { try { - if (controller.signal.aborted) return clientClosedResponse(); + if (controller.signal.aborted) return clientClosedResponse(route); const message = await events.result(); if (message.stopReason === "aborted" || message.stopReason === "error") { const errorMessage = @@ -390,17 +366,25 @@ async function handleFormatEndpoint( error: errorMessage, peer, }); - return json(message.stopReason === "aborted" ? 499 : 502, { error: errorMessage }); + if (message.stopReason === "aborted") { + return route.module.formatError(499, "request_aborted", errorMessage); + } + const classified = classifyGatewayError(new Error(errorMessage)); + return route.module.formatError(classified.status, classified.type, 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(route); + const classified = classifyGatewayError(error); + logger.warn("auth-gateway non-streaming aborted", { + format: route.label, + error: classified.message, + peer, + }); + return route.module.formatError(classified.status, classified.type, classified.message); } } - if (controller.signal.aborted) return clientClosedResponse(); + if (controller.signal.aborted) return clientClosedResponse(route); const sseStream = route.module.encodeStream(events, parsed.modelId, parsed.options); return new Response(sseStream, { @@ -409,17 +393,22 @@ async function handleFormatEndpoint( "Content-Type": "text/event-stream; charset=utf-8", "Cache-Control": "no-cache", Connection: "keep-alive", + // Disable proxy buffering (nginx and ingress controllers honor this). + // Without it the SSE stream gets held until the buffer flushes, which + // stalls the long-thinking-budget calls we exist to support. + "X-Accel-Buffering": "no", }, }); } /** - * 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). + * Snapshot of `GET /v1/usage` — `fetchUsageReports` already caches reports at + * a 5-minute per-credential TTL (with jitter, plus last-good fallback on + * failure) 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?.()) ?? []; +async function handleUsage(storage: AuthStorage, signal: AbortSignal): Promise { + const reports = (await storage.fetchUsageReports?.({ signal })) ?? []; // 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. @@ -450,34 +439,42 @@ export function startAuthGateway(opts: AuthGatewayBootOptions): AuthGatewayServe const url = new URL(req.url); const pathname = url.pathname; const peer = resolvePeer(req); + // CORS preflight is always answered without auth — browsers send + // preflights pre-authentication and a 401 here breaks the actual + // request before the bearer is ever attached. + if (req.method === "OPTIONS") { + return new Response(null, { status: 204, headers: corsHeaders(req) }); + } try { if (req.method === "GET" && pathname === "/healthz") { - return json(200, { ok: true, version }); + return withCors(json(200, { ok: true, version }), req); } if (!isAuthorized(req, tokens)) { logger.info("auth-gateway request unauthorized", { method: req.method, path: pathname, peer }); - return json(401, { error: "unauthorized" }); + return withCors(json(401, { error: "unauthorized" }), req); } - // 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 + // Aggregated usage — backed by AuthStorage's 5-min per-credential 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); + return withCors(await handleUsage(opts.storage, req.signal), req); } // Provider-format dispatch. const formatRoute = FORMAT_ROUTES[pathname]; if (formatRoute && req.method === "POST") { - return await handleFormatEndpoint(formatRoute, opts, req, peer); + return withCors(await handleFormatEndpoint(formatRoute, opts, req, peer), req); } // Model catalog. if (req.method === "GET" && pathname === "/v1/models") { - return handleModelsList(opts); + return withCors(handleModelsList(opts), req); } - return json(404, { error: `No route: ${req.method} ${pathname}` }); + // Route-table miss: no format module to defer to, so we emit a + // plain JSON 404 rather than guessing at a protocol-specific envelope. + return withCors(json(404, { error: `No route: ${req.method} ${pathname}` }), req); } catch (error) { logger.error("auth-gateway handler crashed", { method: req.method, @@ -485,9 +482,12 @@ export function startAuthGateway(opts: AuthGatewayBootOptions): AuthGatewayServe peer, error: String(error), }); - return json(500, { error: "internal error" }); + return withCors(json(500, { error: "internal error" }), req); } }, + // Max-out Bun's idle timeout. Long thinking-budget calls can sit idle + // for minutes before the first token arrives; the default kills them. + idleTimeout: 255, }); const boundHost = server.hostname ?? bind.hostname; diff --git a/packages/ai/src/auth-gateway/types.ts b/packages/ai/src/auth-gateway/types.ts index 584f6e970..34e5c4b3e 100644 --- a/packages/ai/src/auth-gateway/types.ts +++ b/packages/ai/src/auth-gateway/types.ts @@ -18,33 +18,78 @@ export const DEFAULT_AUTH_GATEWAY_BIND = "127.0.0.1:4000"; export type AuthGatewayToolChoice = "auto" | "none" | "required" | { name: string }; export interface AuthGatewayParsedRequestOptions { + // ── Sampling ────────────────────────────────────────────────────────── maxOutputTokens?: number; temperature?: number; topP?: number; topK?: number; + /** OpenAI nucleus-min sampling (`min_p`). */ + minP?: number; + /** Anthropic `stop_sequences` / OpenAI `stop`. */ stopSequences?: string[]; + /** OpenAI `presence_penalty`. */ + presencePenalty?: number; + /** OpenAI `frequency_penalty`. */ + frequencyPenalty?: number; + /** OpenRouter / vLLM `repetition_penalty`. */ + repetitionPenalty?: number; + /** OpenAI deterministic-sampling `seed`. */ + seed?: number; + /** OpenAI `logit_bias` map (token id → bias). */ + logitBias?: Record; + /** OpenAI `response_format` (text | json_object | json_schema). Opaque passthrough. */ + responseFormat?: unknown; + + // ── Tools ───────────────────────────────────────────────────────────── toolChoice?: AuthGatewayToolChoice; + /** OpenAI `parallel_tool_calls`. */ + parallelToolCalls?: boolean; + + // ── Reasoning ───────────────────────────────────────────────────────── /** 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. + * Explicit Anthropic `thinking.budget_tokens`. Mirrors Rust's + * `resolve_thinking_budget`: pins onto whichever effort the client + * requested (defaulting to High when unspecified). Preferred over the + * removed legacy single-number `thinkingBudget` for new code. */ - thinkingBudget?: number; + explicitThinkingBudgetTokens?: number; + /** Per-effort thinking budget map. */ + thinkingBudgets?: Partial>; /** Suppress the provider's reasoning summary stream. */ hideThinkingSummary?: boolean; + + // ── Service / routing ───────────────────────────────────────────────── /** 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; + /** OpenAI Responses `prompt_cache_key`; bridges to pi-ai `sessionId`. */ + promptCacheKey?: string; + /** OpenAI Responses `previous_response_id` for response chaining. */ + previousResponseId?: string; + /** OpenAI / abuse-tracking `user` field. */ + user?: string; + + // ── Passthrough ─────────────────────────────────────────────────────── /** - * Provider-specific request controls that need server-side routing support - * but aren't yet first-class on this interface. + * Provider-specific metadata. Anthropic uses `metadata.user_id`; OpenRouter + * carries routing hints; xAI uses `search_parameters`; OpenAI accepts a + * free-form bag. The gateway forwards as-is. + */ + metadata?: Record; + /** + * Captured allow-listed passthrough headers (anthropic-beta, + * anthropic-version, openai-organization, openai-project, openai-beta, + * x-stainless-*). Keys are lowercased. + */ + headers?: Record; + /** + * Escape hatch for provider-specific request controls that don't yet have a + * first-class field. Prefer adding a typed field over widening this. */ extra?: Record; } @@ -57,13 +102,19 @@ export interface AuthGatewayParsedRequest { } export interface AuthGatewayFormatModule { - parseRequest(body: unknown): AuthGatewayParsedRequest; + parseRequest(body: unknown, headers?: Headers): AuthGatewayParsedRequest; encodeResponse(message: AssistantMessage, requestedModelId: string): Record; encodeStream( events: AssistantMessageEventStream, requestedModelId: string, options?: AuthGatewayParsedRequestOptions, ): ReadableStream; + /** + * Emit a protocol-specific error envelope. OpenAI returns + * `{ error: { message, type } }`; Anthropic returns + * `{ type: "error", error: { type, message } }`. + */ + formatError(status: number, type: string, message: string): Response; } export interface AuthGatewayServerOptions { diff --git a/packages/ai/src/auth-storage.ts b/packages/ai/src/auth-storage.ts index db6e4768d..de038c3e7 100644 --- a/packages/ai/src/auth-storage.ts +++ b/packages/ai/src/auth-storage.ts @@ -145,11 +145,16 @@ export interface AuthCredentialStore { * implements this against the broker; SQLite stores leave it undefined. * * Precedence: `AuthStorageOptions.refreshOAuthCredential` > this hook > local. + * + * `signal` propagates the agent's cancel (ESC, request abort, …) all the + * way to the broker fetch so a hung connection can't strand the caller + * for `timeoutMs * (maxRetries + 1)`. */ refreshOAuthCredential?( provider: Provider, credentialId: number, credential: OAuthCredential, + signal?: AbortSignal, ): Promise; /** * Optional store-supplied aggregate usage fetch. When present, `AuthStorage` @@ -158,8 +163,26 @@ export interface AuthCredentialStore { * isn't rate-limited like a heavy residential client). * * Precedence: `AuthStorageOptions.fetchUsageReports` > this hook > local fan-out. + * + * `signal` propagates the agent's cancel down to the broker fetch. */ - fetchUsageReports?(): Promise; + fetchUsageReports?(signal?: AbortSignal): Promise; + /** + * Optional store-supplied per-credential usage report lookup. When present, + * `AuthStorage` consults this before its own per-credential upstream fetch + * (`#getUsageReport`). `RemoteAuthCredentialStore` implements this against + * the broker's aggregate `/v1/usage` (one coalesced round-trip shared across + * all callers) so multi-credential ranking on the client never hits the + * upstream provider's rate-limited usage endpoint from the laptop IP. + * + * Returning `null` is authoritative — `AuthStorage` does NOT fall back to + * the local fetch path. The store hook owns the decision, since falling + * back would re-introduce the per-IP rate-limit problem the broker exists + * to avoid. + * + * `signal` propagates the agent's cancel down to the broker fetch. + */ + getUsageReport?(provider: Provider, credential: OAuthCredential, signal?: AbortSignal): Promise; } // ───────────────────────────────────────────────────────────────────────────── @@ -214,6 +237,7 @@ export type AuthStorageOptions = { provider: Provider, credentialId: number, credential: OAuthCredential, + signal?: AbortSignal, ) => Promise; /** * Human-readable description of the credential store backing this @@ -235,7 +259,7 @@ export type AuthStorageOptions = { * Implementations may return null when no usage data is available; the * AuthStorage caller surfaces that to its own consumer unchanged. */ - fetchUsageReports?: () => Promise; + fetchUsageReports?: (signal?: AbortSignal) => Promise; }; // ───────────────────────────────────────────────────────────────────────────── @@ -315,6 +339,12 @@ type UsageRequestDescriptor = { type AuthApiKeyOptions = { baseUrl?: string; modelId?: string; + /** + * Caller's cancel signal. Threaded into any broker-bound OAuth refresh so + * `ESC` / request abort actually kills a hung broker fetch instead of + * stranding the caller for `timeoutMs * (maxRetries + 1)`. + */ + signal?: AbortSignal; }; function requiresOpenAICodexProModel(provider: string, modelId: string | undefined): boolean { @@ -362,6 +392,33 @@ function parseUsageCacheEntry(raw: string): UsageCacheEntry | undefined { } } +/** + * Race `promise` against `signal`, rejecting only this caller when the signal + * fires. The underlying promise keeps running so other awaiters on the same + * single-flight fetch aren't punished by a peer's cancel. + */ +function raceUsageWithSignal(promise: Promise, signal: AbortSignal | undefined): Promise { + if (!signal) return promise; + if (signal.aborted) return Promise.reject(new Error("usage fetch aborted")); + return new Promise((resolve, reject) => { + const onAbort = (): void => { + signal.removeEventListener("abort", onAbort); + reject(new Error("usage fetch aborted")); + }; + signal.addEventListener("abort", onAbort, { once: true }); + promise.then( + value => { + signal.removeEventListener("abort", onAbort); + resolve(value); + }, + err => { + signal.removeEventListener("abort", onAbort); + reject(err); + }, + ); + }); +} + // ───────────────────────────────────────────────────────────────────────────── // Usage Cache (backed by AuthCredentialStore) // ───────────────────────────────────────────────────────────────────────────── @@ -425,7 +482,7 @@ export class AuthStorage { #rankingStrategyResolver?: (provider: Provider) => CredentialRankingStrategy | undefined; #usageCache: UsageCache; #usageRequestInFlight: Map> = new Map(); - #usageReportsInFlight: Map> = new Map(); + #usageReportsInFlight: Map> = new Map(); #usageFetch: typeof fetch; #usageRequestTimeoutMs: number; #usageLogger?: UsageLogger; @@ -1520,6 +1577,7 @@ export class AuthStorage { request.provider, refreshableCredential, refreshableCredentialId, + timeoutSignal, ); const refreshedCredential = this.#mergeRefreshedUsageCredential(request.credential, refreshed); this.#persistRefreshedUsageCredential(request.provider, request.credential, refreshedCredential); @@ -1788,8 +1846,16 @@ export class AuthStorage { async #getUsageReport( provider: Provider, credential: OAuthCredential, - options?: { baseUrl?: string; timeoutMs?: number }, + options?: { baseUrl?: string; timeoutMs?: number; signal?: AbortSignal }, ): Promise { + // Store-level hook (e.g. `RemoteAuthCredentialStore`) is authoritative + // when present: the broker already aggregates usage from a less-throttled + // IP, and falling back to the local per-credential fetch would defeat the + // whole point of routing through it. + const storeHook = this.#store.getUsageReport?.bind(this.#store); + if (storeHook) { + return storeHook(provider, credential, options?.signal); + } return this.#fetchUsageCached( this.#buildUsageRequestForOauth(provider, credential, options?.baseUrl), options?.timeoutMs ?? this.#usageRequestTimeoutMs, @@ -1798,6 +1864,8 @@ export class AuthStorage { async fetchUsageReports(options?: { baseUrlResolver?: (provider: Provider) => string | undefined; + /** Caller's cancel signal; only rejects this caller, never the shared upstream fetch. */ + signal?: AbortSignal; }): Promise { // Caller override > store-level hook > local per-credential fan-out. // `RemoteAuthCredentialStore` implements the store hook so a gateway @@ -1805,7 +1873,21 @@ export class AuthStorage { // needing the caller to wire it explicitly. const override = this.#fetchUsageReportsOverride ?? this.#store.fetchUsageReports?.bind(this.#store); if (override) { - return override(); + // Reuse the in-flight map so concurrent callers (widget poll + format + // dispatch + credential selection) coalesce into one upstream call. + // Each caller's `signal` only cancels THAT caller's await; the + // shared upstream fetch runs to completion so peers aren't punished. + const OVERRIDE_KEY = "__override__"; + let shared = this.#usageReportsInFlight.get(OVERRIDE_KEY); + if (!shared) { + // Don't forward the caller signal into the shared fetch — first caller's + // abort would otherwise cancel the upstream for every peer. + shared = override().finally(() => { + this.#usageReportsInFlight.delete(OVERRIDE_KEY); + }); + this.#usageReportsInFlight.set(OVERRIDE_KEY, shared); + } + return raceUsageWithSignal(shared, options?.signal); } if (!this.#usageProviderResolver) return null; @@ -1877,7 +1959,7 @@ export class AuthStorage { async markUsageLimitReached( provider: string, sessionId: string | undefined, - options?: { retryAfterMs?: number; baseUrl?: string }, + options?: { retryAfterMs?: number; baseUrl?: string; signal?: AbortSignal }, ): Promise { const sessionCredential = this.#getSessionCredential(provider, sessionId); if (!sessionCredential) return false; @@ -1998,14 +2080,23 @@ 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 => + const timeoutSignal = Promise.withResolvers(); + // `Bun.sleep` keeps the event loop alive even after Promise.race resolves, + // which leaks a 7.5–15s timer per credential-selection call. Use an unref'd + // timer so the timeout doesn't pin the process and clear it on the happy + // path so memory drops immediately. + const timer = setTimeout(() => timeoutSignal.resolve(null), usageTimeout); + (timer as { unref?: () => void }).unref?.(); + const usageResults = await Promise.race([usagePromise, timeoutSignal.promise]).then(result => { + clearTimeout(timer); + return ( 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]; @@ -2131,6 +2222,7 @@ export class AuthStorage { provider, candidate.selection.credential, credentialId, + options?.signal, ); candidate.selection.credential = { ...candidate.selection.credential, @@ -2176,6 +2268,7 @@ export class AuthStorage { provider: Provider, credential: OAuthCredential, credentialId: number | undefined, + signal?: AbortSignal, ): Promise { if (Date.now() < credential.expires) return credential; let refreshPromise: Promise; @@ -2185,7 +2278,7 @@ export class AuthStorage { const storeRefresh = this.#store.refreshOAuthCredential?.bind(this.#store); const overrideRefresh = this.#refreshOAuthCredentialOverride ?? storeRefresh; if (overrideRefresh && credentialId !== undefined) { - refreshPromise = overrideRefresh(provider, credentialId, credential); + refreshPromise = overrideRefresh(provider, credentialId, credential, signal); } else { const customProvider = getOAuthProvider(provider); if (customProvider) { @@ -2198,17 +2291,29 @@ export class AuthStorage { } } // Bound the refresh so a slow/hanging token endpoint cannot stall credential selection. + // Caller-driven abort jumps the gun on the timeout — the agent's ESC must + // take priority over the floor timeout. let timeout: NodeJS.Timeout | undefined; - const timeoutPromise = new Promise((_, reject) => { + let onAbort: (() => void) | undefined; + const cancellationPromise = new Promise((_, reject) => { timeout = setTimeout( () => reject(new Error(`OAuth token refresh timed out for provider: ${provider}`)), DEFAULT_OAUTH_REFRESH_TIMEOUT_MS, ); + if (signal) { + if (signal.aborted) { + reject(new Error("OAuth token refresh aborted by caller")); + return; + } + onAbort = () => reject(new Error("OAuth token refresh aborted by caller")); + signal.addEventListener("abort", onAbort, { once: true }); + } }); try { - return await Promise.race([refreshPromise, timeoutPromise]); + return await Promise.race([refreshPromise, cancellationPromise]); } finally { if (timeout) clearTimeout(timeout); + if (signal && onAbort) signal.removeEventListener("abort", onAbort); } } @@ -2276,6 +2381,7 @@ export class AuthStorage { provider, selection.credential, this.#getStoredCredentials(provider)[selection.index]?.id, + options?.signal, ); const apiKey = customProvider.getApiKey ? customProvider.getApiKey(refreshedCredentials) @@ -2519,7 +2625,7 @@ export class AuthStorage { * Returns the redacted snapshot entry for the refreshed row. * Throws when no OAuth credential with that id is loaded. */ - async forceRefreshCredentialById(id: number): Promise { + async forceRefreshCredentialById(id: number, signal?: AbortSignal): Promise { for (const [provider, entries] of this.#data) { const index = entries.findIndex(entry => entry.id === id); if (index === -1) continue; @@ -2530,7 +2636,7 @@ export class AuthStorage { // Pass a clone with expires=0 so the cached not-yet-expired short-circuit // in #refreshOAuthCredential doesn't suppress the requested refresh. const stale: OAuthCredential = { ...target.credential, expires: 0 }; - const refreshed = await this.#refreshOAuthCredential(provider as Provider, stale, id); + const refreshed = await this.#refreshOAuthCredential(provider as Provider, stale, id, signal); const updated: OAuthCredential = { type: "oauth", access: refreshed.access, diff --git a/packages/ai/src/providers/anthropic-messages-server-schema.ts b/packages/ai/src/providers/anthropic-messages-server-schema.ts index 767aacaca..09ab75283 100644 --- a/packages/ai/src/providers/anthropic-messages-server-schema.ts +++ b/packages/ai/src/providers/anthropic-messages-server-schema.ts @@ -38,6 +38,22 @@ export const base64ImageSourceSchema = z.object({ media_type: z.string().min(1), }); +export const urlImageSourceSchema = z.object({ + type: z.literal("url"), + url: z.url(), +}); + +export const fileImageSourceSchema = z.object({ + type: z.literal("file"), + file_id: z.string().min(1), +}); + +export const imageSourceSchema = z.discriminatedUnion("type", [ + base64ImageSourceSchema, + urlImageSourceSchema, + fileImageSourceSchema, +]); + const textBlockSchema = z.object({ type: z.literal("text"), text: z.string(), @@ -46,7 +62,7 @@ const textBlockSchema = z.object({ const imageBlockSchema = z.object({ type: z.literal("image"), - source: base64ImageSourceSchema, + source: imageSourceSchema, cache_control: cacheControlSchema.optional(), }); @@ -54,11 +70,13 @@ const thinkingBlockSchema = z.object({ type: z.literal("thinking"), thinking: z.string(), signature: z.string().optional(), + cache_control: cacheControlSchema.optional(), }); const redactedThinkingBlockSchema = z.object({ type: z.literal("redacted_thinking"), data: z.string(), + cache_control: cacheControlSchema.optional(), }); const toolUseBlockSchema = z.object({ @@ -66,6 +84,7 @@ const toolUseBlockSchema = z.object({ id: z.string().min(1), name: z.string().min(1), input: z.record(z.string(), z.unknown()).optional(), + cache_control: cacheControlSchema.optional(), }); const toolResultContentBlockSchema = z.discriminatedUnion("type", [textBlockSchema, imageBlockSchema]); @@ -78,6 +97,12 @@ const toolResultBlockSchema = z.object({ cache_control: cacheControlSchema.optional(), }); +// Catch-all for content block variants Anthropic ships that the gateway doesn't +// natively understand (server_tool_use, web_search_tool_result, mcp_*, +// container_upload, code_execution_*, document, …). The walker flattens these +// to a text placeholder so legitimate Anthropic clients don't get rejected. +const unknownContentBlockSchema = z.object({ type: z.string() }).loose(); + // ─── System ──────────────────────────────────────────────────────────────── const systemBlockSchema = z.object({ @@ -90,13 +115,19 @@ export const systemSchema = z.union([z.string(), z.array(systemBlockSchema)]).op // ─── Messages ────────────────────────────────────────────────────────────── -const userContentBlockSchema = z.discriminatedUnion("type", [textBlockSchema, imageBlockSchema, toolResultBlockSchema]); +const userContentBlockSchema = z.union([ + z.discriminatedUnion("type", [textBlockSchema, imageBlockSchema, toolResultBlockSchema]), + unknownContentBlockSchema, +]); -const assistantContentBlockSchema = z.discriminatedUnion("type", [ - textBlockSchema, - thinkingBlockSchema, - redactedThinkingBlockSchema, - toolUseBlockSchema, +const assistantContentBlockSchema = z.union([ + z.discriminatedUnion("type", [ + textBlockSchema, + thinkingBlockSchema, + redactedThinkingBlockSchema, + toolUseBlockSchema, + ]), + unknownContentBlockSchema, ]); export const userMessageSchema = z.object({ @@ -122,20 +153,18 @@ export const toolSchema = z.object({ // ─── 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", - }); +// `disable_parallel_tool_use` is accepted on every variant; the walker maps it +// onto `options.parallelToolCalls = !disable_parallel_tool_use`. +export const toolChoiceSchema = z.discriminatedUnion("type", [ + z.object({ type: z.literal("auto"), disable_parallel_tool_use: z.boolean().optional() }), + z.object({ type: z.literal("any"), disable_parallel_tool_use: z.boolean().optional() }), + z.object({ type: z.literal("none"), disable_parallel_tool_use: z.boolean().optional() }), + z.object({ + type: z.literal("tool"), + name: z.string().min(1), + disable_parallel_tool_use: z.boolean().optional(), + }), +]); // ─── Thinking ────────────────────────────────────────────────────────────── @@ -175,11 +204,10 @@ export const anthropicMessagesRequestSchema = z.object({ 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(), + // Anthropic clients commonly send `metadata: { user_id }`; the walker + // surfaces it on `options.metadata` for downstream provider forwarding. + metadata: z.record(z.string(), z.unknown()).optional(), + // Spec fields that the gateway tolerates but doesn't translate yet. container: z.unknown().optional(), context_management: z.unknown().optional(), mcp_servers: z.unknown().optional(), diff --git a/packages/ai/src/providers/anthropic-messages-server.ts b/packages/ai/src/providers/anthropic-messages-server.ts index 63bf05153..e0aeb19be 100644 --- a/packages/ai/src/providers/anthropic-messages-server.ts +++ b/packages/ai/src/providers/anthropic-messages-server.ts @@ -1,3 +1,5 @@ +import { logger } from "@oh-my-pi/pi-utils"; +import { captureRequestHeaders, resolvePromptCacheKey } from "../auth-gateway/http"; import type { AssistantMessage, AssistantMessageEventStream, @@ -38,6 +40,43 @@ export type { ParsedRequest }; type ImageContentPart = { type: "image"; data: string; mimeType: string }; +// Dedup noise from unknown-block-type warnings. Module-scoped so the warn +// fires once per (category, type) pair across the lifetime of the process. +const WARNED_UNKNOWN_BLOCK_TYPES = new Set(); +function warnUnknownBlockType(category: "user" | "assistant", blockType: string): void { + const key = `${category}:${blockType}`; + if (WARNED_UNKNOWN_BLOCK_TYPES.has(key)) return; + WARNED_UNKNOWN_BLOCK_TYPES.add(key); + logger.warn("anthropic-messages: unknown content block flattened to text placeholder", { + category, + blockType, + }); +} + +// pi-ai's `ImageContent` only carries base64 + mimeType. When the inbound +// uses `url` or `file_id` sources we surface a text placeholder so the +// downstream provider still sees a sane history; warn once per source kind. +const WARNED_NON_BASE64_IMAGE_SOURCES = new Set(); +function warnNonBase64ImageSource(sourceType: string): void { + if (WARNED_NON_BASE64_IMAGE_SOURCES.has(sourceType)) return; + WARNED_NON_BASE64_IMAGE_SOURCES.add(sourceType); + logger.warn("anthropic-messages: image source surfaced as text placeholder (pi-ai ImageContent lacks URL channel)", { + sourceType, + }); +} + +// Compact, log-safe stringification for unknown content blocks. Keeps the +// placeholder informative without dumping multi-KB structures into history. +function describeUnknownBlock(block: { type: string }): string { + try { + const json = JSON.stringify(block); + if (json !== undefined && json.length <= 200) return `[${block.type}: ${json}]`; + } catch { + // fall through + } + return `[${block.type}]`; +} + function buildSystemPrompt(raw: AnthropicSystem): string[] | undefined { if (raw === undefined) return undefined; if (typeof raw === "string") return raw.length > 0 ? [raw] : undefined; @@ -91,13 +130,30 @@ function walkUserContent( 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"); + // SDK's typed source covers base64+url; our schema also accepts the + // forward-compat `file` variant. Narrow against a widened shape so + // every variant is handled at runtime regardless of SDK lag. + const source = block.source as { + type: string; + data?: string; + media_type?: string; + url?: string; + file_id?: string; + }; + if (source.type === "base64" && source.data && source.media_type) { + userParts.push({ type: "image", data: source.data, mimeType: source.media_type }); + } else { + warnNonBase64ImageSource(source.type); + const ref = + source.type === "url" ? (source.url ?? "") : source.type === "file" ? (source.file_id ?? "") : ""; + userParts.push({ type: "text", text: `[image: ${ref}]` }); } + } else if (block.type === "tool_result") { + // Anthropic permits tool_result blocks to follow plain text/image + // siblings in the same user message. pi-ai's history is a flat + // sequence of typed messages, so flush the accumulated parts as a + // separate UserMessage before emitting the ToolResultMessage. + flush(); messages.push({ role: "toolResult", toolCallId: block.tool_use_id, @@ -107,6 +163,13 @@ function walkUserContent( isError: block.is_error === true, timestamp, }); + } else { + // Unknown variant (server_tool_use, mcp_*, document, web_search_tool_result, + // container_upload, code_execution_*, …). Flatten to a text placeholder + // so the downstream provider still gets a coherent transcript. + const unknown = block as { type: string }; + warnUnknownBlockType("user", unknown.type); + userParts.push({ type: "text", text: describeUnknownBlock(unknown) }); } } flush(); @@ -143,6 +206,14 @@ function walkAssistantContent( arguments: block.input ?? {}, }); break; + default: { + // Unknown assistant variant (server_tool_use, mcp_tool_use, …). + // Flatten to a text placeholder; warn once per unknown type. + const unknown = block as { type: string }; + warnUnknownBlockType("assistant", unknown.type); + out.push({ type: "text", text: describeUnknownBlock(unknown) }); + break; + } } } return out; @@ -215,7 +286,7 @@ function deriveCacheRetention(data: { return strongest; } -export function parseRequest(body: unknown): ParsedRequest { +export function parseRequest(body: unknown, headers?: Headers): ParsedRequest { const parsed = anthropicMessagesRequestSchema.safeParse(body); if (!parsed.success) { throw new Error(`anthropic-messages: ${parsed.error.message}`); @@ -251,23 +322,48 @@ export function parseRequest(body: unknown): ParsedRequest { if (data.stop_sequences) options.stopSequences = data.stop_sequences; const toolChoice = mapToolChoice(data.tool_choice as AnthropicToolChoice | undefined); if (toolChoice !== undefined) options.toolChoice = toolChoice; + // `disable_parallel_tool_use === true` means the client wants the model to + // emit at most one tool call per turn; map to pi-ai's negated boolean. + // Leave undefined when the field is absent or explicitly `false` so we + // don't override provider defaults. + if (data.tool_choice?.disable_parallel_tool_use === true) { + options.parallelToolCalls = false; + } if (data.thinking) { switch (data.thinking.type) { case "enabled": - options.thinkingBudget = data.thinking.budget_tokens; + options.explicitThinkingBudgetTokens = 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; + options.explicitThinkingBudgetTokens = data.thinking.budget_tokens; } break; } } const cacheRetention = deriveCacheRetention(data); if (cacheRetention !== undefined) options.cacheRetention = cacheRetention; + // Anthropic clients commonly send `metadata: { user_id }`; forward verbatim + // so downstream providers (and our anthropic-passthrough fast-path) can + // preserve abuse-tracking signal. + if (data.metadata !== undefined) { + options.metadata = data.metadata as Record; + } + const cacheKey = resolvePromptCacheKey(body, headers); + if (cacheKey !== undefined) options.promptCacheKey = cacheKey; + // Allow-listed header capture. The gateway's `handleFormatEndpoint` + // already merges its own pre-capture under whatever the parser sets, but + // we populate here too so direct callers of `parseRequest` (tests, custom + // wrappers) see the same surface. `anthropic-version` is the most + // load-bearing — some downstream Anthropic-API targets reject requests + // missing it. + if (headers) { + const captured = captureRequestHeaders(headers); + if (Object.keys(captured).length > 0) options.headers = captured; + } return { modelId: data.model, @@ -364,6 +460,9 @@ export function encodeResponse(message: AssistantMessage, requestedModelId: stri model: requestedModelId, content: encodeContentBlocks(message), stop_reason: mapStopReasonOut(message.stopReason), + // TODO: surface the matched stop sequence once pi-ai's + // `AssistantMessage.stopReason` carries the matched string. Intentionally + // `null` for now (Anthropic schema allows it). stop_sequence: null, usage: encodeUsage(message), }; @@ -409,6 +508,8 @@ export function encodeStream( model: requestedModelId, content: [], stop_reason: null, + // TODO: same as encodeResponse — surface matched stop sequence + // once pi-ai propagates it. stop_sequence: null, usage: encodeUsage(partial), }, @@ -522,6 +623,8 @@ export function encodeStream( controller.enqueue( sseFrame("message_delta", { type: "message_delta", + // TODO: surface matched stop sequence once pi-ai + // propagates it on the `done` event. delta: { stop_reason: mapStopReasonOut(ev.reason), stop_sequence: null }, usage: encodeUsage(ev.message), }), @@ -556,3 +659,19 @@ export function encodeStream( }, }); } + +// --------------------------------------------------------------------------- +// Error envelope +// --------------------------------------------------------------------------- + +/** + * Anthropic error envelope: `{ type: "error", error: { type, message } }`. + * See https://docs.anthropic.com/en/api/errors. Returned as a `Response` so + * the gateway can hand it straight back to the client without extra wrapping. + */ +export function formatError(status: number, type: string, message: string): Response { + return new Response(JSON.stringify({ type: "error", error: { type, message } }), { + status, + headers: { "Content-Type": "application/json" }, + }); +} diff --git a/packages/ai/src/providers/anthropic.ts b/packages/ai/src/providers/anthropic.ts index 5b472f066..7d11a1289 100644 --- a/packages/ai/src/providers/anthropic.ts +++ b/packages/ai/src/providers/anthropic.ts @@ -15,6 +15,7 @@ import { isEnoent, isRetryableError, isUnexpectedSocketCloseMessage, + logger, readSseEvents, } from "@oh-my-pi/pi-utils"; import { hasOpus47ApiRestrictions, mapEffortToAnthropicAdaptiveEffort } from "../model-thinking"; @@ -204,6 +205,9 @@ type AnthropicSamplingParams = MessageCreateParamsStreaming & { top_k?: number; }; +const ANTHROPIC_STOP_SEQUENCES_MAX = 4; +let warnedStopSequencesTrim = false; + /** * Adaptive thinking `display` is supported starting with Claude Opus 4.7. * Older adaptive-thinking models (Opus 4.6, Sonnet 4.6+) reject the field. @@ -1781,6 +1785,18 @@ function buildParams( if (options?.topK !== undefined) { params.top_k = options.topK; } + if (options?.stopSequences?.length) { + const seqs = options.stopSequences; + if (seqs.length > ANTHROPIC_STOP_SEQUENCES_MAX && !warnedStopSequencesTrim) { + warnedStopSequencesTrim = true; + logger.warn("anthropic: stop_sequences exceeds 4; extra entries dropped", { + received: seqs.length, + kept: ANTHROPIC_STOP_SEQUENCES_MAX, + }); + } + params.stop_sequences = + seqs.length > ANTHROPIC_STOP_SEQUENCES_MAX ? seqs.slice(0, ANTHROPIC_STOP_SEQUENCES_MAX) : seqs; + } // Opus 4.7+ rejects non-default sampling parameters with 400 error. if (hasOpus47ApiRestrictions(model.id)) { diff --git a/packages/ai/src/providers/openai-chat-server-schema.ts b/packages/ai/src/providers/openai-chat-server-schema.ts index 0e79f4bcb..727c1f833 100644 --- a/packages/ai/src/providers/openai-chat-server-schema.ts +++ b/packages/ai/src/providers/openai-chat-server-schema.ts @@ -1,8 +1,11 @@ /** * 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. + * the shapes the gateway translation layer understands. Unknown fields on + * permissive objects are accepted-and-stripped (via `z.unknown()` passthroughs + * or `.loose()`) so the official OpenAI SDK — which sends a growing pile of + * non-strict defaults (e.g. `stream_options.include_obfuscation`) — does not + * trip 400s on shapes we simply ignore. */ import type { ChatCompletionContentPart, @@ -22,15 +25,62 @@ export const textPartSchema = z.object({ }); /** - * OpenAI documents `image_url` as either `{ url: string }` or — older clients — - * a bare string. Accept both shapes; downstream we extract a URL. + * OpenAI documents `image_url` as either `{ url: string, detail?: ... }` or — + * older clients — a bare string. Accept both shapes; downstream we extract a + * URL. `detail` is accepted for forward-compat but currently dropped (pi-ai's + * `ImageContent` has no detail field — TODO: plumb through if/when added). */ export const imagePartSchema = z.object({ type: z.literal("image_url"), - image_url: z.union([z.string(), z.object({ url: z.string() })]), + image_url: z.union([ + z.string(), + z.object({ + url: z.string(), + detail: z.enum(["auto", "low", "high"]).optional(), + }), + ]), }); -export const userContentPartSchema = z.union([textPartSchema, imagePartSchema]); +/** OpenAI audio input block (gpt-4o-audio). Accepted; currently dropped downstream. */ +export const inputAudioPartSchema = z.object({ + type: z.literal("input_audio"), + input_audio: z.object({ + data: z.string(), + format: z.enum(["wav", "mp3"]), + }), +}); + +/** OpenAI file input block (file_search / vision-document). Accepted; currently dropped downstream. */ +export const filePartSchema = z.object({ + type: z.literal("file"), + file: z.object({ + file_id: z.string().optional(), + filename: z.string().optional(), + file_data: z.string().optional(), + }), +}); + +/** Replayed assistant refusal block. Accepted; currently dropped downstream. */ +export const refusalPartSchema = z.object({ + type: z.literal("refusal"), + refusal: z.string(), +}); + +/** + * Forward-compat catch-all for unknown content-part types. Matches every other + * `{ type: string, ... }` object so a new OpenAI block kind does not 400 the + * whole request; the walker ignores parts whose `type` it does not know. + */ +export const unknownPartSchema = z.object({ type: z.string() }).loose(); + +export const userContentPartSchema = z.union([ + textPartSchema, + imagePartSchema, + inputAudioPartSchema, + filePartSchema, + refusalPartSchema, + unknownPartSchema, +]); // ─── Tool calls / tools ───────────────────────────────────────────────────── @@ -49,6 +99,8 @@ export const toolSchema = z.object({ name: z.string().min(1), description: z.string().optional(), parameters: z.record(z.string(), z.unknown()).optional(), + /** OpenAI structured-output strict mode. Accepted, not enforced upstream. */ + strict: z.boolean().optional(), }), }); @@ -62,6 +114,12 @@ export const toolChoiceSchema = z.union([ type: z.literal("function"), function: z.object({ name: z.string().min(1) }), }), + // Anthropic-style `{ type: 'tool', name }` — translated to the OpenAI + // function shape in the walker. + z.object({ + type: z.literal("tool"), + name: z.string().min(1), + }), ]); // ─── Messages ─────────────────────────────────────────────────────────────── @@ -95,25 +153,40 @@ export const toolMessageSchema = z.object({ tool_call_id: z.string().optional(), }); +/** + * Legacy `function` role (pre-tools API). Translated to a `tool` role + * canonical message in the walker so downstream providers see one shape. + */ +export const functionMessageSchema = z.object({ + role: z.literal("function"), + name: z.string(), + content: z.string().nullable(), +}); + export const messageSchema = z.discriminatedUnion("role", [ systemMessageSchema, developerMessageSchema, userMessageSchema, assistantMessageSchema, toolMessageSchema, + functionMessageSchema, ]); // ─── Stream options ───────────────────────────────────────────────────────── -export const streamOptionsSchema = z - .object({ - include_usage: z.boolean().optional(), - }) - .strict(); +/** + * Permissive: the official OpenAI SDK sets `include_obfuscation: false` by + * default. We only consume `include_usage`, so unknown keys are silently + * stripped rather than 400'd. + */ +export const streamOptionsSchema = z.object({ + include_usage: z.boolean().optional(), +}); // ─── Stop sequences ───────────────────────────────────────────────────────── -export const stopSchema = z.union([z.string(), z.array(z.string())]); +// OpenAI rejects > 4 stop strings; mirror that at the gateway. +export const stopSchema = z.union([z.string(), z.array(z.string()).max(4)]); // ─── Top-level request ────────────────────────────────────────────────────── @@ -129,14 +202,32 @@ export const openaiChatRequestSchema = z.object({ 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. + + // ── Typed first-class passthroughs (now consumed by the walker) ──────── 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(), + reasoning_effort: z.enum(["minimal", "low", "medium", "high", "xhigh"]).optional(), + parallel_tool_calls: z.boolean().optional(), + service_tier: z.enum(["auto", "default", "flex", "scale", "priority"]).optional(), + metadata: z.record(z.string(), z.unknown()).optional(), + + // ── Accept-and-ignore passthroughs ───────────────────────────────────── + // Forward acceptance only: validating these would 400 on shapes the + // gateway has no opinion on. The downstream provider does the real check. + logprobs: z.unknown().optional(), + top_logprobs: z.unknown().optional(), + prediction: z.unknown().optional(), + modalities: z.unknown().optional(), + audio: z.unknown().optional(), + store: z.unknown().optional(), + prompt_cache_key: z.unknown().optional(), + safety_identifier: z.unknown().optional(), + n: z.unknown().optional(), + web_search_options: z.unknown().optional(), }); /** diff --git a/packages/ai/src/providers/openai-chat-server.ts b/packages/ai/src/providers/openai-chat-server.ts index f2a66f9df..2dabd9d02 100644 --- a/packages/ai/src/providers/openai-chat-server.ts +++ b/packages/ai/src/providers/openai-chat-server.ts @@ -1,4 +1,5 @@ import { randomUUID } from "node:crypto"; +import { resolvePromptCacheKey } from "../auth-gateway/http"; /** * Parsed inbound OpenAI chat-completions request, ready to feed into pi-ai * `stream(model, context, options)`. @@ -10,6 +11,7 @@ import type { Context, ImageContent, Message, + ServiceTier, StopReason, TextContent, Tool, @@ -28,11 +30,26 @@ import { export type { ParsedRequest }; +type ReasoningEffort = NonNullable; + +function isReasoningEffort(value: unknown): value is ReasoningEffort { + return value === "minimal" || value === "low" || value === "medium" || value === "high" || value === "xhigh"; +} + +function isServiceTier(value: unknown): value is ServiceTier { + return value === "auto" || value === "default" || value === "flex" || value === "scale" || value === "priority"; +} + // --------------------------------------------------------------------------- // parseRequest // --------------------------------------------------------------------------- -export function parseRequest(body: unknown): ParsedRequest { +export function parseRequest(body: unknown, headers?: Headers): ParsedRequest { + // Header capture is centralized in `auth-gateway/server.ts` (allow-listed + // headers like openai-organization/openai-project/openai-beta/x-stainless-* + // land on `options.headers` automatically). We consult `headers` here too + // for `resolvePromptCacheKey` to pull a cache identity out of inbound + // vendor-neutral headers when the body doesn't carry one. const parsed = openaiChatRequestSchema.safeParse(body); if (!parsed.success) { throw new Error(`openai-chat: ${parsed.error.message}`); @@ -67,8 +84,16 @@ export function parseRequest(body: unknown): ParsedRequest { ); break; case "tool": - messages.push(buildToolMessage(m.content, m.tool_call_id, now)); + pushToolResultMessages(messages, m.content, m.tool_call_id, undefined, now); break; + case "function": { + // Legacy `function` role (pre-tools API): the message carries the tool's + // name on `name` and its output on `content`. Translate to a canonical + // `toolResult` with a synthetic id (no original id on the wire). + const fn = m as { role: "function"; name: string; content: string | null }; + pushToolResultMessages(messages, fn.content ?? "", undefined, fn.name, now); + break; + } } } @@ -83,39 +108,50 @@ export function parseRequest(body: unknown): ParsedRequest { // 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); + // Schema accepts the Anthropic-style {type:'tool', name} variant that the SDK + // union doesn't model; the normalizer collapses it to a plain name lookup. + const toolChoice = normalizeToolChoice(data.tool_choice as Parameters[0]); const includeStreamingUsage = data.stream_options?.include_usage === true; + // `includeStreamingUsage` is the one genuinely-opaque flag — the streaming + // encoder reads it later off `options.extra`. Everything else now lives on + // a typed field; `extra` stays undefined when only typed values are set. 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; } + const options: ParsedRequest["options"] = {}; + if (maxOutputTokens !== undefined) options.maxOutputTokens = maxOutputTokens; + if (data.temperature !== undefined) options.temperature = data.temperature; + if (data.top_p !== undefined) options.topP = data.top_p; + if (stopSequences) options.stopSequences = stopSequences; + if (toolChoice !== undefined) options.toolChoice = toolChoice; + if (data.presence_penalty !== undefined) options.presencePenalty = data.presence_penalty; + if (data.frequency_penalty !== undefined) options.frequencyPenalty = data.frequency_penalty; + if (data.seed !== undefined) options.seed = data.seed; + if (data.logit_bias !== undefined) options.logitBias = data.logit_bias; + if (data.user !== undefined) options.user = data.user; + if (data.response_format !== undefined) options.responseFormat = data.response_format; + if (data.parallel_tool_calls !== undefined) options.parallelToolCalls = data.parallel_tool_calls; + if (data.reasoning_effort !== undefined && isReasoningEffort(data.reasoning_effort)) { + options.reasoning = data.reasoning_effort; + } + if (data.service_tier !== undefined && isServiceTier(data.service_tier)) { + options.serviceTier = data.service_tier; + } + if (data.metadata !== undefined) options.metadata = data.metadata; + const cacheKey = resolvePromptCacheKey(body, headers); + if (cacheKey !== undefined) options.promptCacheKey = cacheKey; + if (hasExtra) options.extra = extra; + 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 } : {}), - }, + options, }; } @@ -141,6 +177,9 @@ function parseUserLikeContent( continue; } if (part.type !== "image_url") continue; + // input_audio / file / refusal / unknown-type parts are accepted by the + // schema for forward-compat but dropped here — pi-ai's canonical user + // content only models text and image today. const url = typeof part.image_url === "string" ? part.image_url : part.image_url.url; const decoded = decodeDataUri(url); if (decoded) { @@ -215,21 +254,63 @@ function buildAssistantMessage( }; } -function buildToolMessage( - content: string | OpenAIChatContentPart[] | undefined, +/** + * Walk a wire `tool` (or legacy `function`) message into canonical messages. + * Tool-result content may carry images alongside text; pi-ai's + * `ToolResultMessage` accepts both, but most downstream providers ignore + * images on tool results. To mirror Rust's `encode_messages` behavior we + * keep text inside the tool-result message and hoist any image parts into a + * follow-up `user` message so they still reach the model. + */ +function pushToolResultMessages( + messages: Message[], + content: string | OpenAIChatContentPart[] | undefined | null, toolCallId: string | undefined, + toolName: string | undefined, now: number, -): ToolResultMessage { - return { +): void { + const textParts: TextContent[] = []; + const imageParts: ImageContent[] = []; + + if (typeof content === "string") { + if (content.length > 0) textParts.push({ type: "text", text: content }); + } else if (Array.isArray(content)) { + for (const part of content) { + if (part.type === "text") { + textParts.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) { + imageParts.push({ type: "image", data: decoded.data, mimeType: decoded.mimeType }); + } else { + // No fetcher available; degrade gracefully to a text placeholder. + textParts.push({ type: "text", text: `[image: ${url}]` }); + } + } + } + + const toolMsg: ToolResultMessage = { 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) }], + // OpenAI's `tool` role omits the tool name on the wire; the legacy + // `function` role supplies it. Downstream providers tolerate empty. + toolName: toolName ?? "", + content: textParts.length > 0 ? textParts : [{ type: "text", text: "" }], isError: false, timestamp: now, }; + messages.push(toolMsg); + + if (imageParts.length > 0) { + messages.push({ + role: "user", + content: imageParts, + timestamp: now, + }); + } } function buildTools(tools: OpenAIChatTool[]): Tool[] | undefined { @@ -255,7 +336,15 @@ function normalizeStop(value: string | string[] | undefined): string[] | undefin 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 }; + if (typeof value === "object" && value !== null) { + // OpenAI canonical: { type: 'function', function: { name } } + if ("function" in value && value.function) return { name: value.function.name }; + // Anthropic-style passthrough (schema-allowed): { type: 'tool', name } + const anthropicLike = value as unknown as { type?: string; name?: string }; + if (anthropicLike.type === "tool" && typeof anthropicLike.name === "string") { + return { name: anthropicLike.name }; + } + } return undefined; } @@ -264,12 +353,19 @@ function normalizeToolChoice(value: OpenAIChatToolChoice | undefined): ParsedReq // --------------------------------------------------------------------------- export function encodeResponse(message: AssistantMessage, requestedModelId: string): Record { - const { text, toolCalls } = flattenAssistant(message); + const { text, reasoning, toolCalls } = flattenAssistant(message); const responseMessage: Record = { role: "assistant", content: text.length > 0 ? text : null, + // pi-ai does not surface real refusals yet; emit `null` so SDKs that + // probe `.refusal` see the documented field shape rather than missing. + refusal: null, }; + if (reasoning.length > 0) { + // DeepSeek-style / o-series reasoning channel. + responseMessage.reasoning_content = reasoning; + } if (toolCalls.length > 0) { responseMessage.tool_calls = toolCalls.map(tc => ({ id: tc.id, @@ -283,11 +379,15 @@ export function encodeResponse(message: AssistantMessage, requestedModelId: stri object: "chat.completion", created: Math.floor(Date.now() / 1000), model: requestedModelId, + // Real OpenAI always emits this key, even when the value is null. Mirror + // the contract so probing SDKs do not throw on a missing field. + system_fingerprint: null, choices: [ { index: 0, message: responseMessage, finish_reason: mapFinishReason(message.stopReason, toolCalls.length > 0), + logprobs: null, }, ], usage: buildUsage(message), @@ -296,29 +396,45 @@ export function encodeResponse(message: AssistantMessage, requestedModelId: stri function buildUsage(message: AssistantMessage): Record { const promptTokens = message.usage.input + message.usage.cacheRead + message.usage.cacheWrite; - return { + const usage: Record = { prompt_tokens: promptTokens, completion_tokens: message.usage.output, total_tokens: promptTokens + message.usage.output, prompt_tokens_details: { cached_tokens: message.usage.cacheRead }, }; + if (message.usage.reasoningTokens !== undefined) { + usage.completion_tokens_details = { reasoning_tokens: message.usage.reasoningTokens }; + } + return usage; } -function flattenAssistant(message: AssistantMessage): { text: string; toolCalls: ToolCall[] } { +function flattenAssistant(message: AssistantMessage): { + text: string; + reasoning: string; + toolCalls: ToolCall[]; +} { let text = ""; + let reasoning = ""; const toolCalls: ToolCall[] = []; for (const part of message.content) { switch (part.type) { case "text": text += part.text; break; + case "thinking": + reasoning += part.thinking; + break; + case "redactedThinking": + // Opaque blob — surface verbatim on the reasoning channel so the + // concatenation round-trips through clients that just echo it. + reasoning += part.data; + break; case "toolCall": toolCalls.push(part); break; - // thinking / redactedThinking: dropped — openai chat-completions has no reasoning channel. } } - return { text, toolCalls }; + return { text, reasoning, toolCalls }; } function isOnlyRaw(args: Record): boolean { @@ -341,6 +457,8 @@ function stringifyArgs(args: Record): string { function mapFinishReason(reason: StopReason, hasToolCalls: boolean): string { if (reason === "toolUse" || (hasToolCalls && reason === "stop")) return "tool_calls"; if (reason === "length") return "length"; + // pi-ai's StopReason does not currently carry a content-filter signal; + // when it does, map it to "content_filter" here. return "stop"; } @@ -367,7 +485,8 @@ export function encodeStream( object: "chat.completion.chunk", created, model: requestedModelId, - choices: [{ index: 0, delta, finish_reason: finishReason }], + system_fingerprint: null, + choices: [{ index: 0, delta, finish_reason: finishReason, logprobs: null }], ...(includeUsage ? { usage: null } : {}), }); @@ -381,6 +500,7 @@ export function encodeStream( object: "chat.completion.chunk", created, model: requestedModelId, + system_fingerprint: null, choices: [], usage: buildUsage(message), }); @@ -406,6 +526,14 @@ export function encodeStream( } break; + case "thinking_delta": + // DeepSeek-style / o-series reasoning channel. Clients that don't + // understand it ignore the unknown delta key. + if (event.delta.length > 0) { + writeSse(controller, baseChunk({ reasoning_content: event.delta }, null)); + } + break; + case "toolcall_start": { hasToolCalls = true; const idx = nextToolIndex++; @@ -463,7 +591,7 @@ export function encodeStream( return; } - // Drop start / *_start / *_end / thinking_* — chat-completions wire only + // Drop start / *_start / *_end — chat-completions wire only // surfaces deltas and the terminal finish_reason. default: break; @@ -482,3 +610,19 @@ export function encodeStream( }, }); } + +// --------------------------------------------------------------------------- +// formatError +// --------------------------------------------------------------------------- + +/** + * OpenAI chat-completions error envelope: + * `{ error: { message, type } }` + * Matches the shape the official SDK auto-parses into `APIError`. + */ +export function formatError(status: number, type: string, message: string): Response { + return new Response(JSON.stringify({ error: { message, type } }), { + status, + headers: { "Content-Type": "application/json" }, + }); +} diff --git a/packages/ai/src/providers/openai-completions.ts b/packages/ai/src/providers/openai-completions.ts index 26626dc7d..8d044b013 100644 --- a/packages/ai/src/providers/openai-completions.ts +++ b/packages/ai/src/providers/openai-completions.ts @@ -1084,6 +1084,13 @@ function buildParams( if (options?.repetitionPenalty !== undefined) { params.repetition_penalty = options.repetitionPenalty; } + if (options?.stopSequences?.length) { + const seqs = options.stopSequences; + params.stop = seqs.length === 1 ? seqs[0] : seqs.slice(0, 4); + } + if (options?.frequencyPenalty !== undefined) { + params.frequency_penalty = options.frequencyPenalty; + } if (shouldSendServiceTier(options?.serviceTier, model.provider)) { params.service_tier = options.serviceTier; } diff --git a/packages/ai/src/providers/openai-responses-server-schema.ts b/packages/ai/src/providers/openai-responses-server-schema.ts index b2d870d3e..144853b6b 100644 --- a/packages/ai/src/providers/openai-responses-server-schema.ts +++ b/packages/ai/src/providers/openai-responses-server-schema.ts @@ -1,9 +1,11 @@ /** * 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. + * gateway. Mirrors https://platform.openai.com/docs/api-reference/responses. + * + * Unsupported / opaque controls (background/include/metadata/prompt/…) are + * accepted as `z.unknown().optional()` so we silently ignore rather than 400. + * Real clients (codex, openai-python, llm-git) routinely send these and a 400 + * is a worse outcome than dropping them on the floor. */ import type { EasyInputMessage, @@ -17,18 +19,46 @@ import type { } from "openai/resources/responses/responses"; import * as z from "zod/v4"; -// ─── Input items ──────────────────────────────────────────────────────────── +// ─── Input content blocks ─────────────────────────────────────────────────── const inputTextSchema = z.object({ type: z.literal("input_text"), text: z.string(), }); +const plainTextSchema = z.object({ + type: z.literal("text"), + text: z.string(), +}); + +const inputImageBlockSchema = z + .object({ + type: z.literal("input_image"), + detail: z.enum(["auto", "low", "high"]).optional(), + image_url: z.string().optional(), + file_id: z.string().optional(), + }) + .refine(v => typeof v.image_url === "string" || typeof v.file_id === "string", { + message: "input_image requires at least one of `image_url` or `file_id`", + }); + +const inputFileBlockSchema = z.object({ + type: z.literal("input_file"), + file_id: z.string().optional(), + filename: z.string().optional(), + file_data: z.string().optional(), +}); + const outputTextSchema = z.object({ type: z.literal("output_text"), text: z.string(), }); +const outputRefusalSchema = z.object({ + type: z.literal("refusal"), + refusal: z.string(), +}); + const summaryTextSchema = z.object({ type: z.literal("summary_text"), text: z.string(), @@ -39,13 +69,15 @@ const reasoningTextSchema = z.object({ text: z.string(), }); -const plainTextSchema = z.object({ - type: z.literal("text"), - text: z.string(), -}); +const inputContentBlockSchema = z.union([ + inputTextSchema, + plainTextSchema, + inputImageBlockSchema, + inputFileBlockSchema, +]); +const outputContentBlockSchema = z.union([outputTextSchema, plainTextSchema, outputRefusalSchema]); -const inputContentBlockSchema = z.union([inputTextSchema, plainTextSchema]); -const outputContentBlockSchema = z.union([outputTextSchema, plainTextSchema]); +// ─── Input items ──────────────────────────────────────────────────────────── const userMessageItemSchema = z.object({ type: z.literal("message").optional(), @@ -83,7 +115,24 @@ const functionCallItemSchema = z.object({ const functionCallOutputItemSchema = z.object({ type: z.literal("function_call_output"), call_id: z.string().min(1), - output: z.string().optional(), + // Codex CLI replays multimodal tool results in array form (text + refusal). + output: z.union([z.string(), z.array(outputContentBlockSchema)]).optional(), +}); + +const customToolCallItemSchema = z.object({ + type: z.literal("custom_tool_call"), + id: z.string().optional(), + call_id: z.string().min(1), + name: z.string().min(1), + // Raw input string — NOT JSON.stringified. apply_patch flow streams a + // freeform body and reading it as JSON would corrupt it. + input: z.string(), +}); + +const customToolCallOutputItemSchema = z.object({ + type: z.literal("custom_tool_call_output"), + call_id: z.string().min(1), + output: z.string(), }); /** @@ -98,8 +147,10 @@ export const inputItemSchema = z.union([ reasoningItemSchema, functionCallItemSchema, functionCallOutputItemSchema, + customToolCallItemSchema, + customToolCallOutputItemSchema, // Tolerated but not bridged (file_search_call, web_search_call, …). - z.object({ type: z.string() }), + z.object({ type: z.string() }).loose(), ]); // Variant types alias the canonical SDK union members so the walker can @@ -112,6 +163,13 @@ export type OpenAIResponsesReasoningItem = ResponseReasoningItem; export type OpenAIResponsesFunctionCallItem = ResponseFunctionToolCall; export type OpenAIResponsesFunctionCallOutputItem = ResponseInputItem.FunctionCallOutput; +/** Inferred shape of the custom tool call input item (no canonical SDK alias). */ +export type OpenAIResponsesCustomToolCallItem = z.infer; +export type OpenAIResponsesCustomToolCallOutputItem = z.infer; +export type OpenAIResponsesInputImageBlock = z.infer; +export type OpenAIResponsesInputFileBlock = z.infer; +export type OpenAIResponsesOutputRefusalBlock = z.infer; + // ─── Tools ────────────────────────────────────────────────────────────────── export const toolSchema = z.object({ @@ -122,14 +180,30 @@ export const toolSchema = z.object({ 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(), -}); +// Built-in / hosted tool entries (web_search_preview, file_search, …) — accepted +// but skipped by the walker. +const builtinToolSchema = z + .object({ + type: z.string(), + }) + .loose(); // ─── Tool choice ──────────────────────────────────────────────────────────── +const hostedToolType = z.enum([ + "web_search_preview", + "file_search", + "computer_use_preview", + "code_interpreter", + "image_generation", + "mcp", +]); + +const allowedToolEntrySchema = z.object({ + type: z.string(), + name: z.string().optional(), +}); + export const toolChoiceSchema = z.union([ z.literal("auto"), z.literal("none"), @@ -138,13 +212,31 @@ export const toolChoiceSchema = z.union([ type: z.literal("function"), name: z.string().min(1), }), + // Codex apply_patch flow. + z.object({ + type: z.literal("custom"), + name: z.string().min(1), + }), + // Hosted-tool selection (no extra fields). + z.object({ + type: hostedToolType, + }), + // `allowed_tools` — walker treats as auto. + z.object({ + type: z.literal("allowed_tools"), + mode: z.enum(["auto", "required"]), + tools: z.array(allowedToolEntrySchema), + }), ]); // ─── Reasoning config ─────────────────────────────────────────────────────── export const reasoningConfigSchema = z.object({ effort: z.string().optional(), - summary: z.string().optional(), + // `none` maps to hideThinkingSummary; auto/concise/detailed mean "show + // summary". pi-ai has no per-level plumbing for the latter — walker logs + // once and treats them as default. + summary: z.enum(["auto", "concise", "detailed", "none"]).optional(), }); // ─── Stop ─────────────────────────────────────────────────────────────────── @@ -153,12 +245,6 @@ 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(), @@ -174,18 +260,21 @@ export const openaiResponsesRequestSchema = z.object({ store: z.boolean().optional(), previous_response_id: z.string().optional(), parallel_tool_calls: z.boolean().optional(), + prompt_cache_key: z.string().optional(), + metadata: z.unknown().optional(), + user: z.string().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"), + frequency_penalty: z.number().optional(), + // Accepted-but-ignored: include `reasoning.encrypted_content` is the canonical + // way to request reasoning replay — silently accept and drop. + background: z.unknown().optional(), + include: z.unknown().optional(), + prompt: z.unknown().optional(), + safety_identifier: z.unknown().optional(), + text: z.unknown().optional(), + top_logprobs: z.unknown().optional(), + truncation: z.unknown().optional(), }); /** diff --git a/packages/ai/src/providers/openai-responses-server.ts b/packages/ai/src/providers/openai-responses-server.ts index bfc4c980c..7fe25b9be 100644 --- a/packages/ai/src/providers/openai-responses-server.ts +++ b/packages/ai/src/providers/openai-responses-server.ts @@ -7,12 +7,10 @@ * * 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 { logger } from "@oh-my-pi/pi-utils"; +import { resolvePromptCacheKey } from "../auth-gateway/http"; import type { AuthGatewayParsedRequest as ParsedRequest } from "../auth-gateway/types"; import type { AssistantMessage, @@ -24,9 +22,20 @@ import type { Tool, ToolCall, } from "../types"; +import { + type OpenAIResponsesFunctionCallItem, + type OpenAIResponsesFunctionCallOutputItem, + type OpenAIResponsesInputContent, + type OpenAIResponsesOutputContent, + type OpenAIResponsesReasoningItem, + type OpenAIResponsesTool, + openaiResponsesRequestSchema, +} from "./openai-responses-server-schema"; export type { ParsedRequest }; +// ─── narrow guards ────────────────────────────────────────────────────────── + function isReasoningEffort(value: unknown): value is NonNullable { return value === "minimal" || value === "low" || value === "medium" || value === "high" || value === "xhigh"; } @@ -35,7 +44,15 @@ function isServiceTier(value: unknown): value is NonNullable { + return typeof v === "object" && v !== null && !Array.isArray(v); +} + +function asString(v: unknown): string | undefined { + return typeof v === "string" ? v : undefined; +} + +// ─── id helpers ───────────────────────────────────────────────────────────── function uuidNoDashes(): string { return crypto.randomUUID().replace(/-/g, ""); @@ -57,59 +74,124 @@ 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 makeCustomCallId(): string { + return `ctc_${uuidNoDashes()}`; } -function asString(v: unknown): string | undefined { - return typeof v === "string" ? v : undefined; -} +// ─── once-only warnings ───────────────────────────────────────────────────── +// Module-scoped so we don't spam logs once per turn. -// ─── inbound parser ───────────────────────────────────────────────────────── +let warnedImageNotSupported = false; +let warnedFileNotSupported = false; +let warnedReasoningSummaryLevel = false; + +// ─── inbound parser helpers ───────────────────────────────────────────────── 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(""); + // Prefer `summary[]` — mirrors real OpenAI and the openai-responses provider + // which writes the surfaced reasoning summary into `summary[].text`. + const fromSummary = (item.summary ?? []).map(c => c.text).join(""); + if (fromSummary) return fromSummary; + return (item.content ?? []).map(c => c.text).join(""); } -function inputTextOf(blocks: OpenAIResponsesInputContent[] | string | undefined): string | TextContent[] { +type InputBlockUnion = + | { type: "input_text"; text: string } + | { type: "text"; text: string } + | { type: "input_image"; detail?: "auto" | "low" | "high"; image_url?: string; file_id?: string } + | { type: "input_file"; file_id?: string; filename?: string; file_data?: string }; + +/** + * Walk an input message's content array and produce pi-ai's `TextContent[]`. + * `input_image`/`input_file` blocks become bracketed text placeholders since + * pi-ai's `ImageContent` only carries inline base64 data and we have no + * resolver for OpenAI `image_url` / `file_id` references. Logs once per kind. + */ +function inputContentParts(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 }); + for (const raw of blocks) { + const block = raw as InputBlockUnion; + if (block.type === "input_text" || block.type === "text") { + parts.push({ type: "text", text: block.text }); + } else if (block.type === "input_image") { + if (!warnedImageNotSupported) { + warnedImageNotSupported = true; + logger.warn("openai-responses-server: input_image dropped (no pi-ai bridge for image_url/file_id)", { + hasUrl: typeof block.image_url === "string", + hasFileId: typeof block.file_id === "string", + }); + } + const ref = block.image_url ?? block.file_id ?? "?"; + parts.push({ type: "text", text: `[image: ${ref}]` }); + } else if (block.type === "input_file") { + if (!warnedFileNotSupported) { + warnedFileNotSupported = true; + logger.warn("openai-responses-server: input_file dropped (no pi-ai bridge for file_id/file_data)", { + hasFileId: typeof block.file_id === "string", + hasFileData: typeof block.file_data === "string", + }); + } + const ref = block.file_id ?? block.filename ?? "?"; + parts.push({ type: "text", text: `[file: ${ref}]` }); + } } return parts.length === 1 ? parts[0].text : parts; } +type OutputBlockUnion = + | { type: "output_text"; text: string } + | { type: "text"; text: string } + | { type: "refusal"; refusal: string }; + 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 }); + for (const raw of blocks) { + const block = raw as OutputBlockUnion; + if (block.type === "output_text" || block.type === "text") { + out.push({ type: "text", text: block.text }); + } else if (block.type === "refusal") { + // Preserve the refusal reason so history replay still carries it. + out.push({ type: "text", text: `[refusal: ${block.refusal}]` }); + } } return out; } -function mapToolChoice(value: OpenAIResponsesToolChoice | undefined): ParsedRequest["options"]["toolChoice"] { +// The schema accepts a much wider tool_choice union than the SDK type so the +// walker narrows against the local schema shape. +type ParsedToolChoice = + | "auto" + | "none" + | "required" + | { type: "function"; name: string } + | { type: "custom"; name: string } + | { + type: + | "web_search_preview" + | "file_search" + | "computer_use_preview" + | "code_interpreter" + | "image_generation" + | "mcp"; + } + | { type: "allowed_tools"; mode: "auto" | "required"; tools: Array<{ type: string; name?: string }> }; + +function mapToolChoice(value: ParsedToolChoice | 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 }; + if ("type" in value) { + // `custom` (codex apply_patch) and `function` both resolve to the same + // pi-ai shape: pi-ai's dispatcher matches `Tool.name` AND `customWireName`, + // so passing the wire name works for either. + if (value.type === "function" || value.type === "custom") return { name: value.name }; + // Hosted tools + allowed_tools — we don't surface these to pi-ai; fall + // back to letting the model pick a tool freely. + return "auto"; + } return undefined; } @@ -117,7 +199,7 @@ function buildTools(tools: Array | undef if (!tools) return undefined; const out: Tool[] = []; for (const t of tools) { - // Skip non-function tools (web_search_call, file_search_call, …). + // Skip non-function tools (web_search, file_search, …). if (t.type !== "function") continue; const fn = t as Extract; const tool: Tool = { @@ -155,7 +237,32 @@ function ensureAssistantPlaceholder(messages: Message[], modelId: string, now: n return placeholder; } -export function parseRequest(body: unknown): ParsedRequest { +/** Flatten a function_call_output array form (text + refusal) into a single string. */ +function flattenFunctionOutputArray(blocks: readonly unknown[]): string { + const parts: string[] = []; + for (const raw of blocks) { + if (!isObj(raw)) continue; + const t = raw.type; + if (t === "output_text" || t === "text") { + const text = asString(raw.text); + if (text) parts.push(text); + } else if (t === "refusal") { + const refusal = asString(raw.refusal); + if (refusal) parts.push(`[refusal: ${refusal}]`); + } + } + return parts.join(""); +} + +// ─── parseRequest ─────────────────────────────────────────────────────────── + +export function parseRequest(body: unknown, headers?: Headers): ParsedRequest { + // Header capture is centralized in `auth-gateway/server.ts` (the + // allow-listed set lands on `options.headers` automatically). We also + // consult `headers` here to populate `options.promptCacheKey` when the + // client signals a cache identity outside the body — see the + // `resolvePromptCacheKey` call further down. + const parsed = openaiResponsesRequestSchema.safeParse(body); if (!parsed.success) { throw new Error(`openai-responses: ${parsed.error.message}`); @@ -183,14 +290,14 @@ export function parseRequest(body: unknown): ParsedRequest { }; switch (msg.role) { case "system": { - const text = inputTextOf(msg.content as OpenAIResponsesInputContent[] | string | undefined); + const text = inputContentParts(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); + const content = inputContentParts(msg.content as OpenAIResponsesInputContent[] | string | undefined); messages.push({ role: msg.role, content, timestamp: now }); break; } @@ -235,8 +342,8 @@ export function parseRequest(body: unknown): ParsedRequest { const argsRaw = call.arguments ?? "{}"; let args: Record; try { - const parsed: unknown = JSON.parse(argsRaw); - args = isObj(parsed) ? parsed : {}; + const parsedArgs: unknown = JSON.parse(argsRaw); + args = isObj(parsedArgs) ? parsedArgs : {}; } catch { throw new Error(`openai-responses: function_call ${call.call_id} has invalid JSON arguments`); } @@ -250,26 +357,49 @@ export function parseRequest(body: unknown): ParsedRequest { ensureAssistantPlaceholder(messages, data.model, now).content.push(toolCall); continue; } + if (effectiveType === "custom_tool_call") { + const call = item as { id?: string; call_id: string; name: string; input: string }; + // Custom tools carry a raw input string. We stash it in `arguments.input` + // matching pi-ai's openai-responses-shared convention, and tag the call + // with `customWireName` so encoders re-emit it as `custom_tool_call`. + const toolCall: ToolCall = { + type: "toolCall", + id: call.call_id, + name: call.name, + arguments: { input: call.input ?? "" }, + customWireName: call.name, + ...(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; - } + const toolName = findToolNameById(messages, output.call_id); + const text = + typeof output.output === "string" + ? output.output + : Array.isArray(output.output) + ? flattenFunctionOutputArray(output.output) + : ""; messages.push({ role: "toolResult", toolCallId: output.call_id, toolName, - content: [{ type: "text", text: typeof output.output === "string" ? output.output : "" }], + content: [{ type: "text", text }], + isError: false, + timestamp: now, + }); + continue; + } + if (effectiveType === "custom_tool_call_output") { + const output = item as { call_id: string; output: string }; + const toolName = findToolNameById(messages, output.call_id); + messages.push({ + role: "toolResult", + toolCallId: output.call_id, + toolName, + content: [{ type: "text", text: output.output ?? "" }], isError: false, timestamp: now, }); @@ -292,25 +422,41 @@ export function parseRequest(body: unknown): ParsedRequest { if (data.stop !== undefined && data.stop !== null) { options.stopSequences = typeof data.stop === "string" ? [data.stop] : data.stop; } - const toolChoice = mapToolChoice(data.tool_choice); + const toolChoice = mapToolChoice(data.tool_choice as ParsedToolChoice | undefined); 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") { + // OpenAI summary: `none` → suppress; `auto`/`concise`/`detailed` → request + // visible summary. pi-ai has no per-level plumbing — log once and let the + // provider default kick in. + if (data.reasoning?.summary === "none") { options.hideThinkingSummary = true; + } else if ( + data.reasoning?.summary === "auto" || + data.reasoning?.summary === "concise" || + data.reasoning?.summary === "detailed" + ) { + if (!warnedReasoningSummaryLevel) { + warnedReasoningSummaryLevel = true; + logger.debug("openai-responses-server: reasoning.summary level not differentiated", { + level: data.reasoning.summary, + }); + } } 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. + if (data.frequency_penalty !== undefined) options.frequencyPenalty = data.frequency_penalty; + if (data.parallel_tool_calls !== undefined) options.parallelToolCalls = data.parallel_tool_calls; + const cacheKey = resolvePromptCacheKey(body, headers); + if (cacheKey !== undefined) options.promptCacheKey = cacheKey; + if (data.previous_response_id !== undefined) options.previousResponseId = data.previous_response_id; + if (data.user !== undefined) options.user = data.user; + if (isObj(data.metadata)) options.metadata = data.metadata; + // `store` is a stateful-storage hint that omp's gateway doesn't honour; + // silently accepted by the schema. No typed slot — drop. return { modelId: data.model, @@ -320,25 +466,61 @@ export function parseRequest(body: unknown): ParsedRequest { }; } +function findToolNameById(messages: Message[], callId: string): string { + 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 === callId) return c.name; + } + } + return ""; +} + +// ─── formatError ──────────────────────────────────────────────────────────── + +export function formatError(status: number, type: string, message: string): Response { + return new Response(JSON.stringify({ error: { message, type } }), { + status, + headers: { "Content-Type": "application/json" }, + }); +} + // ─── 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 }>; + summary: Array<{ type: "summary_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 MessageOutputItem = { + type: "message"; + id: string; + role: "assistant"; + status: "completed"; + content: Array<{ type: "output_text"; text: string; annotations: never[] }>; +}; + +type FunctionCallOutputItem = { + type: "function_call"; + id: string; + call_id: string; + name: string; + arguments: string; + status: "completed"; +}; + +type CustomToolCallOutputItem = { + type: "custom_tool_call"; + id: string; + call_id: string; + name: string; + input: string; + status: "completed"; +}; + +type OutputItem = ReasoningOutputItem | MessageOutputItem | FunctionCallOutputItem | CustomToolCallOutputItem; type ResponseStatus = "completed" | "in_progress" | "failed" | "incomplete"; @@ -349,22 +531,29 @@ function responseStatusForStopReason(message: AssistantMessage): ResponseStatus } function buildReasoningItem(part: ThinkingContent): ReasoningOutputItem { + const baseId = part.itemId ?? makeReasoningId(); 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; + const sigParsed: unknown = JSON.parse(part.thinkingSignature); + if (isObj(sigParsed) && sigParsed.type === "reasoning") { + const id = part.itemId ?? asString(sigParsed.id) ?? makeReasoningId(); + // Preserve any extra fields (encrypted_content, …) the original carried, + // but normalize the summary into the canonical `{type, text}[]` shape. + const merged: Record = { ...sigParsed, type: "reasoning", id }; + merged.summary = [{ type: "summary_text", text: part.thinking }]; + // `content[]` is the encrypted/raw side-channel; leave whatever was + // already there. If absent, omit — real OpenAI only emits `content[]` + // when `include=['reasoning.encrypted_content']` is set. + return merged as ReasoningOutputItem; } } catch { - // Not a serialized Responses reasoning item; fall back to raw thinking text. + // Not a serialized Responses reasoning item; fall through to fresh build. } } return { type: "reasoning", - id: part.itemId ?? makeReasoningId(), - summary: [], - content: [{ type: "reasoning_text", text: part.thinking }], + id: baseId, + summary: [{ type: "summary_text", text: part.thinking }], }; } @@ -372,9 +561,9 @@ 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); + const sigParsed: unknown = JSON.parse(part.thinkingSignature); + if (isObj(sigParsed)) { + const id = asString(sigParsed.id); if (id) return id; } } catch { @@ -390,7 +579,7 @@ function reasoningItemId(part: ThinkingContent): string { */ function buildOutputItems(message: AssistantMessage): OutputItem[] { const out: OutputItem[] = []; - let pendingMessage: Extract | null = null; + let pendingMessage: MessageOutputItem | null = null; const flushMessage = () => { if (pendingMessage) { out.push(pendingMessage); @@ -415,17 +604,28 @@ function buildOutputItems(message: AssistantMessage): OutputItem[] { 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", - }); + if (part.customWireName) { + const rawInput = typeof part.arguments?.input === "string" ? (part.arguments.input as string) : ""; + out.push({ + type: "custom_tool_call", + id: part.thoughtSignature ?? makeCustomCallId(), + call_id: part.id, + name: part.customWireName, + input: rawInput, + status: "completed", + }); + } else { + out.push({ + type: "function_call", + id: part.thoughtSignature ?? makeFuncCallId(), + call_id: part.id, + name: part.name, + arguments: JSON.stringify(part.arguments ?? {}), + status: "completed", + }); + } } - // RedactedThinking is silently dropped — no direct Responses wire representation. + // RedactedThinking / Image are silently dropped — no direct Responses wire representation. } flushMessage(); return out; @@ -501,6 +701,8 @@ interface OpenFunctionCall { callId: string; name: string; argsText: string; + /** Set when the underlying ToolCall is a custom-tool emission. */ + customWireName?: string; } type OpenItem = OpenMessage | OpenReasoning | OpenFunctionCall; @@ -529,6 +731,16 @@ export function encodeStream( const state: { open: OpenItem | null } = { open: null }; const finishedItems: OutputItem[] = []; + const responseSnapshot = (status: ResponseStatus, output: OutputItem[] | []) => ({ + id: responseId, + object: "response", + created_at: createdAt, + status, + model: requestedModelId, + output, + usage: null, + }); + const openMessage = (): OpenMessage => { const itemId = makeMsgId(); const item = { @@ -557,10 +769,18 @@ export function encodeStream( const item = { type: "reasoning" as const, id: itemId, - summary: [] as never[], - content: [] as Array<{ type: "reasoning_text"; text: string }>, + summary: [] as Array<{ type: "summary_text"; text: string }>, }; emit("response.output_item.added", { output_index: outputIndex, item }); + // Open the summary part. Real OpenAI streams summary text in the + // canonical `reasoning_summary_*` lifecycle; pi-ai's own decoder + // reads `summary[].text` from the eventual `output_item.done`. + emit("response.reasoning_summary_part.added", { + item_id: itemId, + output_index: outputIndex, + summary_index: 0, + part: { type: "summary_text", text: "" }, + }); const next: OpenReasoning = { kind: "reasoning", itemId, outputIndex, reasoningText: "" }; state.open = next; return next; @@ -569,19 +789,41 @@ export function encodeStream( 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 customWireName: string | undefined = + tc && typeof tc.customWireName === "string" && tc.customWireName.length > 0 + ? tc.customWireName + : undefined; + const isCustom = customWireName !== undefined; + const itemId = tc?.thoughtSignature ?? (isCustom ? makeCustomCallId() : 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", - }; + const name = customWireName ?? tc?.name ?? ""; + const item = isCustom + ? { + type: "custom_tool_call" as const, + id: itemId, + call_id: callId, + name, + input: "", + status: "in_progress", + } + : { + 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: "" }; + const next: OpenFunctionCall = { + kind: "function_call", + itemId, + outputIndex, + callId, + name, + argsText: "", + ...(isCustom ? { customWireName } : {}), + }; state.open = next; return next; }; @@ -589,8 +831,6 @@ export function encodeStream( 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, @@ -607,38 +847,57 @@ export function encodeStream( content: state.open.content, }); } else if (state.open.kind === "reasoning") { + const summary = [{ type: "summary_text" as const, text: state.open.reasoningText ?? "" }]; const item = { type: "reasoning", id: state.open.itemId, - summary: [], - content: [{ type: "reasoning_text", text: state.open.reasoningText ?? "" }], + summary, }; 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 ?? "" }], + summary, }); } 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", - }); + const text = state.open.argsText ?? ""; + if (state.open.customWireName) { + const item = { + type: "custom_tool_call", + id: state.open.itemId, + call_id: state.open.callId ?? "", + name: state.open.customWireName, + input: text, + status: "completed", + }; + emit("response.output_item.done", { output_index: state.open.outputIndex, item }); + finishedItems.push({ + type: "custom_tool_call", + id: state.open.itemId, + call_id: state.open.callId ?? "", + name: state.open.customWireName, + input: text, + status: "completed", + }); + } else { + const item = { + type: "function_call", + id: state.open.itemId, + call_id: state.open.callId ?? "", + name: state.open.name ?? "", + arguments: text, + 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: text, + status: "completed", + }); + } } outputIndex++; state.open = null; @@ -652,20 +911,24 @@ export function encodeStream( switch (ev.type) { case "start": { createdAt = Math.floor((ev.partial.timestamp || Date.now()) / 1000); + // response.created — initial envelope. 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, - }, + response: responseSnapshot("in_progress", []), + }), + ), + ); + // response.in_progress — mirrors real OpenAI; some clients gate + // on it before reading items. + controller.enqueue( + encoder.encode( + sseEvent("response.in_progress", { + type: "response.in_progress", + sequence_number: seq(), + response: responseSnapshot("in_progress", []), }), ), ); @@ -701,6 +964,9 @@ export function encodeStream( delta: ev.delta, logprobs: [], }); + // TODO: when pi-ai surfaces output_text annotations + // (web_search citations, …), emit + // `response.output_text.annotation.added` here. break; } case "text_end": { @@ -734,10 +1000,10 @@ export function encodeStream( if (!state.open || state.open.kind !== "reasoning") break; const cur: OpenReasoning = state.open; cur.reasoningText += ev.delta; - emit("response.reasoning_text.delta", { + emit("response.reasoning_summary_text.delta", { item_id: cur.itemId, output_index: cur.outputIndex, - content_index: 0, + summary_index: 0, delta: ev.delta, }); break; @@ -747,12 +1013,18 @@ export function encodeStream( const cur: OpenReasoning = state.open; const text = ev.content ?? cur.reasoningText; cur.reasoningText = text; - emit("response.reasoning_text.done", { + emit("response.reasoning_summary_text.done", { item_id: cur.itemId, output_index: cur.outputIndex, - content_index: 0, + summary_index: 0, text, }); + emit("response.reasoning_summary_part.done", { + item_id: cur.itemId, + output_index: cur.outputIndex, + summary_index: 0, + part: { type: "summary_text", text }, + }); closeOpen(); break; } @@ -765,29 +1037,56 @@ export function encodeStream( 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, - }); + if (cur.customWireName) { + emit("response.custom_tool_call_input.delta", { + item_id: cur.itemId, + output_index: cur.outputIndex, + delta: ev.delta, + }); + } else { + 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, - }); + // Promote possibly-late info from the canonical ToolCall. + const tc = ev.toolCall; + if (tc.customWireName && !cur.customWireName) cur.customWireName = tc.customWireName; + if (tc.thoughtSignature) cur.itemId = tc.thoughtSignature; + cur.callId = tc.id; + cur.name = cur.customWireName ?? tc.name; + if (cur.customWireName) { + // Custom tool: raw input string. Streamed deltas accumulated + // the wire-level body; fall back to `arguments.input` from + // the finalized ToolCall when nothing streamed (rare). + const rawInput = + cur.argsText || + (typeof tc.arguments?.input === "string" ? (tc.arguments.input as string) : ""); + cur.argsText = rawInput; + emit("response.custom_tool_call_input.done", { + item_id: cur.itemId, + output_index: cur.outputIndex, + input: rawInput, + name: cur.name, + }); + } else { + // Standard JSON tool: arguments object on the omp side, the + // wire wants the JSON string the model emitted (= streamed deltas). + const argsJson = cur.argsText || JSON.stringify(tc.arguments ?? {}); + cur.argsText = argsJson; + emit("response.function_call_arguments.done", { + item_id: cur.itemId, + output_index: cur.outputIndex, + arguments: argsJson, + name: cur.name, + }); + } closeOpen(); break; } @@ -810,12 +1109,7 @@ export function encodeStream( type: "response.failed", sequence_number: seq(), response: { - id: responseId, - object: "response", - created_at: createdAt, - status: "failed", - model: requestedModelId, - output: finishedItems, + ...responseSnapshot("failed", finishedItems), error: { message: failureMessage.errorMessage ?? "stream failed" }, }, }), diff --git a/packages/ai/src/providers/openai-responses.ts b/packages/ai/src/providers/openai-responses.ts index 0e7cbb64a..63bbcfedc 100644 --- a/packages/ai/src/providers/openai-responses.ts +++ b/packages/ai/src/providers/openai-responses.ts @@ -171,6 +171,7 @@ type OpenAIResponsesSamplingParams = ResponseCreateParamsStreaming & { min_p?: number; presence_penalty?: number; repetition_penalty?: number; + stream_options?: { include_obfuscation?: boolean }; }; /** @@ -404,9 +405,14 @@ function buildParams( prompt_cache_key: promptCacheKey, prompt_cache_retention: promptCacheKey ? getPromptCacheRetention(model.baseUrl, cacheRetention) : undefined, store: false, + stream_options: model.provider === "openai" ? { include_obfuscation: false } : undefined, }; applyCommonResponsesSamplingParams(params, options, model.provider); + // TODO: openai responses has no top-level `stop`/`stop_sequences`; surface via reasoning.stop? + // `StreamOptions.stopSequences` is intentionally dropped for this provider. + // TODO: openai responses has no top-level `frequency_penalty` field as of the current SDK; + // `StreamOptions.frequencyPenalty` is intentionally dropped for this provider. if (context.tools) { params.tools = convertTools(context.tools, supportsStrictMode(model), model); diff --git a/packages/ai/src/types.ts b/packages/ai/src/types.ts index e94f0493e..5aa9785ac 100644 --- a/packages/ai/src/types.ts +++ b/packages/ai/src/types.ts @@ -220,6 +220,18 @@ export interface StreamOptions { minP?: number; presencePenalty?: number; repetitionPenalty?: number; + /** + * Stop sequences. Anthropic encodes as `stop_sequences` (array, max 4); + * OpenAI chat-completions encodes as `stop` (string or array of up to 4); + * OpenAI Responses API has no `stop` field today (silently dropped by the + * provider when present). + */ + stopSequences?: string[]; + /** + * Frequency penalty (OpenAI). Penalizes new tokens based on existing frequency + * in the text so far. Range -2.0 to 2.0. Parallel to {@link presencePenalty}. + */ + frequencyPenalty?: number; maxTokens?: number; signal?: AbortSignal; apiKey?: string; diff --git a/packages/ai/src/utils/parse-bind.ts b/packages/ai/src/utils/parse-bind.ts new file mode 100644 index 000000000..e55905e49 --- /dev/null +++ b/packages/ai/src/utils/parse-bind.ts @@ -0,0 +1,54 @@ +/** + * Shared `host:port` parser used by the auth-broker and auth-gateway boot + * paths. Centralized so the two servers can't drift on what they accept (the + * gateway used to silently allow empty hostnames; this fixes it). + */ + +export interface ParsedBind { + hostname: string; + 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; +} + +/** + * Parse a `host:port` (or bare `port`, which assumes loopback) string. + * + * Accepts: + * - `"4000"` → `127.0.0.1:4000` + * - `"0.0.0.0:4000"` → as written + * - `"[::1]:4000"` → as written (brackets retained, Bun handles them) + * + * Rejects: + * - empty input + * - empty hostname (`":4000"`) + * - non-integer / out-of-range port + */ +export 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: parsePort(trimmed, raw) }; + } + const lastColon = trimmed.lastIndexOf(":"); + if (lastColon < 0) { + throw new Error(`Invalid bind '${raw}'; expected 'host:port' or 'port'.`); + } + const hostPart = trimmed.slice(0, lastColon); + const portPart = trimmed.slice(lastColon + 1); + if (hostPart.length === 0) { + throw new Error(`Invalid bind '${raw}'; host must not be empty.`); + } + return { hostname: hostPart, port: parsePort(portPart, raw) }; +} diff --git a/packages/ai/test/auth-gateway-anthropic-caching.test.ts b/packages/ai/test/auth-gateway-anthropic-caching.test.ts new file mode 100644 index 000000000..36921fa74 --- /dev/null +++ b/packages/ai/test/auth-gateway-anthropic-caching.test.ts @@ -0,0 +1,184 @@ +/** + * E2E test: exercise an Anthropic conversation through a live auth-gateway and + * assert prompt caching round-trips. Defends against regressions where the + * gateway either strips `cache_control` markers, places them on the wrong + * block, drops them from the upstream wire, or fails to surface + * `cache_creation_input_tokens` / `cache_read_input_tokens` in the response. + * + * Skips unless a local gateway is reachable at the default `127.0.0.1:4000` + * (override via `OMP_E2E_GATEWAY_URL`) AND the bearer token file exists at + * `~/.omp/auth-gateway.token`. + * + * To run: `bun --cwd packages/ai test test/auth-gateway-anthropic-caching.test.ts` + * with the gateway live (`omp auth-gateway serve` or pm2). + */ +import { describe, expect, it } from "bun:test"; +import * as os from "node:os"; +import * as path from "node:path"; +import { isEnoent } from "@oh-my-pi/pi-utils"; + +interface AnthropicUsage { + input_tokens: number; + output_tokens: number; + cache_creation_input_tokens?: number; + cache_read_input_tokens?: number; +} + +interface AnthropicResponse { + type?: string; + stop_reason?: string; + content?: Array<{ type: string; text?: string }>; + usage: AnthropicUsage; + error?: { type: string; message: string }; +} + +const GATEWAY_URL = Bun.env.OMP_E2E_GATEWAY_URL ?? "http://127.0.0.1:4000"; +const TOKEN_PATH = path.join(os.homedir(), ".omp", "auth-gateway.token"); +const MODEL = Bun.env.OMP_E2E_ANTHROPIC_MODEL ?? "claude-sonnet-4-5"; + +async function checkGatewayAvailable(): Promise<{ ok: boolean; token?: string; reason?: string }> { + let token: string; + try { + token = (await Bun.file(TOKEN_PATH).text()).trim(); + } catch (err) { + if (isEnoent(err)) return { ok: false, reason: `no token at ${TOKEN_PATH}` }; + throw err; + } + if (!token) return { ok: false, reason: `empty token at ${TOKEN_PATH}` }; + try { + const res = await fetch(`${GATEWAY_URL}/healthz`, { signal: AbortSignal.timeout(2_000) }); + if (!res.ok) return { ok: false, reason: `healthz returned ${res.status}` }; + } catch (err) { + const msg = err instanceof Error ? err.message : String(err); + return { ok: false, reason: `healthz unreachable: ${msg}` }; + } + return { ok: true, token }; +} + +const gateway = await checkGatewayAvailable(); + +// Build a system prompt that comfortably exceeds Anthropic's 1024-token cache +// floor for Sonnet. Using a deterministic repeated paragraph so cache keys are +// stable across runs of this test. +const SYSTEM_PARAGRAPH = ` +You are a precise assistant participating in an automated end-to-end test of +the omp auth-gateway's Anthropic prompt-caching pipeline. The same system +prompt will be reused across two turns; the gateway must place a cache +breakpoint on the final system block so that the second request hits the +ephemeral cache instead of being re-tokenized from scratch. Always respond +with extreme brevity: a single short word or phrase, never more than five +tokens. Do not add filler, do not add explanations, do not add punctuation +beyond what is strictly necessary. If asked to confirm something, respond +with "yes". If asked to deny, respond with "no". If asked to repeat your +previous reply, repeat it verbatim. Reasoning, hedging, and conversational +preamble are strictly forbidden. This block is intentionally verbose so the +caching threshold is comfortably cleared on every run; please disregard the +verbosity itself and follow the brevity rule above. +`.trim(); + +const SYSTEM_TEXT = Array.from({ length: 12 }, () => SYSTEM_PARAGRAPH).join("\n\n"); + +interface MessageBlock { + role: "user" | "assistant"; + content: string | Array<{ type: string; text?: string }>; +} + +async function callGateway(body: unknown, token: string): Promise { + const res = await fetch(`${GATEWAY_URL}/v1/messages`, { + method: "POST", + headers: { + "Content-Type": "application/json", + Authorization: `Bearer ${token}`, + "anthropic-version": "2023-06-01", + }, + body: JSON.stringify(body), + }); + const text = await res.text(); + let parsed: AnthropicResponse; + try { + parsed = JSON.parse(text) as AnthropicResponse; + } catch { + throw new Error(`gateway returned non-JSON (status=${res.status}): ${text.slice(0, 200)}`); + } + if (parsed.error) { + throw new Error(`gateway error: ${parsed.error.type}: ${parsed.error.message}`); + } + return parsed; +} + +function extractAssistantText(res: AnthropicResponse): string { + const block = res.content?.find(c => c.type === "text"); + return block?.text ?? ""; +} + +describe.skipIf(!gateway.ok)("auth-gateway: anthropic prompt caching e2e", () => { + if (!gateway.ok) { + // Surface the skip reason once so a quick rerun with `-v` shows it. + console.warn(`[skip] anthropic caching e2e: ${gateway.reason}`); + return; + } + const token = gateway.token; + if (!token) throw new Error("invariant: token must be present when gateway.ok is true"); + + it("writes the system prefix to ephemeral cache on turn 1 and reads it on turn 2", async () => { + // Per-run nonce ensures we always start with a cold cache. The bytes + // before the breakpoint must be unique to this run; otherwise a + // previously-warm Anthropic cache entry hits on turn 1 and we lose the + // ability to assert "first turn writes, second turn reads" cleanly. + const nonce = `${Date.now().toString(36)}-${crypto.randomUUID()}`; + const systemTextWithNonce = `${SYSTEM_TEXT}\n\n[run-nonce: ${nonce}]`; + const system = [{ type: "text", text: systemTextWithNonce, cache_control: { type: "ephemeral" } }]; + + // ── Turn 1 ─────────────────────────────────────────────────────── + const turn1Messages: MessageBlock[] = [{ role: "user", content: "Respond with the single word: alpha" }]; + const turn1 = await callGateway( + { + model: MODEL, + max_tokens: 32, + system, + messages: turn1Messages, + }, + token, + ); + + const turn1Text = extractAssistantText(turn1); + expect(turn1Text.length).toBeGreaterThan(0); + + // Anthropic populates cache_creation_input_tokens with the size of the + // content written to the cache. Above the 1024-token floor this MUST + // be > 0 on the first turn or the gateway stripped our cache_control. + const turn1Created = turn1.usage.cache_creation_input_tokens ?? 0; + const turn1Read = turn1.usage.cache_read_input_tokens ?? 0; + expect(turn1Created).toBeGreaterThan(0); + // First turn cannot hit the cache (nothing to read yet). + expect(turn1Read).toBe(0); + + // ── Turn 2: append assistant + new user, re-send with same system ── + const turn2Messages: MessageBlock[] = [ + ...turn1Messages, + { role: "assistant", content: turn1Text }, + { role: "user", content: "Respond with the single word: beta" }, + ]; + const turn2 = await callGateway( + { + model: MODEL, + max_tokens: 32, + system, + messages: turn2Messages, + }, + token, + ); + + const turn2Text = extractAssistantText(turn2); + expect(turn2Text.length).toBeGreaterThan(0); + + // Second turn MUST read from the cache populated by turn 1. If + // cache_read_input_tokens is 0 the gateway either dropped the marker, + // rewrote the cached prefix bytes, or routed the request without + // Anthropic's cache-aware OAuth headers. + const turn2Read = turn2.usage.cache_read_input_tokens ?? 0; + expect(turn2Read).toBeGreaterThan(0); + // The cache read should cover at least the system block we wrote. + expect(turn2Read).toBeGreaterThanOrEqual(turn1Created); + }, 60_000); +}); diff --git a/packages/ai/test/auth-gateway-anthropic-messages.test.ts b/packages/ai/test/auth-gateway-anthropic-messages.test.ts index 1a9834d55..deff20cea 100644 --- a/packages/ai/test/auth-gateway-anthropic-messages.test.ts +++ b/packages/ai/test/auth-gateway-anthropic-messages.test.ts @@ -118,7 +118,7 @@ describe("anthropic-messages parseRequest", () => { 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.explicitThinkingBudgetTokens).toBe(2048); expect(parsed.options.extra).toBeUndefined(); expect(parsed.context.tools).toHaveLength(1); @@ -198,22 +198,24 @@ describe("anthropic-messages parseRequest", () => { 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("splits user text/image blocks into a separate UserMessage before a tool_result", () => { + const parsed = parseRequest({ + model: "m", + max_tokens: 8, + messages: [ + { + role: "user", + content: [ + { type: "text", text: "preface text" }, + { type: "tool_result", tool_use_id: "t1", content: "ok" }, + ], + }, + ], + }); + // Expect a flush before the tool result: user("preface text") then toolResult(t1). + expect(parsed.context.messages).toHaveLength(2); + expect(parsed.context.messages[0]).toMatchObject({ role: "user", content: "preface text" }); + expect(parsed.context.messages[1]!.role).toBe("toolResult"); }); it("rejects missing required fields and unsupported request controls", () => { @@ -222,9 +224,8 @@ describe("anthropic-messages parseRequest", () => { 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. + // `metadata` is tolerated permissively and surfaced on options for + // downstream forwarding (Anthropic clients ship `metadata.user_id`). const withMetadata = parseRequest({ model: "m", max_tokens: 1, @@ -232,6 +233,7 @@ describe("anthropic-messages parseRequest", () => { metadata: { user_id: "u_1" }, }); expect(withMetadata.options.extra).toBeUndefined(); + expect(withMetadata.options.metadata).toEqual({ user_id: "u_1" }); }); }); diff --git a/packages/ai/test/auth-gateway-anthropic-to-codex-caching.test.ts b/packages/ai/test/auth-gateway-anthropic-to-codex-caching.test.ts new file mode 100644 index 000000000..e96d13257 --- /dev/null +++ b/packages/ai/test/auth-gateway-anthropic-to-codex-caching.test.ts @@ -0,0 +1,193 @@ +/** + * E2E test: send an Anthropic Messages request to an OPENAI CODEX backend + * through the auth-gateway and assert prompt caching survives the + * cross-protocol translate path in the other direction. + * + * Pipeline under test: + * client → POST /v1/messages (Anthropic shape, cache_control markers) + * → anthropic-messages parser → omp Context (cacheRetention derived) + * → pi-ai openai-codex-responses provider + * → upstream Codex (ChatGPT-subscription Responses API) + * → assistant stream → anthropic-messages encoder + * → Anthropic-shape response with cache_read_input_tokens carrying + * Codex's cached_tokens (mapped via usage.cacheRead) + * + * Regression surface: the inbound parser strips cache_control hints into + * `cacheRetention`, but the codex provider doesn't consume `cacheRetention` + * directly — caching only works if pi-ai's codex transport reaches Codex + * with an effective cache identity (prompt_cache_key from sessionId, or + * implicit session reuse). If that path breaks, this test catches it. + * + * Skips unless a local gateway is reachable at the default `127.0.0.1:4000` + * (override via `OMP_E2E_GATEWAY_URL`) AND the bearer token file exists at + * `~/.omp/auth-gateway.token`. + * + * To run: `bun --cwd packages/ai test test/auth-gateway-anthropic-to-codex-caching.test.ts` + */ +import { describe, expect, it } from "bun:test"; +import * as os from "node:os"; +import * as path from "node:path"; +import { isEnoent } from "@oh-my-pi/pi-utils"; + +interface AnthropicUsage { + input_tokens: number; + output_tokens: number; + cache_creation_input_tokens?: number; + cache_read_input_tokens?: number; +} + +interface AnthropicResponse { + type?: string; + stop_reason?: string; + content?: Array<{ type: string; text?: string }>; + usage: AnthropicUsage; + error?: { type: string; message: string }; +} + +const GATEWAY_URL = Bun.env.OMP_E2E_GATEWAY_URL ?? "http://127.0.0.1:4000"; +const TOKEN_PATH = path.join(os.homedir(), ".omp", "auth-gateway.token"); +const MODEL = Bun.env.OMP_E2E_CODEX_MODEL ?? "gpt-5.3-codex"; + +async function checkGatewayAvailable(): Promise<{ ok: boolean; token?: string; reason?: string }> { + let token: string; + try { + token = (await Bun.file(TOKEN_PATH).text()).trim(); + } catch (err) { + if (isEnoent(err)) return { ok: false, reason: `no token at ${TOKEN_PATH}` }; + throw err; + } + if (!token) return { ok: false, reason: `empty token at ${TOKEN_PATH}` }; + try { + const res = await fetch(`${GATEWAY_URL}/healthz`, { signal: AbortSignal.timeout(2_000) }); + if (!res.ok) return { ok: false, reason: `healthz returned ${res.status}` }; + } catch (err) { + const msg = err instanceof Error ? err.message : String(err); + return { ok: false, reason: `healthz unreachable: ${msg}` }; + } + return { ok: true, token }; +} + +const gateway = await checkGatewayAvailable(); + +// Long deterministic instructions, repeated to clear Codex's 1024-token +// cache floor with headroom. +const SYSTEM_PARAGRAPH = ` +You are a precise assistant participating in an automated end-to-end test of +the omp auth-gateway's cross-protocol prompt-caching pipeline. The request +arrives over the Anthropic Messages wire format but is fulfilled by an +OpenAI Codex backend, so the gateway must preserve the cached prefix across +the translation. Always respond with extreme brevity: a single short word or +phrase, never more than five tokens. Do not add filler, do not add +explanations, do not add punctuation beyond what is strictly necessary. If +asked to confirm, respond "yes". If asked to deny, respond "no". If asked +to repeat a previous reply, repeat it verbatim. Reasoning, hedging, and +conversational preamble are strictly forbidden. This block is intentionally +verbose so the caching threshold is comfortably cleared on every run; +disregard the verbosity itself and follow the brevity rule above. +`.trim(); + +const SYSTEM_TEXT = Array.from({ length: 12 }, () => SYSTEM_PARAGRAPH).join("\n\n"); + +interface MessageBlock { + role: "user" | "assistant"; + content: string | Array<{ type: string; text?: string }>; +} + +async function callGateway(body: unknown, token: string): Promise { + const res = await fetch(`${GATEWAY_URL}/v1/messages`, { + method: "POST", + headers: { + "Content-Type": "application/json", + Authorization: `Bearer ${token}`, + "anthropic-version": "2023-06-01", + }, + body: JSON.stringify(body), + }); + const text = await res.text(); + let parsed: AnthropicResponse; + try { + parsed = JSON.parse(text) as AnthropicResponse; + } catch { + throw new Error(`gateway returned non-JSON (status=${res.status}): ${text.slice(0, 200)}`); + } + if (parsed.error) { + throw new Error(`gateway error: ${parsed.error.type}: ${parsed.error.message}`); + } + return parsed; +} + +function extractAssistantText(res: AnthropicResponse): string { + const block = res.content?.find(c => c.type === "text"); + return block?.text ?? ""; +} + +describe.skipIf(!gateway.ok)("auth-gateway: anthropic-messages → openai-codex caching e2e", () => { + if (!gateway.ok) { + console.warn(`[skip] anthropic→codex caching e2e: ${gateway.reason}`); + return; + } + const token = gateway.token; + if (!token) throw new Error("invariant: token must be present when gateway.ok is true"); + + it("caches the system prefix across a cross-protocol translate (messages→codex)", async () => { + // Prepend nonce so prefix-tree caching (chunk-level) starts cold on + // every run. Appending wouldn't help — earlier chunks in the prefix + // would still match warm cache entries from prior runs. + const nonce = `${Date.now().toString(36)}-${crypto.randomUUID()}`; + const systemWithNonce = `[run-nonce: ${nonce}]\n\n${SYSTEM_TEXT}`; + const system = [{ type: "text", text: systemWithNonce, cache_control: { type: "ephemeral" } }]; + + // ── Turn 1 ─────────────────────────────────────────────────────── + const turn1Messages: MessageBlock[] = [{ role: "user", content: "Respond with the single word: alpha" }]; + const turn1 = await callGateway( + { + model: MODEL, + max_tokens: 32, + system, + messages: turn1Messages, + }, + token, + ); + + const turn1Text = extractAssistantText(turn1); + expect(turn1Text.length).toBeGreaterThan(0); + + // First turn cannot hit the cache (nonce ensures cold start). + const turn1Read = turn1.usage.cache_read_input_tokens ?? 0; + expect(turn1Read).toBe(0); + // Confirm the prefix actually crossed the 1024-token caching floor. + expect(turn1.usage.input_tokens).toBeGreaterThan(1024); + + // ── Turn 2 ─────────────────────────────────────────────────────── + const turn2Messages: MessageBlock[] = [ + ...turn1Messages, + { role: "assistant", content: turn1Text }, + { role: "user", content: "Respond with the single word: beta" }, + ]; + const turn2 = await callGateway( + { + model: MODEL, + max_tokens: 32, + system, + messages: turn2Messages, + }, + token, + ); + + const turn2Text = extractAssistantText(turn2); + expect(turn2Text.length).toBeGreaterThan(0); + + // Second turn MUST read the cached prefix. If cache_read_input_tokens + // is 0, one of: + // - anthropic-messages parser stripped the cache_control hint and + // downstream lost the cache-retention signal; + // - the codex provider didn't surface a stable cache identity to + // Codex (no prompt_cache_key, no session reuse, etc.); + // - the anthropic-messages encoder forgot to map pi-ai's + // `usage.cacheRead` to `cache_read_input_tokens` on the wire. + const turn2Read = turn2.usage.cache_read_input_tokens ?? 0; + expect(turn2Read).toBeGreaterThan(0); + // Cached read should cover at least the system block we sent. + expect(turn2Read).toBeGreaterThan(1024); + }, 90_000); +}); diff --git a/packages/ai/test/auth-gateway-cache-key.test.ts b/packages/ai/test/auth-gateway-cache-key.test.ts new file mode 100644 index 000000000..4f6a55bca --- /dev/null +++ b/packages/ai/test/auth-gateway-cache-key.test.ts @@ -0,0 +1,70 @@ +import { describe, expect, it } from "bun:test"; +import { resolvePromptCacheKey } from "../src/auth-gateway/http"; + +describe("resolvePromptCacheKey", () => { + it("prefers body.prompt_cache_key over everything else", () => { + const headers = new Headers({ "x-prompt-cache-key": "from-header" }); + expect( + resolvePromptCacheKey( + { + prompt_cache_key: "from-body", + metadata: { session_id: "from-metadata" }, + }, + headers, + ), + ).toBe("from-body"); + }); + + it("falls back to body.metadata.session_id when prompt_cache_key absent", () => { + expect(resolvePromptCacheKey({ metadata: { session_id: "from-metadata" } }, undefined)).toBe("from-metadata"); + }); + + it("falls back to body.metadata.conversation_id", () => { + expect(resolvePromptCacheKey({ metadata: { conversation_id: "conv-1" } }, undefined)).toBe("conv-1"); + }); + + it("prefers explicit metadata.prompt_cache_key over session/conversation ids", () => { + expect( + resolvePromptCacheKey( + { metadata: { prompt_cache_key: "meta-pck", session_id: "sid", conversation_id: "cid" } }, + undefined, + ), + ).toBe("meta-pck"); + }); + + it("falls back to x-prompt-cache-key header when body lacks anything", () => { + expect(resolvePromptCacheKey({}, new Headers({ "x-prompt-cache-key": "hdr-pck" }))).toBe("hdr-pck"); + }); + + it("falls back to codex session_id / conversation_id headers", () => { + expect(resolvePromptCacheKey({}, new Headers({ session_id: "codex-sid" }))).toBe("codex-sid"); + expect(resolvePromptCacheKey({}, new Headers({ conversation_id: "codex-cid" }))).toBe("codex-cid"); + }); + + it("falls back to vendor-neutral x-session-id / x-conversation-id headers", () => { + expect(resolvePromptCacheKey({}, new Headers({ "x-session-id": "x-sid" }))).toBe("x-sid"); + expect(resolvePromptCacheKey({}, new Headers({ "x-conversation-id": "x-cid" }))).toBe("x-cid"); + }); + + it("returns undefined when nothing resolvable is present", () => { + expect(resolvePromptCacheKey({}, new Headers())).toBeUndefined(); + expect(resolvePromptCacheKey({}, undefined)).toBeUndefined(); + expect(resolvePromptCacheKey(null, undefined)).toBeUndefined(); + expect(resolvePromptCacheKey("not-an-object", undefined)).toBeUndefined(); + }); + + it("ignores empty string body fields and empty header values", () => { + expect(resolvePromptCacheKey({ prompt_cache_key: "" }, new Headers({ "x-prompt-cache-key": "fallback" }))).toBe( + "fallback", + ); + }); + + it("ignores non-string body fields", () => { + expect( + resolvePromptCacheKey( + { prompt_cache_key: 123, metadata: { session_id: { nested: "wrong-type" } } }, + new Headers({ "x-session-id": "hdr-sid" }), + ), + ).toBe("hdr-sid"); + }); +}); diff --git a/packages/ai/test/auth-gateway-cross-protocol-caching.test.ts b/packages/ai/test/auth-gateway-cross-protocol-caching.test.ts new file mode 100644 index 000000000..b2736140a --- /dev/null +++ b/packages/ai/test/auth-gateway-cross-protocol-caching.test.ts @@ -0,0 +1,202 @@ +/** + * E2E test: send an OpenAI Responses request to an ANTHROPIC backend through + * the auth-gateway and assert that prompt caching still works across the + * cross-protocol translate path. This is the canonical mixed-format use case + * — clients targeting `/v1/responses` should keep their caching benefits + * regardless of which credential the model resolves to. + * + * Pipeline under test: + * client → POST /v1/responses (OpenAI shape) + * → openai-responses parser → omp Context + * → pi-ai anthropic provider (auto cache_control via cacheRetention) + * → upstream Anthropic (Messages API) + * → assistant stream → openai-responses encoder + * → OpenAI Responses-shape response with input_tokens_details.cached_tokens + * carrying Anthropic's cache_read_input_tokens + * + * The cross-protocol path is exactly where regressions tend to hide: the + * inbound parser silently strips info that the outbound provider needs, the + * encoder forgets to surface a usage subfield, or the per-turn message rebuild + * mutates the cached prefix bytes. + * + * Skips unless a local gateway is reachable at the default `127.0.0.1:4000` + * (override via `OMP_E2E_GATEWAY_URL`) AND the bearer token file exists at + * `~/.omp/auth-gateway.token`. + * + * To run: `bun --cwd packages/ai test test/auth-gateway-cross-protocol-caching.test.ts` + * with the gateway live (`omp auth-gateway serve` or pm2). + */ +import { describe, expect, it } from "bun:test"; +import * as os from "node:os"; +import * as path from "node:path"; +import { isEnoent } from "@oh-my-pi/pi-utils"; + +interface OpenAIResponsesUsage { + input_tokens: number; + output_tokens: number; + input_tokens_details?: { cached_tokens?: number }; + output_tokens_details?: { reasoning_tokens?: number }; + total_tokens?: number; +} + +interface OpenAIResponse { + status?: string; + output?: Array<{ + type: string; + content?: Array<{ type: string; text?: string }>; + }>; + usage: OpenAIResponsesUsage; + error?: { type?: string; message: string }; +} + +const GATEWAY_URL = Bun.env.OMP_E2E_GATEWAY_URL ?? "http://127.0.0.1:4000"; +const TOKEN_PATH = path.join(os.homedir(), ".omp", "auth-gateway.token"); +const MODEL = Bun.env.OMP_E2E_ANTHROPIC_MODEL ?? "claude-sonnet-4-5"; + +async function checkGatewayAvailable(): Promise<{ ok: boolean; token?: string; reason?: string }> { + let token: string; + try { + token = (await Bun.file(TOKEN_PATH).text()).trim(); + } catch (err) { + if (isEnoent(err)) return { ok: false, reason: `no token at ${TOKEN_PATH}` }; + throw err; + } + if (!token) return { ok: false, reason: `empty token at ${TOKEN_PATH}` }; + try { + const res = await fetch(`${GATEWAY_URL}/healthz`, { signal: AbortSignal.timeout(2_000) }); + if (!res.ok) return { ok: false, reason: `healthz returned ${res.status}` }; + } catch (err) { + const msg = err instanceof Error ? err.message : String(err); + return { ok: false, reason: `healthz unreachable: ${msg}` }; + } + return { ok: true, token }; +} + +const gateway = await checkGatewayAvailable(); + +// Long deterministic instructions, repeated to clear Anthropic's 1024-token +// cache floor for Sonnet. +const INSTRUCTIONS_PARAGRAPH = ` +You are a precise assistant participating in an automated end-to-end test of +the omp auth-gateway's cross-protocol prompt-caching pipeline. The request +arrives over the OpenAI Responses wire format but is fulfilled by an +Anthropic backend, so the gateway must preserve the cached prefix across the +translation. Always respond with extreme brevity: a single short word or +phrase, never more than five tokens. Do not add filler, do not add +explanations, do not add punctuation beyond what is strictly necessary. If +asked to confirm, respond "yes". If asked to deny, respond "no". If asked to +repeat a previous reply, repeat it verbatim. Reasoning, hedging, and +conversational preamble are strictly forbidden. This block is intentionally +verbose so the caching threshold is comfortably cleared on every run; +disregard the verbosity itself and follow the brevity rule above. +`.trim(); + +const INSTRUCTIONS = Array.from({ length: 12 }, () => INSTRUCTIONS_PARAGRAPH).join("\n\n"); + +interface ResponseInputMessage { + role: "user" | "assistant" | "developer" | "system"; + content: string | Array<{ type: string; text?: string }>; +} + +async function callGateway(body: unknown, token: string): Promise { + const res = await fetch(`${GATEWAY_URL}/v1/responses`, { + method: "POST", + headers: { + "Content-Type": "application/json", + Authorization: `Bearer ${token}`, + }, + body: JSON.stringify(body), + }); + const text = await res.text(); + let parsed: OpenAIResponse; + try { + parsed = JSON.parse(text) as OpenAIResponse; + } catch { + throw new Error(`gateway returned non-JSON (status=${res.status}): ${text.slice(0, 200)}`); + } + if (parsed.error) { + throw new Error(`gateway error: ${parsed.error.type ?? "unknown"}: ${parsed.error.message}`); + } + return parsed; +} + +function extractAssistantText(res: OpenAIResponse): string { + for (const item of res.output ?? []) { + if (item.type !== "message") continue; + const block = item.content?.find(c => c.type === "output_text"); + if (block?.text) return block.text; + } + return ""; +} + +describe.skipIf(!gateway.ok)("auth-gateway: openai-responses → anthropic caching e2e", () => { + if (!gateway.ok) { + console.warn(`[skip] cross-protocol caching e2e: ${gateway.reason}`); + return; + } + const token = gateway.token; + if (!token) throw new Error("invariant: token must be present when gateway.ok is true"); + + it("caches the instructions prefix across a cross-protocol translate (responses→anthropic)", async () => { + // Prepend nonce so prefix-tree caching (chunk-level) starts cold on every run. + const nonce = `${Date.now().toString(36)}-${crypto.randomUUID()}`; + const instructionsWithNonce = `[run-nonce: ${nonce}]\n\n${INSTRUCTIONS}`; + + // ── Turn 1 ─────────────────────────────────────────────────────── + const turn1Input: ResponseInputMessage[] = [{ role: "user", content: "Respond with the single word: alpha" }]; + const turn1 = await callGateway( + { + model: MODEL, + max_output_tokens: 64, + instructions: instructionsWithNonce, + input: turn1Input, + }, + token, + ); + + const turn1Text = extractAssistantText(turn1); + expect(turn1Text.length).toBeGreaterThan(0); + + // First turn cannot hit the cache (nothing to read yet thanks to the nonce). + const turn1Cached = turn1.usage.input_tokens_details?.cached_tokens ?? 0; + expect(turn1Cached).toBe(0); + // Confirm the request actually crossed the 1024-token caching floor; + // otherwise no cache entry gets created and turn 2 can't possibly read. + expect(turn1.usage.input_tokens).toBeGreaterThan(1024); + + // ── Turn 2: append assistant + new user, re-send with same instructions ── + const turn2Input: ResponseInputMessage[] = [ + ...turn1Input, + { role: "assistant", content: turn1Text }, + { role: "user", content: "Respond with the single word: beta" }, + ]; + const turn2 = await callGateway( + { + model: MODEL, + max_output_tokens: 64, + instructions: instructionsWithNonce, + input: turn2Input, + }, + token, + ); + + const turn2Text = extractAssistantText(turn2); + expect(turn2Text.length).toBeGreaterThan(0); + + // Second turn MUST hit the cache populated by turn 1. The Anthropic + // provider auto-places cache markers via the default `short` retention, + // and the openai-responses encoder maps Anthropic's + // cache_read_input_tokens → input_tokens_details.cached_tokens. + // If cached_tokens is 0, one of: + // - openai-responses parser stripped per-turn content into different + // bytes (so the cache prefix moved), + // - the anthropic provider failed to apply cache_control markers, + // - the encoder forgot to surface the cached-tokens subfield. + const turn2Cached = turn2.usage.input_tokens_details?.cached_tokens ?? 0; + expect(turn2Cached).toBeGreaterThan(0); + // The cached prefix should cover at least the instructions block we + // established on turn 1 — sanity check that we're not catching a + // trivial overlap. + expect(turn2Cached).toBeGreaterThan(1024); + }, 90_000); +}); diff --git a/packages/ai/test/auth-gateway-openai-chat.test.ts b/packages/ai/test/auth-gateway-openai-chat.test.ts index ef138343f..31659ace9 100644 --- a/packages/ai/test/auth-gateway-openai-chat.test.ts +++ b/packages/ai/test/auth-gateway-openai-chat.test.ts @@ -136,10 +136,8 @@ describe("auth-gateway openai-chat: parseRequest", () => { 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, - }); + expect(parsed.options.responseFormat).toEqual({ type: "json_object" }); + expect(parsed.options.extra).toEqual({ includeStreamingUsage: true }); }); it("rejects missing required fields", () => { diff --git a/packages/ai/test/auth-gateway-openai-responses-caching.test.ts b/packages/ai/test/auth-gateway-openai-responses-caching.test.ts new file mode 100644 index 000000000..021033ac5 --- /dev/null +++ b/packages/ai/test/auth-gateway-openai-responses-caching.test.ts @@ -0,0 +1,202 @@ +/** + * E2E test: exercise an OpenAI Responses conversation through a live + * auth-gateway and assert automatic prompt caching round-trips. OpenAI + * Responses caches prefixes ≥1024 tokens automatically — no explicit + * `cache_control` markers — so the bug surface is "did we keep the prefix + * byte-identical across the two turns" and "did we surface + * input_tokens_details.cached_tokens in the response usage block". + * + * Skips unless a local gateway is reachable at the default `127.0.0.1:4000` + * (override via `OMP_E2E_GATEWAY_URL`) AND the bearer token file exists at + * `~/.omp/auth-gateway.token`. + * + * To run: `bun --cwd packages/ai test test/auth-gateway-openai-responses-caching.test.ts` + * with the gateway live (`omp auth-gateway serve` or pm2). + */ +import { describe, expect, it } from "bun:test"; +import * as os from "node:os"; +import * as path from "node:path"; +import { isEnoent } from "@oh-my-pi/pi-utils"; + +interface OpenAIResponsesUsage { + input_tokens: number; + output_tokens: number; + input_tokens_details?: { cached_tokens?: number }; + output_tokens_details?: { reasoning_tokens?: number }; + total_tokens?: number; +} + +interface OpenAIResponse { + status?: string; + output?: Array<{ + type: string; + content?: Array<{ type: string; text?: string }>; + }>; + usage: OpenAIResponsesUsage; + error?: { type?: string; message: string }; +} + +const GATEWAY_URL = Bun.env.OMP_E2E_GATEWAY_URL ?? "http://127.0.0.1:4000"; +const TOKEN_PATH = path.join(os.homedir(), ".omp", "auth-gateway.token"); +// `gpt-5.3-codex` is the model we've verified the ChatGPT-subscription Codex +// backend accepts; older or higher-tier ids 4xx with "model not supported". +const MODEL = Bun.env.OMP_E2E_OPENAI_RESPONSES_MODEL ?? "gpt-5.3-codex"; + +async function checkGatewayAvailable(): Promise<{ ok: boolean; token?: string; reason?: string }> { + let token: string; + try { + token = (await Bun.file(TOKEN_PATH).text()).trim(); + } catch (err) { + if (isEnoent(err)) return { ok: false, reason: `no token at ${TOKEN_PATH}` }; + throw err; + } + if (!token) return { ok: false, reason: `empty token at ${TOKEN_PATH}` }; + try { + const res = await fetch(`${GATEWAY_URL}/healthz`, { signal: AbortSignal.timeout(2_000) }); + if (!res.ok) return { ok: false, reason: `healthz returned ${res.status}` }; + } catch (err) { + const msg = err instanceof Error ? err.message : String(err); + return { ok: false, reason: `healthz unreachable: ${msg}` }; + } + return { ok: true, token }; +} + +const gateway = await checkGatewayAvailable(); + +// Long deterministic instructions, repeated to clear OpenAI's 1024-token +// automatic-caching floor with plenty of headroom. +const INSTRUCTIONS_PARAGRAPH = ` +You are a precise assistant participating in an automated end-to-end test of +the omp auth-gateway's OpenAI Responses prompt-caching pipeline. The same +instructions block will be reused across two turns; OpenAI automatically +caches identical prefixes ≥1024 tokens, so the second turn must see the +same prefix bytes as the first or the cache misses silently. Always respond +with extreme brevity: a single short word or phrase, never more than five +tokens. Do not add filler, do not add explanations, do not add punctuation +beyond what is strictly necessary. If asked to confirm something, respond +with "yes". If asked to deny, respond with "no". If asked to repeat your +previous reply, repeat it verbatim. Reasoning, hedging, and conversational +preamble are strictly forbidden. This block is intentionally verbose so the +caching threshold is comfortably cleared on every run; please disregard the +verbosity itself and follow the brevity rule above. +`.trim(); + +const INSTRUCTIONS = Array.from({ length: 12 }, () => INSTRUCTIONS_PARAGRAPH).join("\n\n"); + +interface ResponseInputMessage { + role: "user" | "assistant" | "developer" | "system"; + content: string | Array<{ type: string; text?: string }>; +} + +async function callGateway(body: unknown, token: string): Promise { + const res = await fetch(`${GATEWAY_URL}/v1/responses`, { + method: "POST", + headers: { + "Content-Type": "application/json", + Authorization: `Bearer ${token}`, + }, + body: JSON.stringify(body), + }); + const text = await res.text(); + let parsed: OpenAIResponse; + try { + parsed = JSON.parse(text) as OpenAIResponse; + } catch { + throw new Error(`gateway returned non-JSON (status=${res.status}): ${text.slice(0, 200)}`); + } + if (parsed.error) { + throw new Error(`gateway error: ${parsed.error.type ?? "unknown"}: ${parsed.error.message}`); + } + return parsed; +} + +function extractAssistantText(res: OpenAIResponse): string { + for (const item of res.output ?? []) { + if (item.type !== "message") continue; + const block = item.content?.find(c => c.type === "output_text"); + if (block?.text) return block.text; + } + return ""; +} + +describe.skipIf(!gateway.ok)("auth-gateway: openai-responses prompt caching e2e", () => { + if (!gateway.ok) { + console.warn(`[skip] openai-responses caching e2e: ${gateway.reason}`); + return; + } + const token = gateway.token; + if (!token) throw new Error("invariant: token must be present when gateway.ok is true"); + + it("automatically caches the instructions prefix across two turns", async () => { + // Per-run nonce ensures we always start with a cold cache. The bytes + // before the cacheable prefix boundary must be unique to this run; + // otherwise a previously-warm cache entry silently hits on turn 1 + // and we lose the ability to assert "first turn cold, second turn warm". + const nonce = `${Date.now().toString(36)}-${crypto.randomUUID()}`; + // Prepend (not append) — OpenAI caches at prefix-tree granularity, so the + // first chunk must differ across runs to guarantee a cold start. + const instructionsWithNonce = `[run-nonce: ${nonce}]\n\n${INSTRUCTIONS}`; + // Stable per-run cache key. The ChatGPT-subscription Codex backend + // only coalesces prefixes across requests when an explicit + // `prompt_cache_key` is set — caching is opt-in there, unlike public + // OpenAI Responses which caches automatically. Reusing the same key + // across both turns is the contract that makes turn 2 hit. + const cacheKey = `omp-e2e-${nonce}`; + + // ── Turn 1 ─────────────────────────────────────────────────────── + const turn1Input: ResponseInputMessage[] = [{ role: "user", content: "Respond with the single word: alpha" }]; + const turn1 = await callGateway( + { + model: MODEL, + max_output_tokens: 64, + instructions: instructionsWithNonce, + prompt_cache_key: cacheKey, + input: turn1Input, + }, + token, + ); + + const turn1Text = extractAssistantText(turn1); + + expect(turn1Text.length).toBeGreaterThan(0); + + // First turn cannot hit the cache (nothing to read yet thanks to the nonce). + const turn1Cached = turn1.usage.input_tokens_details?.cached_tokens ?? 0; + expect(turn1Cached).toBe(0); + // Confirm the request actually crossed the 1024-token caching floor; + // otherwise OpenAI never registers a cache entry and turn 2 can't + // possibly read. + expect(turn1.usage.input_tokens).toBeGreaterThan(1024); + + // ── Turn 2: append assistant + new user, re-send with same instructions ── + const turn2Input: ResponseInputMessage[] = [ + ...turn1Input, + { role: "assistant", content: turn1Text }, + { role: "user", content: "Respond with the single word: beta" }, + ]; + const turn2 = await callGateway( + { + prompt_cache_key: cacheKey, + model: MODEL, + max_output_tokens: 64, + instructions: instructionsWithNonce, + input: turn2Input, + }, + token, + ); + + const turn2Text = extractAssistantText(turn2); + expect(turn2Text.length).toBeGreaterThan(0); + + // Second turn MUST hit the cache populated by turn 1. If + // cached_tokens is 0, the gateway either mutated the prefix bytes + // between turns or failed to surface input_tokens_details from the + // upstream usage block. + const turn2Cached = turn2.usage.input_tokens_details?.cached_tokens ?? 0; + expect(turn2Cached).toBeGreaterThan(0); + // The cached prefix should cover at least the instructions block we + // established on turn 1 — sanity check that we're not catching a + // trivial 64-token overlap. + expect(turn2Cached).toBeGreaterThan(1024); + }, 90_000); +}); diff --git a/packages/ai/test/auth-gateway-openai-responses.test.ts b/packages/ai/test/auth-gateway-openai-responses.test.ts index d08bc038a..fe3a21801 100644 --- a/packages/ai/test/auth-gateway-openai-responses.test.ts +++ b/packages/ai/test/auth-gateway-openai-responses.test.ts @@ -59,8 +59,7 @@ describe("openai-responses parseRequest", () => { const reasoningItem = { type: "reasoning", id: "rs_abc", - summary: [], - content: [{ type: "reasoning_text", text: "The user wants arithmetic." }], + summary: [{ type: "summary_text", text: "The user wants arithmetic." }], }; const parsed = parseRequest({ model: "gpt-5.3-codex-spark", @@ -229,8 +228,7 @@ describe("openai-responses encodeResponse", () => { const reasoningItem = { type: "reasoning", id: "rs_signed", - summary: [], - content: [{ type: "reasoning_text", text: "thinking aloud" }], + summary: [{ type: "summary_text", text: "thinking aloud" }], }; const message: AssistantMessage = { role: "assistant", @@ -328,7 +326,7 @@ describe("openai-responses encodeResponse", () => { }); describe("openai-responses encodeStream", () => { - it("emits response.created, reasoning_text.delta, output_text.delta, function_call_arguments.delta, response.completed, [DONE]", async () => { + it("emits response.created, reasoning_summary_text.delta, output_text.delta, function_call_arguments.delta, response.completed, [DONE]", async () => { const stream = new AssistantMessageEventStream(); const partial: AssistantMessage = { @@ -416,8 +414,8 @@ describe("openai-responses encodeStream", () => { // 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 idxReasoningDelta = names.indexOf("response.reasoning_summary_text.delta"); + const idxReasoningDone = names.indexOf("response.reasoning_summary_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"); @@ -439,7 +437,7 @@ describe("openai-responses encodeStream", () => { expect(idxArgsDone).toBeGreaterThan(idxArgsDelta); expect(idxCompleted).toBeGreaterThan(idxArgsDone); - // reasoning_text.delta must carry item_id matching the signature, and output_index 0. + // reasoning_summary_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); diff --git a/packages/ai/test/remote-auth-store.test.ts b/packages/ai/test/remote-auth-store.test.ts index 0929c72d9..ae22162cf 100644 --- a/packages/ai/test/remote-auth-store.test.ts +++ b/packages/ai/test/remote-auth-store.test.ts @@ -115,4 +115,58 @@ describe("RemoteAuthCredentialStore + AuthStorage integration", () => { expect(() => remoteStore.deleteAuthCredentialsForProvider("anthropic", "x")).toThrow(/read-only/); remoteStore.close(); }); + + test("getUsageReport coalesces parallel callers and matches by identity", async () => { + const brokerClient = new AuthBrokerClient({ url: handle!.url, token }); + const remoteStore = new RemoteAuthCredentialStore({ + client: brokerClient, + initialSnapshot: { generatedAt: 0, credentials: [] }, + }); + + const reportForA = { + provider: "anthropic" as const, + fetchedAt: Date.now(), + limits: [], + metadata: { email: "a@example.com" }, + }; + const reportForB = { + provider: "anthropic" as const, + fetchedAt: Date.now(), + limits: [], + metadata: { email: "b@example.com" }, + }; + const fetchSpy = vi + .spyOn(brokerClient, "fetchUsage") + .mockResolvedValue({ generatedAt: Date.now(), reports: [reportForA, reportForB] }); + + const credA = { + type: "oauth" as const, + access: "ax", + refresh: REMOTE_REFRESH_SENTINEL, + expires: Date.now() + 60_000, + email: "a@example.com", + }; + const credB = { ...credA, email: "b@example.com" }; + + const [resA, resB] = await Promise.all([ + remoteStore.getUsageReport("anthropic", credA), + remoteStore.getUsageReport("anthropic", credB), + ]); + // Parallel callers share a single broker round-trip. + expect(fetchSpy).toHaveBeenCalledTimes(1); + expect(resA?.metadata?.email).toBe("a@example.com"); + expect(resB?.metadata?.email).toBe("b@example.com"); + + // Cached on the second call — still one fetch total. + const cached = await remoteStore.getUsageReport("anthropic", credA); + expect(cached?.metadata?.email).toBe("a@example.com"); + expect(fetchSpy).toHaveBeenCalledTimes(1); + + // Unknown provider → null, no extra fetch. + const miss = await remoteStore.getUsageReport("openai-codex", credA); + expect(miss).toBeNull(); + expect(fetchSpy).toHaveBeenCalledTimes(1); + + remoteStore.close(); + }); }); diff --git a/packages/coding-agent/CHANGELOG.md b/packages/coding-agent/CHANGELOG.md index 133cdf667..2e3faf001 100644 --- a/packages/coding-agent/CHANGELOG.md +++ b/packages/coding-agent/CHANGELOG.md @@ -1,7 +1,6 @@ # Changelog ## [Unreleased] - ### Breaking Changes - Renamed the embedded-documentation internal URL scheme from `pi://` to `omp://`. `OmpProtocolHandler` replaces `PiProtocolHandler`; update any external references accordingly. @@ -10,21 +9,22 @@ ### Added +- Added optional backend push for the auto-QA grievance database (`dev.autoqaPush.enabled`, `dev.autoqaPush.endpoint`, `dev.autoqaPush.token`; env overrides `PI_AUTO_QA_PUSH`, `PI_AUTO_QA_PUSH_URL`, `PI_AUTO_QA_PUSH_TOKEN`). When enabled, every `report_tool_issue` call schedules a background flush that `POST`s pending rows to the configured endpoint and deletes them on HTTP 2xx. Each push carries a stable per-install UUID (`installId`) generated on first use and persisted at `~/.omp/install-id` via `getInstallId()` (new export from `@oh-my-pi/pi-utils`), so the receiver can dedup retries across host renames and `autoqa.db` wipes. Single-flight, 5s request timeout, 30s in-memory cooldown after failure, and a row-id watermark so rows inserted during an in-flight push survive and ship next time. Tool execution remains non-blocking and never throws. - `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`. - - `login [--via=user@host] [--dry-run]` — drives the OAuth flow locally or via SSH `-L` tunnel into a remote broker (callback ports pinned per provider). - - `logout ` — disables every credential for the given provider in the local SQLite store. - - `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`). +- `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`. +- `login [--via=user@host] [--dry-run]` — drives the OAuth flow locally or via SSH `-L` tunnel into a remote broker (callback ports pinned per provider). +- `logout ` — disables every credential for the given provider in the local SQLite store. +- `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. +- `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, 5-min per-credential cache), `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 @@ -32,6 +32,8 @@ ### Fixed +- Fixed `omp auth-broker migrate` to skip local placeholder `` API credentials (not real keys) when exporting to a remote broker +- Fixed `auth-gateway` token initialization to avoid clobbering an existing token when multiple processes initialize it concurrently - 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. diff --git a/packages/coding-agent/src/cli/auth-broker-cli.ts b/packages/coding-agent/src/cli/auth-broker-cli.ts index c4973c1bf..2876bfad2 100644 --- a/packages/coding-agent/src/cli/auth-broker-cli.ts +++ b/packages/coding-agent/src/cli/auth-broker-cli.ts @@ -549,6 +549,19 @@ async function runMigrate(flags: AuthBrokerCommandArgs["flags"]): Promise const plannedApiKeyProviders = new Set(); try { for (const row of localStore.listAuthCredentials()) { + // Skip placeholder sentinels that pi-ai treats as "authenticated via + // out-of-band mechanism" (Bedrock/Vertex ``). They + // aren't real keys and uploading them would store garbage on the + // broker. Mirrors the env-var path's guard below. + if (row.credential.type === "api_key" && row.credential.key === "") { + skipped.push({ + source: "local-sqlite", + provider: row.provider, + identity: "(api key)", + reason: "placeholder sentinel '' is not a real key", + }); + continue; + } const identity = credentialIdentity(row.provider, row.credential); if (row.credential.type === "oauth" && flags.includeOauth !== true) { skipped.push({ diff --git a/packages/coding-agent/src/cli/auth-gateway-cli.ts b/packages/coding-agent/src/cli/auth-gateway-cli.ts index 28b369d78..cf800eeca 100644 --- a/packages/coding-agent/src/cli/auth-gateway-cli.ts +++ b/packages/coding-agent/src/cli/auth-gateway-cli.ts @@ -77,6 +77,29 @@ async function writeToken(token: string): Promise { } } +/** + * Atomically create the token file, refusing to clobber an existing one. + * Returns `true` on success, `false` when the file already existed (so the + * caller re-reads it instead of racing another concurrent `ensureToken`). + */ +async function createTokenExclusive(token: string): Promise { + const file = getTokenFilePath(); + await fs.mkdir(path.dirname(file), { recursive: true, mode: 0o700 }); + try { + // `wx` = O_CREAT | O_EXCL — fails with EEXIST if the file is already there. + await fs.writeFile(file, token, { flag: "wx", mode: 0o600 }); + } catch (err) { + if ((err as NodeJS.ErrnoException).code === "EEXIST") return false; + throw err; + } + try { + await fs.chmod(file, 0o600); + } catch { + // Best-effort (e.g. Windows). + } + return true; +} + function generateToken(): string { return crypto.randomBytes(32).toString("base64url"); } @@ -85,6 +108,12 @@ async function ensureToken(): Promise { const existing = await readToken(); if (existing) return existing; const token = generateToken(); + if (await createTokenExclusive(token)) return token; + // Another concurrent invocation won the create race; read what they wrote. + const fromRace = await readToken(); + if (fromRace) return fromRace; + // File existed-then-disappeared between EEXIST and read; last resort, write + // our generated token unconditionally so callers don't see an empty string. await writeToken(token); return token; } diff --git a/packages/coding-agent/src/session/agent-session.ts b/packages/coding-agent/src/session/agent-session.ts index f3d8b01a1..22444b124 100644 --- a/packages/coding-agent/src/session/agent-session.ts +++ b/packages/coding-agent/src/session/agent-session.ts @@ -8149,11 +8149,12 @@ export class AgentSession { }; } - async fetchUsageReports(): Promise { + async fetchUsageReports(signal?: AbortSignal): Promise { const authStorage = this.#modelRegistry.authStorage; if (!authStorage.fetchUsageReports) return null; return authStorage.fetchUsageReports({ baseUrlResolver: provider => this.#modelRegistry.getProviderBaseUrl?.(provider), + signal, }); }