Keep runtime provider overrides on overlay refresh
This commit is contained in:
@@ -1672,6 +1672,16 @@ export class ModelRegistry {
|
||||
compat: override.compat ? mergeCompat(baseOverride?.compat, override.compat) : baseOverride?.compat,
|
||||
};
|
||||
}
|
||||
#applyProviderTransportOverride<T extends { baseUrl?: string; headers?: Record<string, string> }>(
|
||||
entry: T,
|
||||
override: Pick<ProviderOverride, "baseUrl" | "headers">,
|
||||
): T {
|
||||
return {
|
||||
...entry,
|
||||
baseUrl: override.baseUrl ?? entry.baseUrl,
|
||||
headers: override.headers ? { ...entry.headers, ...override.headers } : entry.headers,
|
||||
};
|
||||
}
|
||||
#applyModelOverrides(models: Model<Api>[], overrides: Map<string, Map<string, ModelOverride>>): Model<Api>[] {
|
||||
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();
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
|
||||
|
||||
Reference in New Issue
Block a user