feat(ai): add Radius gateway support

This commit is contained in:
Armin Ronacher
2026-07-14 11:01:21 +02:00
parent 0e6909f050
commit 961fa6c142
16 changed files with 1483 additions and 11 deletions
+4
View File
@@ -0,0 +1,4 @@
import type { ProviderStreams } from "../types.ts";
import { lazyApi } from "./lazy.ts";
export const piMessagesApi = (): ProviderStreams => lazyApi(() => import("./pi-messages.ts"));
+436
View File
@@ -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 `<baseUrl>/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<PiMessagesStopReason, "stop" | "length" | "toolUse">;
usage: PiMessagesUsage;
responseId?: string;
rewrite?: PiMessagesRewriteImpact;
}
| {
type: "error";
reason: Extract<PiMessagesStopReason, "aborted" | "error">;
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<string, unknown>;
constructor(message: string, code: string | undefined, diagnosticDetails: Record<string, unknown>) {
super(message);
this.name = "PiMessagesResponseError";
this.code = code;
this.diagnosticDetails = diagnosticDetails;
}
}
function isRecord(value: unknown): value is Record<string, unknown> {
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<number, string>();
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<ToolCall["arguments"]>(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<Uint8Array>): AsyncGenerator<PiMessagesEvent> {
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,
});
};
+6
View File
@@ -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<Api, ReturnType<typeof getApiProvider>>();
+1
View File
@@ -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",
+1
View File
@@ -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";
+11 -6
View File
@@ -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<TProvider extends KnownProvider, TModelId extends keyof (typeof MODELS)[TProvider]>(
export function getBuiltinModel<TProvider extends BuiltinProvider, TModelId extends keyof (typeof MODELS)[TProvider]>(
provider: TProvider,
modelId: TModelId,
): Model<BuiltinModelApi<TProvider, TModelId>> {
@@ -53,11 +58,11 @@ export function getBuiltinModel<TProvider extends KnownProvider, TModelId extend
return models?.[modelId as string] as Model<BuiltinModelApi<TProvider, TModelId>>;
}
export function getBuiltinProviders(): KnownProvider[] {
return Object.keys(MODELS) as KnownProvider[];
export function getBuiltinProviders(): BuiltinProvider[] {
return Object.keys(MODELS) as BuiltinProvider[];
}
export function getBuiltinModels<TProvider extends KnownProvider>(
export function getBuiltinModels<TProvider extends BuiltinProvider>(
provider: TProvider,
): Model<BuiltinModelApi<TProvider, keyof (typeof MODELS)[TProvider]>>[] {
const models = MODELS[provider] as Record<string, Model<Api>> | undefined;
+5 -1
View File
@@ -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;
}
/**
+16
View File
@@ -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<string, OAuthProviderInterface>(
+557
View File
@@ -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<Api>["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<string, unknown> {
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<RadiusGatewayConfig> {
const headers: Record<string, string> = { 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<RadiusOAuthConfig> {
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<OAuthResponseError> {
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<OAuthCredentials> {
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<string | null>;
close(): void;
};
function startOAuthCallbackServer(
expectedState: string,
signal: AbortSignal | undefined,
): Promise<OAuthCallbackServer> {
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<string | null>((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<OAuthCredentials> {
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<DeviceAuthorizationResponse> {
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<DeviceAuthorizationResponse>;
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<OAuthCredentials> {
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<OAuthCredentials>({
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<RadiusOAuthCredentials> {
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<OAuthCredentials> {
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<OAuthCredentials> {
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<Api>[], credentials: OAuthCredentials): Model<Api>[] {
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<Api>,
);
return [...models, ...added];
},
};
}
+243
View File
@@ -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<string, string>;
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<void>((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<string, string> | 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");
});
});
@@ -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<string, Map<string, ModelOverride>>,
): Model<Api>[] {
return getProviders().flatMap((provider) => {
const models = getModels(provider as KnownProvider) as Model<Api>[];
const models = getModels(provider as BuiltinProvider) as Model<Api>[];
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<Api>[];
const builtIn = getModels(providerName as BuiltinProvider) as Model<Api>[];
if (builtIn.length === 0) return undefined;
const defaults = { api: builtIn[0].api, baseUrl: builtIn[0].baseUrl };
builtInDefaultsCache.set(providerName, defaults);
@@ -18,6 +18,7 @@ export const defaultModelPerProvider: Record<KnownProvider, string> = {
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",
@@ -23,6 +23,7 @@ export const BUILT_IN_PROVIDER_DISPLAY_NAMES: Record<string, string> = {
"opencode-go": "OpenCode Go",
openai: "OpenAI",
openrouter: "OpenRouter",
radius: "Radius",
together: "Together AI",
"vercel-ai-gateway": "Vercel AI Gateway",
xai: "xAI",
+34
View File
@@ -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, ""),
}),
);
}
+146
View File
@@ -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<string, unknown>, 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');
});
});
+3
View File
@@ -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