diff --git a/pi-extension/index.ts b/pi-extension/index.ts index 4dcc6be..32635b0 100644 --- a/pi-extension/index.ts +++ b/pi-extension/index.ts @@ -3,6 +3,7 @@ import { type ExtensionAPI, type ExtensionContext, } from "@mariozechner/pi-coding-agent"; +import { getModel } from "@earendil-works/pi-ai"; import { existsSync, readFileSync } from "node:fs"; import { dirname, join } from "node:path"; import { fileURLToPath } from "node:url"; @@ -59,6 +60,15 @@ function isIntegrationBaseUrl(baseUrl: string | undefined, info: IntegrationProv type CurrentModel = NonNullable; +function builtInModelCapabilities(provider: string, modelID: string) { + const model = getModel(provider, modelID); + if (!model) return undefined; + return { + reasoning: model.reasoning, + thinkingLevelMap: model.thinkingLevelMap, + }; +} + let procEnvCache: Map | null = null; // envValue mirrors pi-ai's Bun compiled-binary workaround for sandboxed Linux @@ -118,7 +128,12 @@ export default async function (pi: ExtensionAPI) { const integrationNames = discovered.integrations.map((integration) => integration.name); const routeLabel = integrationProviderDisplayName(integrationNames); const availableIntegrationsLabel = integrationPromptAvailabilityLabel(integrationNames); - const integrationInfos = providerInfosFromIntegrationCatalogs(discovered.integrations, pricingCatalog); + const integrationInfos = providerInfosFromIntegrationCatalogs( + discovered.integrations, + pricingCatalog, + console.warn, + builtInModelCapabilities, + ); if (discovered.found && integrationInfos.size === 0 && !disabled) { console.warn(`[pi-exe-dev] LLM integration discovered, but no supported models were available`); } diff --git a/pi-extension/integration_catalog.test.ts b/pi-extension/integration_catalog.test.ts index 508aae3..0ebdcf0 100644 --- a/pi-extension/integration_catalog.test.ts +++ b/pi-extension/integration_catalog.test.ts @@ -229,6 +229,110 @@ test("namespaces reflected xAI models and never creates routes from pricing meta ]); }); +test("inherits built-in thinking capabilities before aliasing integration models", () => { + const thinkingLevelMap = { + off: "none", + minimal: null, + low: "low", + medium: "medium", + high: "high", + xhigh: "xhigh", + max: "max", + }; + const lookups: string[] = []; + const infos = providerInfosFromIntegrationCatalogs( + [ + { + name: "llm", + baseURL: "https://llm.int.exe.xyz", + catalog: { schema_version: 1, models: [openAIGPTModel("chatgpt")] }, + }, + ], + undefined, + () => {}, + (provider, modelID) => { + lookups.push(`${provider}/${modelID}`); + return { reasoning: true, thinkingLevelMap }; + }, + ); + + assert.deepEqual(lookups, ["openai/gpt-5.5"]); + const model = infos.get("exe-dev-openai")?.config.models?.[0]; + assert.equal(model?.id, "gpt-5.5@llm"); + assert.equal(model?.reasoning, true); + assert.deepEqual(model?.thinkingLevelMap, thinkingLevelMap); +}); + +test("preserves explicit false from built-in capabilities over pricing metadata", () => { + const pricingCatalog: Catalog = { + schemaVersion: 1, + providers: [ + { + id: "openai", + path: "openai/v1", + models: [ + { + id: "gpt-5.5", + reasoning: true, + cost: { input: 1, output: 2, cacheRead: 0.1, cacheWrite: 0.2 }, + }, + ], + }, + ], + }; + const infos = providerInfosFromIntegrationCatalogs( + [ + { + name: "llm", + baseURL: "https://llm.int.exe.xyz", + catalog: { schema_version: 1, models: [openAIGPTModel("managed")] }, + }, + ], + pricingCatalog, + () => {}, + () => ({ reasoning: false, thinkingLevelMap: undefined }), + ); + + const model = infos.get("exe-dev-openai")?.config.models?.[0]; + assert.equal(model?.reasoning, false); + assert.equal("thinkingLevelMap" in (model ?? {}), false); +}); + +test("retains pricing-catalog reasoning for unknown built-in models", () => { + const pricingCatalog: Catalog = { + schemaVersion: 1, + providers: [ + { + id: "xai", + path: "xai/v1", + models: [ + { + id: "grok-4.5", + reasoning: true, + cost: { input: 2, output: 6, cacheRead: 0.5, cacheWrite: 0 }, + }, + ], + }, + ], + }; + const infos = providerInfosFromIntegrationCatalogs( + [ + { + name: "llm", + baseURL: "https://llm.int.exe.xyz", + catalog: { schema_version: 1, models: [xaiGrokModel()] }, + }, + ], + pricingCatalog, + () => {}, + () => undefined, + ); + + const model = infos.get("exe-dev-xai")?.config.models?.[0]; + assert.equal(model?.reasoning, true); + assert.equal("thinkingLevelMap" in (model ?? {}), false); +}); + test("routes arbitrary providers by client protocol priority", () => { const infos = providerInfosFromIntegrationCatalogs( [ diff --git a/pi-extension/integration_catalog.ts b/pi-extension/integration_catalog.ts index c57d8d4..297cc28 100644 --- a/pi-extension/integration_catalog.ts +++ b/pi-extension/integration_catalog.ts @@ -108,6 +108,11 @@ type CompatBag = { export type JSONFetcher = (url: string) => Promise; +export type ModelCapabilityLookup = ( + provider: string, + modelID: string, +) => Pick | undefined; + // Both the integration catalog and the bundled pricing sidecar currently use // schema version 1. Unknown shapes are ignored rather than becoming routes. export const SCHEMA_VERSION = 1; @@ -322,6 +327,7 @@ function configFromIntegrationModel( integration: DiscoveredIntegration, model: IntegrationModel, fallback: CatalogModel | undefined, + lookupCapabilities: ModelCapabilityLookup, ): ProviderModelConfig | undefined { if (!validIntegrationProviderID(model.provider)) return undefined; const adapter = integrationAPIAdapter(model); @@ -330,13 +336,17 @@ function configFromIntegrationModel( const modelID = integrationModelID(model); if (!modelID) return undefined; + const capabilities = lookupCapabilities(model.provider, modelID); const compat = sanitizeCompat(fallback?.compat, model.provider, modelID); return { id: modelID, name: model.name || fallback?.name || modelID, api: adapter.piAPI, baseUrl, - reasoning: fallback?.reasoning ?? false, + reasoning: capabilities?.reasoning ?? fallback?.reasoning ?? false, + ...(capabilities?.thinkingLevelMap !== undefined + ? { thinkingLevelMap: capabilities.thinkingLevelMap } + : {}), input: inputModalities(model, fallback), contextWindow: model.limits?.context_window ?? fallback?.contextWindow ?? 128000, maxTokens: model.limits?.max_output_tokens ?? fallback?.maxTokens ?? 4096, @@ -381,6 +391,7 @@ export function providerInfosFromIntegrationCatalogs( integrations: DiscoveredIntegration[], pricingCatalog: Catalog | undefined, warn: WarnFn = console.warn, + lookupCapabilities: ModelCapabilityLookup = () => undefined, ): Map { const costs = costCatalogIndex(pricingCatalog); const grouped = new Map(); @@ -393,7 +404,7 @@ export function providerInfosFromIntegrationCatalogs( const modelID = integrationModelID(model); if (!modelID) continue; const fallback = fallbackCatalogModel(costs, provider, modelID, model.id); - const config = configFromIntegrationModel(integration, model, fallback); + const config = configFromIntegrationModel(integration, model, fallback, lookupCapabilities); if (!config) continue; if (!fallback) warnedMissingPricing.add(`${provider}/${modelID}`);