diff --git a/src/plugin-handlers/provider-config-handler.test.ts b/src/plugin-handlers/provider-config-handler.test.ts new file mode 100644 index 000000000..421dd5d0e --- /dev/null +++ b/src/plugin-handlers/provider-config-handler.test.ts @@ -0,0 +1,84 @@ +/// + +import { describe, expect, test } from "bun:test" +import { applyProviderConfig } from "./provider-config-handler" +import { createModelCacheState } from "../plugin-state" +import { clearVisionCapableModelsCache, readVisionCapableModelsCache } from "../shared/vision-capable-models-cache" + +describe("applyProviderConfig", () => { + test("caches vision-capable models from modalities and capabilities", () => { + // given + const modelCacheState = createModelCacheState() + const visionCapableModelsCache = modelCacheState.visionCapableModelsCache + if (!visionCapableModelsCache) { + throw new Error("visionCapableModelsCache should be initialized") + } + const config = { + provider: { + rundao: { + models: { + "public/qwen3.5-397b": { + modalities: { + input: ["text", "image"], + }, + }, + "public/text-only": { + modalities: { + input: ["text"], + }, + }, + }, + }, + google: { + models: { + "gemini-3-flash": { + capabilities: { + input: { + image: true, + }, + }, + }, + }, + }, + }, + } satisfies Record + + // when + applyProviderConfig({ config, modelCacheState }) + + // then + expect(Array.from(visionCapableModelsCache.keys())).toEqual([ + "rundao/public/qwen3.5-397b", + "google/gemini-3-flash", + ]) + expect(readVisionCapableModelsCache()).toEqual([ + { providerID: "rundao", modelID: "public/qwen3.5-397b" }, + { providerID: "google", modelID: "gemini-3-flash" }, + ]) + }) + + test("clears stale vision-capable models when provider config changes", () => { + // given + const modelCacheState = createModelCacheState() + const visionCapableModelsCache = modelCacheState.visionCapableModelsCache + if (!visionCapableModelsCache) { + throw new Error("visionCapableModelsCache should be initialized") + } + visionCapableModelsCache.set("stale/old-model", { + providerID: "stale", + modelID: "old-model", + }) + + // when + applyProviderConfig({ + config: { provider: {} }, + modelCacheState, + }) + + // then + expect(visionCapableModelsCache.size).toBe(0) + expect(readVisionCapableModelsCache()).toEqual([]) + }) +}) + +clearVisionCapableModelsCache() diff --git a/src/plugin-handlers/provider-config-handler.ts b/src/plugin-handlers/provider-config-handler.ts index 75964d20b..cc01508bd 100644 --- a/src/plugin-handlers/provider-config-handler.ts +++ b/src/plugin-handlers/provider-config-handler.ts @@ -1,10 +1,31 @@ -import type { ModelCacheState } from "../plugin-state"; +import type { ModelCacheState, VisionCapableModel } from "../plugin-state"; +import { setVisionCapableModelsCache } from "../shared/vision-capable-models-cache" type ProviderConfig = { options?: { headers?: Record }; - models?: Record; + models?: Record; }; +type ProviderModelConfig = { + limit?: { context?: number }; + modalities?: { + input?: string[]; + }; + capabilities?: { + input?: { + image?: boolean; + }; + }; +} + +function supportsImageInput(modelConfig: ProviderModelConfig | undefined): boolean { + if (modelConfig?.modalities?.input?.includes("image")) { + return true + } + + return modelConfig?.capabilities?.input?.image === true +} + export function applyProviderConfig(params: { config: Record; modelCacheState: ModelCacheState; @@ -17,6 +38,12 @@ export function applyProviderConfig(params: { params.modelCacheState.anthropicContext1MEnabled = anthropicBeta?.includes("context-1m") ?? false; + const visionCapableModelsCache = params.modelCacheState.visionCapableModelsCache + ?? new Map() + params.modelCacheState.visionCapableModelsCache = visionCapableModelsCache + visionCapableModelsCache.clear() + setVisionCapableModelsCache(visionCapableModelsCache) + if (!providers) return; for (const [providerID, providerConfig] of Object.entries(providers)) { @@ -24,6 +51,13 @@ export function applyProviderConfig(params: { if (!models) continue; for (const [modelID, modelConfig] of Object.entries(models)) { + if (supportsImageInput(modelConfig)) { + visionCapableModelsCache.set( + `${providerID}/${modelID}`, + { providerID, modelID }, + ) + } + const contextLimit = modelConfig?.limit?.context; if (!contextLimit) continue;