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.
This commit is contained in:
@@ -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<string, { module: FormatModule; label: string }> = {
|
||||
*
|
||||
* 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<Response> {
|
||||
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);
|
||||
|
||||
@@ -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<keyof SimpleStreamOptions>([
|
||||
"signal",
|
||||
"apiKey",
|
||||
"fetch",
|
||||
"onPayload",
|
||||
"onResponse",
|
||||
"onSseEvent",
|
||||
"execHandlers",
|
||||
"cursorExecHandlers",
|
||||
"cursorOnToolResult",
|
||||
"providerSessionState",
|
||||
]);
|
||||
|
||||
function buildWireOptions(options: SimpleStreamOptions | undefined): Record<string, unknown> {
|
||||
if (!options) return {};
|
||||
const wire: Record<string, unknown> = {};
|
||||
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<Error> {
|
||||
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<Api>): 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<Api>, apiKey: string | undefined): Record<string, string> {
|
||||
const headers: Record<string, string> = {
|
||||
"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<TApi extends Api>(
|
||||
model: Model<TApi>,
|
||||
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<Api>);
|
||||
const fetchImpl = options?.fetch ?? globalThis.fetch;
|
||||
const headers = buildHeaders(model as Model<Api>, 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<AssistantMessageEvent>(
|
||||
response.body as ReadableStream<Uint8Array>,
|
||||
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<Api>);
|
||||
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<Api>): 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(),
|
||||
};
|
||||
}
|
||||
@@ -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<keyof SimpleStreamOptions> = 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<string, unknown>;
|
||||
|
||||
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<string, unknown>;
|
||||
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<string, unknown>;
|
||||
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<string, unknown>;
|
||||
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<Uint8Array> {
|
||||
return new ReadableStream<Uint8Array>({
|
||||
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",
|
||||
},
|
||||
});
|
||||
}
|
||||
@@ -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<TApi extends Api>(
|
||||
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);
|
||||
|
||||
@@ -767,6 +767,21 @@ export interface Model<TApi extends Api = any> {
|
||||
contextWindow: number;
|
||||
maxTokens: number;
|
||||
headers?: Record<string, string>;
|
||||
/**
|
||||
* 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). */
|
||||
|
||||
@@ -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<AssistantMessage> }).result = async () => final;
|
||||
return stream;
|
||||
}
|
||||
|
||||
async function collectSse(stream: ReadableStream<Uint8Array>): Promise<string[]> {
|
||||
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>): 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<string, unknown>
|
||||
>;
|
||||
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<AssistantMessage> }).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" } });
|
||||
});
|
||||
});
|
||||
@@ -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<Uint8Array> {
|
||||
return new ReadableStream<Uint8Array>({
|
||||
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> = {}): 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">> = {}): 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<AssistantMessageEvent>,
|
||||
): Promise<AssistantMessageEvent[]> {
|
||||
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<string, string>;
|
||||
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<string, string>;
|
||||
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<Uint8Array>({
|
||||
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<typeof spyOn>).mock.calls.length).toBe(0);
|
||||
});
|
||||
|
||||
it("cancels the response body when the caller aborts mid-stream", async () => {
|
||||
let cancelReason: unknown;
|
||||
const blockedBody = new ReadableStream<Uint8Array>({
|
||||
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);
|
||||
});
|
||||
});
|
||||
@@ -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<string>();
|
||||
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[] {
|
||||
|
||||
@@ -242,13 +242,14 @@ export const ModelsConfigFile = new ConfigFile<ModelsConfig>("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<string, string>;
|
||||
apiKey?: string;
|
||||
authHeader?: boolean;
|
||||
compat?: Model<Api>["compat"];
|
||||
transport?: Model<Api>["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<T extends { baseUrl?: string; headers?: Record<string, string> }>(
|
||||
entry: T,
|
||||
override: Pick<ProviderOverride, "baseUrl" | "headers" | "authHeader" | "apiKey">,
|
||||
override: Pick<ProviderOverride, "baseUrl" | "headers" | "authHeader" | "apiKey" | "transport">,
|
||||
): 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<Api>[]): Model<Api>[] {
|
||||
@@ -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<string, string>;
|
||||
compat?: Model<Api>["compat"];
|
||||
authHeader?: boolean;
|
||||
/** Streaming transport override — see {@link Model.transport}. */
|
||||
transport?: Model<Api>["transport"];
|
||||
oauth?: {
|
||||
name: string;
|
||||
login(callbacks: OAuthLoginCallbacks): Promise<OAuthCredentials | string>;
|
||||
|
||||
@@ -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({
|
||||
|
||||
@@ -121,6 +121,10 @@ export async function startCpuProfile(): Promise<ProfilerSession> {
|
||||
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 {
|
||||
|
||||
Reference in New Issue
Block a user