From fc6076e9bd3c438fd46ae333ead2ada0fbe9cb5a Mon Sep 17 00:00:00 2001 From: Zakhar Kogan <36503576+zaharkogan@users.noreply.github.com> Date: Sun, 22 Mar 2026 16:08:13 +0200 Subject: [PATCH] Fix custom model precedence across load and refresh --- .../coding-agent/src/config/model-registry.ts | 207 +++++++++++---- .../coding-agent/test/model-registry.test.ts | 239 +++++++++++++++++- 2 files changed, 396 insertions(+), 50 deletions(-) diff --git a/packages/coding-agent/src/config/model-registry.ts b/packages/coding-agent/src/config/model-registry.ts index 251dc2964..79a37904b 100644 --- a/packages/coding-agent/src/config/model-registry.ts +++ b/packages/coding-agent/src/config/model-registry.ts @@ -356,7 +356,7 @@ export interface ProviderDiscoveryState { /** Result of loading custom models from models.json */ interface CustomModelsResult { - models?: Model[]; + models?: CustomModelOverlay[]; overrides?: Map; modelOverrides?: Map>; keylessProviders?: Set; @@ -551,6 +551,24 @@ interface CustomModelBuildOptions { useDefaults: boolean; } +type CustomModelOverlay = { + id: string; + provider: string; + api: Api; + baseUrl: string; + name?: string; + reasoning?: boolean; + thinking?: ThinkingConfig; + input?: ("text" | "image")[]; + cost?: { input: number; output: number; cacheRead: number; cacheWrite: number }; + contextWindow?: number; + maxTokens?: number; + headers?: Record; + compat?: Model["compat"]; + contextPromotionTarget?: string; + premiumMultiplier?: number; +}; + function mergeCustomModelHeaders( providerHeaders: Record | undefined, modelHeaders: Record | undefined, @@ -567,6 +585,68 @@ function mergeCustomModelHeaders( return headers; } +function buildCustomModelOverlay( + providerName: string, + providerBaseUrl: string, + providerApi: Api | undefined, + providerHeaders: Record | undefined, + providerApiKey: string | undefined, + authHeader: boolean | undefined, + providerCompat: Model["compat"] | undefined, + modelDef: CustomModelDefinitionLike, +): CustomModelOverlay | undefined { + const api = modelDef.api ?? providerApi; + if (!api) return undefined; + return { + id: modelDef.id, + provider: providerName, + api, + baseUrl: modelDef.baseUrl ?? providerBaseUrl, + name: modelDef.name, + reasoning: modelDef.reasoning, + thinking: modelDef.thinking as ThinkingConfig | undefined, + input: modelDef.input as ("text" | "image")[] | undefined, + cost: modelDef.cost, + contextWindow: modelDef.contextWindow, + maxTokens: modelDef.maxTokens, + headers: mergeCustomModelHeaders(providerHeaders, modelDef.headers, authHeader, providerApiKey), + compat: mergeCompat(providerCompat, modelDef.compat), + contextPromotionTarget: modelDef.contextPromotionTarget, + premiumMultiplier: modelDef.premiumMultiplier, + }; +} + +function applyStandaloneCustomModelPolicies(model: CustomModelOverlay): CustomModelOverlay { + if (model.id !== "gpt-5.4" || model.provider === "github-copilot" || model.contextWindow !== undefined) { + return model; + } + return { ...model, contextWindow: 1_000_000 }; +} + +function finalizeCustomModel(model: CustomModelOverlay, options: CustomModelBuildOptions): Model { + const resolvedModel = options.useDefaults ? applyStandaloneCustomModelPolicies(model) : model; + const cost = + resolvedModel.cost ?? (options.useDefaults ? { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 } : undefined); + const input = resolvedModel.input ?? (options.useDefaults ? ["text"] : undefined); + return enrichModelThinking({ + id: resolvedModel.id, + name: resolvedModel.name ?? (options.useDefaults ? resolvedModel.id : undefined), + api: resolvedModel.api, + provider: resolvedModel.provider, + baseUrl: resolvedModel.baseUrl, + reasoning: resolvedModel.reasoning ?? (options.useDefaults ? false : undefined), + thinking: resolvedModel.thinking, + input: input as ("text" | "image")[], + cost, + contextWindow: resolvedModel.contextWindow ?? (options.useDefaults ? 128000 : undefined), + maxTokens: resolvedModel.maxTokens ?? (options.useDefaults ? 16384 : undefined), + headers: resolvedModel.headers, + compat: resolvedModel.compat, + contextPromotionTarget: resolvedModel.contextPromotionTarget, + premiumMultiplier: resolvedModel.premiumMultiplier, + } as Model); +} + function buildCustomModel( providerName: string, providerBaseUrl: string, @@ -578,28 +658,18 @@ function buildCustomModel( modelDef: CustomModelDefinitionLike, options: CustomModelBuildOptions, ): Model | undefined { - const api = modelDef.api ?? providerApi; - if (!api) return undefined; - const withDefaults = options.useDefaults; - const cost = modelDef.cost ?? (withDefaults ? { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 } : undefined); - const input = modelDef.input ?? (withDefaults ? ["text"] : undefined); - return enrichModelThinking({ - id: modelDef.id, - name: modelDef.name ?? (withDefaults ? modelDef.id : undefined), - api, - provider: providerName, - baseUrl: modelDef.baseUrl ?? providerBaseUrl, - reasoning: modelDef.reasoning ?? (withDefaults ? false : undefined), - thinking: modelDef.thinking as ThinkingConfig | undefined, - input: input as ("text" | "image")[], - cost, - contextWindow: modelDef.contextWindow ?? (withDefaults ? 128000 : undefined), - maxTokens: modelDef.maxTokens ?? (withDefaults ? 16384 : undefined), - headers: mergeCustomModelHeaders(providerHeaders, modelDef.headers, authHeader, providerApiKey), - compat: mergeCompat(providerCompat, modelDef.compat), - contextPromotionTarget: modelDef.contextPromotionTarget, - premiumMultiplier: modelDef.premiumMultiplier, - } as Model); + const model = buildCustomModelOverlay( + providerName, + providerBaseUrl, + providerApi, + providerHeaders, + providerApiKey, + authHeader, + providerCompat, + modelDef, + ); + if (!model) return undefined; + return finalizeCustomModel(model, options); } /** @@ -610,6 +680,8 @@ export class ModelRegistry { #customProviderApiKeys: Map = new Map(); #keylessProviders: Set = new Set(); #discoverableProviders: DiscoveryProviderConfig[] = []; + #customModelOverlays: CustomModelOverlay[] = []; + #providerOverrides: Map = new Map(); #modelOverrides: Map> = new Map(); #configError: ConfigError | undefined = undefined; #modelsConfigFile: ConfigFile; @@ -676,6 +748,7 @@ export class ModelRegistry { this.#customProviderApiKeys.clear(); this.#keylessProviders.clear(); this.#discoverableProviders = []; + this.#providerOverrides.clear(); this.#modelOverrides.clear(); this.#configError = undefined; this.#providerDiscoveryStates.clear(); @@ -703,14 +776,17 @@ export class ModelRegistry { this.#configError = configError; this.#keylessProviders = keylessProviders; this.#discoverableProviders = discoverableProviders; + this.#customModelOverlays = customModels; + this.#providerOverrides = overrides; this.#modelOverrides = modelOverrides; this.#addImplicitDiscoverableProviders(configuredProviders); - const builtInModels = this.#loadBuiltInModels(overrides, modelOverrides); - const cachedDiscoveries = this.#loadCachedDiscoverableModels(); - const combined = this.#mergeCustomModels(builtInModels, [...customModels, ...cachedDiscoveries]); + const builtInModels = this.#applyHardcodedModelPolicies(this.#loadBuiltInModels(overrides, modelOverrides)); + const cachedDiscoveries = this.#applyHardcodedModelPolicies(this.#loadCachedDiscoverableModels()); + const resolvedDefaults = this.#mergeResolvedModels(builtInModels, cachedDiscoveries); + const combined = this.#mergeCustomModels(resolvedDefaults, this.#customModelOverlays); - this.#models = this.#applyHardcodedModelPolicies(combined); + this.#models = this.#applyModelOverrides(combined, this.#modelOverrides); } /** Load built-in models, applying provider and per-model overrides */ @@ -742,15 +818,50 @@ export class ModelRegistry { }); } + #mergeResolvedModels(baseModels: Model[], replacementModels: Model[]): Model[] { + const merged = [...baseModels]; + for (const replacementModel of replacementModels) { + const existingIndex = merged.findIndex( + m => m.provider === replacementModel.provider && m.id === replacementModel.id, + ); + if (existingIndex >= 0) { + merged[existingIndex] = replacementModel; + } else { + merged.push(replacementModel); + } + } + return merged; + } + /** Merge custom models with built-in, replacing by provider+id match */ - #mergeCustomModels(builtInModels: Model[], customModels: Model[]): Model[] { + #mergeCustomModels(builtInModels: Model[], customModels: CustomModelOverlay[]): Model[] { const merged = [...builtInModels]; for (const customModel of customModels) { const existingIndex = merged.findIndex(m => m.provider === customModel.provider && m.id === customModel.id); if (existingIndex >= 0) { - merged[existingIndex] = customModel; + merged[existingIndex] = enrichModelThinking({ + ...merged[existingIndex], + id: customModel.id, + provider: customModel.provider, + api: customModel.api, + baseUrl: customModel.baseUrl, + name: customModel.name ?? merged[existingIndex].name, + reasoning: customModel.reasoning ?? merged[existingIndex].reasoning, + thinking: customModel.thinking ?? merged[existingIndex].thinking, + input: customModel.input ?? merged[existingIndex].input, + cost: customModel.cost ?? merged[existingIndex].cost, + contextWindow: customModel.contextWindow ?? merged[existingIndex].contextWindow, + maxTokens: customModel.maxTokens ?? merged[existingIndex].maxTokens, + headers: customModel.headers + ? { ...merged[existingIndex].headers, ...customModel.headers } + : merged[existingIndex].headers, + compat: mergeCompat(merged[existingIndex].compat, customModel.compat), + contextPromotionTarget: + customModel.contextPromotionTarget ?? merged[existingIndex].contextPromotionTarget, + premiumMultiplier: customModel.premiumMultiplier ?? merged[existingIndex].premiumMultiplier, + } as Model); } else { - merged.push(customModel); + merged.push(finalizeCustomModel(customModel, { useDefaults: true })); } } return merged; @@ -935,22 +1046,31 @@ export class ModelRegistry { if (discovered.length === 0) { return; } - const merged = this.#mergeCustomModels( - this.#models, + const discoveredModels = this.#applyHardcodedModelPolicies( discovered.map(model => { - const existing = - this.find(model.provider, model.id) ?? - this.#models.find(candidate => candidate.provider === model.provider); - return existing + const existing = this.find(model.provider, model.id); + if (existing) { + return { + ...model, + baseUrl: existing.baseUrl, + headers: existing.headers ? { ...existing.headers, ...model.headers } : model.headers, + }; + } + const providerOverride = this.#providerOverrides.get(model.provider); + return providerOverride ? { ...model, - baseUrl: existing.baseUrl, - headers: existing.headers ? { ...existing.headers, ...model.headers } : model.headers, + baseUrl: providerOverride.baseUrl ?? model.baseUrl, + headers: providerOverride.headers + ? { ...model.headers, ...providerOverride.headers } + : model.headers, } : model; }), ); - this.#models = this.#applyHardcodedModelPolicies(this.#applyModelOverrides(merged, this.#modelOverrides)); + const resolved = this.#mergeResolvedModels(this.#models, discoveredModels); + const combined = this.#mergeCustomModels(resolved, this.#customModelOverlays); + this.#models = this.#applyModelOverrides(combined, this.#modelOverrides); } async #discoverProviderModels( @@ -1454,8 +1574,8 @@ export class ModelRegistry { }); } - #parseModels(config: ModelsConfig): Model[] { - const models: Model[] = []; + #parseModels(config: ModelsConfig): CustomModelOverlay[] { + const models: CustomModelOverlay[] = []; for (const [providerName, providerConfig] of Object.entries(config.providers)) { const modelDefs = providerConfig.models ?? []; @@ -1464,7 +1584,7 @@ export class ModelRegistry { this.#customProviderApiKeys.set(providerName, providerConfig.apiKey); } for (const modelDef of modelDefs) { - const model = buildCustomModel( + const model = buildCustomModelOverlay( providerName, providerConfig.baseUrl!, providerConfig.api as Api | undefined, @@ -1473,7 +1593,6 @@ export class ModelRegistry { providerConfig.authHeader, providerConfig.compat, modelDef as CustomModelDefinitionLike, - { useDefaults: true }, ); if (!model) continue; models.push(model); @@ -1635,7 +1754,7 @@ export class ModelRegistry { config.authHeader, config.compat, modelDef as CustomModelDefinitionLike, - { useDefaults: false }, + { useDefaults: true }, ); if (!model) { throw new Error(`Provider ${providerName}, model ${modelDef.id}: no "api" specified.`); diff --git a/packages/coding-agent/test/model-registry.test.ts b/packages/coding-agent/test/model-registry.test.ts index 77478029c..c40cc8e92 100644 --- a/packages/coding-agent/test/model-registry.test.ts +++ b/packages/coding-agent/test/model-registry.test.ts @@ -50,7 +50,13 @@ describe("ModelRegistry", () => { /** Create minimal provider config */ function providerConfig( baseUrl: string, - models: Array<{ id: string; name?: string; reasoning?: boolean; thinking?: ThinkingConfig }>, + models: Array<{ + id: string; + name?: string; + reasoning?: boolean; + thinking?: ThinkingConfig; + contextWindow?: number; + }>, api: string = "anthropic-messages", ) { return { @@ -64,7 +70,7 @@ describe("ModelRegistry", () => { thinking: m.thinking, input: ["text"], cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 }, - contextWindow: 100000, + contextWindow: m.contextWindow ?? 100000, maxTokens: 8000, })), }; @@ -88,6 +94,19 @@ describe("ModelRegistry", () => { fs.writeFileSync(modelsJsonPath, JSON.stringify({ providers })); } + function mockOpenAiCompatibleModels(url: string, modelIds: string[]) { + return hookFetch(input => { + const requestUrl = String(input); + if (requestUrl === url) { + return new Response(JSON.stringify({ data: modelIds.map(id => ({ id })) }), { + status: 200, + headers: { "Content-Type": "application/json" }, + }); + } + throw new Error(`Unexpected URL: ${requestUrl}`); + }); + } + describe("baseUrl override (no custom models)", () => { test("overriding baseUrl keeps all built-in models", () => { writeRawModelsJson({ @@ -422,13 +441,221 @@ describe("ModelRegistry", () => { expect(registry.find("openai", "gpt-5.4")?.contextWindow).toBe(1_000_000); }); - test("custom gpt-5.4 replacement also applies the hardcoded context window policy", () => { - writeModelsJson({ - openai: providerConfig("https://my-proxy.example.com/v1", [{ id: "gpt-5.4" }], "openai-responses"), + test("custom gpt-5.4 replacement keeps the hardcoded context window when contextWindow is omitted", () => { + writeRawModelsJson({ + openai: { + baseUrl: "https://my-proxy.example.com/v1", + apiKey: "TEST_KEY", + api: "openai-responses", + models: [{ id: "gpt-5.4" }], + }, }); const registry = new ModelRegistry(authStorage, modelsJsonPath); - expect(registry.find("openai", "gpt-5.4")?.contextWindow).toBe(1_000_000); + const model = registry.find("openai", "gpt-5.4"); + expect(model?.contextWindow).toBe(1_000_000); + expect(model?.baseUrl).toBe("https://my-proxy.example.com/v1"); + }); + + test("custom-only gpt-5.4 provider keeps the hardcoded context window when contextWindow is omitted", () => { + writeRawModelsJson({ + "my-proxy": { + baseUrl: "https://my-proxy.example.com/v1", + apiKey: "TEST_KEY", + api: "openai-responses", + models: [{ id: "gpt-5.4" }], + }, + }); + + const registry = new ModelRegistry(authStorage, modelsJsonPath); + const model = registry.find("my-proxy", "gpt-5.4"); + expect(model?.contextWindow).toBe(1_000_000); + expect(model?.baseUrl).toBe("https://my-proxy.example.com/v1"); + }); + + test("custom gpt-5.4 replacement preserves its explicit context window", () => { + writeModelsJson({ + openai: providerConfig( + "https://my-proxy.example.com/v1", + [{ id: "gpt-5.4", contextWindow: 256000 }], + "openai-responses", + ), + }); + + const registry = new ModelRegistry(authStorage, modelsJsonPath); + expect(registry.find("openai", "gpt-5.4")?.contextWindow).toBe(256000); + }); + + test("modelOverrides can still patch a custom gpt-5.4 replacement", () => { + writeRawModelsJson({ + openai: { + baseUrl: "https://my-proxy.example.com/v1", + apiKey: "TEST_KEY", + api: "openai-responses", + models: [ + { + id: "gpt-5.4", + name: "gpt-5.4", + reasoning: false, + input: ["text"], + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 }, + contextWindow: 256000, + maxTokens: 128000, + }, + ], + modelOverrides: { + "gpt-5.4": { + contextWindow: 512000, + }, + }, + }, + }); + + const registry = new ModelRegistry(authStorage, modelsJsonPath); + expect(registry.find("openai", "gpt-5.4")?.contextWindow).toBe(512000); + }); + + test("discoverable bundled replacement survives refresh", async () => { + writeModelsJson({ + openai: providerConfig( + "https://my-proxy.example.com/v1", + [{ id: "gpt-5.4", name: "Proxy GPT-5.4", contextWindow: 256000 }], + "openai-responses", + ), + }); + const registry = new ModelRegistry(authStorage, modelsJsonPath); + expect(registry.find("openai", "gpt-5.4")?.name).toBe("Proxy GPT-5.4"); + expect(registry.find("openai", "gpt-5.4")?.contextWindow).toBe(256000); + + using _hook = mockOpenAiCompatibleModels("https://my-proxy.example.com/v1/models", ["gpt-5.4"]); + await registry.refreshProvider("openai", "online"); + + const model = registry.find("openai", "gpt-5.4"); + expect(model?.name).toBe("Proxy GPT-5.4"); + expect(model?.contextWindow).toBe(256000); + expect(model?.baseUrl).toBe("https://my-proxy.example.com/v1"); + }); + + test("discoverable custom-only gpt-5.4 survives refresh", async () => { + writeRawModelsJson({ + "custom-local": { + baseUrl: "http://127.0.0.1:8080", + apiKey: "TEST_KEY", + api: "openai-responses", + discovery: { type: "llama.cpp" }, + models: [{ id: "gpt-5.4" }], + }, + }); + const registry = new ModelRegistry(authStorage, modelsJsonPath); + expect(registry.find("custom-local", "gpt-5.4")?.contextWindow).toBe(1_000_000); + + using _hook = mockOpenAiCompatibleModels("http://127.0.0.1:8080/models", ["gpt-5.4"]); + await registry.refreshProvider("custom-local", "online"); + + const model = registry.find("custom-local", "gpt-5.4"); + expect(model?.contextWindow).toBe(1_000_000); + expect(model?.baseUrl).toBe("http://127.0.0.1:8080"); + }); + + test("discoverable custom compat survives refresh", async () => { + writeRawModelsJson({ + openai: { + baseUrl: "https://my-proxy.example.com/v1", + apiKey: "TEST_KEY", + api: "openai-responses", + models: [ + { + id: "gpt-5.4", + compat: { + extraBody: { source: "proxy" }, + }, + }, + ], + }, + }); + const registry = new ModelRegistry(authStorage, modelsJsonPath); + expect(registry.find("openai", "gpt-5.4")?.compat?.extraBody).toEqual({ source: "proxy" }); + + using _hook = mockOpenAiCompatibleModels("https://my-proxy.example.com/v1/models", ["gpt-5.4"]); + await registry.refreshProvider("openai", "online"); + + expect(registry.find("openai", "gpt-5.4")?.compat?.extraBody).toEqual({ source: "proxy" }); + }); + + test("modelOverrides still apply after discoverable refresh", async () => { + writeRawModelsJson({ + openai: { + baseUrl: "https://my-proxy.example.com/v1", + apiKey: "TEST_KEY", + api: "openai-responses", + models: [ + { + id: "gpt-5.4", + contextWindow: 256000, + }, + ], + modelOverrides: { + "gpt-5.4": { + contextWindow: 512000, + }, + }, + }, + }); + const registry = new ModelRegistry(authStorage, modelsJsonPath); + expect(registry.find("openai", "gpt-5.4")?.contextWindow).toBe(512000); + + using _hook = mockOpenAiCompatibleModels("https://my-proxy.example.com/v1/models", ["gpt-5.4"]); + await registry.refreshProvider("openai", "online"); + + expect(registry.find("openai", "gpt-5.4")?.contextWindow).toBe(512000); + }); + + test("newly discovered ids inherit provider fields, not another model's custom fields", async () => { + writeRawModelsJson({ + openai: { + baseUrl: "https://provider.example.com/v1", + headers: { "X-Provider": "provider" }, + apiKey: "TEST_KEY", + api: "openai-responses", + models: [ + { + id: "gpt-5.4", + baseUrl: "https://special.example.com/v1", + headers: { "X-Model": "special" }, + }, + ], + }, + }); + const registry = new ModelRegistry(authStorage, modelsJsonPath); + expect(registry.find("openai", "gpt-5.4")?.baseUrl).toBe("https://special.example.com/v1"); + + using _hook = mockOpenAiCompatibleModels("https://provider.example.com/v1/models", ["gpt-5.4", "gpt-5.5"]); + await registry.refreshProvider("openai", "online"); + + const discovered = registry.find("openai", "gpt-5.5"); + expect(discovered?.baseUrl).toBe("https://provider.example.com/v1"); + expect(discovered?.headers?.["X-Provider"]).toBe("provider"); + expect(discovered?.headers?.["X-Model"]).toBeUndefined(); + }); + + test("provider compat overlays preserve bundled model compat", () => { + writeRawModelsJson({ + "minimax-code": { + baseUrl: "https://proxy.example.com/v1", + apiKey: "TEST_KEY", + api: "openai-completions", + compat: { + extraBody: { source: "proxy" }, + }, + models: [{ id: "MiniMax-M2.5" }], + }, + }); + + const registry = new ModelRegistry(authStorage, modelsJsonPath); + const model = registry.find("minimax-code", "MiniMax-M2.5"); + expect(model?.compat?.thinkingFormat).toBe("zai"); + expect(model?.compat?.reasoningContentField).toBe("reasoning_content"); + expect(model?.compat?.extraBody).toEqual({ source: "proxy" }); }); test("removing custom models from models.json keeps built-in provider models", async () => {