diff --git a/packages/ai/.changes/prime-inference-model-catalog.md b/packages/ai/.changes/prime-inference-model-catalog.md new file mode 100644 index 0000000000..3470623079 --- /dev/null +++ b/packages/ai/.changes/prime-inference-model-catalog.md @@ -0,0 +1 @@ +- Added live Prime Inference model names, pricing, limits, modalities, and reasoning support to the bundled catalog. diff --git a/packages/ai/scripts/generate-models.ts b/packages/ai/scripts/generate-models.ts index 04c09654c7..5a942ab3a7 100644 --- a/packages/ai/scripts/generate-models.ts +++ b/packages/ai/scripts/generate-models.ts @@ -1,12 +1,16 @@ #!/usr/bin/env tsx -import { readFileSync, writeFileSync } from "fs"; -import { homedir } from "os"; +import { writeFileSync } from "fs"; import { dirname, join } from "path"; import { fileURLToPath } from "url"; import { getAnthropicCacheCosts } from "../src/cache-pricing.js"; import { COPILOT_CLIENT_HEADERS } from "../src/copilot-client-version.js"; import { getOpenRouterReasoningCapabilities } from "../src/openrouter-reasoning.js"; +import { + isPrivatePrimeInferenceModelId, + parsePrimeInferenceModelCatalog, + type PrimeInferenceCatalogEntry, +} from "../src/prime-inference-model-catalog.js"; import { CLOUDFLARE_AI_GATEWAY_ANTHROPIC_BASE_URL, CLOUDFLARE_AI_GATEWAY_COMPAT_BASE_URL, @@ -21,6 +25,7 @@ import { type OpenAICompletionsCompat, } from "../src/types.js"; import { MODELS as EXISTING_MODELS } from "../src/models.generated.js"; +import { renderModelsFile } from "./render-models.js"; const __filename = fileURLToPath(import.meta.url); const __dirname = dirname(__filename); @@ -116,15 +121,6 @@ const PRIME_INFERENCE_COMPAT: OpenAICompletionsCompat = { maxTokensField: "max_tokens", supportsStrictMode: false, }; -interface PrimeInferenceCatalogEntry { - id: string; - input: number; - output: number; - contextWindow?: number; - maxTokens?: number; - reasoning?: boolean; -} - interface PrimeInferenceModelMetadata { contextWindow?: number; maxTokens?: number; @@ -132,13 +128,9 @@ interface PrimeInferenceModelMetadata { name?: string; } -// The full Prime Inference catalog is registered (minus raw/duplicate variants). -// Prime's /models endpoint publishes pricing only, so context/output limits and -// modalities are read from OpenRouter's public catalog, used here purely as a -// published spec sheet for the same upstream models — requests always go to -// Prime's own baseUrl. Entries below override those specs where the Prime route -// enforces a different limit (verified against the live API) or fill gaps for -// models OpenRouter does not list or leaves incomplete. +// Prime's /models endpoint is authoritative for route metadata. OpenRouter and +// these overrides only fill gaps for older or incomplete endpoint entries; +// requests always go to Prime's own baseUrl. const PRIME_INFERENCE_MODEL_METADATA: Record = { // These routes accept 200k, checked against the live API 2026-07-08. The // other Claude routes take the full window their spec lists. @@ -218,22 +210,6 @@ const PRIME_INFERENCE_OPENROUTER_ALIASES: Record = { const PRIME_INFERENCE_DEFAULT_CONTEXT_WINDOW = 128000; const PRIME_INFERENCE_DEFAULT_MAX_TOKENS = 8192; -// Raw checkpoints and duplicate routes that would clutter the picker: BF16 -// exports, fine-tune outputs, zai-org/ and HF-cased twins of canonical ids. -function isPrimeInferenceRawVariant(modelId: string): boolean { - const id = modelId.toLowerCase(); - if (id.endsWith("-bf16") || id.includes(":")) { - return true; - } - const vendor = modelId.split("/")[0] ?? ""; - return vendor === "zai-org" || vendor !== vendor.toLowerCase(); -} - -function isPrimeInferencePrivateModel(modelId: string): boolean { - const id = modelId.toLowerCase(); - return id.startsWith("internal/") || id.startsWith("dev/"); -} - const OPENAI_RESPONSES_NONE_REASONING_MODELS = new Set([ "gpt-5.1", "gpt-5.2", @@ -385,51 +361,6 @@ function getOptionalNumber(value: unknown): number | undefined { return typeof value === "number" && Number.isFinite(value) ? value : undefined; } -function getOptionalBoolean(value: unknown): boolean | undefined { - return typeof value === "boolean" ? value : undefined; -} - -function readPrimeCliConfig(): Record { - try { - const parsed = JSON.parse(readFileSync(join(homedir(), ".prime", "config.json"), "utf8")); - return isRecord(parsed) ? parsed : {}; - } catch { - return {}; - } -} - -function getPrimeInferenceConfigValue( - envName: "PRIME_API_KEY" | "PRIME_TEAM_ID", - config: Record, - configKeys: readonly string[], -): string | undefined { - const fromEnv = process.env[envName]?.trim(); - if (fromEnv) { - return fromEnv; - } - - for (const key of configKeys) { - const value = config[key]; - if (typeof value === "string" && value.trim()) { - return value.trim(); - } - } - - return undefined; -} - -function getPrimeInferenceHeaders(apiKey: string | undefined, teamId: string | undefined): Record | undefined { - const headers: Record = {}; - if (apiKey) { - headers.Authorization = `Bearer ${apiKey}`; - } - if (teamId) { - headers["X-Prime-Team-ID"] = teamId; - } - - return Object.keys(headers).length > 0 ? headers : undefined; -} - function getPrimeInferenceCacheCosts(modelId: string, inputCost: number): { cacheRead: number; cacheWrite: number } { return modelId.toLowerCase().startsWith("anthropic/") ? getAnthropicCacheCosts(inputCost, "5m") @@ -439,7 +370,7 @@ function getPrimeInferenceCacheCosts(modelId: string, inputCost: number): { cach function getExistingPrimeInferenceModels(): Model<"openai-completions">[] { const models = EXISTING_MODELS["prime-inference"] as unknown as Record>; return Object.values(models) - .filter((model) => !isPrimeInferenceRawVariant(model.id) && !isPrimeInferencePrivateModel(model.id)) + .filter((model) => !isPrivatePrimeInferenceModelId(model.id)) .map((model) => ({ ...model, input: [...model.input], @@ -486,20 +417,6 @@ function refreshPrimeInferenceAliasLimits( }); } -function includesCatalogCapability(value: unknown, capabilities: readonly string[]): boolean { - if (!Array.isArray(value)) { - return false; - } - - return value.some((item) => { - if (typeof item !== "string") { - return false; - } - const normalized = item.toLowerCase(); - return capabilities.some((capability) => normalized.includes(capability)); - }); -} - function getPrimeInferenceDisplayName(modelId: string): string { const rawName = modelId.split("/").at(-1) ?? modelId; return rawName @@ -513,29 +430,6 @@ function getPrimeInferenceDisplayName(modelId: string): string { .join(" "); } -function getPrimeInferenceCatalogReasoning(item: Record): boolean | undefined { - const metadata = isRecord(item.metadata) ? item.metadata : {}; - const direct = - getOptionalBoolean(item.reasoning) ?? - getOptionalBoolean(item.supports_reasoning) ?? - getOptionalBoolean(item.supportsReasoning) ?? - getOptionalBoolean(metadata.reasoning) ?? - getOptionalBoolean(metadata.supports_reasoning) ?? - getOptionalBoolean(metadata.supportsReasoning); - if (direct !== undefined) { - return direct; - } - - return includesCatalogCapability(item.supported_parameters, ["reasoning", "thinking"]) || - includesCatalogCapability(item.capabilities, ["reasoning", "thinking"]) || - includesCatalogCapability(item.tags, ["reasoning", "thinking"]) || - includesCatalogCapability(metadata.supported_parameters, ["reasoning", "thinking"]) || - includesCatalogCapability(metadata.capabilities, ["reasoning", "thinking"]) || - includesCatalogCapability(metadata.tags, ["reasoning", "thinking"]) - ? true - : undefined; -} - function isPrimeInferenceReasoningModel(modelId: string, catalogReasoning?: boolean): boolean { if (catalogReasoning !== undefined) { return catalogReasoning; @@ -572,37 +466,6 @@ function getPrimeInferenceCompat(modelId: string): OpenAICompletionsCompat { return PRIME_INFERENCE_COMPAT; } -function parsePrimeInferenceCatalog(data: unknown): PrimeInferenceCatalogEntry[] { - if (!isRecord(data) || !Array.isArray(data.data)) { - return []; - } - - return data.data.flatMap((item): PrimeInferenceCatalogEntry[] => { - if (!isRecord(item) || typeof item.id !== "string") { - return []; - } - - const pricing = isRecord(item.pricing) ? item.pricing : {}; - const input = getOptionalNumber(pricing.input_usd_per_mtok); - const output = getOptionalNumber(pricing.output_usd_per_mtok); - if (input === undefined || output === undefined) { - return []; - } - - const limit = isRecord(item.limit) ? item.limit : {}; - return [ - { - id: item.id, - input, - output, - contextWindow: getOptionalNumber(item.context_window ?? item.contextWindow ?? limit.context), - maxTokens: getOptionalNumber(item.max_tokens ?? item.maxTokens ?? limit.output), - reasoning: getPrimeInferenceCatalogReasoning(item), - }, - ]; - }); -} - interface PrimeInferenceOpenRouterMetadata { contextWindow?: number; maxTokens?: number; @@ -651,17 +514,12 @@ function getPrimeInferenceOpenRouterMetadata( } async function fetchPrimeInferenceModels(): Promise[]> { - const primeConfig = readPrimeCliConfig(); - const apiKey = getPrimeInferenceConfigValue("PRIME_API_KEY", primeConfig, ["api_key", "apiKey"]); - const teamId = getPrimeInferenceConfigValue("PRIME_TEAM_ID", primeConfig, ["team_id", "teamId", "teamID"]); let catalog: PrimeInferenceCatalogEntry[] = []; try { - console.log("Fetching models from Prime Inference API..."); - const response = await fetch(`${PRIME_INFERENCE_BASE_URL}/models`, { - headers: getPrimeInferenceHeaders(apiKey, teamId), - }); - catalog = parsePrimeInferenceCatalog(await response.json()); + console.log("Fetching public models from Prime Inference API..."); + const response = await fetch(`${PRIME_INFERENCE_BASE_URL}/models`); + catalog = parsePrimeInferenceModelCatalog(await response.json()); } catch (error) { console.error("Failed to fetch Prime Inference models:", error); } @@ -680,7 +538,7 @@ async function fetchPrimeInferenceModels(): Promise[ } const catalogModels = catalog - .filter((entry) => !isPrimeInferenceRawVariant(entry.id) && !isPrimeInferencePrivateModel(entry.id)) + .filter((entry) => !isPrivatePrimeInferenceModelId(entry.id)) .map((entry) => createPrimeInferenceModel( entry, @@ -689,6 +547,10 @@ async function fetchPrimeInferenceModels(): Promise[ ), ); let snapshotModels = getExistingPrimeInferenceModels(); + if (catalog.length > 0 && catalogModels.length < Math.ceil(snapshotModels.length * 0.5)) { + console.error("Prime Inference catalog is severely truncated; keeping snapshot models"); + return snapshotModels; + } if (catalog.length > 0) { const liveIds = new Set(catalogModels.map((model) => model.id.toLowerCase())); snapshotModels = snapshotModels.filter((model) => liveIds.has(model.id.toLowerCase())); @@ -704,8 +566,12 @@ function createPrimeInferenceModel( override: PrimeInferenceModelMetadata | undefined, openRouter: PrimeInferenceOpenRouterMetadata | undefined, ): Model<"openai-completions"> { - const vision = override?.vision ?? openRouter?.vision ?? false; - const cacheCosts = getPrimeInferenceCacheCosts(entry.id, entry.input); + const vision = entry.vision ?? override?.vision ?? openRouter?.vision ?? false; + const fallbackCacheCosts = getPrimeInferenceCacheCosts(entry.id, entry.input); + const cacheCosts = { + cacheRead: entry.cacheRead ?? fallbackCacheCosts.cacheRead, + cacheWrite: entry.cacheWrite ?? fallbackCacheCosts.cacheWrite, + }; const contextWindow = entry.contextWindow ?? override?.contextWindow ?? @@ -721,7 +587,7 @@ function createPrimeInferenceModel( return { id: entry.id, ...(PRIME_INFERENCE_FEATURED_MODELS.has(entry.id.toLowerCase()) ? { featured: true } : {}), - name: override?.name ?? getPrimeInferenceDisplayName(entry.id), + name: entry.name ?? override?.name ?? getPrimeInferenceDisplayName(entry.id), api: "openai-completions", provider: "prime-inference", baseUrl: PRIME_INFERENCE_BASE_URL, @@ -2394,7 +2260,7 @@ async function generateModels() { } // Group by provider and deduplicate by model ID - const providers: Record>> = {}; + const providers: Record>> = {}; for (const model of allModels) { if (!providers[model.provider]) { providers[model.provider] = {}; @@ -2406,63 +2272,9 @@ async function generateModels() { } } - // Generate TypeScript file - let output = `// This file is auto-generated by scripts/generate-models.ts -// Do not edit manually - run 'npm run generate-models' to update - -import type { Model } from "./types.js"; - -export const MODELS = { -`; - - // Generate provider sections (sorted for deterministic output) - const sortedProviderIds = Object.keys(providers).sort(); - for (const providerId of sortedProviderIds) { - const models = providers[providerId]; - output += `\t${JSON.stringify(providerId)}: {\n`; - - const sortedModelIds = Object.keys(models).sort(); - for (const modelId of sortedModelIds) { - const model = models[modelId]; - output += `\t\t"${model.id}": {\n`; - output += `\t\t\tid: "${model.id}",\n`; - output += `\t\t\tname: "${model.name}",\n`; - output += `\t\t\tapi: "${model.api}",\n`; - output += `\t\t\tprovider: "${model.provider}",\n`; - if (model.baseUrl !== undefined) { - output += `\t\t\tbaseUrl: "${model.baseUrl}",\n`; - } - if (model.headers) { - output += `\t\t\theaders: ${JSON.stringify(model.headers)},\n`; - } - if (model.compat) { - output += ` compat: ${JSON.stringify(model.compat)}, -`; - } - output += `\t\t\treasoning: ${model.reasoning},\n`; - if (model.thinkingLevelMap) { - output += `\t\t\tthinkingLevelMap: ${JSON.stringify(model.thinkingLevelMap)},\n`; - } - output += `\t\t\tinput: [${model.input.map(i => `"${i}"`).join(", ")}],\n`; - output += `\t\t\tcost: {\n`; - output += `\t\t\t\tinput: ${model.cost.input},\n`; - output += `\t\t\t\toutput: ${model.cost.output},\n`; - output += `\t\t\t\tcacheRead: ${model.cost.cacheRead},\n`; - output += `\t\t\t\tcacheWrite: ${model.cost.cacheWrite},\n`; - output += `\t\t\t},\n`; - output += `\t\t\tcontextWindow: ${model.contextWindow},\n`; - output += `\t\t\tmaxTokens: ${model.maxTokens},\n`; - if (model.featured) { - output += `\t\t\tfeatured: true,\n`; - } - output += `\t\t} satisfies Model<"${model.api}">,\n`; - } - - output += `\t},\n`; - } - - output += `} as const; -`; + // Generate TypeScript file. JSON string literals prevent remote catalog + // text from becoming executable source code. + const output = renderModelsFile(providers); // Write file writeFileSync(join(packageRoot, "src/models.generated.ts"), output); diff --git a/packages/ai/scripts/render-models.ts b/packages/ai/scripts/render-models.ts new file mode 100644 index 0000000000..d906425ee6 --- /dev/null +++ b/packages/ai/scripts/render-models.ts @@ -0,0 +1,49 @@ +import type { Api, Model } from "../src/types.js"; + +export function renderModelsFile(providers: Record>>): string { + let output = `// This file is auto-generated by scripts/generate-models.ts +// Do not edit manually - run 'npm run generate-models' to update + +import type { Model } from "./types.js"; + +export const MODELS = { +`; + + for (const providerId of Object.keys(providers).sort()) { + const models = providers[providerId]; + output += `\t${JSON.stringify(providerId)}: {\n`; + + for (const modelId of Object.keys(models).sort()) { + const model = models[modelId]; + output += `\t\t${JSON.stringify(model.id)}: {\n`; + output += `\t\t\tid: ${JSON.stringify(model.id)},\n`; + output += `\t\t\tname: ${JSON.stringify(model.name)},\n`; + output += `\t\t\tapi: ${JSON.stringify(model.api)},\n`; + output += `\t\t\tprovider: ${JSON.stringify(model.provider)},\n`; + if (model.baseUrl !== undefined) { + output += `\t\t\tbaseUrl: ${JSON.stringify(model.baseUrl)},\n`; + } + if (model.headers) output += `\t\t\theaders: ${JSON.stringify(model.headers)},\n`; + if (model.compat) output += `\t\t\tcompat: ${JSON.stringify(model.compat)},\n`; + output += `\t\t\treasoning: ${model.reasoning},\n`; + if (model.thinkingLevelMap) { + output += `\t\t\tthinkingLevelMap: ${JSON.stringify(model.thinkingLevelMap)},\n`; + } + output += `\t\t\tinput: [${model.input.map((input) => JSON.stringify(input)).join(", ")}],\n`; + output += `\t\t\tcost: {\n`; + output += `\t\t\t\tinput: ${model.cost.input},\n`; + output += `\t\t\t\toutput: ${model.cost.output},\n`; + output += `\t\t\t\tcacheRead: ${model.cost.cacheRead},\n`; + output += `\t\t\t\tcacheWrite: ${model.cost.cacheWrite},\n`; + output += `\t\t\t},\n`; + output += `\t\t\tcontextWindow: ${model.contextWindow},\n`; + output += `\t\t\tmaxTokens: ${model.maxTokens},\n`; + if (model.featured) output += `\t\t\tfeatured: true,\n`; + output += `\t\t} satisfies Model<${JSON.stringify(model.api)}>,\n`; + } + + output += `\t},\n`; + } + + return `${output}} as const;\n`; +} diff --git a/packages/ai/src/index.ts b/packages/ai/src/index.ts index 388dffd6e2..7b59b84597 100644 --- a/packages/ai/src/index.ts +++ b/packages/ai/src/index.ts @@ -5,6 +5,7 @@ export * from "./api-registry.js"; export * from "./env-api-keys.js"; export * from "./log.js"; export * from "./models.js"; +export * from "./prime-inference-model-catalog.js"; export type { BedrockOptions, BedrockThinkingDisplay } from "./providers/amazon-bedrock.js"; export type { AnthropicEffort, AnthropicOptions, AnthropicThinkingDisplay } from "./providers/anthropic.js"; export type { AzureOpenAIResponsesOptions } from "./providers/azure-openai-responses.js"; diff --git a/packages/ai/src/prime-inference-model-catalog.ts b/packages/ai/src/prime-inference-model-catalog.ts new file mode 100644 index 0000000000..4ca8cb8967 --- /dev/null +++ b/packages/ai/src/prime-inference-model-catalog.ts @@ -0,0 +1,89 @@ +export interface PrimeInferenceCatalogEntry { + id: string; + name?: string; + input: number; + output: number; + cacheRead?: number; + cacheWrite?: number; + contextWindow?: number; + maxTokens?: number; + vision?: boolean; + reasoning?: boolean; +} + +function isRecord(value: unknown): value is Record { + return typeof value === "object" && value !== null && !Array.isArray(value); +} + +function nonNegativeNumber(value: unknown): number | undefined { + return typeof value === "number" && Number.isFinite(value) && value >= 0 ? value : undefined; +} + +function positiveInteger(value: unknown): number | undefined { + return typeof value === "number" && Number.isInteger(value) && value > 0 ? value : undefined; +} + +export function isPrivatePrimeInferenceModelId(modelId: string): boolean { + const normalizedId = modelId.toLowerCase(); + return normalizedId.startsWith("internal/") || normalizedId.startsWith("dev/") || normalizedId.includes(":"); +} + +export function parsePrimeInferenceModelCatalog( + value: unknown, + options: { allowEmpty?: boolean } = {}, +): PrimeInferenceCatalogEntry[] { + if (!isRecord(value) || !Array.isArray(value.data)) throw new Error("Invalid Prime Inference model catalog"); + const models: PrimeInferenceCatalogEntry[] = []; + const seen = new Set(); + for (const item of value.data) { + if (!isRecord(item) || typeof item.id !== "string" || !item.id || item.id.length > 1_024) continue; + if (seen.has(item.id)) throw new Error(`Duplicate Prime Inference model ${item.id}`); + const pricing = isRecord(item.pricing) ? item.pricing : {}; + const input = nonNegativeNumber(pricing.input_usd_per_mtok); + const output = nonNegativeNumber(pricing.output_usd_per_mtok); + if (input === undefined || output === undefined) continue; + + const name = typeof item.display_name === "string" ? item.display_name.trim() : ""; + const specs = isRecord(item.specs) ? item.specs : {}; + const modalities = isRecord(specs.modalities) ? specs.modalities : {}; + const inputModalities = + Array.isArray(modalities.input) && modalities.input.every((modality) => typeof modality === "string") + ? modalities.input + : undefined; + const outputModalities = + Array.isArray(modalities.output) && modalities.output.every((modality) => typeof modality === "string") + ? modalities.output + : undefined; + const contextWindow = positiveInteger(specs.context_window); + const maxTokens = positiveInteger(specs.max_output_tokens); + const reasoning = typeof specs.supports_reasoning === "boolean" ? specs.supports_reasoning : undefined; + const hasSpecs = + contextWindow !== undefined && + maxTokens !== undefined && + reasoning !== undefined && + inputModalities !== undefined && + outputModalities !== undefined; + const cacheRead = nonNegativeNumber(pricing.cache_read_usd_per_mtok); + const cacheWrite = nonNegativeNumber(pricing.cache_write_usd_per_mtok); + + seen.add(item.id); + models.push({ + id: item.id, + ...(name ? { name } : {}), + input, + output, + ...(cacheRead !== undefined ? { cacheRead } : {}), + ...(cacheWrite !== undefined ? { cacheWrite } : {}), + ...(hasSpecs + ? { + contextWindow, + maxTokens: Math.min(maxTokens, contextWindow), + vision: inputModalities.includes("image"), + reasoning, + } + : {}), + }); + } + if (models.length === 0 && !options.allowEmpty) throw new Error("Prime Inference model catalog is empty"); + return models; +} diff --git a/packages/ai/test/prime-inference-model-catalog.test.ts b/packages/ai/test/prime-inference-model-catalog.test.ts new file mode 100644 index 0000000000..98ead74f89 --- /dev/null +++ b/packages/ai/test/prime-inference-model-catalog.test.ts @@ -0,0 +1,74 @@ +import { describe, expect, test } from "vitest"; +import { parsePrimeInferenceModelCatalog } from "../src/prime-inference-model-catalog.js"; + +function response(...data: unknown[]) { + return { object: "list", data }; +} + +describe("Prime Inference model catalog", () => { + test("parses pricing and model specs", () => { + const [model] = parsePrimeInferenceModelCatalog( + response({ + id: "vendor/model", + display_name: "Model Name", + pricing: { + input_usd_per_mtok: 1, + output_usd_per_mtok: 2, + cache_read_usd_per_mtok: 0.1, + cache_write_usd_per_mtok: 1.25, + }, + specs: { + context_window: 200_000, + max_output_tokens: 64_000, + modalities: { input: ["text", "image", "file"], output: ["text"] }, + supports_reasoning: true, + }, + }), + ); + expect(model).toEqual({ + id: "vendor/model", + name: "Model Name", + input: 1, + output: 2, + cacheRead: 0.1, + cacheWrite: 1.25, + contextWindow: 200_000, + maxTokens: 64_000, + vision: true, + reasoning: true, + }); + }); + + test("keeps priced entries without complete specs for bundled fallback", () => { + expect( + parsePrimeInferenceModelCatalog( + response( + { + id: "vendor/no-specs", + pricing: { input_usd_per_mtok: 1, output_usd_per_mtok: 2 }, + specs: null, + }, + { + id: "vendor/partial-specs", + pricing: { input_usd_per_mtok: 3, output_usd_per_mtok: 4 }, + specs: { + context_window: 100_000, + max_output_tokens: 10_000, + modalities: { output: ["text"] }, + supports_reasoning: false, + }, + }, + ), + ), + ).toEqual([ + { id: "vendor/no-specs", input: 1, output: 2 }, + { id: "vendor/partial-specs", input: 3, output: 4 }, + ]); + }); + + test("rejects empty and duplicate catalogs", () => { + expect(() => parsePrimeInferenceModelCatalog(response())).toThrow(/empty/); + const model = { id: "duplicate", pricing: { input_usd_per_mtok: 1, output_usd_per_mtok: 2 } }; + expect(() => parsePrimeInferenceModelCatalog(response(model, model))).toThrow(/duplicate/i); + }); +}); diff --git a/packages/ai/test/prime-inference-models.test.ts b/packages/ai/test/prime-inference-models.test.ts index cd79d4a177..58c3a03b8f 100644 --- a/packages/ai/test/prime-inference-models.test.ts +++ b/packages/ai/test/prime-inference-models.test.ts @@ -42,13 +42,10 @@ describe("Prime Inference models", () => { ); }); - it("skips private, raw, and duplicate catalog variants", () => { + it("excludes private routes from the bundled public catalog", () => { const modelIds = getModels("prime-inference").map((model) => model.id); - expect(modelIds.filter((id) => id.startsWith("internal/"))).toEqual([]); - expect(modelIds).not.toContain("zai-org/GLM-4.7"); - expect(modelIds).not.toContain("nvidia/NVIDIA-Nemotron-3-Nano-30B-A3B-BF16"); - expect(modelIds).not.toContain("Qwen/Qwen3.5-4B"); + expect(modelIds.filter((id) => id.startsWith("internal/") || id.startsWith("dev/"))).toEqual([]); expect(modelIds.filter((id) => id.includes(":"))).toEqual([]); }); @@ -94,8 +91,7 @@ describe("Prime Inference models", () => { expect(model.reasoning).toBe(true); expect(getSupportedThinkingLevels(model)).toEqual(["off", "low", "high", "max"]); expect(model.input).toEqual(["text", "image"]); - expect(model.contextWindow).toBe(1048576); - expect(model.maxTokens).toBe(1048576); + expect(model.maxTokens).toBeLessThanOrEqual(model.contextWindow); expect(model.cost.input).toBe(provider === "prime-inference" ? 3.45 : 3); expect(model.cost.output).toBe(provider === "prime-inference" ? 17.25 : 15); } @@ -111,8 +107,7 @@ describe("Prime Inference models", () => { const nemotronSuper = getModel("prime-inference", "nvidia/nemotron-3-super-120b-a12b"); expect(nemotronSuper.reasoning).toBe(true); expect(nemotronSuper.input).toEqual(["text"]); - expect(nemotronSuper.contextWindow).toBe(262144); - expect(nemotronSuper.maxTokens).toBe(4096); + expect(nemotronSuper.maxTokens).toBeLessThanOrEqual(nemotronSuper.contextWindow); const maverick = getModel("prime-inference", "meta-llama/llama-4-maverick"); expect(maverick.contextWindow).toBe(1048576); @@ -207,7 +202,6 @@ describe("Prime Inference models", () => { expect(getModel("prime-inference", "anthropic/claude-sonnet-4.6").contextWindow).toBe(1000000); expect(getModel("prime-inference", "anthropic/claude-sonnet-5").contextWindow).toBe(1000000); expect(getModel("prime-inference", "anthropic/claude-haiku-4.5").contextWindow).toBe(200000); - expect(getModel("prime-inference", "anthropic/claude-sonnet-4.5").contextWindow).toBe(200000); }); it("resolves PRIME_API_KEY from the environment", () => { diff --git a/packages/ai/test/render-models.test.ts b/packages/ai/test/render-models.test.ts new file mode 100644 index 0000000000..b0fbc5f457 --- /dev/null +++ b/packages/ai/test/render-models.test.ts @@ -0,0 +1,30 @@ +import { describe, expect, test } from "vitest"; +import { renderModelsFile } from "../scripts/render-models.js"; + +describe("generated model serialization", () => { + test("escapes remote strings before writing TypeScript source", () => { + const id = 'vendor/model";\nexport const injected = true; //'; + const name = 'Model "name"\nwith a newline'; + const output = renderModelsFile({ + "prime-inference": { + [id]: { + id, + name, + api: "openai-completions", + provider: "prime-inference", + baseUrl: "https://api.pinference.ai/api/v1", + reasoning: false, + input: ["text"], + cost: { input: 1, output: 2, cacheRead: 0, cacheWrite: 0 }, + contextWindow: 128_000, + maxTokens: 8_192, + }, + }, + }); + + expect(output).toContain(`\t\t${JSON.stringify(id)}: {`); + expect(output).toContain(`\t\t\tid: ${JSON.stringify(id)},`); + expect(output).toContain(`\t\t\tname: ${JSON.stringify(name)},`); + expect(output).not.toContain(`name: "${name}",`); + }); +}); diff --git a/packages/coding-agent/.changes/prime-inference-model-catalog.md b/packages/coding-agent/.changes/prime-inference-model-catalog.md new file mode 100644 index 0000000000..e8317479a5 --- /dev/null +++ b/packages/coding-agent/.changes/prime-inference-model-catalog.md @@ -0,0 +1 @@ +- Added live refreshes for public and authorized private Prime Inference models while retaining bundled and cached fallbacks. diff --git a/packages/coding-agent/docs/providers.md b/packages/coding-agent/docs/providers.md index 5c7ff364bb..1aaa6a2957 100644 --- a/packages/coding-agent/docs/providers.md +++ b/packages/coding-agent/docs/providers.md @@ -1,6 +1,6 @@ # Providers -Prime Agent supports subscription-based providers via OAuth and API key providers via environment variables or the auth file. Its built-in model catalog is updated with each Prime Agent release. +Prime Agent supports subscription-based providers via OAuth and API key providers via environment variables or the auth file. Models for external providers are bundled with each release. Prime Inference models refresh from its `/models` endpoint, with the bundled list and a validated disk cache as fallbacks. Set `PI_OFFLINE=1` to skip network refreshes. ## Table of Contents diff --git a/packages/coding-agent/src/core/model-registry.ts b/packages/coding-agent/src/core/model-registry.ts index 787fce9386..359a1ff7a1 100644 --- a/packages/coding-agent/src/core/model-registry.ts +++ b/packages/coding-agent/src/core/model-registry.ts @@ -16,6 +16,7 @@ import { type OAuthProviderInterface, type OpenAICompletionsCompat, type OpenAIResponsesCompat, + parsePrimeInferenceModelCatalog, registerApiProvider, resetApiProviders, type SimpleStreamOptions, @@ -31,7 +32,13 @@ import { getAgentDir } from "../config.js"; import type { AuthSourceToken, AuthStatus, AuthStorage } from "./auth-storage.js"; import { PRIME_INFERENCE_PROVIDER_ID } from "./prime-inference-auth.js"; import { - fetchAuthorizedPrivatePrimeInferenceModelIds, + buildPrimeInferenceModels, + mergePrimeInferenceModels, + readCachedPrimeInferenceModels, + refreshPrimeInferenceModels, +} from "./prime-inference-model-catalog.js"; +import { + fetchAuthorizedPrivatePrimeInferenceModels, getPrivatePrimeInferenceModels, isPrivatePrimeInferenceModel, } from "./prime-inference-models.js"; @@ -416,7 +423,7 @@ const PRIVATE_PRIME_BACKGROUND_REFRESH_TIMEOUT_MS = 3_000; interface PrivatePrimeAuthorizationCache { fingerprint: string; - modelIds: Set; + models: Model<"openai-completions">[]; refreshedAt: number; } @@ -441,10 +448,12 @@ export class ModelRegistry { private modelRequestHeaders: Map> = new Map(); private registeredProviders: Map = new Map(); private authorizedPrivatePrimeInferenceModelIds = new Set(); + private authorizedPrivatePrimeInferenceModels: Model<"openai-completions">[] = []; private authorizedPrivatePrimeInferenceTeamId: string | undefined; private explicitPrivatePrimeInferenceModelIds = new Set(); private openAICodexModelsCache: { authFingerprint: string; modelIds: Set; refreshedAt: number } | undefined; private backgroundPrivatePrimeAuthorization: { fingerprint: string; promise: Promise } | undefined; + private livePrimeInferenceModels: Model<"openai-completions">[] | undefined; private loadError: string | undefined = undefined; /** Re-register dynamic OAuth providers (e.g. user MCP servers) after refresh() resets the registry. */ @@ -496,13 +505,28 @@ export class ModelRegistry { registerBuiltinMcpOAuthProviders(); this.onOAuthProvidersReset?.(); + this.reloadModelsAfterCatalogChange(); + } + + private reloadModelsAfterCatalogChange(): void { this.loadModels(); + this.reapplyRegisteredProviders(); + } + private reapplyRegisteredProviders(): void { for (const [providerName, config] of this.registeredProviders.entries()) { this.applyProviderConfig(providerName, config); } } + private primeInferenceCatalogCachePath(): string | undefined { + return this.modelsJsonPath ? join(dirname(this.modelsJsonPath), "prime-inference-models-cache.json") : undefined; + } + + private bundledPrimeInferenceModels(): Model<"openai-completions">[] { + return getModels(PRIME_INFERENCE_PROVIDER_ID) as Model<"openai-completions">[]; + } + /** * Get any error from loading models.json (undefined if no error). */ @@ -525,7 +549,20 @@ export class ModelRegistry { this.explicitPrivatePrimeInferenceModelIds = new Set( customModels.filter(isPrivatePrimeInferenceModel).map((model) => model.id), ); - const builtInModels = [...this.loadBuiltInModels(overrides, modelOverrides), ...getPrivatePrimeInferenceModels()]; + const cachePath = this.primeInferenceCatalogCachePath(); + this.livePrimeInferenceModels ??= cachePath + ? readCachedPrimeInferenceModels(cachePath, this.bundledPrimeInferenceModels()) + : undefined; + const privateModels = new Map( + [...getPrivatePrimeInferenceModels(), ...this.authorizedPrivatePrimeInferenceModels].map((model) => [ + model.id, + model, + ]), + ); + const builtInModels = [ + ...this.loadBuiltInModels(overrides, modelOverrides, this.livePrimeInferenceModels), + ...privateModels.values(), + ]; let combined = this.mergeCustomModels(builtInModels, customModels); for (const oauthProvider of this.authStorage.getOAuthProviders()) { @@ -542,30 +579,24 @@ export class ModelRegistry { private loadBuiltInModels( overrides: Map, modelOverrides: Map>, + livePrimeInferenceModels?: Model<"openai-completions">[], ): Model[] { - return getProviders().flatMap((provider) => { - const models = getModels(provider as KnownProvider) as Model[]; - const providerOverride = overrides.get(provider); - const perModelOverrides = modelOverrides.get(provider); - - return models.map((m) => { - let model = m; - - if (providerOverride) { - model = { - ...model, - baseUrl: providerOverride.baseUrl ?? model.baseUrl, - compat: mergeCompat(model.compat, providerOverride.compat), - }; - } - - const modelOverride = perModelOverrides?.get(m.id); - if (modelOverride) { - model = applyModelOverride(model, modelOverride); - } + const bundledModels = getProviders().flatMap((provider) => getModels(provider as KnownProvider) as Model[]); + return mergePrimeInferenceModels(bundledModels, livePrimeInferenceModels).map((model) => { + const providerOverride = overrides.get(model.provider); + const perModelOverrides = modelOverrides.get(model.provider); + let configuredModel = model; + + if (providerOverride) { + configuredModel = { + ...configuredModel, + baseUrl: providerOverride.baseUrl ?? configuredModel.baseUrl, + compat: mergeCompat(configuredModel.compat, providerOverride.compat), + }; + } - return model; - }); + const modelOverride = perModelOverrides?.get(model.id); + return modelOverride ? applyModelOverride(configuredModel, modelOverride) : configuredModel; }); } @@ -773,10 +804,24 @@ export class ModelRegistry { }); } + /** + * Reload local state and private authorization. Public Prime Inference models + * return from the disk/bundled fallback immediately and refresh in the background. + */ async refreshAvailableModels(): Promise[]> { const previousPrivateModelIds = new Set(this.authorizedPrivatePrimeInferenceModelIds); const previousTeamId = this.authorizedPrivatePrimeInferenceTeamId; this.refresh(); + const cachePath = this.primeInferenceCatalogCachePath(); + if (cachePath) { + void refreshPrimeInferenceModels(cachePath, this.bundledPrimeInferenceModels(), { + offline: isOfflineModeEnabled(), + }).then((models) => { + if (!models) return; + this.livePrimeInferenceModels = models; + this.reloadModelsAfterCatalogChange(); + }); + } await this.refreshPrivatePrimeInferenceAuthorization(previousPrivateModelIds, previousTeamId); return this.getAvailable(); } @@ -790,34 +835,41 @@ export class ModelRegistry { const teamId = teamHeaders?.["X-Prime-Team-ID"]; if (!apiKey || !teamHeaders || !teamId) { this.authorizedPrivatePrimeInferenceModelIds.clear(); + this.authorizedPrivatePrimeInferenceModels = []; this.authorizedPrivatePrimeInferenceTeamId = undefined; + this.reloadModelsAfterCatalogChange(); return; } const fingerprint = privatePrimeAuthorizationFingerprint(apiKey, teamId); const cached = this.readPrivatePrimeAuthorizationCache(); if (cached?.fingerprint === fingerprint) { - // Serve the persisted authorization decision so startup and model lists - // don't block on the network. A stale cache refreshes in the background - // and the updated ids apply to subsequent lookups in this process. - this.authorizedPrivatePrimeInferenceModelIds = new Set(cached.modelIds); + // Serve the credential-scoped cache so startup and model lists don't + // block on the network. Stale entries refresh in the background. + this.authorizedPrivatePrimeInferenceModels = cached.models; + this.authorizedPrivatePrimeInferenceModelIds = new Set(cached.models.map((model) => model.id)); this.authorizedPrivatePrimeInferenceTeamId = teamId; + this.reloadModelsAfterCatalogChange(); const cacheIsFresh = Date.now() - cached.refreshedAt < PRIVATE_PRIME_AUTHORIZATION_CACHE_TTL_MS; - if (cacheIsFresh || isOfflineModeEnabled()) { - return; - } + if (isOfflineModeEnabled() || cacheIsFresh) return; this.startBackgroundPrivatePrimeAuthorizationRefresh(apiKey, teamHeaders, teamId, fingerprint); return; } if (isOfflineModeEnabled()) { this.authorizedPrivatePrimeInferenceModelIds.clear(); + this.authorizedPrivatePrimeInferenceModels = []; this.authorizedPrivatePrimeInferenceTeamId = undefined; + this.reloadModelsAfterCatalogChange(); return; } - let authorizedIds: Set | undefined; + let authorizedModels: Model<"openai-completions">[] | undefined; try { - authorizedIds = await fetchAuthorizedPrivatePrimeInferenceModelIds(apiKey, teamHeaders); + authorizedModels = await fetchAuthorizedPrivatePrimeInferenceModels( + apiKey, + teamHeaders, + new Set((this.livePrimeInferenceModels ?? this.bundledPrimeInferenceModels()).map((model) => model.id)), + ); } catch { // Fall back to the previous authorization below. } @@ -825,16 +877,20 @@ export class ModelRegistry { if ((await this.currentPrivatePrimeAuthorizationFingerprint()) !== fingerprint) { return; } - if (authorizedIds) { - this.authorizedPrivatePrimeInferenceModelIds = authorizedIds; + if (authorizedModels) { + this.authorizedPrivatePrimeInferenceModels = authorizedModels; + this.authorizedPrivatePrimeInferenceModelIds = new Set(authorizedModels.map((model) => model.id)); this.authorizedPrivatePrimeInferenceTeamId = teamId; - this.writePrivatePrimeAuthorizationCache({ fingerprint, modelIds: authorizedIds, refreshedAt: Date.now() }); + this.reloadModelsAfterCatalogChange(); + this.writePrivatePrimeAuthorizationCache({ fingerprint, models: authorizedModels, refreshedAt: Date.now() }); } else if (teamId === previousTeamId) { this.authorizedPrivatePrimeInferenceModelIds = previousPrivateModelIds; this.authorizedPrivatePrimeInferenceTeamId = teamId; } else { this.authorizedPrivatePrimeInferenceModelIds.clear(); + this.authorizedPrivatePrimeInferenceModels = []; this.authorizedPrivatePrimeInferenceTeamId = undefined; + this.reloadModelsAfterCatalogChange(); } } @@ -855,18 +911,25 @@ export class ModelRegistry { } const run = async () => { try { - const authorizedIds = await fetchAuthorizedPrivatePrimeInferenceModelIds( + const authorizedModels = await fetchAuthorizedPrivatePrimeInferenceModels( apiKey, teamHeaders, + new Set((this.livePrimeInferenceModels ?? this.bundledPrimeInferenceModels()).map((model) => model.id)), undefined, PRIVATE_PRIME_BACKGROUND_REFRESH_TIMEOUT_MS, ); if ((await this.currentPrivatePrimeAuthorizationFingerprint()) !== fingerprint) { return; } - this.authorizedPrivatePrimeInferenceModelIds = authorizedIds; + this.authorizedPrivatePrimeInferenceModels = authorizedModels; + this.authorizedPrivatePrimeInferenceModelIds = new Set(authorizedModels.map((model) => model.id)); this.authorizedPrivatePrimeInferenceTeamId = teamId; - this.writePrivatePrimeAuthorizationCache({ fingerprint, modelIds: authorizedIds, refreshedAt: Date.now() }); + this.reloadModelsAfterCatalogChange(); + this.writePrivatePrimeAuthorizationCache({ + fingerprint, + models: authorizedModels, + refreshedAt: Date.now(), + }); } catch { // Keep the cached authorization. } @@ -896,25 +959,28 @@ export class ModelRegistry { private readPrivatePrimeAuthorizationCache(): PrivatePrimeAuthorizationCache | undefined { const cachePath = this.privatePrimeAuthorizationCachePath(); - if (!cachePath) { - return undefined; - } + if (!cachePath) return undefined; try { - const parsed = JSON.parse(readFileSync(cachePath, "utf8")) as Partial< - Omit & { modelIds: string[] } - >; + const parsed = JSON.parse(readFileSync(cachePath, "utf8")) as { + fingerprint?: unknown; + data?: unknown; + refreshedAt?: unknown; + }; if ( typeof parsed.fingerprint !== "string" || - !Array.isArray(parsed.modelIds) || + !Array.isArray(parsed.data) || typeof parsed.refreshedAt !== "number" ) { return undefined; } - return { - fingerprint: parsed.fingerprint, - modelIds: new Set(parsed.modelIds), - refreshedAt: parsed.refreshedAt, - }; + const entries = parsePrimeInferenceModelCatalog({ data: parsed.data }, { allowEmpty: true }).filter((entry) => + isPrivatePrimeInferenceModel({ provider: PRIME_INFERENCE_PROVIDER_ID, id: entry.id }), + ); + const models = buildPrimeInferenceModels(getPrivatePrimeInferenceModels(), entries, { + includePrivate: true, + minimumModels: 0, + }); + return { fingerprint: parsed.fingerprint, models: models ?? [], refreshedAt: parsed.refreshedAt }; } catch { return undefined; } @@ -922,12 +988,32 @@ export class ModelRegistry { private writePrivatePrimeAuthorizationCache(cache: PrivatePrimeAuthorizationCache): void { const cachePath = this.privatePrimeAuthorizationCachePath(); - if (!cachePath) { - return; - } + if (!cachePath) return; + const data = cache.models.map((model) => ({ + id: model.id, + display_name: model.name, + pricing: { + input_usd_per_mtok: model.cost.input, + output_usd_per_mtok: model.cost.output, + cache_read_usd_per_mtok: model.cost.cacheRead, + cache_write_usd_per_mtok: model.cost.cacheWrite, + }, + specs: { + context_window: model.contextWindow, + max_output_tokens: model.maxTokens, + modalities: { input: model.input, output: ["text"] }, + supports_reasoning: model.reasoning, + }, + })); try { const tmpPath = `${cachePath}.${process.pid}.tmp`; - writeFileSync(tmpPath, JSON.stringify({ ...cache, modelIds: [...cache.modelIds] }), { mode: 0o600 }); + writeFileSync( + tmpPath, + JSON.stringify({ fingerprint: cache.fingerprint, data, refreshedAt: cache.refreshedAt }), + { + mode: 0o600, + }, + ); renameSync(tmpPath, cachePath); } catch { // A failed cache write only requires a later refetch. diff --git a/packages/coding-agent/src/core/prime-inference-model-catalog.ts b/packages/coding-agent/src/core/prime-inference-model-catalog.ts new file mode 100644 index 0000000000..f229b9b612 --- /dev/null +++ b/packages/coding-agent/src/core/prime-inference-model-catalog.ts @@ -0,0 +1,176 @@ +import { Buffer } from "node:buffer"; +import { existsSync, readFileSync, renameSync, unlinkSync, writeFileSync } from "node:fs"; +import { + type Api, + isPrivatePrimeInferenceModelId, + type Model, + type OpenAICompletionsCompat, + type PrimeInferenceCatalogEntry, + parsePrimeInferenceModelCatalog, +} from "@earendil-works/pi-ai"; + +export const PRIME_INFERENCE_BASE_URL = "https://api.pinference.ai/api/v1"; +const FETCH_TIMEOUT_MS = 5_000; +const MAX_RESPONSE_BYTES = 2 * 1024 * 1024; +const MIN_CATALOG_COVERAGE = 0.5; +const pendingRefreshes = new Map[] | undefined>>(); + +const DEFAULT_COMPAT: OpenAICompletionsCompat = { + supportsStore: false, + supportsDeveloperRole: false, + // The endpoint does not yet describe reasoning controls. Do not send an + // unconfirmed reasoning_effort parameter for models without a bundled template. + supportsReasoningEffort: false, + maxTokensField: "max_tokens", + supportsStrictMode: false, +}; + +function cacheCosts(entry: PrimeInferenceCatalogEntry, template?: Model<"openai-completions">) { + const anthropic = entry.id.toLowerCase().startsWith("anthropic/"); + return { + cacheRead: entry.cacheRead ?? template?.cost.cacheRead ?? (anthropic ? entry.input * 0.1 : 0), + cacheWrite: entry.cacheWrite ?? template?.cost.cacheWrite ?? (anthropic ? entry.input * 1.25 : 0), + }; +} + +export function buildPrimeInferenceModels( + bundledModels: readonly Model<"openai-completions">[], + entries: readonly PrimeInferenceCatalogEntry[], + options: { includePrivate?: boolean; minimumModels?: number } = {}, +): Model<"openai-completions">[] | undefined { + const bundled = new Map(bundledModels.map((model) => [model.id.toLowerCase(), model])); + const models: Model<"openai-completions">[] = []; + for (const entry of entries) { + if (!options.includePrivate && isPrivatePrimeInferenceModelId(entry.id)) continue; + const template = bundled.get(entry.id.toLowerCase()); + if (!template && (!entry.contextWindow || !entry.maxTokens || entry.reasoning === undefined)) continue; + const contextWindow = entry.contextWindow ?? template?.contextWindow ?? 0; + const maxTokens = Math.min(entry.maxTokens ?? template?.maxTokens ?? 0, contextWindow); + models.push({ + id: entry.id, + name: entry.name ?? template?.name ?? entry.id, + api: "openai-completions", + provider: "prime-inference", + baseUrl: PRIME_INFERENCE_BASE_URL, + reasoning: entry.reasoning ?? template?.reasoning ?? false, + ...(template?.thinkingLevelMap ? { thinkingLevelMap: { ...template.thinkingLevelMap } } : {}), + input: (entry.vision ?? template?.input.includes("image")) ? ["text", "image"] : ["text"], + cost: { input: entry.input, output: entry.output, ...cacheCosts(entry, template) }, + contextWindow, + maxTokens, + ...(template?.featured ? { featured: true } : {}), + compat: structuredClone(template?.compat ?? DEFAULT_COMPAT), + }); + } + const minimumModels = options.minimumModels ?? Math.ceil(bundledModels.length * MIN_CATALOG_COVERAGE); + const coveredBundledModels = models.filter((model) => bundled.has(model.id.toLowerCase())).length; + return coveredBundledModels >= minimumModels ? models : undefined; +} + +export function mergePrimeInferenceModels( + bundledModels: readonly Model[], + livePrimeInferenceModels?: readonly Model<"openai-completions">[], +): Model[] { + if (!livePrimeInferenceModels) return [...bundledModels]; + return [...bundledModels.filter((model) => model.provider !== "prime-inference"), ...livePrimeInferenceModels]; +} + +export function readCachedPrimeInferenceModels( + cachePath: string, + bundledModels: readonly Model<"openai-completions">[], +): Model<"openai-completions">[] | undefined { + if (!existsSync(cachePath)) return undefined; + try { + return buildPrimeInferenceModels( + bundledModels, + parsePrimeInferenceModelCatalog(JSON.parse(readFileSync(cachePath, "utf8")) as unknown), + ); + } catch { + return undefined; + } +} + +function writeCache(cachePath: string, value: unknown): void { + const temporaryPath = `${cachePath}.${process.pid}.${Date.now()}.tmp`; + try { + writeFileSync(temporaryPath, JSON.stringify(value), { encoding: "utf8", mode: 0o600 }); + renameSync(temporaryPath, cachePath); + } catch { + // The bundled catalog remains available when the cache cannot be persisted. + } finally { + try { + if (existsSync(temporaryPath)) unlinkSync(temporaryPath); + } catch { + // Ignore cache cleanup failures. + } + } +} + +export class PrimeInferenceCatalogRequestError extends Error { + constructor(readonly status: number) { + super(`Prime Inference model catalog request failed with status ${status}`); + } +} + +async function readResponse(response: Response): Promise { + if (!response.ok) throw new PrimeInferenceCatalogRequestError(response.status); + const contentLength = Number(response.headers.get("content-length")); + if (Number.isFinite(contentLength) && contentLength > MAX_RESPONSE_BYTES) throw new Error("Response is too large"); + if (!response.body) throw new Error("Response body is empty"); + const reader = response.body.getReader(); + const chunks: Uint8Array[] = []; + let bytesRead = 0; + try { + while (true) { + const { done, value } = await reader.read(); + if (done) break; + bytesRead += value.byteLength; + if (bytesRead > MAX_RESPONSE_BYTES) { + await reader.cancel().catch(() => {}); + throw new Error("Response is too large"); + } + chunks.push(value); + } + } finally { + reader.releaseLock(); + } + return JSON.parse(Buffer.concat(chunks, bytesRead).toString("utf8")) as unknown; +} + +export async function fetchPrimeInferenceModelCatalog( + options: { fetchFn?: typeof fetch; headers?: Record; timeoutMs?: number; allowEmpty?: boolean } = {}, +): Promise<{ payload: unknown; entries: PrimeInferenceCatalogEntry[] }> { + const response = await (options.fetchFn ?? fetch)(`${PRIME_INFERENCE_BASE_URL}/models`, { + headers: { accept: "application/json", ...options.headers }, + signal: AbortSignal.timeout(options.timeoutMs ?? FETCH_TIMEOUT_MS), + }); + const payload = await readResponse(response); + return { payload, entries: parsePrimeInferenceModelCatalog(payload, { allowEmpty: options.allowEmpty }) }; +} + +export async function refreshPrimeInferenceModels( + cachePath: string, + bundledModels: readonly Model<"openai-completions">[], + options: { fetchFn?: typeof fetch; offline?: boolean } = {}, +): Promise[] | undefined> { + const cached = readCachedPrimeInferenceModels(cachePath, bundledModels); + if (options.offline) return cached; + const existing = pendingRefreshes.get(cachePath); + if (existing) return existing; + const promise = (async () => { + try { + const { payload, entries } = await fetchPrimeInferenceModelCatalog({ fetchFn: options.fetchFn }); + const models = buildPrimeInferenceModels(bundledModels, entries); + if (!models) return cached; + writeCache(cachePath, payload); + return models; + } catch { + return cached; + } + })(); + pendingRefreshes.set(cachePath, promise); + void promise.finally(() => { + if (pendingRefreshes.get(cachePath) === promise) pendingRefreshes.delete(cachePath); + }); + return promise; +} diff --git a/packages/coding-agent/src/core/prime-inference-models.ts b/packages/coding-agent/src/core/prime-inference-models.ts index c9af39adfa..15c229c3fb 100644 --- a/packages/coding-agent/src/core/prime-inference-models.ts +++ b/packages/coding-agent/src/core/prime-inference-models.ts @@ -1,6 +1,13 @@ -import type { Model } from "@earendil-works/pi-ai"; +import { isPrivatePrimeInferenceModelId, type Model } from "@earendil-works/pi-ai"; +import { + buildPrimeInferenceModels, + fetchPrimeInferenceModelCatalog, + PRIME_INFERENCE_BASE_URL, + PrimeInferenceCatalogRequestError, +} from "./prime-inference-model-catalog.js"; + +export { PRIME_INFERENCE_BASE_URL }; -export const PRIME_INFERENCE_BASE_URL = "https://api.pinference.ai/api/v1"; const PRIVATE_MODEL_REFRESH_TIMEOUT_MS = 10_000; const PRIVATE_PRIME_INFERENCE_MODELS: readonly Model<"openai-completions">[] = [ @@ -24,7 +31,7 @@ const PRIVATE_PRIME_INFERENCE_MODELS: readonly Model<"openai-completions">[] = [ ]; export function isPrivatePrimeInferenceModel(model: Pick, "provider" | "id">): boolean { - return model.provider === "prime-inference" && model.id.startsWith("internal/"); + return model.provider === "prime-inference" && isPrivatePrimeInferenceModelId(model.id); } export function getPrivatePrimeInferenceModels(): Model<"openai-completions">[] { @@ -36,42 +43,46 @@ export function getPrivatePrimeInferenceModels(): Model<"openai-completions">[] })); } -export async function fetchAuthorizedPrivatePrimeInferenceModelIds( +export async function fetchAuthorizedPrivatePrimeInferenceModels( apiKey: string, teamHeaders: Record, + publicModelIds: ReadonlySet, fetchFn: typeof fetch = fetch, timeoutMs: number = PRIVATE_MODEL_REFRESH_TIMEOUT_MS, -): Promise> { - if (!teamHeaders["X-Prime-Team-ID"]) { - return new Set(); - } - - const response = await fetchFn(`${PRIME_INFERENCE_BASE_URL}/models`, { - headers: { - Authorization: `Bearer ${apiKey}`, - ...teamHeaders, - }, - signal: AbortSignal.timeout(timeoutMs), - }); - if (response.status === 401 || response.status === 403) { - return new Set(); +): Promise[]> { + if (!teamHeaders["X-Prime-Team-ID"]) return []; + try { + const { payload, entries } = await fetchPrimeInferenceModelCatalog({ + fetchFn, + timeoutMs, + allowEmpty: true, + headers: { ...teamHeaders, Authorization: `Bearer ${apiKey}` }, + }); + const publicIds = new Set([...publicModelIds].map((id) => id.toLowerCase())); + const bundledPrivateModels = getPrivatePrimeInferenceModels(); + const bundledById = new Map(bundledPrivateModels.map((model) => [model.id.toLowerCase(), model])); + const entriesById = new Map(entries.map((entry) => [entry.id.toLowerCase(), entry])); + const data = + payload && typeof payload === "object" && "data" in payload && Array.isArray(payload.data) ? payload.data : []; + const privateEntries = data.flatMap((item) => { + if (!item || typeof item !== "object" || !("id" in item) || typeof item.id !== "string") return []; + const id = item.id.toLowerCase(); + if (publicIds.has(id) || !isPrivatePrimeInferenceModelId(id)) return []; + const parsed = entriesById.get(id); + if (parsed) return [parsed]; + const template = bundledById.get(id); + return template ? [{ id: item.id, input: template.cost.input, output: template.cost.output }] : []; + }); + return ( + buildPrimeInferenceModels(bundledPrivateModels, privateEntries, { + includePrivate: true, + minimumModels: 0, + }) ?? [] + ); + } catch (error) { + if (error instanceof PrimeInferenceCatalogRequestError && (error.status === 401 || error.status === 403)) { + return []; + } + throw error; } - if (!response.ok) { - throw new Error(`Prime Inference model catalog request failed with status ${response.status}`); - } - - const payload = (await response.json()) as unknown; - if (!payload || typeof payload !== "object" || !("data" in payload) || !Array.isArray(payload.data)) { - throw new Error("Prime Inference model catalog response is invalid"); - } - - const knownPrivateIds = new Set(PRIVATE_PRIME_INFERENCE_MODELS.map((model) => model.id)); - return new Set( - payload.data.flatMap((entry) => { - if (!entry || typeof entry !== "object" || !("id" in entry) || typeof entry.id !== "string") { - return []; - } - return knownPrivateIds.has(entry.id) ? [entry.id] : []; - }), - ); } diff --git a/packages/coding-agent/test/model-registry.test.ts b/packages/coding-agent/test/model-registry.test.ts index 0afca0a61f..d9014a1602 100644 --- a/packages/coding-agent/test/model-registry.test.ts +++ b/packages/coding-agent/test/model-registry.test.ts @@ -2,9 +2,9 @@ import { existsSync, mkdirSync, readFileSync, rmSync, writeFileSync } from "node import { tmpdir } from "node:os"; import { join } from "node:path"; import type { AnthropicMessagesCompat, Api, Context, Model, OpenAICompletionsCompat } from "@earendil-works/pi-ai"; -import { getApiProvider } from "@earendil-works/pi-ai"; +import { getApiProvider, getModels } from "@earendil-works/pi-ai"; import { getOAuthProvider, registerOAuthProvider } from "@earendil-works/pi-ai/oauth"; -import { afterEach, beforeEach, describe, expect, test } from "vitest"; +import { afterEach, beforeEach, describe, expect, test, vi } from "vitest"; import { AuthStorage } from "../src/core/auth-storage.js"; import { ModelRegistry, type ProviderConfigInput } from "../src/core/model-registry.js"; @@ -21,6 +21,7 @@ describe("ModelRegistry", () => { }); afterEach(() => { + vi.unstubAllGlobals(); if (tempDir && existsSync(tempDir)) { rmSync(tempDir, { recursive: true }); } @@ -603,6 +604,96 @@ describe("ModelRegistry", () => { }); }); + describe("live Prime Inference models", () => { + test("loads the cache without replacing external providers and applies local overrides", () => { + const bundled = getModels("prime-inference") as Model<"openai-completions">[]; + const catalogEntries = bundled.map((model) => ({ + id: model.id, + display_name: `Live ${model.name}`, + pricing: { input_usd_per_mtok: model.cost.input, output_usd_per_mtok: model.cost.output }, + specs: { + context_window: model.contextWindow, + max_output_tokens: model.maxTokens, + modalities: { input: model.input, output: ["text"] }, + supports_reasoning: model.reasoning, + }, + })); + catalogEntries.push({ + id: "test/live-added", + display_name: "Live Added", + pricing: { input_usd_per_mtok: 1, output_usd_per_mtok: 2 }, + specs: { + context_window: 200_000, + max_output_tokens: 20_000, + modalities: { input: ["text"], output: ["text"] }, + supports_reasoning: false, + }, + }); + writeFileSync( + join(tempDir, "prime-inference-models-cache.json"), + JSON.stringify({ object: "list", data: catalogEntries }), + ); + writeRawModelsJson({ + "prime-inference": { + baseUrl: "https://local-proxy.example.com/v1", + modelOverrides: { "test/live-added": { name: "Local Added", contextWindow: 123_456 } }, + }, + }); + + const registry = ModelRegistry.create(authStorage, modelsJsonPath); + expect(registry.find("prime-inference", "test/live-added")).toMatchObject({ + name: "Local Added", + baseUrl: "https://local-proxy.example.com/v1", + contextWindow: 123_456, + cost: { input: 1, output: 2 }, + }); + expect(getModelsForProvider(registry, "openrouter")).toHaveLength(getModels("openrouter").length); + }); + + test("restores cached authorized deployment metadata without waiting for the network", async () => { + const privateRoute = { + id: "vendor/model:deployment", + display_name: "Private Deployment", + pricing: { input_usd_per_mtok: 1, output_usd_per_mtok: 2 }, + specs: { + context_window: 200_000, + max_output_tokens: 20_000, + modalities: { input: ["text"], output: ["text"] }, + supports_reasoning: false, + }, + }; + authStorage.set("prime-inference", { + type: "api_key", + key: "prime-key", + primeTeam: { teamId: "research-team", name: "Research" }, + }); + vi.stubGlobal( + "fetch", + vi.fn( + async (_url: string | URL | Request, init?: RequestInit) => + new Response( + JSON.stringify({ data: new Headers(init?.headers).has("Authorization") ? [privateRoute] : [] }), + ), + ), + ); + const firstRegistry = ModelRegistry.create(authStorage, modelsJsonPath); + expect( + (await firstRegistry.refreshAvailableModels()).find((model) => model.id === privateRoute.id), + ).toMatchObject({ name: "Private Deployment", contextWindow: 200_000 }); + + vi.stubGlobal( + "fetch", + vi.fn(async () => { + throw new Error("offline"); + }), + ); + const restoredRegistry = ModelRegistry.create(authStorage, modelsJsonPath); + expect( + (await restoredRegistry.refreshAvailableModels()).find((model) => model.id === privateRoute.id), + ).toMatchObject({ name: "Private Deployment", contextWindow: 200_000 }); + }); + }); + describe("modelOverrides (per-model customization)", () => { test("model override applies to a single built-in model", () => { writeRawModelsJson({ diff --git a/packages/coding-agent/test/prime-inference-model-catalog.test.ts b/packages/coding-agent/test/prime-inference-model-catalog.test.ts new file mode 100644 index 0000000000..e5a033a00e --- /dev/null +++ b/packages/coding-agent/test/prime-inference-model-catalog.test.ts @@ -0,0 +1,195 @@ +import { mkdtempSync, readFileSync, rmSync } from "node:fs"; +import { tmpdir } from "node:os"; +import { join } from "node:path"; +import type { Model } from "@earendil-works/pi-ai"; +import { afterEach, describe, expect, test, vi } from "vitest"; +import { + buildPrimeInferenceModels, + mergePrimeInferenceModels, + PRIME_INFERENCE_BASE_URL, + refreshPrimeInferenceModels, +} from "../src/core/prime-inference-model-catalog.js"; +import { + fetchAuthorizedPrivatePrimeInferenceModels, + isPrivatePrimeInferenceModel, +} from "../src/core/prime-inference-models.js"; + +const directories: string[] = []; +const model = (id: string, provider = "prime-inference"): Model<"openai-completions"> => ({ + id, + name: `Bundled ${id}`, + api: "openai-completions", + provider, + baseUrl: provider === "prime-inference" ? PRIME_INFERENCE_BASE_URL : "https://example.com/v1", + reasoning: true, + thinkingLevelMap: { high: "high" }, + input: ["text"], + cost: { input: 9, output: 10, cacheRead: 0.9, cacheWrite: 11.25 }, + contextWindow: 100_000, + maxTokens: 10_000, + featured: true, + compat: { supportsDeveloperRole: false, maxTokensField: "max_tokens" }, +}); + +const entry = (id: string, overrides: Record = {}) => ({ + id, + input: 1, + output: 2, + contextWindow: 200_000, + maxTokens: 20_000, + vision: true, + reasoning: false, + ...overrides, +}); + +const payloadEntry = ( + id: string, + specs: unknown = { + context_window: 200_000, + max_output_tokens: 20_000, + modalities: { input: ["text", "image"], output: ["text"] }, + supports_reasoning: false, + }, +) => ({ + id, + display_name: `Live ${id}`, + pricing: { input_usd_per_mtok: 1, output_usd_per_mtok: 2 }, + specs, +}); + +const response = (...data: unknown[]) => new Response(JSON.stringify({ object: "list", data })); + +afterEach(() => { + for (const directory of directories.splice(0)) rmSync(directory, { recursive: true, force: true }); +}); + +describe("Prime Inference model catalog", () => { + test("uses live metadata while retaining bundled client compatibility", () => { + const [live] = buildPrimeInferenceModels( + [model("vendor/model")], + [entry("vendor/model", { name: "Live Name", cacheRead: 0.1, cacheWrite: 1.25, maxTokens: 250_000 })], + ) ?? [undefined]; + expect(live).toMatchObject({ + id: "vendor/model", + name: "Live Name", + baseUrl: PRIME_INFERENCE_BASE_URL, + api: "openai-completions", + provider: "prime-inference", + reasoning: false, + input: ["text", "image"], + cost: { input: 1, output: 2, cacheRead: 0.1, cacheWrite: 1.25 }, + contextWindow: 200_000, + maxTokens: 200_000, + thinkingLevelMap: { high: "high" }, + featured: true, + compat: { supportsDeveloperRole: false, maxTokensField: "max_tokens" }, + }); + expect(live).not.toHaveProperty("headers"); + }); + + test("adds complete new models and skips incomplete unknown models", () => { + const models = + buildPrimeInferenceModels( + [model("bundled")], + [entry("new/complete"), { id: "new/incomplete", input: 1, output: 2 }], + { minimumModels: 0 }, + ) ?? []; + expect(models.map(({ id }) => id)).toEqual(["new/complete"]); + }); + + test("retains bundled specs when an existing live entry has none", () => { + const [live] = buildPrimeInferenceModels( + [model("vendor/model")], + [{ id: "vendor/model", name: "Renamed", input: 1, output: 2 }], + ) ?? [undefined]; + expect(live).toMatchObject({ name: "Renamed", contextWindow: 100_000, maxTokens: 10_000, reasoning: true }); + }); + + test("filters private routes and measures coverage against bundled models", () => { + const bundled = [model("one"), model("two"), model("three")]; + expect( + buildPrimeInferenceModels(bundled, [ + entry("internal/private"), + entry("dev/private"), + entry("poolside/model:deployment"), + entry("one"), + ]), + ).toBeUndefined(); + expect( + buildPrimeInferenceModels(bundled, [entry("new/one"), entry("new/two"), entry("new/three")]), + ).toBeUndefined(); + }); + + test("requires authorization for private prefixes and deployment routes", () => { + for (const id of ["internal/model", "INTERNAL/model", "dev/model", "vendor/model:deployment"]) { + expect(isPrivatePrimeInferenceModel(model(id))).toBe(true); + } + expect(isPrivatePrimeInferenceModel(model("public/model"))).toBe(false); + expect(isPrivatePrimeInferenceModel(model("vendor/model:deployment", "openrouter"))).toBe(false); + }); + + test("replaces only the Prime Inference provider list", () => { + const external = model("external", "openrouter"); + const live = model("live"); + expect(mergePrimeInferenceModels([external, model("removed")], [live])).toEqual([external, live]); + }); + + test("caches valid responses and falls back to the cache", async () => { + const directory = mkdtempSync(join(tmpdir(), "prime-models-")); + directories.push(directory); + const cachePath = join(directory, "cache.json"); + const bundled = [model("vendor/model")]; + const fetched = await refreshPrimeInferenceModels(cachePath, bundled, { + fetchFn: vi.fn(async () => response(payloadEntry("vendor/model"))), + }); + expect(fetched?.[0]?.name).toBe("Live vendor/model"); + expect(JSON.parse(readFileSync(cachePath, "utf8")).data).toHaveLength(1); + const fallback = await refreshPrimeInferenceModels(cachePath, bundled, { + fetchFn: vi.fn(async () => { + throw new Error("offline"); + }), + }); + expect(fallback?.[0]?.name).toBe("Live vendor/model"); + }); + + test("uses authenticated responses only for private routes with complete metadata", async () => { + const fetchFn = vi.fn(async (_url: string | URL | Request, init?: RequestInit) => { + expect(new Headers(init?.headers).get("Authorization")).toBe("Bearer secret"); + expect(new Headers(init?.headers).get("X-Prime-Team-ID")).toBe("team"); + return response( + payloadEntry("public/model"), + payloadEntry("internal/model"), + payloadEntry("dev/model"), + payloadEntry("poolside/model:deployment"), + payloadEntry("internal/incomplete", null), + ); + }); + const models = await fetchAuthorizedPrivatePrimeInferenceModels( + "secret", + { "X-Prime-Team-ID": "team" }, + new Set(["public/model"]), + fetchFn, + ); + expect(models.map(({ id }) => id)).toEqual(["internal/model", "dev/model", "poolside/model:deployment"]); + }); + + test("uses bundled metadata to authorize an existing private route", async () => { + const models = await fetchAuthorizedPrivatePrimeInferenceModels( + "secret", + { "X-Prime-Team-ID": "team" }, + new Set(), + vi.fn(async () => response({ id: "internal/glm-5.2-fast" })), + ); + expect(models.map(({ id }) => id)).toEqual(["internal/glm-5.2-fast"]); + }); + + test("treats rejected authenticated requests as no private access", async () => { + const models = await fetchAuthorizedPrivatePrimeInferenceModels( + "bad", + { "X-Prime-Team-ID": "team" }, + new Set(), + vi.fn(async () => new Response(null, { status: 403 })), + ); + expect(models).toEqual([]); + }); +}); diff --git a/packages/coding-agent/test/suite/regressions/4645-internal-glm.test.ts b/packages/coding-agent/test/suite/regressions/4645-internal-glm.test.ts index 5a87100a7f..1e23b26c78 100644 --- a/packages/coding-agent/test/suite/regressions/4645-internal-glm.test.ts +++ b/packages/coding-agent/test/suite/regressions/4645-internal-glm.test.ts @@ -52,6 +52,7 @@ describe("ENG-4645 internal GLM configuration", () => { }); expect(fetchMock).toHaveBeenCalledWith("https://api.pinference.ai/api/v1/models", { headers: { + accept: "application/json", Authorization: "Bearer prime-key", "X-Prime-Team-ID": "engineering-team", },