From cb05764eac0761b3eaa6fa0812e5963e9fd5d04f Mon Sep 17 00:00:00 2001 From: can1357 Date: Sun, 17 May 2026 03:05:57 +0200 Subject: [PATCH] feat(ai): added pi-native auth-gateway transport - Added POST /v1/pi/stream endpoint that accepts canonical Context directly, skipping wire-format translation layers. - Added client-side streamPiNative dispatch activated via Model.transport = "pi-native". - Propagated transport field through model registry, provider overrides, and models.yml schema. - Refactored deriveSessionId to accept explicit arguments instead of ParsedFormatRequest. --- packages/ai/src/auth-gateway/server.ts | 156 ++++++++- packages/ai/src/providers/pi-native-client.ts | 228 +++++++++++++ packages/ai/src/providers/pi-native-server.ts | 210 ++++++++++++ packages/ai/src/stream.ts | 13 +- packages/ai/src/types.ts | 15 + .../ai/test/auth-gateway-pi-native.test.ts | 280 ++++++++++++++++ packages/ai/test/pi-native-client.test.ts | 313 ++++++++++++++++++ .../src/config/model-equivalence.ts | 41 ++- .../coding-agent/src/config/model-registry.ts | 29 +- .../src/config/models-config-schema.ts | 8 + packages/coding-agent/src/debug/profiler.ts | 4 + 11 files changed, 1285 insertions(+), 12 deletions(-) create mode 100644 packages/ai/src/providers/pi-native-client.ts create mode 100644 packages/ai/src/providers/pi-native-server.ts create mode 100644 packages/ai/test/auth-gateway-pi-native.test.ts create mode 100644 packages/ai/test/pi-native-client.test.ts diff --git a/packages/ai/src/auth-gateway/server.ts b/packages/ai/src/auth-gateway/server.ts index 79eb5a85f..3b713f7a4 100644 --- a/packages/ai/src/auth-gateway/server.ts +++ b/packages/ai/src/auth-gateway/server.ts @@ -22,8 +22,9 @@ import { Effort } from "../model-thinking"; import * as anthropicMessages from "../providers/anthropic-messages-server"; import * as openaiChat from "../providers/openai-chat-server"; import * as openaiResponses from "../providers/openai-responses-server"; +import * as piNative from "../providers/pi-native-server"; import { streamSimple } from "../stream"; -import type { Api, AssistantMessageEventStream, Model, SimpleStreamOptions } from "../types"; +import type { Api, AssistantMessageEventStream, Context, Model, SimpleStreamOptions } from "../types"; import { parseBind } from "../utils/parse-bind"; import { captureRequestHeaders, corsHeaders, isAuthorized, json, resolvePeer, withCors } from "./http"; import type { @@ -89,8 +90,7 @@ const FORMAT_ROUTES: Record = { * * Anthropic-backed requests ignore `sessionId`; the key is harmless there. */ -function deriveSessionId(parsed: ParsedFormatRequest): string { - const { modelId, context } = parsed; +function deriveSessionId(modelId: string, context: Context): string { const parts: string[] = [modelId]; if (context.systemPrompt && context.systemPrompt.length > 0) { parts.push(context.systemPrompt.join("\n\n")); @@ -144,7 +144,7 @@ function buildStreamOptions(parsed: ParsedFormatRequest, api: Api, signal: Abort // 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); + opts.sessionId = options.promptCacheKey ?? deriveSessionId(parsed.modelId, parsed.context); if (options.thinkingBudgets) { opts.thinkingBudgets = { ...(opts.thinkingBudgets ?? {}), ...options.thinkingBudgets }; } @@ -401,6 +401,148 @@ async function handleFormatEndpoint( }); } +/** + * Pi-native fast path: `POST /v1/pi/stream`. Accepts the canonical pi-ai + * `Context` directly (no wire-format round-trip) and emits a bandwidth-shrunk + * event stream matching `pi-agent`'s `streamProxy`. Skips the OpenAI / + * Anthropic / Responses translation layers — those exist to bridge foreign + * SDKs (llm-git, anthropic-sdk, openai-sdk), and bridging back to pi-native + * just to bridge forward again is wasted work. + * + * Every other gateway concern (bearer auth, model resolve, credential fetch, + * abort mirroring, codex temperature/topP strip, prefix-cache key derivation, + * Claude-Code OAuth shaping inside `streamSimple`) still applies — only + * `parseRequest`/`encodeResponse`/`encodeStream` differ from the format-endpoint + * path. + */ +async function handlePiNative(bootOpts: AuthGatewayBootOptions, req: Request, peer: string): Promise { + const controller = mirrorRequestAbort(req); + const aborted = (): Response => piNative.formatError(499, "request_aborted", "client closed request"); + if (controller.signal.aborted) return aborted(); + + let body: unknown; + try { + body = await req.json(); + } catch (error) { + if (controller.signal.aborted) return aborted(); + return piNative.formatError(400, "invalid_request_error", `Invalid JSON body: ${String(error)}`); + } + if (controller.signal.aborted) return aborted(); + + let parsed: piNative.PiNativeParsedRequest; + try { + parsed = piNative.parseRequest(body, req.headers); + } catch (error) { + if (controller.signal.aborted) return aborted(); + const message = error instanceof Error ? error.message : String(error); + return piNative.formatError(400, "invalid_request_error", message); + } + + const model = bootOpts.resolveModel(parsed.modelId); + if (!model) { + return piNative.formatError(404, "invalid_request_error", `Unknown model: ${parsed.modelId}`); + } + + let apiKey: string | undefined; + try { + apiKey = await bootOpts.storage.getApiKey(model.provider, undefined, { + modelId: model.id, + signal: controller.signal, + }); + } catch (error) { + if (controller.signal.aborted) return aborted(); + const classified = classifyGatewayError(error); + logger.warn("auth-gateway getApiKey threw", { provider: model.provider, peer, error: classified.message }); + return piNative.formatError(classified.status, classified.type, classified.message); + } + if (controller.signal.aborted) return aborted(); + if (!apiKey) { + return piNative.formatError( + 401, + "authentication_error", + `No credential available for provider ${model.provider}`, + ); + } + + // Build the SimpleStreamOptions actually handed to `streamSimple`. We + // trust the client's options (already allow-listed by `parseRequest`) and + // only inject server-controlled fields. The codex temperature/topP strip + // matches `buildStreamOptions` — Codex rejects them with a 400. + const streamOpts: SimpleStreamOptions = { ...parsed.options, apiKey, signal: controller.signal }; + if (model.api === "openai-codex-responses") { + delete streamOpts.temperature; + delete streamOpts.topP; + } + // Merge gateway-captured passthrough headers under the client's own + // headers — the client's values win when they collide. + const captured = captureRequestHeaders(req.headers); + streamOpts.headers = { ...captured, ...(streamOpts.headers ?? {}) }; + // Cache identity: explicit `sessionId` wins, then derive a stable key + // from model + system + tools + first message so Codex prefix caching + // engages on the same logical conversation across turns. + streamOpts.sessionId ??= deriveSessionId(parsed.modelId, parsed.context); + + logger.info("auth-gateway request", { + format: "pi-native", + model: parsed.modelId, + resolvedProvider: model.provider, + resolvedModel: model.id, + stream: parsed.stream, + peer, + }); + + let events: AssistantMessageEventStream; + try { + if (controller.signal.aborted) return aborted(); + events = streamSimple(model, parsed.context, streamOpts); + } catch (error) { + const classified = classifyGatewayError(error); + logger.warn("auth-gateway streamSimple threw", { format: "pi-native", error: classified.message, peer }); + return piNative.formatError(classified.status, classified.type, classified.message); + } + + if (!parsed.stream) { + try { + if (controller.signal.aborted) return aborted(); + const message = await events.result(); + if (message.stopReason === "aborted" || message.stopReason === "error") { + const errorMessage = + message.errorMessage ?? + (message.stopReason === "aborted" ? "Request was aborted" : "Upstream request failed"); + logger.warn("auth-gateway non-streaming failed", { + format: "pi-native", + reason: message.stopReason, + error: errorMessage, + peer, + }); + if (message.stopReason === "aborted") { + return piNative.formatError(499, "request_aborted", errorMessage); + } + const classified = classifyGatewayError(new Error(errorMessage)); + return piNative.formatError(classified.status, classified.type, errorMessage); + } + return json(200, { message }); + } catch (error) { + if (controller.signal.aborted) return aborted(); + const classified = classifyGatewayError(error); + logger.warn("auth-gateway non-streaming aborted", { format: "pi-native", error: classified.message, peer }); + return piNative.formatError(classified.status, classified.type, classified.message); + } + } + if (controller.signal.aborted) return aborted(); + + const sseStream = piNative.encodeStream(events); + return new Response(sseStream, { + status: 200, + headers: { + "Content-Type": "text/event-stream; charset=utf-8", + "Cache-Control": "no-cache", + Connection: "keep-alive", + "X-Accel-Buffering": "no", + }, + }); +} + /** * Snapshot of `GET /v1/usage` — `fetchUsageReports` already caches reports at * a 5-minute per-credential TTL (with jitter, plus last-good fallback on @@ -467,6 +609,12 @@ export function startAuthGateway(opts: AuthGatewayBootOptions): AuthGatewayServe return withCors(await handleFormatEndpoint(formatRoute, opts, req, peer), req); } + // Pi-native fast path. Same auth + provider plumbing as the + // foreign-wire routes, just without the wire-format translation. + if (req.method === "POST" && pathname === "/v1/pi/stream") { + return withCors(await handlePiNative(opts, req, peer), req); + } + // Model catalog. if (req.method === "GET" && pathname === "/v1/models") { return withCors(handleModelsList(opts), req); diff --git a/packages/ai/src/providers/pi-native-client.ts b/packages/ai/src/providers/pi-native-client.ts new file mode 100644 index 000000000..b5df79636 --- /dev/null +++ b/packages/ai/src/providers/pi-native-client.ts @@ -0,0 +1,228 @@ +/** + * Client half of the pi-native auth-gateway protocol. + * + * Dispatches a {@link streamSimple}-shaped request to an `omp auth-gateway` + * via `POST /v1/pi/stream`, reads the SSE event stream back, and pushes the + * parsed events into a local {@link AssistantMessageEventStream} — the same + * stream type every other provider client produces. Callers downstream of + * `streamSimple` cannot tell whether the events came from a real provider + * SDK or from a gateway hop; they consume `AssistantMessageEvent`s either + * way. + * + * Activated when a {@link Model} has `transport: "pi-native"` set; the + * dispatch hook lives in `streamSimple()` (see `../stream.ts`). Used by + * containerized omp deployments (robomp slots, the swarm extension) that + * route every LLM call through a credential-holding sidecar so the slot + * itself stays credential-free. + */ +import { readSseJson } from "@oh-my-pi/pi-utils"; +import type { + Api, + AssistantMessage, + AssistantMessageEvent, + AssistantMessageEventStream as AssistantMessageEventStreamType, + Context, + Model, + SimpleStreamOptions, +} from "../types"; +import { AssistantMessageEventStream } from "../utils/event-stream"; + +/** + * Fields that must not cross the wire — either non-serializable (functions, + * `AbortSignal`, the provider-session `Map`) or server-controlled + * (`apiKey`, which the gateway injects from its own credential store; the + * client's `apiKey` is the gateway *bearer*, sent in the `Authorization` + * header rather than the request body). + */ +const NON_WIRE_KEYS = new Set([ + "signal", + "apiKey", + "fetch", + "onPayload", + "onResponse", + "onSseEvent", + "execHandlers", + "cursorExecHandlers", + "cursorOnToolResult", + "providerSessionState", +]); + +function buildWireOptions(options: SimpleStreamOptions | undefined): Record { + if (!options) return {}; + const wire: Record = {}; + for (const [k, v] of Object.entries(options)) { + if (v === undefined) continue; + if (NON_WIRE_KEYS.has(k as keyof SimpleStreamOptions)) continue; + wire[k] = v; + } + return wire; +} + +async function decodeGatewayError(response: Response): Promise { + const status = response.status; + let body: unknown; + try { + body = await response.json(); + } catch { + body = await response.text().catch(() => ""); + } + if (typeof body === "object" && body !== null && "error" in body) { + const err = (body as { error: unknown }).error; + if (typeof err === "object" && err !== null) { + const message = (err as { message?: unknown }).message; + const type = (err as { type?: unknown }).type; + const out = new Error(typeof message === "string" ? message : `auth-gateway ${status}`); + (out as { status?: number; type?: string }).status = status; + if (typeof type === "string") (out as { type?: string }).type = type; + return out; + } + } + const text = typeof body === "string" ? body : JSON.stringify(body); + const err = new Error(`auth-gateway ${status}: ${text || response.statusText}`); + (err as { status?: number }).status = status; + return err; +} + +/** + * Resolve the `/v1/pi/stream` endpoint URL from the model's `baseUrl`. + * Trims a trailing slash so concatenation can't double-slash; throws when + * the baseUrl is missing (transport=pi-native without a gateway target is + * a configuration error, not a runtime recoverable one). + */ +function resolveStreamUrl(model: Model): string { + if (!model.baseUrl) { + throw new Error( + `pi-native transport requires \`baseUrl\` on model ${model.id} (set it on the provider config in models.yml)`, + ); + } + return `${model.baseUrl.replace(/\/+$/, "")}/v1/pi/stream`; +} + +function buildHeaders(model: Model, apiKey: string | undefined): Record { + const headers: Record = { + "Content-Type": "application/json", + Accept: "text/event-stream", + ...(model.headers ?? {}), + }; + if (apiKey && !headers.Authorization) { + headers.Authorization = `Bearer ${apiKey}`; + } + return headers; +} + +/** + * Stream a turn through an `omp auth-gateway` over the pi-native protocol. + * + * The returned {@link AssistantMessageEventStream} receives each parsed + * `AssistantMessageEvent` verbatim from the gateway; the terminal `done` / + * `error` event resolves `.result()` automatically via the base class's + * completion check. Non-streaming consumers just call `.result()` and pay + * for SSE framing they don't use — that overhead is dominated by provider + * latency, so we always stream rather than maintaining a parallel + * non-streaming path. + */ +export function streamPiNative( + model: Model, + context: Context, + options?: SimpleStreamOptions, +): AssistantMessageEventStreamType { + const stream = new AssistantMessageEventStream(); + + void (async () => { + const signal = options?.signal; + // Abort propagation: cancel the response body when the caller's signal + // fires. Mirror `streamProxy`'s shape — explicit listener + finally + // cleanup — so we don't leak listeners on the long-running case. + let response: Response | null = null; + const onAbort = (): void => { + const body = response?.body; + if (body) body.cancel("Request aborted by caller").catch(() => {}); + }; + if (signal) { + if (signal.aborted) { + stream.fail(signal.reason instanceof Error ? signal.reason : new Error(String(signal.reason ?? "aborted"))); + return; + } + signal.addEventListener("abort", onAbort, { once: true }); + } + + try { + const url = resolveStreamUrl(model as Model); + const fetchImpl = options?.fetch ?? globalThis.fetch; + const headers = buildHeaders(model as Model, options?.apiKey); + const body = JSON.stringify({ + modelId: model.id, + context, + options: buildWireOptions(options), + stream: true, + }); + + response = await fetchImpl(url, { method: "POST", headers, body, signal }); + if (!response.ok) { + stream.fail(await decodeGatewayError(response)); + return; + } + if (!response.body) { + stream.fail(new Error("auth-gateway returned empty body")); + return; + } + + let sawTerminal = false; + for await (const event of readSseJson( + response.body as ReadableStream, + signal, + )) { + if (event.type === "done" || event.type === "error") sawTerminal = true; + stream.push(event); + // `stream.push` resolves `.result()` on `done`/`error`; subsequent + // pushes are silently dropped by the base class. We still iterate + // to drain any trailing bytes from the wire so the underlying TCP + // stream closes cleanly. + } + + if (!sawTerminal) { + // SSE closed before a terminal event reached us — synthesize one + // so awaiters of `.result()` resolve instead of hanging forever. + // Matches the gateway's own defensive fallback in + // `pi-native-server.encodeStream`. + const aborted = signal?.aborted === true; + const partial = makeSyntheticAssistant(model as Model); + if (aborted) { + partial.stopReason = "aborted"; + partial.errorMessage = "stream closed without terminal event"; + stream.push({ type: "error", reason: "aborted", error: partial }); + } else { + partial.stopReason = "stop"; + stream.push({ type: "done", reason: "stop", message: partial }); + } + } + stream.end(); + } catch (err) { + stream.fail(err); + } finally { + if (signal) signal.removeEventListener("abort", onAbort); + } + })(); + + return stream; +} + +function makeSyntheticAssistant(model: Model): AssistantMessage { + return { + role: "assistant", + content: [], + api: model.api, + provider: model.provider, + model: model.id, + usage: { + input: 0, + output: 0, + cacheRead: 0, + cacheWrite: 0, + totalTokens: 0, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, + }, + stopReason: "stop", + timestamp: Date.now(), + }; +} diff --git a/packages/ai/src/providers/pi-native-server.ts b/packages/ai/src/providers/pi-native-server.ts new file mode 100644 index 000000000..2be0e9f7e --- /dev/null +++ b/packages/ai/src/providers/pi-native-server.ts @@ -0,0 +1,210 @@ +/** + * Pi-native wire format for the auth-gateway. + * + * Where the OpenAI / Anthropic / Responses route modules translate foreign + * wire shapes through pi-ai's canonical {@link Context}, this module accepts + * the canonical shape *directly* — for clients that already speak pi-ai + * (containerized omp, the swarm extension, robomp's sidecar auth-gateway). + * Skipping the wire-format → Context → wire-format round-trip cuts + * per-request CPU but, more importantly, avoids the quantization that those + * translations impose on first-class pi-ai fields (service tier, cache + * markers, thinking budgets, tool-choice variants, …). + * + * The streaming wire is {@link AssistantMessageEvent} serialized verbatim and + * SSE-framed. Same type pi-ai already produces internally; the client feeds + * each parsed event straight into `AssistantMessageEventStream.push()` with + * no translation. Including `partial: AssistantMessage` on every delta is + * O(N²) in turn length on the wire — acceptable for the loopback / sidecar + * topology this transport is designed for; provider latency dominates the + * actual cost. + * + * Endpoint contract: + * POST /v1/pi/stream + * body: { modelId, context, options?, stream? } // `stream` defaults to true + * 200 SSE: stream of `AssistantMessageEvent` (terminated by `data: [DONE]`) + * 200 JSON (stream=false): { message: AssistantMessage } + * 4xx/5xx: { error: { type, message } } + */ +import type { AssistantMessageEventStream, Context, SimpleStreamOptions } from "../types"; + +export interface PiNativeParsedRequest { + modelId: string; + context: Context; + options: SimpleStreamOptions; + stream: boolean; +} +/** + * Subset of {@link SimpleStreamOptions} accepted from the wire. Function-valued + * fields (`fetch`, `onPayload`, `onResponse`, `onSseEvent`, exec handlers, the + * provider-session map) and gateway-owned controls (`apiKey`, `signal`) are + * intentionally absent — those are server-side concerns. Anything outside this + * allow-list is dropped silently rather than 400ing, so clients can forward + * `SimpleStreamOptions` from older / newer omp builds without per-version + * conditionals. + */ +const ALLOWED_OPTION_KEYS: ReadonlySet = new Set([ + "temperature", + "topP", + "topK", + "minP", + "presencePenalty", + "frequencyPenalty", + "repetitionPenalty", + "stopSequences", + "maxTokens", + "cacheRetention", + "headers", + "initiatorOverride", + "maxRetryDelayMs", + "metadata", + "sessionId", + "streamFirstEventTimeoutMs", + "streamIdleTimeoutMs", + "reasoning", + "disableReasoning", + "hideThinkingSummary", + "thinkingBudgets", + "toolChoice", + "serviceTier", + "kimiApiFormat", + "syntheticApiFormat", + "preferWebsockets", +] as const satisfies readonly (keyof SimpleStreamOptions)[]); + +// --------------------------------------------------------------------------- +// parseRequest +// --------------------------------------------------------------------------- + +/** + * Parse a pi-native request body. Validation is intentionally minimal — only + * the shape the gateway itself reads is checked (`modelId`, `context.messages` + * array, options is an object). Everything downstream is the canonical pi-ai + * type surface; mis-shaped values surface as a `502 upstream_error` from + * `streamSimple` rather than being re-validated here. + * + * Accepts both `{ modelId: string }` and `{ model: { id: string } }` so the + * existing `streamProxy` client (which sends the full Model object) can target + * the gateway with only a URL swap. + */ +export function parseRequest(body: unknown, _headers?: Headers): PiNativeParsedRequest { + if (typeof body !== "object" || body === null || Array.isArray(body)) { + throw new Error("Request body must be a JSON object"); + } + const obj = body as Record; + + let modelId: string | undefined; + if (typeof obj.modelId === "string" && obj.modelId.length > 0) { + modelId = obj.modelId; + } else if (typeof obj.model === "string" && obj.model.length > 0) { + modelId = obj.model; + } else if (typeof obj.model === "object" && obj.model !== null) { + const m = obj.model as Record; + if (typeof m.id === "string" && m.id.length > 0) modelId = m.id; + } + if (!modelId) throw new Error("Missing `modelId` (or `model.id`) field"); + + const context = obj.context; + if (typeof context !== "object" || context === null || Array.isArray(context)) { + throw new Error("Missing `context` object"); + } + const ctxObj = context as Record; + if (!Array.isArray(ctxObj.messages)) { + throw new Error("`context.messages` must be an array"); + } + if (ctxObj.systemPrompt !== undefined && !Array.isArray(ctxObj.systemPrompt)) { + throw new Error("`context.systemPrompt` must be an array of strings when present"); + } + if (ctxObj.tools !== undefined && !Array.isArray(ctxObj.tools)) { + throw new Error("`context.tools` must be an array when present"); + } + + const options: SimpleStreamOptions = {}; + const rawOpts = obj.options; + if (typeof rawOpts === "object" && rawOpts !== null && !Array.isArray(rawOpts)) { + const optsBag = options as Record; + for (const [k, v] of Object.entries(rawOpts)) { + if (v === undefined || v === null) continue; + if (!ALLOWED_OPTION_KEYS.has(k as keyof SimpleStreamOptions)) continue; + optsBag[k] = v; + } + } + + // `stream` defaults to true — pi-native clients overwhelmingly stream, and + // matching `streamProxy`'s implicit-stream behavior avoids a one-flag papercut. + const stream = typeof obj.stream === "boolean" ? obj.stream : true; + + return { + modelId, + context: context as Context, + options, + stream, + }; +} +// --------------------------------------------------------------------------- +// encodeStream (SSE) +// --------------------------------------------------------------------------- + +const SSE_ENCODER = new TextEncoder(); +const SSE_DONE = SSE_ENCODER.encode("data: [DONE]\n\n"); + +/** + * Ship every {@link AssistantMessageEvent} verbatim, SSE-framed. + * + * No per-event re-shaping: the pi-native client is pi-ai itself, so the + * canonical event type IS the wire type. Including the rolling + * `partial: AssistantMessage` on every delta is quadratic in turn length + * on the wire, but for the loopback / sidecar topology this transport + * targets (containerized omp → host gateway, robomp slot → omp-auth-gateway + * sidecar) the bandwidth cost is negligible compared to provider latency — + * and the client gets to feed the events straight into its existing + * `AssistantMessageEventStream.push()` plumbing with zero translation. + */ +export function encodeStream(events: AssistantMessageEventStream): ReadableStream { + return new ReadableStream({ + async start(controller) { + try { + for await (const event of events) { + controller.enqueue(SSE_ENCODER.encode(`data: ${JSON.stringify(event)}\n\n`)); + if (event.type === "done" || event.type === "error") break; + } + controller.enqueue(SSE_DONE); + controller.close(); + } catch (err) { + // Best-effort error envelope so the client iterator resolves + // instead of hanging on the dropped connection. Shape matches the + // canonical `error` event minus the unrecoverable `error: + // AssistantMessage` payload (we don't have a usable one here). + const message = err instanceof Error ? err.message : String(err); + controller.enqueue( + SSE_ENCODER.encode( + `data: ${JSON.stringify({ type: "error", reason: "error", errorMessage: message })}\n\n`, + ), + ); + controller.enqueue(SSE_DONE); + controller.close(); + } + }, + }); +} + +// --------------------------------------------------------------------------- +// formatError +// --------------------------------------------------------------------------- + +/** + * Pi-native error envelope: + * `{ error: { type, message } }` + * + * Mirrors OpenAI's outer shape (which clients/SDKs already parse) without the + * provider-specific status taxonomy — pi-native callers consume `type` + * directly. + */ +export function formatError(status: number, type: string, message: string): Response { + return new Response(JSON.stringify({ error: { type, message } }), { + status, + headers: { + "Content-Type": "application/json; charset=utf-8", + "Cache-Control": "no-store", + }, + }); +} diff --git a/packages/ai/src/stream.ts b/packages/ai/src/stream.ts index ced7ee158..566a91282 100644 --- a/packages/ai/src/stream.ts +++ b/packages/ai/src/stream.ts @@ -40,6 +40,7 @@ import { streamOpenAICompletions, streamOpenAIResponses, } from "./providers/register-builtins"; +import { streamPiNative } from "./providers/pi-native-client"; import { isSyntheticModel, streamSynthetic } from "./providers/synthetic"; import type { Api, @@ -278,7 +279,17 @@ export function streamSimple( context: Context, options?: SimpleStreamOptions, ): AssistantMessageEventStream { - // Check custom API registry first (extension-provided APIs) + // Pi-native transport short-circuits the per-provider dispatch entirely: + // the gateway resolves provider + credential server-side, so we don't + // need an `apiKey` from `getEnvApiKey` here — `options.apiKey` carries + // the gateway bearer instead. Comes BEFORE the custom-API check so + // extension-registered APIs can't accidentally override a configured + // pi-native transport. + if (model.transport === "pi-native") { + return streamPiNative(model, context, options); + } + + // Check custom API registry (extension-provided APIs) const customApiProvider = getCustomApi(model.api); if (customApiProvider) { return customApiProvider.streamSimple(model, context, options); diff --git a/packages/ai/src/types.ts b/packages/ai/src/types.ts index 5aa9785ac..096b221bb 100644 --- a/packages/ai/src/types.ts +++ b/packages/ai/src/types.ts @@ -767,6 +767,21 @@ export interface Model { contextWindow: number; maxTokens: number; headers?: Record; + /** + * Streaming transport override. When `"pi-native"`, `streamSimple` routes + * the request to the model's `baseUrl` via the auth-gateway's + * `POST /v1/pi/stream` endpoint instead of dispatching the per-API + * provider client. The `baseUrl` must point at an `omp auth-gateway` + * (or compatible) host; `headers.Authorization` (or `apiKey` resolved by + * the registry) carries the gateway bearer. + * + * Used by containerized omp installs (e.g. robomp slots) to route every + * LLM call through a sidecar gateway that holds the real provider + * credentials. The model's other metadata (pricing, context window, + * thinking config, …) still resolves locally; only the streaming + * dispatch is redirected. + */ + transport?: "pi-native"; /** Hint that websocket transport should be preferred when supported by the provider implementation. */ preferWebsockets?: boolean; /** Preferred model to switch to when context promotion is triggered (model id or provider/id). */ diff --git a/packages/ai/test/auth-gateway-pi-native.test.ts b/packages/ai/test/auth-gateway-pi-native.test.ts new file mode 100644 index 000000000..7c10a65a7 --- /dev/null +++ b/packages/ai/test/auth-gateway-pi-native.test.ts @@ -0,0 +1,280 @@ +import { describe, expect, it } from "bun:test"; +import { Effort } from "../src/model-thinking"; +import { encodeStream, formatError, parseRequest } from "../src/providers/pi-native-server"; +import type { + AssistantMessage, + AssistantMessageEvent, + AssistantMessageEventStream, + Context, + Usage, +} from "../src/types"; + +function makeEventStream(events: AssistantMessageEvent[], final: AssistantMessage): AssistantMessageEventStream { + async function* iter() { + for (const e of events) yield e; + } + const stream = iter() as unknown as AssistantMessageEventStream; + (stream as { result(): Promise }).result = async () => final; + return stream; +} + +async function collectSse(stream: ReadableStream): Promise { + const reader = stream.getReader(); + const decoder = new TextDecoder(); + let buf = ""; + for (;;) { + const { value, done } = await reader.read(); + if (done) break; + buf += decoder.decode(value, { stream: true }); + } + buf += decoder.decode(); + return buf.split("\n\n").filter(s => s.length > 0); +} + +function parseSseLine(line: string): unknown { + const stripped = line.replace(/^data: /, ""); + if (stripped === "[DONE]") return "[DONE]"; + return JSON.parse(stripped); +} + +const ZERO_USAGE: Usage = { + input: 0, + output: 0, + cacheRead: 0, + cacheWrite: 0, + totalTokens: 0, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, +}; + +function baseAssistant(overrides?: Partial): AssistantMessage { + return { + role: "assistant", + content: [], + api: "anthropic-messages", + provider: "anthropic", + model: "claude-sonnet-4-5", + usage: ZERO_USAGE, + stopReason: "stop", + timestamp: 0, + ...overrides, + }; +} + +const baseContext: Context = { + systemPrompt: ["you are helpful"], + messages: [{ role: "user", content: "hi", timestamp: 0 }], +}; + +describe("pi-native parseRequest", () => { + it("accepts modelId + context and returns canonical shape", () => { + const parsed = parseRequest({ + modelId: "claude-sonnet-4-5", + context: baseContext, + options: { temperature: 0.5, reasoning: Effort.High }, + stream: false, + }); + expect(parsed.modelId).toBe("claude-sonnet-4-5"); + expect(parsed.context).toEqual(baseContext); + expect(parsed.options.temperature).toBe(0.5); + expect(parsed.options.reasoning).toBe(Effort.High); + expect(parsed.stream).toBe(false); + }); + + it("falls back to model.id when modelId is absent (streamProxy compat)", () => { + const parsed = parseRequest({ + model: { id: "claude-opus-4-1", provider: "anthropic", api: "anthropic-messages" }, + context: baseContext, + }); + expect(parsed.modelId).toBe("claude-opus-4-1"); + }); + + it("accepts top-level string `model` as the id (extra compat)", () => { + const parsed = parseRequest({ + model: "gpt-5", + context: baseContext, + }); + expect(parsed.modelId).toBe("gpt-5"); + }); + + it("defaults stream to true when omitted", () => { + const parsed = parseRequest({ modelId: "x", context: baseContext }); + expect(parsed.stream).toBe(true); + }); + + it("drops server-controlled and unknown option keys", () => { + const parsed = parseRequest({ + modelId: "x", + context: baseContext, + options: { + temperature: 0.2, + apiKey: "should-be-stripped", + signal: {}, + fetch: () => {}, + onPayload: () => {}, + onResponse: () => {}, + onSseEvent: () => {}, + execHandlers: {}, + providerSessionState: new Map(), + notARealField: "ignored", + }, + }); + expect(parsed.options).toEqual({ temperature: 0.2 }); + expect("apiKey" in parsed.options).toBe(false); + expect("signal" in parsed.options).toBe(false); + expect("fetch" in parsed.options).toBe(false); + expect("onPayload" in parsed.options).toBe(false); + expect("notARealField" in parsed.options).toBe(false); + }); + + it("preserves headers, metadata, sessionId, thinkingBudgets", () => { + const parsed = parseRequest({ + modelId: "x", + context: baseContext, + options: { + headers: { "x-foo": "bar" }, + metadata: { user_id: "u" }, + sessionId: "explicit-session", + thinkingBudgets: { high: 8192 }, + stopSequences: ["\n\n"], + toolChoice: "required", + serviceTier: "priority", + cacheRetention: "long", + }, + }); + expect(parsed.options.headers).toEqual({ "x-foo": "bar" }); + expect(parsed.options.metadata).toEqual({ user_id: "u" }); + expect(parsed.options.sessionId).toBe("explicit-session"); + expect(parsed.options.thinkingBudgets).toEqual({ high: 8192 }); + expect(parsed.options.stopSequences).toEqual(["\n\n"]); + expect(parsed.options.toolChoice).toBe("required"); + expect(parsed.options.serviceTier).toBe("priority"); + expect(parsed.options.cacheRetention).toBe("long"); + }); + + it("rejects missing required fields", () => { + expect(() => parseRequest({ context: baseContext })).toThrow(/modelId/); + expect(() => parseRequest({ modelId: "x" })).toThrow(/context/); + expect(() => parseRequest({ modelId: "x", context: { systemPrompt: [] } })).toThrow(/messages/); + }); + + it("rejects non-object body", () => { + expect(() => parseRequest(null)).toThrow(); + expect(() => parseRequest("hello")).toThrow(); + expect(() => parseRequest([])).toThrow(); + }); + + it("validates systemPrompt and tools shape", () => { + expect(() => parseRequest({ modelId: "x", context: { systemPrompt: "not array", messages: [] } })).toThrow( + /systemPrompt/, + ); + expect(() => parseRequest({ modelId: "x", context: { messages: [], tools: "not array" } })).toThrow(/tools/); + }); + + it("skips null and undefined option values", () => { + const parsed = parseRequest({ + modelId: "x", + context: baseContext, + options: { temperature: null, topP: undefined, maxTokens: 100 }, + }); + expect("temperature" in parsed.options).toBe(false); + expect("topP" in parsed.options).toBe(false); + expect(parsed.options.maxTokens).toBe(100); + }); +}); +describe("pi-native encodeStream", () => { + it("ships every AssistantMessageEvent verbatim, terminated by [DONE]", async () => { + // Pi-native is omp-talks-to-omp: the client feeds parsed events directly + // into `AssistantMessageEventStream.push()`, so the wire IS the canonical + // event type. No partial-stripping, no per-event re-shaping. + const finalMessage = baseAssistant({ + content: [{ type: "text", text: "hi" }], + usage: { ...ZERO_USAGE, input: 4, output: 2, totalTokens: 6 }, + }); + const partialAfterDelta: AssistantMessage = baseAssistant({ + content: [{ type: "text", text: "hi" }], + }); + const events: AssistantMessageEvent[] = [ + { type: "start", partial: baseAssistant() }, + { type: "text_start", contentIndex: 0, partial: baseAssistant({ content: [{ type: "text", text: "" }] }) }, + { type: "text_delta", contentIndex: 0, delta: "hi", partial: partialAfterDelta }, + { type: "text_end", contentIndex: 0, content: "hi", partial: partialAfterDelta }, + { type: "done", reason: "stop", message: finalMessage }, + ]; + const chunks = await collectSse(encodeStream(makeEventStream(events, finalMessage))); + const parsed = chunks.map(parseSseLine); + + // Every payload is the input event verbatim — partials, signatures, + // usage all intact. Terminator follows `done`/`error`. + expect(parsed.length).toBe(events.length + 1); + for (let i = 0; i < events.length; i++) { + expect(parsed[i]).toEqual(JSON.parse(JSON.stringify(events[i]))); + } + expect(parsed[parsed.length - 1]).toBe("[DONE]"); + }); + + it("preserves the rolling `partial` on every delta (sanity: no shrink)", async () => { + // Guards against an accidental re-introduction of partial-stripping + // optimization. Clients depend on `partial` being present. + const final = baseAssistant({ content: [{ type: "text", text: "abc" }] }); + const events: AssistantMessageEvent[] = [ + { type: "text_delta", contentIndex: 0, delta: "abc", partial: final }, + { type: "done", reason: "stop", message: final }, + ]; + const parsed = (await collectSse(encodeStream(makeEventStream(events, final)))).map(parseSseLine) as Array< + Record + >; + expect(parsed[0]).toHaveProperty("partial"); + expect((parsed[0] as { partial: AssistantMessage }).partial.content).toEqual([{ type: "text", text: "abc" }]); + }); + + it("stops streaming after a terminal `done` and emits [DONE] once", async () => { + const final = baseAssistant(); + const events: AssistantMessageEvent[] = [ + { type: "done", reason: "stop", message: final }, + // This trailing event must NOT reach the wire — terminal events end + // the stream so the client iterator resolves cleanly. + { type: "text_delta", contentIndex: 0, delta: "ghost", partial: final }, + ]; + const parsed = (await collectSse(encodeStream(makeEventStream(events, final)))).map(parseSseLine); + expect(parsed.length).toBe(2); + expect((parsed[0] as { type: string }).type).toBe("done"); + expect(parsed[1]).toBe("[DONE]"); + }); + + it("forwards `error` events verbatim, then closes with [DONE]", async () => { + const errored = baseAssistant({ + stopReason: "error", + errorMessage: "upstream blew up", + usage: { ...ZERO_USAGE, input: 3 }, + }); + const events: AssistantMessageEvent[] = [{ type: "error", reason: "error", error: errored }]; + const parsed = (await collectSse(encodeStream(makeEventStream(events, errored)))).map(parseSseLine); + expect(parsed[0]).toEqual({ type: "error", reason: "error", error: JSON.parse(JSON.stringify(errored)) }); + expect(parsed[1]).toBe("[DONE]"); + }); + + it("emits a synthetic error envelope when the source iterator throws", async () => { + // Source-stream failures (network drop after `streamSimple` returned) + // must not hang the client. We surface a minimal `error` event followed + // by `[DONE]` so the iterator on the other end resolves. + const broken = (async function* () { + yield { type: "start", partial: baseAssistant() } satisfies AssistantMessageEvent; + throw new Error("connection reset"); + })() as unknown as AssistantMessageEventStream; + (broken as { result(): Promise }).result = async () => baseAssistant(); + + const parsed = (await collectSse(encodeStream(broken))).map(parseSseLine); + expect((parsed[0] as { type: string }).type).toBe("start"); + expect(parsed[1]).toEqual({ type: "error", reason: "error", errorMessage: "connection reset" }); + expect(parsed[2]).toBe("[DONE]"); + }); +}); + +describe("pi-native formatError", () => { + it("emits { error: { type, message } } with the given status", async () => { + const res = formatError(401, "authentication_error", "no credential"); + expect(res.status).toBe(401); + expect(res.headers.get("Content-Type")).toBe("application/json; charset=utf-8"); + expect(await res.json()).toEqual({ error: { type: "authentication_error", message: "no credential" } }); + }); +}); diff --git a/packages/ai/test/pi-native-client.test.ts b/packages/ai/test/pi-native-client.test.ts new file mode 100644 index 000000000..c4b61477f --- /dev/null +++ b/packages/ai/test/pi-native-client.test.ts @@ -0,0 +1,313 @@ +import { afterEach, describe, expect, it, mock, spyOn } from "bun:test"; +import { streamPiNative } from "../src/providers/pi-native-client"; +import type { AssistantMessage, AssistantMessageEvent, Context, FetchImpl, Model } from "../src/types"; + +function sseBytes(events: AssistantMessageEvent[]): Uint8Array { + const encoder = new TextEncoder(); + const parts: Uint8Array[] = []; + for (const event of events) { + parts.push(encoder.encode(`data: ${JSON.stringify(event)}\n\n`)); + } + parts.push(encoder.encode("data: [DONE]\n\n")); + const total = parts.reduce((n, p) => n + p.byteLength, 0); + const out = new Uint8Array(total); + let offset = 0; + for (const part of parts) { + out.set(part, offset); + offset += part.byteLength; + } + return out; +} + +function fakeBody(bytes: Uint8Array): ReadableStream { + return new ReadableStream({ + start(controller) { + controller.enqueue(bytes); + controller.close(); + }, + }); +} + +function fakeResponse(events: AssistantMessageEvent[], init: ResponseInit = {}): Response { + return new Response(fakeBody(sseBytes(events)), { + status: 200, + headers: { "Content-Type": "text/event-stream" }, + ...init, + }); +} + +function baseAssistant(overrides: Partial = {}): AssistantMessage { + return { + role: "assistant", + content: [], + api: "anthropic-messages", + provider: "anthropic", + model: "claude-sonnet-4-5", + usage: { + input: 0, + output: 0, + cacheRead: 0, + cacheWrite: 0, + totalTokens: 0, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, + }, + stopReason: "stop", + timestamp: 0, + ...overrides, + }; +} + +function fakeModel(overrides: Partial> = {}): Model<"anthropic-messages"> { + return { + id: "claude-sonnet-4-5", + name: "Claude Sonnet 4.5", + api: "anthropic-messages", + provider: "anthropic", + baseUrl: "http://llm-gateway.internal:4000", + reasoning: true, + input: ["text"], + cost: { input: 3, output: 15, cacheRead: 0.3, cacheWrite: 3.75 }, + contextWindow: 200000, + maxTokens: 64000, + transport: "pi-native", + ...overrides, + }; +} + +const baseContext: Context = { + systemPrompt: ["you are helpful"], + messages: [{ role: "user", content: "hi", timestamp: 0 }], +}; + +async function collectEvents( + stream: AsyncIterable, +): Promise { + const out: AssistantMessageEvent[] = []; + for await (const event of stream) out.push(event); + return out; +} + +afterEach(() => { + mock.restore(); +}); + +describe("streamPiNative request shape", () => { + it("POSTs `{modelId, context, options, stream:true}` to `${baseUrl}/v1/pi/stream`", async () => { + const final = baseAssistant(); + const captured: { url?: string; init?: RequestInit } = {}; + const fetchImpl: FetchImpl = (async (input, init) => { + captured.url = typeof input === "string" ? input : input.toString(); + captured.init = init; + return fakeResponse([{ type: "done", reason: "stop", message: final }]); + }) as FetchImpl; + + const stream = streamPiNative(fakeModel(), baseContext, { + apiKey: "gw-bearer", + fetch: fetchImpl, + temperature: 0.7, + }); + await stream.result(); + + expect(captured.url).toBe("http://llm-gateway.internal:4000/v1/pi/stream"); + expect(captured.init?.method).toBe("POST"); + const headers = captured.init?.headers as Record; + expect(headers["Content-Type"]).toBe("application/json"); + expect(headers.Accept).toBe("text/event-stream"); + expect(headers.Authorization).toBe("Bearer gw-bearer"); + + const body = JSON.parse(captured.init?.body as string); + expect(body.modelId).toBe("claude-sonnet-4-5"); + expect(body.context).toEqual(baseContext); + expect(body.stream).toBe(true); + expect(body.options.temperature).toBe(0.7); + }); + + it("strips non-wire fields (signal, apiKey, fetch, callbacks) from `options`", async () => { + // `apiKey` must ride in the Authorization header, never the body — sending + // it twice would let a logged request leak the gateway bearer. The other + // fields are non-serializable function/runtime handles. + const captured: { init?: RequestInit } = {}; + const fetchImpl: FetchImpl = (async (_input, init) => { + captured.init = init; + return fakeResponse([{ type: "done", reason: "stop", message: baseAssistant() }]); + }) as FetchImpl; + + const controller = new AbortController(); + const stream = streamPiNative(fakeModel(), baseContext, { + apiKey: "gw-bearer", + fetch: fetchImpl, + signal: controller.signal, + onPayload: () => undefined, + onResponse: () => undefined, + onSseEvent: () => undefined, + providerSessionState: new Map(), + maxTokens: 1024, + }); + await stream.result(); + + const body = JSON.parse(captured.init?.body as string); + expect("apiKey" in body.options).toBe(false); + expect("signal" in body.options).toBe(false); + expect("fetch" in body.options).toBe(false); + expect("onPayload" in body.options).toBe(false); + expect("onResponse" in body.options).toBe(false); + expect("onSseEvent" in body.options).toBe(false); + expect("providerSessionState" in body.options).toBe(false); + // And the legitimate options survive + expect(body.options.maxTokens).toBe(1024); + }); + + it("normalizes trailing slashes on `baseUrl` so the endpoint never double-slashes", async () => { + const captured: { url?: string } = {}; + const fetchImpl: FetchImpl = (async (input, _init) => { + captured.url = typeof input === "string" ? input : input.toString(); + return fakeResponse([{ type: "done", reason: "stop", message: baseAssistant() }]); + }) as FetchImpl; + + await streamPiNative( + fakeModel({ baseUrl: "http://llm-gateway.internal:4000///" }), + baseContext, + { apiKey: "k", fetch: fetchImpl }, + ).result(); + expect(captured.url).toBe("http://llm-gateway.internal:4000/v1/pi/stream"); + }); + + it("forwards `model.headers` and lets a caller-supplied Authorization win", async () => { + const captured: { init?: RequestInit } = {}; + const fetchImpl: FetchImpl = (async (_input, init) => { + captured.init = init; + return fakeResponse([{ type: "done", reason: "stop", message: baseAssistant() }]); + }) as FetchImpl; + + await streamPiNative( + fakeModel({ headers: { "x-omp-slot": "robomp-1", Authorization: "Bearer model-wins" } }), + baseContext, + { apiKey: "options-loses", fetch: fetchImpl }, + ).result(); + + const headers = captured.init?.headers as Record; + expect(headers["x-omp-slot"]).toBe("robomp-1"); + expect(headers.Authorization).toBe("Bearer model-wins"); + }); + + it("throws synchronously when `baseUrl` is missing", async () => { + const broken = fakeModel({ baseUrl: "" as unknown as string }); + // The promise the iterator awaits surfaces the error via `.result()`. + const stream = streamPiNative(broken, baseContext, { apiKey: "k" }); + await expect(stream.result()).rejects.toThrow(/baseUrl/); + }); +}); + +describe("streamPiNative event flow", () => { + it("pushes parsed events verbatim and resolves `.result()` on terminal `done`", async () => { + const final = baseAssistant({ + content: [{ type: "text", text: "hi" }], + usage: { + input: 4, + output: 2, + cacheRead: 0, + cacheWrite: 0, + totalTokens: 6, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, + }, + }); + const partial = baseAssistant({ content: [{ type: "text", text: "hi" }] }); + const events: AssistantMessageEvent[] = [ + { type: "start", partial: baseAssistant() }, + { type: "text_delta", contentIndex: 0, delta: "hi", partial }, + { type: "done", reason: "stop", message: final }, + ]; + const fetchImpl: FetchImpl = (async () => fakeResponse(events)) as FetchImpl; + + const stream = streamPiNative(fakeModel(), baseContext, { apiKey: "k", fetch: fetchImpl }); + const seen = await collectEvents(stream); + const result = await stream.result(); + + expect(seen).toEqual(events); + expect(result).toEqual(final); + }); + + it("classifies non-2xx responses into Errors with status + type tags", async () => { + const fetchImpl: FetchImpl = (async () => + new Response(JSON.stringify({ error: { type: "authentication_error", message: "no credential" } }), { + status: 401, + headers: { "Content-Type": "application/json" }, + })) as FetchImpl; + + const stream = streamPiNative(fakeModel(), baseContext, { apiKey: "k", fetch: fetchImpl }); + await expect(stream.result()).rejects.toThrow(/no credential/); + }); + + it("falls back to plain text on a non-JSON error body", async () => { + const fetchImpl: FetchImpl = (async () => + new Response("bad gateway", { status: 502 })) as FetchImpl; + const stream = streamPiNative(fakeModel(), baseContext, { apiKey: "k", fetch: fetchImpl }); + await expect(stream.result()).rejects.toThrow(/502/); + }); + + it("synthesizes a terminal `done` when the SSE stream closes silently", async () => { + // Models the gateway dropping mid-stream — without this synthetic terminator, + // `.result()` would hang forever. + const halfEvents: AssistantMessageEvent[] = [{ type: "start", partial: baseAssistant() }]; + const encoder = new TextEncoder(); + const body = new ReadableStream({ + start(controller) { + for (const e of halfEvents) controller.enqueue(encoder.encode(`data: ${JSON.stringify(e)}\n\n`)); + controller.close(); + }, + }); + const fetchImpl: FetchImpl = (async () => + new Response(body, { status: 200, headers: { "Content-Type": "text/event-stream" } })) as FetchImpl; + + const stream = streamPiNative(fakeModel(), baseContext, { apiKey: "k", fetch: fetchImpl }); + const seen = await collectEvents(stream); + expect(seen.length).toBeGreaterThanOrEqual(2); + expect(seen[seen.length - 1].type).toBe("done"); + + const result = await stream.result(); + expect(result.role).toBe("assistant"); + expect(result.stopReason).toBe("stop"); + }); + + it("fails fast when the caller's signal is already aborted before fetch fires", async () => { + const fetchImpl = spyOn({ fetch: globalThis.fetch }, "fetch") as unknown as FetchImpl; + const controller = new AbortController(); + controller.abort(new Error("pre-aborted")); + + const stream = streamPiNative(fakeModel(), baseContext, { + apiKey: "k", + fetch: fetchImpl, + signal: controller.signal, + }); + + await expect(stream.result()).rejects.toThrow(/pre-aborted/); + // fetch was never called — short-circuit happened in the abort guard + expect((fetchImpl as unknown as ReturnType).mock.calls.length).toBe(0); + }); + + it("cancels the response body when the caller aborts mid-stream", async () => { + let cancelReason: unknown; + const blockedBody = new ReadableStream({ + start() { + // Never enqueues a terminal event — we abort instead. + }, + cancel(reason) { + cancelReason = reason; + }, + }); + const fetchImpl: FetchImpl = (async () => + new Response(blockedBody, { status: 200, headers: { "Content-Type": "text/event-stream" } })) as FetchImpl; + + const controller = new AbortController(); + const stream = streamPiNative(fakeModel(), baseContext, { + apiKey: "k", + fetch: fetchImpl, + signal: controller.signal, + }); + + // Schedule the abort after the request body is in-flight. + setTimeout(() => controller.abort(new Error("operator abort")), 5); + await expect(stream.result()).rejects.toThrow(); + expect(String(cancelReason)).toMatch(/aborted/i); + }); +}); diff --git a/packages/coding-agent/src/config/model-equivalence.ts b/packages/coding-agent/src/config/model-equivalence.ts index 77861ffb0..5a5efdbcb 100644 --- a/packages/coding-agent/src/config/model-equivalence.ts +++ b/packages/coding-agent/src/config/model-equivalence.ts @@ -333,8 +333,45 @@ function selectBestOfficialCandidate(candidates: readonly string[]): string | un if (candidates.length === 0) { return undefined; } - const ranked = [...new Set(candidates)].sort(compareCandidatePreference); - return ranked[0]; + const seen = new Set(); + let bestCandidate: string | undefined; + let bestPenalty = 0; + let bestLength = 0; + for (const candidate of candidates) { + if (seen.has(candidate)) { + continue; + } + seen.add(candidate); + const penalty = getCandidatePenalty(candidate); + const length = candidate.length; + if (bestCandidate === undefined) { + bestCandidate = candidate; + bestPenalty = penalty; + bestLength = length; + continue; + } + if (penalty < bestPenalty) { + bestCandidate = candidate; + bestPenalty = penalty; + bestLength = length; + continue; + } + if (penalty > bestPenalty) { + continue; + } + if (length < bestLength) { + bestCandidate = candidate; + bestLength = length; + continue; + } + if (length > bestLength) { + continue; + } + if (candidate.localeCompare(bestCandidate) < 0) { + bestCandidate = candidate; + } + } + return bestCandidate; } function getWrapperCanonicalCandidates(candidate: string): string[] { diff --git a/packages/coding-agent/src/config/model-registry.ts b/packages/coding-agent/src/config/model-registry.ts index 361ec281b..dbb821c74 100644 --- a/packages/coding-agent/src/config/model-registry.ts +++ b/packages/coding-agent/src/config/model-registry.ts @@ -242,13 +242,14 @@ export const ModelsConfigFile = new ConfigFile("models", ModelsCon }, ); -/** Provider override config (baseUrl, headers, apiKey, compat) without custom models */ +/** Provider override config (baseUrl, headers, apiKey, compat, transport) without custom models */ interface ProviderOverride { baseUrl?: string; headers?: Record; apiKey?: string; authHeader?: boolean; compat?: Model["compat"]; + transport?: Model["transport"]; } interface DiscoveryProviderConfig { @@ -1085,14 +1086,15 @@ export class ModelRegistry { const configuredProviders = new Set(Object.keys(value.providers ?? {})); for (const [providerName, providerConfig] of providerEntries) { - // Always set overrides when baseUrl/headers/apiKey/authHeader/compat/disableStrictTools are present + // Always set overrides when baseUrl/headers/apiKey/authHeader/compat/disableStrictTools/transport are present if ( providerConfig.baseUrl || providerConfig.headers || providerConfig.apiKey || providerConfig.authHeader !== undefined || providerConfig.compat || - providerConfig.disableStrictTools + providerConfig.disableStrictTools || + providerConfig.transport ) { const disableStrictCompat = providerConfig.disableStrictTools ? { disableStrictTools: true } : undefined; overrides.set(providerName, { @@ -1101,6 +1103,7 @@ export class ModelRegistry { apiKey: providerConfig.apiKey, authHeader: providerConfig.authHeader, compat: mergeCompat(providerConfig.compat, disableStrictCompat), + transport: providerConfig.transport, }); } @@ -1192,6 +1195,9 @@ export class ModelRegistry { headers: providerOverride.headers ? { ...model.headers, ...providerOverride.headers } : model.headers, + ...(providerOverride.transport !== undefined + ? { transport: providerOverride.transport } + : {}), } : model; }), @@ -1693,11 +1699,12 @@ export class ModelRegistry { authHeader: override.authHeader ?? baseOverride?.authHeader, headers: override.headers ? { ...(baseOverride?.headers ?? {}), ...override.headers } : baseOverride?.headers, compat: override.compat ? mergeCompat(baseOverride?.compat, override.compat) : baseOverride?.compat, + transport: override.transport ?? baseOverride?.transport, }; } #applyProviderTransportOverride }>( entry: T, - override: Pick, + override: Pick, ): T { const headers = mergeAuthHeader( override.headers ? { ...entry.headers, ...override.headers } : entry.headers, @@ -1708,6 +1715,9 @@ export class ModelRegistry { ...entry, baseUrl: override.baseUrl ?? entry.baseUrl, headers, + // Preserve the model's existing transport when the override omits one; + // providers without a `transport` field keep the default per-API dispatch. + ...(override.transport !== undefined ? { transport: override.transport } : {}), }; } #applyRuntimeProviderOverrides(models: Model[]): Model[] { @@ -2182,12 +2192,19 @@ export class ModelRegistry { return; } - if (config.baseUrl || config.headers || config.apiKey || config.authHeader !== undefined) { + if ( + config.baseUrl || + config.headers || + config.apiKey || + config.authHeader !== undefined || + config.transport !== undefined + ) { const transportOverride = { baseUrl: config.baseUrl, headers: config.headers, apiKey: config.apiKey, authHeader: config.authHeader, + transport: config.transport, }; const nextRuntimeOverride = this.#mergeProviderOverride( this.#runtimeProviderOverrides.get(providerName), @@ -2235,6 +2252,8 @@ export interface ProviderConfigInput { headers?: Record; compat?: Model["compat"]; authHeader?: boolean; + /** Streaming transport override — see {@link Model.transport}. */ + transport?: Model["transport"]; oauth?: { name: string; login(callbacks: OAuthLoginCallbacks): Promise; diff --git a/packages/coding-agent/src/config/models-config-schema.ts b/packages/coding-agent/src/config/models-config-schema.ts index 9b3802bcf..d8a6632d9 100644 --- a/packages/coding-agent/src/config/models-config-schema.ts +++ b/packages/coding-agent/src/config/models-config-schema.ts @@ -151,6 +151,14 @@ const ProviderConfigSchema = z.object({ models: z.array(ModelDefinitionSchema).optional(), modelOverrides: z.record(z.string(), ModelOverrideSchema).optional(), disableStrictTools: z.boolean().optional(), + /** + * Streaming transport override. When set to `"pi-native"`, omp dispatches + * every model under this provider via the auth-gateway's + * `POST /v1/pi/stream` endpoint instead of the per-provider SDK. The + * provider's `baseUrl` must point at a compatible `omp auth-gateway` + * and `apiKey` must carry the gateway bearer. + */ + transport: z.literal("pi-native").optional(), }); const EquivalenceConfigSchema = z.object({ diff --git a/packages/coding-agent/src/debug/profiler.ts b/packages/coding-agent/src/debug/profiler.ts index 38774bc7e..242cb2a2f 100644 --- a/packages/coding-agent/src/debug/profiler.ts +++ b/packages/coding-agent/src/debug/profiler.ts @@ -121,6 +121,10 @@ export async function startCpuProfile(): Promise { session.connect(); await session.post("Profiler.enable"); + // Default CDP interval is 1ms, which mis-attributes await-resumption samples + // to the line after `await` (one sparse sample inherits the entire wait). 100µs + // scatters samples enough to keep CPU vs. async-wait attribution honest. + await session.post("Profiler.setSamplingInterval", { interval: 100 }); await session.post("Profiler.start"); return {