From 961fa6c14228d3c652869256fa2349d1f08e6306 Mon Sep 17 00:00:00 2001 From: Armin Ronacher Date: Tue, 14 Jul 2026 11:01:21 +0200 Subject: [PATCH] feat(ai): add Radius gateway support --- packages/ai/src/api/pi-messages.lazy.ts | 4 + packages/ai/src/api/pi-messages.ts | 436 ++++++++++++++ packages/ai/src/compat.ts | 6 + packages/ai/src/env-api-keys.ts | 1 + packages/ai/src/index.ts | 1 + packages/ai/src/providers/all.ts | 17 +- packages/ai/src/types.ts | 6 +- packages/ai/src/utils/oauth/index.ts | 16 + packages/ai/src/utils/oauth/radius.ts | 557 ++++++++++++++++++ packages/ai/test/pi-messages.test.ts | 243 ++++++++ .../coding-agent/src/core/model-registry.ts | 22 +- .../coding-agent/src/core/model-resolver.ts | 1 + .../src/core/provider-display-names.ts | 1 + packages/coding-agent/src/core/radius.ts | 34 ++ packages/coding-agent/test/radius.test.ts | 146 +++++ test.sh | 3 + 16 files changed, 1483 insertions(+), 11 deletions(-) create mode 100644 packages/ai/src/api/pi-messages.lazy.ts create mode 100644 packages/ai/src/api/pi-messages.ts create mode 100644 packages/ai/src/utils/oauth/radius.ts create mode 100644 packages/ai/test/pi-messages.test.ts create mode 100644 packages/coding-agent/src/core/radius.ts create mode 100644 packages/coding-agent/test/radius.test.ts diff --git a/packages/ai/src/api/pi-messages.lazy.ts b/packages/ai/src/api/pi-messages.lazy.ts new file mode 100644 index 00000000..41af9ee5 --- /dev/null +++ b/packages/ai/src/api/pi-messages.lazy.ts @@ -0,0 +1,4 @@ +import type { ProviderStreams } from "../types.ts"; +import { lazyApi } from "./lazy.ts"; + +export const piMessagesApi = (): ProviderStreams => lazyApi(() => import("./pi-messages.ts")); diff --git a/packages/ai/src/api/pi-messages.ts b/packages/ai/src/api/pi-messages.ts new file mode 100644 index 00000000..0af13f88 --- /dev/null +++ b/packages/ai/src/api/pi-messages.ts @@ -0,0 +1,436 @@ +/** + * pi-messages API implementation. + * + * Streams pi's own message protocol directly to a backend: the request is a + * single POST of `{ model, context, options }` to `/messages`, the + * response is an SSE stream of serialized assistant-message events plus a + * terminal `done`/`error` event. This is the wire protocol spoken by the + * Radius gateway, but any backend implementing it can be used, e.g. via a + * models.json custom provider with `"api": "pi-messages"`. + */ + +import type { + AssistantMessage, + AssistantMessageEvent, + CacheRetention, + Context, + Model, + ProviderEnv, + SimpleStreamOptions, + StreamFunction, + StreamOptions, + ThinkingLevel, + ToolCall, +} from "../types.ts"; +import { appendAssistantMessageDiagnostic, createAssistantMessageDiagnostic } from "../utils/diagnostics.ts"; +import { AssistantMessageEventStream } from "../utils/event-stream.ts"; +import { headersToRecord, providerHeadersToRecord } from "../utils/headers.ts"; +import { parseStreamingJson } from "../utils/json-parse.ts"; +import { getProviderEnvValue } from "../utils/provider-env.ts"; + +export interface PiMessagesOptions extends StreamOptions { + reasoning?: ThinkingLevel; + toolChoice?: "auto" | "none" | "required" | { type: "function"; function: { name: string } }; + /** Ask the backend for debug metadata (e.g. routing response headers). */ + debug?: boolean; +} + +type PiMessagesUsage = AssistantMessage["usage"]; +type PiMessagesStopReason = AssistantMessage["stopReason"]; + +/** Impact summary of a server-side message rewrite (e.g. a gateway policy). */ +export type PiMessagesRewriteImpact = { + policyId: string; + policyVersion: number; + changed: boolean; + tokenCountChange: number; + messageCountChange: number; + systemPromptChanged: boolean; +}; + +/** Serialized assistant-message event as sent by a pi-messages backend. */ +export type PiMessagesEvent = + | { type: "start" } + | { type: "text_start"; contentIndex: number } + | { type: "text_delta"; contentIndex: number; delta: string } + | { type: "text_end"; contentIndex: number; content: string; contentSignature?: string } + | { type: "thinking_start"; contentIndex: number } + | { type: "thinking_delta"; contentIndex: number; delta: string } + | { + type: "thinking_end"; + contentIndex: number; + content: string; + contentSignature?: string; + redacted?: boolean; + } + | { type: "toolcall_start"; contentIndex: number; id: string; toolName: string } + | { type: "toolcall_delta"; contentIndex: number; delta: string } + | { type: "toolcall_end"; contentIndex: number; toolCall: ToolCall } + | { + type: "done"; + reason: Extract; + usage: PiMessagesUsage; + responseId?: string; + rewrite?: PiMessagesRewriteImpact; + } + | { + type: "error"; + reason: Extract; + usage: PiMessagesUsage; + errorMessage?: string; + responseId?: string; + rewrite?: PiMessagesRewriteImpact; + }; + +type PiMessagesErrorBody = { + error?: { + message?: unknown; + code?: unknown; + details?: unknown; + [key: string]: unknown; + }; +}; + +export class PiMessagesResponseError extends Error { + code?: string; + readonly diagnosticDetails: Record; + + constructor(message: string, code: string | undefined, diagnosticDetails: Record) { + super(message); + this.name = "PiMessagesResponseError"; + this.code = code; + this.diagnosticDetails = diagnosticDetails; + } +} + +function isRecord(value: unknown): value is Record { + return typeof value === "object" && value !== null && !Array.isArray(value); +} + +function parsePiMessagesErrorBody(body: string): PiMessagesErrorBody | undefined { + try { + const parsed = JSON.parse(body) as unknown; + return isRecord(parsed) && isRecord(parsed.error) ? (parsed as PiMessagesErrorBody) : undefined; + } catch { + return undefined; + } +} + +function truncateDiagnosticString(value: string): string { + const maxLength = 8192; + return value.length > maxLength ? `${value.slice(0, maxLength)}…` : value; +} + +function formatPiMessagesResponseError( + response: Response, + body: string, + errorBody: PiMessagesErrorBody | undefined, +): string { + const message = typeof errorBody?.error?.message === "string" ? errorBody.error.message : undefined; + const code = typeof errorBody?.error?.code === "string" ? errorBody.error.code : undefined; + const suffix = message ?? body; + const codeSuffix = code ? ` (${code})` : ""; + return `${response.status} ${response.statusText}: ${suffix}${codeSuffix}`; +} + +function createPiMessagesResponseError( + model: Model<"pi-messages">, + url: URL, + response: Response, + body: string, +): PiMessagesResponseError { + const errorBody = parsePiMessagesErrorBody(body); + const code = typeof errorBody?.error?.code === "string" ? errorBody.error.code : undefined; + return new PiMessagesResponseError(formatPiMessagesResponseError(response, body, errorBody), code, { + version: 1, + provider: model.provider, + model: model.id, + url: url.toString(), + status: response.status, + statusText: response.statusText, + error: errorBody?.error, + body: errorBody ? undefined : truncateDiagnosticString(body), + timestampMs: Date.now(), + }); +} + +function createEmptyUsage(): PiMessagesUsage { + return { + input: 0, + output: 0, + cacheRead: 0, + cacheWrite: 0, + totalTokens: 0, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, + }; +} + +function appendRewriteDiagnostic(message: AssistantMessage, rewrite: PiMessagesRewriteImpact | undefined): void { + if (!rewrite) { + return; + } + appendAssistantMessageDiagnostic(message, { + type: "pi_messages_rewrite", + timestamp: Date.now(), + details: { ...rewrite }, + }); +} + +function createEventConverter(model: Model<"pi-messages">) { + const partial: AssistantMessage = { + role: "assistant", + content: [], + api: model.api, + provider: model.provider, + model: model.id, + usage: createEmptyUsage(), + stopReason: "stop", + timestamp: Date.now(), + }; + const toolJson = new Map(); + + return (event: PiMessagesEvent): AssistantMessageEvent => { + switch (event.type) { + case "done": + Object.assign(partial, { + stopReason: event.reason, + usage: event.usage, + responseId: event.responseId, + }); + appendRewriteDiagnostic(partial, event.rewrite); + return { type: "done", reason: event.reason, message: partial }; + case "error": + Object.assign(partial, { + stopReason: event.reason, + usage: event.usage, + errorMessage: event.errorMessage, + responseId: event.responseId, + }); + appendRewriteDiagnostic(partial, event.rewrite); + return { type: "error", reason: event.reason, error: partial }; + case "start": + break; + case "text_start": + partial.content[event.contentIndex] = { type: "text", text: "" }; + break; + case "text_delta": + (partial.content[event.contentIndex] as { text: string }).text += event.delta; + break; + case "text_end": + Object.assign(partial.content[event.contentIndex]!, { + text: event.content, + textSignature: event.contentSignature, + }); + break; + case "thinking_start": + partial.content[event.contentIndex] = { type: "thinking", thinking: "" }; + break; + case "thinking_delta": + (partial.content[event.contentIndex] as { thinking: string }).thinking += event.delta; + break; + case "thinking_end": + Object.assign(partial.content[event.contentIndex]!, { + thinking: event.content, + thinkingSignature: event.contentSignature, + redacted: event.redacted, + }); + break; + case "toolcall_start": + partial.content[event.contentIndex] = { + type: "toolCall", + id: event.id, + name: event.toolName, + arguments: {}, + }; + toolJson.set(event.contentIndex, ""); + break; + case "toolcall_delta": { + const json = `${toolJson.get(event.contentIndex) ?? ""}${event.delta}`; + toolJson.set(event.contentIndex, json); + (partial.content[event.contentIndex] as ToolCall).arguments = + parseStreamingJson(json); + break; + } + case "toolcall_end": + Object.assign(partial.content[event.contentIndex]!, event.toolCall); + toolJson.delete(event.contentIndex); + return { + type: "toolcall_end", + contentIndex: event.contentIndex, + toolCall: partial.content[event.contentIndex] as ToolCall, + partial, + }; + } + + return { ...event, partial } as AssistantMessageEvent; + }; +} + +async function* readPiMessagesEvents(stream: ReadableStream): AsyncGenerator { + const decoder = new TextDecoder(); + const reader = stream.getReader(); + let buffer = ""; + + try { + while (true) { + const { done, value } = await reader.read(); + buffer += done ? decoder.decode() : decoder.decode(value, { stream: true }); + buffer = buffer.replace(/\r\n/g, "\n"); + + let split = buffer.indexOf("\n\n"); + while (split !== -1) { + const event = parsePiMessagesEvent(buffer.slice(0, split)); + if (event) { + yield event; + } + buffer = buffer.slice(split + 2); + split = buffer.indexOf("\n\n"); + } + + if (done) { + break; + } + } + + if (buffer.trim()) { + const event = parsePiMessagesEvent(buffer); + if (event) { + yield event; + } + } + } finally { + reader.releaseLock(); + } +} + +function parsePiMessagesEvent(raw: string): PiMessagesEvent | undefined { + const data = raw + .split("\n") + .find((line) => line.startsWith("data:")) + ?.slice(5) + .trim(); + + return data && data !== "[DONE]" ? (JSON.parse(data) as PiMessagesEvent) : undefined; +} + +function createErrorEvent(model: Model<"pi-messages">, error: unknown, aborted: boolean): AssistantMessageEvent { + const reason = aborted ? "aborted" : "error"; + const assistantMessage: AssistantMessage = { + role: "assistant", + content: [], + api: model.api, + provider: model.provider, + model: model.id, + usage: createEmptyUsage(), + stopReason: reason, + errorMessage: error instanceof Error ? error.message : String(error), + timestamp: Date.now(), + }; + + if (!aborted && error instanceof PiMessagesResponseError) { + appendAssistantMessageDiagnostic( + assistantMessage, + createAssistantMessageDiagnostic("pi_messages_response_failure", error, error.diagnosticDetails), + ); + } + + return { type: "error", reason, error: assistantMessage }; +} + +function resolveCacheRetention(cacheRetention?: CacheRetention, env?: ProviderEnv): CacheRetention | undefined { + if (cacheRetention) { + return cacheRetention; + } + // Backend defaults apply when unset; only the legacy env opt-in is mapped. + return getProviderEnvValue("PI_CACHE_RETENTION", env) === "long" ? "long" : undefined; +} + +export const stream: StreamFunction<"pi-messages", PiMessagesOptions> = ( + model: Model<"pi-messages">, + context: Context, + options?: PiMessagesOptions, +): AssistantMessageEventStream => { + const eventStream = new AssistantMessageEventStream(); + const convertEvent = createEventConverter(model); + + void (async () => { + try { + const apiKey = options?.apiKey; + if (!apiKey) { + throw new Error(`No API key provided for provider "${model.provider}"`); + } + + const url = new URL(`${model.baseUrl.replace(/\/+$/u, "")}/messages`); + if (options?.debug) { + url.searchParams.set("debug", "1"); + } + + let payload: unknown = { + model: model.id, + context, + options: { + temperature: options?.temperature, + maxTokens: options?.maxTokens, + reasoning: options?.reasoning, + cacheRetention: resolveCacheRetention(options?.cacheRetention, options?.env), + sessionId: options?.sessionId, + toolChoice: options?.toolChoice, + }, + }; + const nextPayload = await options?.onPayload?.(payload, model); + if (nextPayload !== undefined) { + payload = nextPayload; + } + + const response = await fetch(url, { + method: "POST", + headers: { + authorization: `Bearer ${apiKey}`, + accept: "text/event-stream", + "content-type": "application/json", + ...providerHeadersToRecord(options?.headers), + }, + body: JSON.stringify(payload), + signal: options?.signal, + }); + + await options?.onResponse?.({ status: response.status, headers: headersToRecord(response.headers) }, model); + + if (!response.ok) { + const body = await response.text(); + throw createPiMessagesResponseError(model, url, response, body); + } + if (!response.body) { + throw new Error(`${model.provider} response has no body`); + } + + for await (const piEvent of readPiMessagesEvents(response.body)) { + const event = convertEvent(piEvent); + eventStream.push(event); + if (event.type === "done" || event.type === "error") { + return; + } + } + + throw new Error(`${model.provider} stream ended without a terminal event`); + } catch (error) { + eventStream.push(createErrorEvent(model, error, options?.signal?.aborted ?? false)); + } + })(); + + return eventStream; +}; + +export const streamSimple: StreamFunction<"pi-messages", SimpleStreamOptions> = ( + model: Model<"pi-messages">, + context: Context, + options?: SimpleStreamOptions, +): AssistantMessageEventStream => { + const extra = options as PiMessagesOptions | undefined; + return stream(model, context, { + ...options, + reasoning: options?.reasoning, + toolChoice: extra?.toolChoice, + debug: extra?.debug, + }); +}; diff --git a/packages/ai/src/compat.ts b/packages/ai/src/compat.ts index eb1151fd..c2b9fc69 100644 --- a/packages/ai/src/compat.ts +++ b/packages/ai/src/compat.ts @@ -19,6 +19,7 @@ export * from "./api/mistral-conversations.lazy.ts"; export * from "./api/openai-codex-responses.lazy.ts"; export * from "./api/openai-completions.lazy.ts"; export * from "./api/openai-responses.lazy.ts"; +export * from "./api/pi-messages.lazy.ts"; export * from "./env-api-keys.ts"; export * from "./image-models.ts"; export * from "./images.ts"; @@ -36,8 +37,12 @@ import { mistralConversationsApi } from "./api/mistral-conversations.lazy.ts"; import { openAICodexResponsesApi } from "./api/openai-codex-responses.lazy.ts"; import { openAICompletionsApi } from "./api/openai-completions.lazy.ts"; import { openAIResponsesApi } from "./api/openai-responses.lazy.ts"; +import { piMessagesApi } from "./api/pi-messages.lazy.ts"; import { getEnvApiKey } from "./env-api-keys.ts"; import { builtinModels, getBuiltinModel, getBuiltinModels, getBuiltinProviders } from "./providers/all.ts"; + +export type { BuiltinProvider } from "./providers/all.ts"; + import { createFauxCore, type FauxProviderRegistration, type RegisterFauxProviderOptions } from "./providers/faux.ts"; import type { Api, @@ -179,6 +184,7 @@ const BUILTIN_APIS: [Api, ProviderStreams][] = [ ["google-vertex", googleVertexApi()], ["mistral-conversations", mistralConversationsApi()], ["bedrock-converse-stream", bedrockConverseStreamApi()], + ["pi-messages", piMessagesApi()], ]; const builtinApiProviderInstances = new Map>(); diff --git a/packages/ai/src/env-api-keys.ts b/packages/ai/src/env-api-keys.ts index 7bd955e1..8687fcc4 100644 --- a/packages/ai/src/env-api-keys.ts +++ b/packages/ai/src/env-api-keys.ts @@ -82,6 +82,7 @@ function getApiKeyEnvVars(provider: string): readonly string[] | undefined { groq: "GROQ_API_KEY", cerebras: "CEREBRAS_API_KEY", xai: "XAI_API_KEY", + radius: "PI_GATEWAY_API_KEY", openrouter: "OPENROUTER_API_KEY", "vercel-ai-gateway": "AI_GATEWAY_API_KEY", zai: "ZAI_API_KEY", diff --git a/packages/ai/src/index.ts b/packages/ai/src/index.ts index 57c448cc..25290c6d 100644 --- a/packages/ai/src/index.ts +++ b/packages/ai/src/index.ts @@ -17,6 +17,7 @@ export type { MistralOptions } from "./api/mistral-conversations.ts"; export type { OpenAICodexResponsesOptions, OpenAICodexWebSocketDebugStats } from "./api/openai-codex-responses.ts"; export type { OpenAICompletionsOptions } from "./api/openai-completions.ts"; export type { OpenAIResponsesOptions } from "./api/openai-responses.ts"; +export type { PiMessagesEvent, PiMessagesOptions, PiMessagesRewriteImpact } from "./api/pi-messages.ts"; export * from "./auth/context.ts"; export * from "./auth/credential-store.ts"; export * from "./auth/helpers.ts"; diff --git a/packages/ai/src/providers/all.ts b/packages/ai/src/providers/all.ts index 85ba0301..ad919016 100644 --- a/packages/ai/src/providers/all.ts +++ b/packages/ai/src/providers/all.ts @@ -1,7 +1,7 @@ import { createImagesModels, type ImagesProvider, type MutableImagesModels } from "../images-models.ts"; import { MODELS } from "../models.generated.ts"; import { type CreateModelsOptions, createModels, type MutableModels, type Provider } from "../models.ts"; -import type { Api, KnownProvider, Model } from "../types.ts"; +import type { Api, Model } from "../types.ts"; import { amazonBedrockProvider } from "./amazon-bedrock.ts"; import { antLingProvider } from "./ant-ling.ts"; import { anthropicProvider } from "./anthropic.ts"; @@ -39,13 +39,18 @@ import { xiaomiTokenPlanSgpProvider } from "./xiaomi-token-plan-sgp.ts"; import { zaiProvider } from "./zai.ts"; import { zaiCodingCnProvider } from "./zai-coding-cn.ts"; +/** Providers present in the generated catalog. `KnownProvider` additionally + * includes purely dynamic providers (e.g. "radius") that have no static + * catalog entry. */ +export type BuiltinProvider = keyof typeof MODELS; + type BuiltinModelApi< - TProvider extends KnownProvider, + TProvider extends BuiltinProvider, TModelId extends keyof (typeof MODELS)[TProvider], > = (typeof MODELS)[TProvider][TModelId] extends { api: infer TApi } ? (TApi extends Api ? TApi : never) : never; /** Typed read of the generated built-in catalog. */ -export function getBuiltinModel( +export function getBuiltinModel( provider: TProvider, modelId: TModelId, ): Model> { @@ -53,11 +58,11 @@ export function getBuiltinModel>; } -export function getBuiltinProviders(): KnownProvider[] { - return Object.keys(MODELS) as KnownProvider[]; +export function getBuiltinProviders(): BuiltinProvider[] { + return Object.keys(MODELS) as BuiltinProvider[]; } -export function getBuiltinModels( +export function getBuiltinModels( provider: TProvider, ): Model>[] { const models = MODELS[provider] as Record> | undefined; diff --git a/packages/ai/src/types.ts b/packages/ai/src/types.ts index 10530fae..69356b22 100644 --- a/packages/ai/src/types.ts +++ b/packages/ai/src/types.ts @@ -7,6 +7,7 @@ import type { MistralOptions } from "./api/mistral-conversations.ts"; import type { OpenAICodexResponsesOptions } from "./api/openai-codex-responses.ts"; import type { OpenAICompletionsOptions } from "./api/openai-completions.ts"; import type { OpenAIResponsesOptions } from "./api/openai-responses.ts"; +import type { PiMessagesOptions } from "./api/pi-messages.ts"; import type { AssistantMessageDiagnostic } from "./utils/diagnostics.ts"; import type { AssistantMessageEventStream } from "./utils/event-stream.ts"; @@ -21,7 +22,8 @@ export type KnownApi = | "anthropic-messages" | "bedrock-converse-stream" | "google-generative-ai" - | "google-vertex"; + | "google-vertex" + | "pi-messages"; export type Api = KnownApi | (string & {}); @@ -38,6 +40,7 @@ export type KnownProvider = | "openai" | "azure-openai-responses" | "openai-codex" + | "radius" | "nvidia" | "deepseek" | "github-copilot" @@ -202,6 +205,7 @@ export interface ApiOptionsMap { "google-vertex": GoogleVertexOptions; "mistral-conversations": MistralOptions; "bedrock-converse-stream": BedrockOptions; + "pi-messages": PiMessagesOptions; } /** diff --git a/packages/ai/src/utils/oauth/index.ts b/packages/ai/src/utils/oauth/index.ts index a57badda..e9e1f0e9 100644 --- a/packages/ai/src/utils/oauth/index.ts +++ b/packages/ai/src/utils/oauth/index.ts @@ -28,21 +28,37 @@ export { refreshOpenAICodexToken, } from "./openai-codex.ts"; +// Radius (pi-messages gateway) +export { + createRadiusOAuthProvider, + DEFAULT_RADIUS_GATEWAY, + type RadiusGatewayConfig, + type RadiusGatewayModel, + type RadiusOAuthCredentials, + type RadiusOAuthProviderOptions, +} from "./radius.ts"; export * from "./types.ts"; // ============================================================================ // Provider Registry // ============================================================================ +import { getProviderEnvValue } from "../provider-env.ts"; import { anthropicOAuthProvider } from "./anthropic.ts"; import { githubCopilotOAuthProvider } from "./github-copilot.ts"; import { openaiCodexOAuthProvider } from "./openai-codex.ts"; +import { createRadiusOAuthProvider, DEFAULT_RADIUS_GATEWAY } from "./radius.ts"; import type { OAuthCredentials, OAuthProviderId, OAuthProviderInfo, OAuthProviderInterface } from "./types.ts"; const BUILT_IN_OAUTH_PROVIDERS: OAuthProviderInterface[] = [ anthropicOAuthProvider, githubCopilotOAuthProvider, openaiCodexOAuthProvider, + createRadiusOAuthProvider({ + id: "radius", + name: "Radius", + gateway: getProviderEnvValue("PI_GATEWAY") || DEFAULT_RADIUS_GATEWAY, + }), ]; const oauthProviderRegistry = new Map( diff --git a/packages/ai/src/utils/oauth/radius.ts b/packages/ai/src/utils/oauth/radius.ts new file mode 100644 index 00000000..077619ec --- /dev/null +++ b/packages/ai/src/utils/oauth/radius.ts @@ -0,0 +1,557 @@ +/** + * Radius gateway OAuth flow and model catalog loading. + * + * Radius is a pi-messages gateway. OAuth endpoints are discovered from the + * gateway (`/v1/oauth`); the model catalog comes from `/v1/config` and is + * cached on the stored credential (`gatewayConfig`) so models are available + * at startup and refreshed whenever the token refreshes. + * + * NOTE: This module uses node:http for the OAuth callback server. + * It is only intended for CLI use, not browser environments. + */ + +// NEVER convert to top-level imports - breaks browser/Vite builds +let _http: typeof import("node:http") | null = null; +if (typeof process !== "undefined" && (process.versions?.node || process.versions?.bun)) { + import("node:http").then((m) => { + _http = m; + }); +} + +import type { Api, Model, ThinkingLevelMap } from "../../types.ts"; +import { pollOAuthDeviceCodeFlow } from "./device-code.ts"; +import { oauthErrorHtml, oauthSuccessHtml } from "./oauth-page.ts"; +import { generatePKCE } from "./pkce.ts"; +import type { OAuthCredentials, OAuthLoginCallbacks, OAuthProviderInterface } from "./types.ts"; + +export const DEFAULT_RADIUS_GATEWAY = "https://radius.pi.dev"; + +const CALLBACK_HOST = "127.0.0.1"; +const CALLBACK_PORT = 1456; +const CALLBACK_PATH = "/oauth/callback"; +const REDIRECT_URI = `http://${CALLBACK_HOST}:${CALLBACK_PORT}${CALLBACK_PATH}`; +const TOKEN_EXPIRY_SKEW_MS = 60_000; +const LOGIN_METHOD_BROWSER = "browser"; +const LOGIN_METHOD_DEVICE_CODE = "device-code"; + +/** Model metadata served by the gateway config endpoint. */ +export type RadiusGatewayModel = { + id: string; + name: string; + reasoning: boolean; + thinkingLevelMap?: ThinkingLevelMap; + input: ("text" | "image")[]; + cost: Model["cost"]; + contextWindow: number; + maxTokens: number; +}; + +export type RadiusGatewayConfig = { + baseUrl: string; + models: RadiusGatewayModel[]; +}; + +export type RadiusOAuthCredentials = OAuthCredentials & { + gatewayConfig?: RadiusGatewayConfig; +}; + +type RadiusOAuthConfig = { + issuer: string; + authorizationEndpoint: string; + tokenEndpoint: string; + deviceAuthorizationEndpoint: string; + deviceAuthorizationEventsEndpoint: string; + verificationEndpoint: string; + clientId: string; + scope: string; + deviceCodeGrantType: string; +}; + +type DeviceAuthorizationResponse = { + device_code: string; + user_code: string; + verification_uri?: string; + verification_uri_complete?: string; + expires_in: number; + interval?: number; +}; + +function isRecord(value: unknown): value is Record { + return typeof value === "object" && value !== null && !Array.isArray(value); +} + +function normalizeRadiusGatewayUrl(value: string): string { + const withScheme = /^https?:\/\//iu.test(value) ? value : `https://${value}`; + return withScheme.replace(/\/+$/u, ""); +} + +// The gateway is a trusted first-party service. The shape checks below only +// guard against version skew and stale credential caches: malformed entries +// are dropped rather than failing the whole catalog, and nested fields (e.g. +// `input` members, `cost` rates) are intentionally not validated in depth. +// Do not turn this into strict validation. +function isRadiusGatewayModel(value: unknown): value is RadiusGatewayModel { + if (!isRecord(value)) { + return false; + } + return ( + typeof value.id === "string" && + typeof value.name === "string" && + typeof value.reasoning === "boolean" && + Array.isArray(value.input) && + isRecord(value.cost) && + typeof value.contextWindow === "number" && + typeof value.maxTokens === "number" + ); +} + +function sanitizeRadiusGatewayConfig(config: unknown): RadiusGatewayConfig | undefined { + if (!isRecord(config)) { + return undefined; + } + const baseUrl = config.baseUrl; + const models = config.models; + if (typeof baseUrl !== "string" || !Array.isArray(models)) { + return undefined; + } + + return { + baseUrl, + models: models.filter(isRadiusGatewayModel).map((model) => ({ ...model })), + }; +} + +function getRadiusCredentialConfig(credentials: OAuthCredentials | undefined): RadiusGatewayConfig | undefined { + return sanitizeRadiusGatewayConfig((credentials as RadiusOAuthCredentials | undefined)?.gatewayConfig); +} + +function truncateHttpBody(body: string): string { + const trimmed = body.trim(); + return trimmed.length > 512 ? `${trimmed.slice(0, 512)}…` : trimmed; +} + +async function loadRadiusGatewayConfig(gateway: string, apiKey?: string): Promise { + const headers: Record = { accept: "application/json" }; + if (apiKey) { + headers.authorization = `Bearer ${apiKey}`; + } + + const response = await fetch(new URL("/v1/config", gateway), { headers }); + if (!response.ok) { + throw new Error( + `Could not load Radius config from ${gateway}: ${response.status}: ${truncateHttpBody(await response.text())}`, + ); + } + + const config = sanitizeRadiusGatewayConfig(await response.json()); + if (!config) { + throw new Error(`Invalid Radius config from ${gateway}`); + } + return config; +} + +async function loadRadiusOAuthConfig(gateway: string): Promise { + const response = await fetch(new URL("/v1/oauth", gateway), { + headers: { accept: "application/json" }, + }); + + if (!response.ok) { + throw new Error( + `Could not load Radius OAuth config from ${gateway}: ${response.status} ${await response.text()}`, + ); + } + + return (await response.json()) as RadiusOAuthConfig; +} + +class OAuthResponseError extends Error { + readonly status: number; + readonly oauthError?: string; + + constructor(status: number, oauthError: string | undefined, description: string | undefined, message: string) { + const detail = oauthError + ? description + ? `${oauthError}: ${description}` + : oauthError + : description || String(status); + super(`${message}: ${detail}`); + this.status = status; + this.oauthError = oauthError; + } +} + +async function readOAuthResponseError(response: Response, message: string): Promise { + const text = await response.text().catch(() => ""); + let oauthError: string | undefined; + let description: string | undefined; + + if (text) { + try { + const data = JSON.parse(text) as { error?: unknown; error_description?: unknown }; + oauthError = typeof data.error === "string" ? data.error : undefined; + description = typeof data.error_description === "string" ? data.error_description : undefined; + } catch { + description = text; + } + } + + return new OAuthResponseError(response.status, oauthError, description, message); +} + +async function requestOAuthToken( + oauth: RadiusOAuthConfig, + body: URLSearchParams, + signal?: AbortSignal, +): Promise { + let response: Response; + try { + response = await fetch(oauth.tokenEndpoint, { + method: "POST", + headers: { accept: "application/json", "content-type": "application/x-www-form-urlencoded" }, + body, + signal, + }); + } catch (error) { + if (signal?.aborted) { + throw new Error("Login cancelled"); + } + throw error; + } + + if (!response.ok) { + throw await readOAuthResponseError(response, "Radius OAuth token request failed"); + } + + const data = (await response.json()) as { + access_token: string; + refresh_token: string; + expires_in: number; + scope?: string; + }; + + return { + access: data.access_token, + refresh: data.refresh_token, + expires: Date.now() + data.expires_in * 1000 - TOKEN_EXPIRY_SKEW_MS, + scope: data.scope, + }; +} + +type OAuthCallbackServer = { + waitForCode(): Promise; + close(): void; +}; + +function startOAuthCallbackServer( + expectedState: string, + signal: AbortSignal | undefined, +): Promise { + if (!_http) { + throw new Error("Radius OAuth is only available in Node.js environments"); + } + + let settle: (code: string | null) => void = () => {}; + let settled = false; + const wait = new Promise((resolve) => { + settle = resolve; + }); + const finish = (code: string | null) => { + if (settled) { + return; + } + settled = true; + signal?.removeEventListener("abort", onAbort); + settle(code); + }; + const onAbort = () => finish(null); + signal?.addEventListener("abort", onAbort, { once: true }); + + const sendPage = (response: import("node:http").ServerResponse, status: number, html: string) => { + response.statusCode = status; + response.setHeader("content-type", "text/html; charset=utf-8"); + response.end(html); + }; + + const server = _http.createServer((request, response) => { + const url = new URL(request.url ?? "/", REDIRECT_URI); + if (url.pathname !== CALLBACK_PATH) { + sendPage(response, 404, oauthErrorHtml("Callback route not found.")); + return; + } + if (url.searchParams.get("state") !== expectedState) { + sendPage(response, 400, oauthErrorHtml("OAuth state mismatch.")); + return; + } + + const error = url.searchParams.get("error"); + if (error) { + sendPage(response, 400, oauthErrorHtml(url.searchParams.get("error_description") ?? error)); + finish(null); + return; + } + + const code = url.searchParams.get("code"); + if (!code) { + sendPage(response, 400, oauthErrorHtml("Missing authorization code.")); + return; + } + + sendPage(response, 200, oauthSuccessHtml("Signed in to Radius. You may now close this page.")); + finish(code); + }); + + return new Promise((resolve) => { + server + .listen(CALLBACK_PORT, CALLBACK_HOST, () => { + resolve({ + waitForCode: () => wait, + close: () => { + finish(null); + server.close(); + }, + }); + }) + .once("error", () => { + finish(null); + resolve({ waitForCode: async () => null, close: () => {} }); + }); + }); +} + +async function loginWithBrowser(oauth: RadiusOAuthConfig, callbacks: OAuthLoginCallbacks): Promise { + const { verifier, challenge } = await generatePKCE(); + const state = crypto.randomUUID(); + const authorizeUrl = new URL(oauth.authorizationEndpoint); + authorizeUrl.search = new URLSearchParams({ + response_type: "code", + client_id: oauth.clientId, + redirect_uri: REDIRECT_URI, + scope: oauth.scope, + code_challenge: challenge, + code_challenge_method: "S256", + handoff: "url", + state, + }).toString(); + + const callbackServer = await startOAuthCallbackServer(state, callbacks.signal); + callbacks.onProgress?.(`Listening for OAuth callback on ${REDIRECT_URI}`); + callbacks.onAuth({ + url: authorizeUrl.toString(), + instructions: "Continue in your browser.", + }); + + try { + const code = await callbackServer.waitForCode(); + if (!code) { + if (callbacks.signal?.aborted) { + throw new Error("Login cancelled"); + } + throw new Error("OAuth callback did not complete."); + } + return await requestOAuthToken( + oauth, + new URLSearchParams({ + grant_type: "authorization_code", + client_id: oauth.clientId, + redirect_uri: REDIRECT_URI, + code, + code_verifier: verifier, + }), + callbacks.signal, + ); + } finally { + callbackServer.close(); + } +} + +async function requestDeviceAuthorization( + oauth: RadiusOAuthConfig, + signal: AbortSignal | undefined, +): Promise { + let response: Response; + try { + response = await fetch(oauth.deviceAuthorizationEndpoint, { + method: "POST", + headers: { accept: "application/json", "content-type": "application/x-www-form-urlencoded" }, + body: new URLSearchParams({ client_id: oauth.clientId, scope: oauth.scope }), + signal, + }); + } catch (error) { + if (signal?.aborted) { + throw new Error("Login cancelled"); + } + throw error; + } + + if (!response.ok) { + throw await readOAuthResponseError(response, "Radius OAuth device authorization failed"); + } + + const data = (await response.json()) as Partial; + if (!data.device_code || !data.user_code || !data.expires_in) { + throw new Error("Radius OAuth device authorization response is missing required fields"); + } + + return { + device_code: data.device_code, + user_code: data.user_code, + verification_uri: data.verification_uri, + verification_uri_complete: data.verification_uri_complete, + expires_in: data.expires_in, + interval: data.interval, + }; +} + +async function loginWithDeviceCode( + oauth: RadiusOAuthConfig, + callbacks: OAuthLoginCallbacks, +): Promise { + const device = await requestDeviceAuthorization(oauth, callbacks.signal); + callbacks.onDeviceCode({ + userCode: device.user_code, + verificationUri: device.verification_uri || oauth.verificationEndpoint, + intervalSeconds: device.interval, + expiresInSeconds: device.expires_in, + }); + + return pollOAuthDeviceCodeFlow({ + intervalSeconds: device.interval, + expiresInSeconds: device.expires_in, + signal: callbacks.signal, + poll: async () => { + try { + const credentials = await requestOAuthToken( + oauth, + new URLSearchParams({ + grant_type: oauth.deviceCodeGrantType, + client_id: oauth.clientId, + device_code: device.device_code, + }), + callbacks.signal, + ); + return { status: "complete", value: credentials }; + } catch (error) { + if (!(error instanceof OAuthResponseError)) { + throw error; + } + switch (error.oauthError) { + case "authorization_pending": + return { status: "pending" }; + case "slow_down": + return { status: "slow_down" }; + case "expired_token": + return { status: "failed", message: "Device authorization expired." }; + case "access_denied": + return { status: "failed", message: "Device authorization was denied." }; + default: + throw error; + } + } + }, + }); +} + +async function attachGatewayConfig( + gateway: string, + credentials: OAuthCredentials, + previous?: OAuthCredentials, +): Promise { + try { + const config = await loadRadiusGatewayConfig(gateway, credentials.access); + return { ...credentials, gatewayConfig: config }; + } catch (error) { + // Keep the previous catalog so models do not vanish on transient + // config failures; the next token refresh retries. + const previousConfig = getRadiusCredentialConfig(previous); + if (previousConfig) { + return { ...credentials, gatewayConfig: previousConfig }; + } + // No catalog to retain (e.g. initial login): fail loudly instead of + // completing a sign-in that would register no models. + throw error; + } +} + +export interface RadiusOAuthProviderOptions { + id: string; + name: string; + gateway: string; +} + +export function createRadiusOAuthProvider(options: RadiusOAuthProviderOptions): OAuthProviderInterface { + const gateway = normalizeRadiusGatewayUrl(options.gateway); + + return { + id: options.id, + name: options.name, + + async login(callbacks: OAuthLoginCallbacks): Promise { + const oauth = await loadRadiusOAuthConfig(gateway); + const loginMethod = await callbacks.onSelect({ + message: `Sign in to ${options.name}:`, + options: [ + { id: LOGIN_METHOD_BROWSER, label: "Sign in with browser (recommended)" }, + { + id: LOGIN_METHOD_DEVICE_CODE, + label: "Sign in with device code (when signing in from another device)", + }, + ], + }); + if (!loginMethod) { + throw new Error("Login cancelled"); + } + + let credentials: OAuthCredentials; + if (loginMethod === LOGIN_METHOD_DEVICE_CODE) { + credentials = await loginWithDeviceCode(oauth, callbacks); + } else if (loginMethod === LOGIN_METHOD_BROWSER) { + credentials = await loginWithBrowser(oauth, callbacks); + } else { + throw new Error(`Unknown ${options.name} sign-in method: ${loginMethod}`); + } + + return attachGatewayConfig(gateway, credentials); + }, + + async refreshToken(credentials: OAuthCredentials): Promise { + const oauth = await loadRadiusOAuthConfig(gateway); + const refreshed = await requestOAuthToken( + oauth, + new URLSearchParams({ + grant_type: "refresh_token", + client_id: oauth.clientId, + refresh_token: credentials.refresh, + }), + ); + return attachGatewayConfig(gateway, refreshed, credentials); + }, + + getApiKey(credentials: OAuthCredentials): string { + return credentials.access; + }, + + modifyModels(models: Model[], credentials: OAuthCredentials): Model[] { + const config = getRadiusCredentialConfig(credentials); + if (!config) { + return models; + } + + // Keep models already registered for this provider (e.g. models.json + // custom entries) and add catalog models that are not present. + const existingIds = new Set(models.filter((model) => model.provider === options.id).map((model) => model.id)); + const added = config.models + .filter((model) => !existingIds.has(model.id)) + .map( + (model) => + ({ + ...model, + api: "pi-messages", + provider: options.id, + baseUrl: config.baseUrl, + }) as Model, + ); + + return [...models, ...added]; + }, + }; +} diff --git a/packages/ai/test/pi-messages.test.ts b/packages/ai/test/pi-messages.test.ts new file mode 100644 index 00000000..a689c9f2 --- /dev/null +++ b/packages/ai/test/pi-messages.test.ts @@ -0,0 +1,243 @@ +import { createServer, type IncomingMessage, type Server, type ServerResponse } from "node:http"; +import type { AddressInfo } from "node:net"; +import { afterEach, describe, expect, it } from "vitest"; +import { type PiMessagesOptions, stream, streamSimple } from "../src/api/pi-messages.ts"; +import type { Api, AssistantMessageEvent, Context, Model } from "../src/types.ts"; + +type RecordedRequest = { + url: string; + headers: IncomingMessage["headers"]; + body: unknown; +}; + +type ResponderOptions = { + status?: number; + headers?: Record; + events?: unknown[]; + rawBody?: string; +}; + +let server: Server | undefined; + +afterEach(() => { + server?.close(); + server = undefined; +}); + +async function startServer(options: ResponderOptions): Promise<{ baseUrl: string; requests: RecordedRequest[] }> { + const requests: RecordedRequest[] = []; + + server = createServer((request: IncomingMessage, response: ServerResponse) => { + const chunks: Buffer[] = []; + request.on("data", (chunk: Buffer) => chunks.push(chunk)); + request.on("end", () => { + const raw = Buffer.concat(chunks).toString("utf-8"); + requests.push({ + url: request.url ?? "", + headers: request.headers, + body: raw ? JSON.parse(raw) : undefined, + }); + + if (options.status && options.status !== 200) { + response.statusCode = options.status; + response.setHeader("content-type", "application/json"); + response.end(options.rawBody ?? "{}"); + return; + } + + response.statusCode = 200; + response.setHeader("content-type", "text/event-stream"); + for (const [name, value] of Object.entries(options.headers ?? {})) { + response.setHeader(name, value); + } + for (const event of options.events ?? []) { + response.write(`data: ${JSON.stringify(event)}\n\n`); + } + response.end(); + }); + }); + + await new Promise((resolve) => { + server!.listen(0, "127.0.0.1", () => resolve()); + }); + + const address = server!.address() as AddressInfo; + return { baseUrl: `http://127.0.0.1:${address.port}/v1`, requests }; +} + +function createModel(baseUrl: string): Model<"pi-messages"> { + return { + id: "auto", + name: "Radius Auto", + api: "pi-messages", + provider: "radius", + baseUrl, + reasoning: false, + input: ["text"], + cost: { input: 1, output: 2, cacheRead: 0.1, cacheWrite: 0.2 }, + contextWindow: 128000, + maxTokens: 16384, + }; +} + +const context: Context = { + messages: [{ role: "user", content: "Hello", timestamp: Date.now() }], +}; + +const usage = { + input: 10, + output: 5, + cacheRead: 0, + cacheWrite: 0, + totalTokens: 15, + cost: { input: 0.1, output: 0.2, cacheRead: 0, cacheWrite: 0, total: 0.3 }, +}; + +describe("pi-messages", () => { + it("streams text and tool calls and resolves the terminal message", async () => { + const { baseUrl, requests } = await startServer({ + events: [ + { type: "start" }, + { type: "text_start", contentIndex: 0 }, + { type: "text_delta", contentIndex: 0, delta: "Hel" }, + { type: "text_delta", contentIndex: 0, delta: "lo" }, + { type: "text_end", contentIndex: 0, content: "Hello" }, + { type: "toolcall_start", contentIndex: 1, id: "call_1", toolName: "read" }, + { type: "toolcall_delta", contentIndex: 1, delta: '{"path":' }, + { type: "toolcall_delta", contentIndex: 1, delta: '"a.txt"}' }, + { + type: "toolcall_end", + contentIndex: 1, + toolCall: { type: "toolCall", id: "call_1", name: "read", arguments: { path: "a.txt" } }, + }, + { type: "done", reason: "toolUse", usage, responseId: "resp_1" }, + ], + }); + const model = createModel(baseUrl); + + const events: AssistantMessageEvent[] = []; + const eventStream = stream(model, context, { + apiKey: "test-key", + sessionId: "session-1", + toolChoice: "auto", + maxTokens: 100, + headers: { "x-custom": "1" }, + }); + for await (const event of eventStream) { + events.push(event); + } + const message = await eventStream.result(); + + expect(message.stopReason).toBe("toolUse"); + expect(message.usage).toEqual(usage); + expect(message.responseId).toBe("resp_1"); + expect(message.model).toBe("auto"); + expect(message.provider).toBe("radius"); + expect(message.content).toEqual([ + { type: "text", text: "Hello", textSignature: undefined }, + { type: "toolCall", id: "call_1", name: "read", arguments: { path: "a.txt" } }, + ]); + expect(events.some((event) => event.type === "text_delta")).toBe(true); + expect(events.filter((event) => event.type === "toolcall_end")).toHaveLength(1); + + expect(requests).toHaveLength(1); + const request = requests[0]; + expect(request.url).toBe("/v1/messages"); + expect(request.headers.authorization).toBe("Bearer test-key"); + expect(request.headers["x-custom"]).toBe("1"); + expect(request.body).toEqual({ + model: "auto", + context, + options: { maxTokens: 100, sessionId: "session-1", toolChoice: "auto" }, + }); + }); + + it("appends debug=1 and reports response headers via onResponse", async () => { + const { baseUrl, requests } = await startServer({ + headers: { "x-pi-gateway-upstream-provider": "anthropic" }, + events: [{ type: "done", reason: "stop", usage }], + }); + const model = createModel(baseUrl); + + let observedHeaders: Record | undefined; + const options: PiMessagesOptions = { + apiKey: "test-key", + debug: true, + onResponse: (response) => { + observedHeaders = response.headers; + }, + }; + const message = await streamSimple(model, context, options).result(); + + expect(message.stopReason).toBe("stop"); + expect(requests[0].url).toBe("/v1/messages?debug=1"); + expect(observedHeaders?.["x-pi-gateway-upstream-provider"]).toBe("anthropic"); + }); + + it("surfaces backend error responses with diagnostics", async () => { + const { baseUrl } = await startServer({ + status: 401, + rawBody: JSON.stringify({ error: { message: "Token expired", code: "unauthorized" } }), + }); + const model = createModel(baseUrl); + + const message = await stream(model, context, { apiKey: "stale" }).result(); + + expect(message.stopReason).toBe("error"); + expect(message.errorMessage).toContain("401"); + expect(message.errorMessage).toContain("Token expired"); + expect(message.errorMessage).toContain("unauthorized"); + expect(message.diagnostics?.[0]?.type).toBe("pi_messages_response_failure"); + expect(message.diagnostics?.[0]?.details?.status).toBe(401); + }); + + it("propagates server-sent error events", async () => { + const { baseUrl } = await startServer({ + events: [{ type: "start" }, { type: "error", reason: "error", usage, errorMessage: "Upstream failed" }], + }); + const model = createModel(baseUrl); + + const message = await stream(model, context, { apiKey: "test-key" }).result(); + + expect(message.stopReason).toBe("error"); + expect(message.errorMessage).toBe("Upstream failed"); + expect(message.usage).toEqual(usage); + }); + + it("errors when no API key is provided", async () => { + const model = createModel("http://127.0.0.1:1/v1"); + + const message = await stream(model, context).result(); + + expect(message.stopReason).toBe("error"); + expect(message.errorMessage).toContain("No API key provided"); + }); + + it("errors when the stream ends without a terminal event", async () => { + const { baseUrl } = await startServer({ + events: [ + { type: "start" }, + { type: "text_start", contentIndex: 0 }, + { type: "text_delta", contentIndex: 0, delta: "partial" }, + ], + }); + const model = createModel(baseUrl); + + const message = await stream(model, context, { apiKey: "test-key" }).result(); + + expect(message.stopReason).toBe("error"); + expect(message.errorMessage).toContain("stream ended without a terminal event"); + }); +}); + +describe("pi-messages api registration", () => { + it("is registered as a builtin api provider", async () => { + const { getApiProvider } = await import("../src/compat.ts"); + expect(getApiProvider("pi-messages")).toBeDefined(); + }); + + it("is a known api usable on models", () => { + const api: Api = "pi-messages"; + expect(api).toBe("pi-messages"); + }); +}); diff --git a/packages/coding-agent/src/core/model-registry.ts b/packages/coding-agent/src/core/model-registry.ts index cb84fe5d..450ea15e 100644 --- a/packages/coding-agent/src/core/model-registry.ts +++ b/packages/coding-agent/src/core/model-registry.ts @@ -6,10 +6,10 @@ import { type AnthropicMessagesCompat, type Api, type AssistantMessageEventStream, + type BuiltinProvider, type Context, getModels, getProviders, - type KnownProvider, type Model, type OAuthProviderInterface, type OpenAICompletionsCompat, @@ -29,6 +29,7 @@ import { stripJsonComments } from "../utils/json.ts"; import { normalizePath } from "../utils/paths.ts"; import type { AuthStatus, AuthStorage } from "./auth-storage.ts"; import { BUILT_IN_PROVIDER_DISPLAY_NAMES } from "./provider-display-names.ts"; +import { registerCustomRadiusOAuthProvider } from "./radius.ts"; import { clearConfigValueCache, getConfigValueEnvVarNames, @@ -224,6 +225,9 @@ const ProviderConfigSchema = Type.Object({ baseUrl: Type.Optional(Type.String({ minLength: 1 })), apiKey: Type.Optional(Type.String({ minLength: 1 })), api: Type.Optional(Type.String({ minLength: 1 })), + /** OAuth flavor spoken by this provider's endpoint. Registers a sign-in + * provider with a dynamic model catalog (e.g. a custom Radius gateway). */ + oauth: Type.Optional(Type.Literal("radius")), headers: Type.Optional(Type.Record(Type.String(), Type.String())), compat: Type.Optional(ProviderCompatSchema), authHeader: Type.Optional(Type.Boolean()), @@ -453,7 +457,7 @@ export class ModelRegistry { modelOverrides: Map>, ): Model[] { return getProviders().flatMap((provider) => { - const models = getModels(provider as KnownProvider) as Model[]; + const models = getModels(provider as BuiltinProvider) as Model[]; const providerOverride = overrides.get(provider); const perModelOverrides = modelOverrides.get(provider); @@ -537,6 +541,12 @@ export class ModelRegistry { }); } + if (providerConfig.oauth === "radius") { + // Must run before the modifyModels loop in loadModels() so the + // credential-cached catalog is injected on this load. + registerCustomRadiusOAuthProvider(providerName, providerConfig.name, providerConfig.baseUrl!); + } + this.storeProviderRequestConfig(providerName, providerConfig); if (providerConfig.modelOverrides) { @@ -568,7 +578,11 @@ export class ModelRegistry { const hasModelOverrides = providerConfig.modelOverrides && Object.keys(providerConfig.modelOverrides).length > 0; - if (models.length === 0) { + if (providerConfig.oauth && !providerConfig.baseUrl) { + throw new Error(`Provider ${providerName}: "baseUrl" is required when "oauth" is set.`); + } + + if (models.length === 0 && !providerConfig.oauth) { // Override-only config: needs baseUrl, headers, compat, modelOverrides, or some combination. if (!providerConfig.baseUrl && !providerConfig.headers && !providerConfig.compat && !hasModelOverrides) { throw new Error( @@ -614,7 +628,7 @@ export class ModelRegistry { const getBuiltInDefaults = (providerName: string): { api: string; baseUrl: string } | undefined => { if (!builtInProviders.has(providerName)) return undefined; if (builtInDefaultsCache.has(providerName)) return builtInDefaultsCache.get(providerName); - const builtIn = getModels(providerName as KnownProvider) as Model[]; + const builtIn = getModels(providerName as BuiltinProvider) as Model[]; if (builtIn.length === 0) return undefined; const defaults = { api: builtIn[0].api, baseUrl: builtIn[0].baseUrl }; builtInDefaultsCache.set(providerName, defaults); diff --git a/packages/coding-agent/src/core/model-resolver.ts b/packages/coding-agent/src/core/model-resolver.ts index 3a3341df..c23fe59b 100644 --- a/packages/coding-agent/src/core/model-resolver.ts +++ b/packages/coding-agent/src/core/model-resolver.ts @@ -18,6 +18,7 @@ export const defaultModelPerProvider: Record = { openai: "gpt-5.5", "azure-openai-responses": "gpt-5.4", "openai-codex": "gpt-5.5", + radius: "auto", nvidia: "nvidia/nemotron-3-super-120b-a12b", deepseek: "deepseek-v4-pro", google: "gemini-3.1-pro-preview", diff --git a/packages/coding-agent/src/core/provider-display-names.ts b/packages/coding-agent/src/core/provider-display-names.ts index d33c3d7d..fd29d713 100644 --- a/packages/coding-agent/src/core/provider-display-names.ts +++ b/packages/coding-agent/src/core/provider-display-names.ts @@ -23,6 +23,7 @@ export const BUILT_IN_PROVIDER_DISPLAY_NAMES: Record = { "opencode-go": "OpenCode Go", openai: "OpenAI", openrouter: "OpenRouter", + radius: "Radius", together: "Together AI", "vercel-ai-gateway": "Vercel AI Gateway", xai: "xAI", diff --git a/packages/coding-agent/src/core/radius.ts b/packages/coding-agent/src/core/radius.ts new file mode 100644 index 00000000..93ac3301 --- /dev/null +++ b/packages/coding-agent/src/core/radius.ts @@ -0,0 +1,34 @@ +/** + * Radius (pi-messages gateway) provider wiring. + * + * The main Radius provider is a built-in OAuth provider in pi-ai; models are + * dynamic, cached on the stored OAuth credential (`gatewayConfig`) and + * injected via the OAuth provider's `modifyModels` hook, so startup, /reload, + * and registry refreshes work without network access. The catalog refreshes + * on login and on every token refresh. + * + * Additional gateways (e.g. a local dev gateway) can be declared in + * models.json with `"oauth": "radius"`; each entry is an independent Radius + * instance with its own credentials and catalog. + */ + +import { createRadiusOAuthProvider, registerOAuthProvider } from "@earendil-works/pi-ai/oauth"; + +export const RADIUS_PROVIDER_ID = "radius"; + +/** + * Register a Radius-style OAuth provider for a custom gateway declared in + * models.json (`"oauth": "radius"`). Runs on every models.json load so the + * registration survives `resetOAuthProviders()` during registry refreshes. + */ +export function registerCustomRadiusOAuthProvider(id: string, name: string | undefined, gateway: string): void { + registerOAuthProvider( + createRadiusOAuthProvider({ + id, + name: name ?? id, + // Tolerate an API base URL: the gateway root is what the OAuth and + // config discovery endpoints hang off. + gateway: gateway.replace(/\/v1\/?$/u, ""), + }), + ); +} diff --git a/packages/coding-agent/test/radius.test.ts b/packages/coding-agent/test/radius.test.ts new file mode 100644 index 00000000..99ba8db1 --- /dev/null +++ b/packages/coding-agent/test/radius.test.ts @@ -0,0 +1,146 @@ +import { existsSync, mkdirSync, rmSync, writeFileSync } from "node:fs"; +import { tmpdir } from "node:os"; +import { join } from "node:path"; +import { getOAuthProvider, resetOAuthProviders } from "@earendil-works/pi-ai/oauth"; +import { afterEach, beforeEach, describe, expect, it } from "vitest"; +import { AuthStorage } from "../src/core/auth-storage.ts"; +import { ModelRegistry } from "../src/core/model-registry.ts"; +import { RADIUS_PROVIDER_ID } from "../src/core/radius.ts"; + +function radiusOAuthCredential(gatewayBaseUrl: string) { + return { + type: "oauth" as const, + access: "access-token", + refresh: "refresh-token", + expires: Date.now() + 60 * 60 * 1000, + gatewayConfig: { + baseUrl: gatewayBaseUrl, + models: [ + { + id: "auto", + name: "Radius Auto", + reasoning: false, + input: ["text"], + cost: { input: 1, output: 2, cacheRead: 0.1, cacheWrite: 0.2 }, + contextWindow: 128000, + maxTokens: 16384, + }, + { + id: "byok/gpt-5.5", + name: "GPT-5.5 (BYOK)", + reasoning: true, + input: ["text", "image"], + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 }, + contextWindow: 256000, + maxTokens: 32000, + }, + ], + }, + }; +} + +let tempDir: string; + +beforeEach(() => { + tempDir = join(tmpdir(), `pi-test-radius-${Date.now()}-${Math.random().toString(36).slice(2)}`); + mkdirSync(tempDir, { recursive: true }); + resetOAuthProviders(); +}); + +afterEach(() => { + if (tempDir && existsSync(tempDir)) { + rmSync(tempDir, { recursive: true }); + } + resetOAuthProviders(); +}); + +describe("radius oauth provider", () => { + it("is registered as a built-in OAuth provider", () => { + expect(getOAuthProvider(RADIUS_PROVIDER_ID)?.name).toBe("Radius"); + }); +}); + +describe("radius models via ModelRegistry", () => { + it("injects catalog models from the stored credential", () => { + const registry = ModelRegistry.inMemory( + AuthStorage.inMemory({ radius: radiusOAuthCredential("https://radius.example.com/v1") }), + ); + + const auto = registry.find(RADIUS_PROVIDER_ID, "auto"); + expect(auto).toBeDefined(); + expect(auto?.api).toBe("pi-messages"); + expect(auto?.baseUrl).toBe("https://radius.example.com/v1"); + expect(auto?.name).toBe("Radius Auto"); + + // byok ids are registered verbatim + const byok = registry.find(RADIUS_PROVIDER_ID, "byok/gpt-5.5"); + expect(byok).toBeDefined(); + expect(byok?.contextWindow).toBe(256000); + + expect(registry.hasConfiguredAuth(auto!)).toBe(true); + expect(registry.getProviderDisplayName(RADIUS_PROVIDER_ID)).toBe("Radius"); + }); + + it("exposes no radius models without credentials", () => { + const registry = ModelRegistry.inMemory(AuthStorage.inMemory()); + + expect(registry.getAll().filter((model) => model.provider === RADIUS_PROVIDER_ID)).toHaveLength(0); + expect(getOAuthProvider(RADIUS_PROVIDER_ID)).toBeDefined(); + }); + + it("keeps radius models across registry refresh", () => { + const registry = ModelRegistry.inMemory( + AuthStorage.inMemory({ radius: radiusOAuthCredential("https://radius.example.com/v1") }), + ); + + registry.refresh(); + + expect(registry.find(RADIUS_PROVIDER_ID, "auto")).toBeDefined(); + }); +}); + +describe("custom radius gateways via models.json", () => { + function createRegistry(providers: Record, authStorage: AuthStorage): ModelRegistry { + const modelsJsonPath = join(tempDir, "models.json"); + writeFileSync(modelsJsonPath, JSON.stringify({ providers })); + return ModelRegistry.create(authStorage, modelsJsonPath); + } + + it("registers an independent radius-style provider", () => { + const registry = createRegistry( + { "radius-dev": { name: "Radius (dev)", baseUrl: "http://localhost:8788", oauth: "radius" } }, + AuthStorage.inMemory({ "radius-dev": radiusOAuthCredential("http://localhost:8788/v1") }), + ); + + expect(registry.getError()).toBeUndefined(); + expect(getOAuthProvider("radius-dev")?.name).toBe("Radius (dev)"); + expect(getOAuthProvider(RADIUS_PROVIDER_ID)?.name).toBe("Radius"); + + // Dev gateway models are injected under the custom provider id only. + const devAuto = registry.find("radius-dev", "auto"); + expect(devAuto).toBeDefined(); + expect(devAuto?.api).toBe("pi-messages"); + expect(devAuto?.baseUrl).toBe("http://localhost:8788/v1"); + expect(registry.find(RADIUS_PROVIDER_ID, "auto")).toBeUndefined(); + + expect(registry.getProviderDisplayName("radius-dev")).toBe("Radius (dev)"); + }); + + it("survives registry refresh", () => { + const registry = createRegistry( + { "radius-dev": { baseUrl: "http://localhost:8788", oauth: "radius" } }, + AuthStorage.inMemory({ "radius-dev": radiusOAuthCredential("http://localhost:8788/v1") }), + ); + + registry.refresh(); + + expect(getOAuthProvider("radius-dev")).toBeDefined(); + expect(registry.find("radius-dev", "auto")).toBeDefined(); + }); + + it("requires baseUrl when oauth is set", () => { + const registry = createRegistry({ "radius-dev": { oauth: "radius" } }, AuthStorage.inMemory()); + + expect(registry.getError()).toContain('"baseUrl" is required when "oauth" is set'); + }); +}); diff --git a/test.sh b/test.sh index b4ea70c7..f2ff5f68 100755 --- a/test.sh +++ b/test.sh @@ -55,6 +55,9 @@ unset XIAOMI_API_KEY unset XIAOMI_TOKEN_PLAN_CN_API_KEY unset XIAOMI_TOKEN_PLAN_AMS_API_KEY unset XIAOMI_TOKEN_PLAN_SGP_API_KEY +unset PI_GATEWAY_API_KEY +unset PI_GATEWAY +unset PI_EXPERIMENTAL unset COPILOT_GITHUB_TOKEN unset GH_TOKEN unset GITHUB_TOKEN