From 5e89736bc9e4128ada2faf5ace1880eddbcdccef Mon Sep 17 00:00:00 2001 From: Aidan Date: Fri, 24 Apr 2026 15:13:26 -0400 Subject: [PATCH] Keep runtime provider overrides on overlay refresh --- .../coding-agent/src/config/model-registry.ts | 29 +++++++++---- .../model-registry-runtime-provider.test.ts | 43 +++++++++++++++++++ 2 files changed, 63 insertions(+), 9 deletions(-) diff --git a/packages/coding-agent/src/config/model-registry.ts b/packages/coding-agent/src/config/model-registry.ts index eb3e44b47..41b31c312 100644 --- a/packages/coding-agent/src/config/model-registry.ts +++ b/packages/coding-agent/src/config/model-registry.ts @@ -1672,6 +1672,16 @@ export class ModelRegistry { compat: override.compat ? mergeCompat(baseOverride?.compat, override.compat) : baseOverride?.compat, }; } + #applyProviderTransportOverride }>( + entry: T, + override: Pick, + ): T { + return { + ...entry, + baseUrl: override.baseUrl ?? entry.baseUrl, + headers: override.headers ? { ...entry.headers, ...override.headers } : entry.headers, + }; + } #applyModelOverrides(models: Model[], overrides: Map>): Model[] { if (overrides.size === 0) return models; return models.map(model => { @@ -2075,18 +2085,19 @@ export class ModelRegistry { } if (config.baseUrl || config.headers) { - const nextRuntimeOverride = this.#mergeProviderOverride(this.#runtimeProviderOverrides.get(providerName), { - baseUrl: config.baseUrl, - headers: config.headers, - }); + const transportOverride = { baseUrl: config.baseUrl, headers: config.headers }; + const nextRuntimeOverride = this.#mergeProviderOverride( + this.#runtimeProviderOverrides.get(providerName), + transportOverride, + ); this.#runtimeProviderOverrides.set(providerName, nextRuntimeOverride); + this.#runtimeModelOverlays = this.#runtimeModelOverlays.map(overlay => { + if (overlay.provider !== providerName) return overlay; + return this.#applyProviderTransportOverride(overlay, transportOverride); + }); this.#models = this.#models.map(m => { if (m.provider !== providerName) return m; - return { - ...m, - baseUrl: config.baseUrl ?? m.baseUrl, - headers: config.headers ? { ...m.headers, ...config.headers } : m.headers, - }; + return this.#applyProviderTransportOverride(m, transportOverride); }); this.#rebuildCanonicalIndex(); } diff --git a/packages/coding-agent/test/model-registry-runtime-provider.test.ts b/packages/coding-agent/test/model-registry-runtime-provider.test.ts index ad119dfa0..02fdd1c0b 100644 --- a/packages/coding-agent/test/model-registry-runtime-provider.test.ts +++ b/packages/coding-agent/test/model-registry-runtime-provider.test.ts @@ -261,6 +261,49 @@ describe("ModelRegistry runtime provider registration", () => { expect(model?.api).toBe("openai-completions"); }); + test("runtime model overlays keep provider overrides across refresh cycles", async () => { + const registry = new ModelRegistry(authStorage, modelsJsonPath); + const runtimeHeader = "X-Runtime-Overlay-Header"; + const overrideBaseUrl = "https://runtime-overridden.example.com/v1"; + const modelId = "runtime-override-survivor"; + + registry.registerProvider( + "runtime-provider", + { + baseUrl: "https://runtime.example.com/v1", + apiKey: "RUNTIME_KEY", + api: "openai-completions", + models: [{ ...baseModel, id: modelId }], + }, + "ext://runtime", + ); + registry.registerProvider( + "runtime-provider", + { baseUrl: overrideBaseUrl, headers: { [runtimeHeader]: "runtime-header" } }, + "ext://runtime", + ); + + const modelAfterOverride = registry.find("runtime-provider", modelId); + expect(modelAfterOverride).toBeDefined(); + expect(modelAfterOverride?.baseUrl).toBe(overrideBaseUrl); + expect(modelAfterOverride?.headers?.[runtimeHeader]).toBe("runtime-header"); + + await registry.refresh("offline"); + const modelAfterRefresh = registry.find("runtime-provider", modelId); + expect(modelAfterRefresh).toBeDefined(); + expect(modelAfterRefresh?.baseUrl).toBe(overrideBaseUrl); + expect(modelAfterRefresh?.headers?.[runtimeHeader]).toBe("runtime-header"); + + await registry.refreshProvider("runtime-provider", "offline"); + const modelAfterProviderRefresh = registry.find("runtime-provider", modelId); + expect(modelAfterProviderRefresh).toBeDefined(); + expect(modelAfterProviderRefresh?.baseUrl).toBe(overrideBaseUrl); + expect(modelAfterProviderRefresh?.headers?.[runtimeHeader]).toBe("runtime-header"); + + registry.clearSourceRegistrations("ext://runtime"); + expect(registry.find("runtime-provider", modelId)).toBeUndefined(); + }); + test("extension-registered API keys survive refresh cycle for auth resolution", async () => { const registry = new ModelRegistry(authStorage, modelsJsonPath);