diff --git a/src/api/providers/__tests__/litellm.test.ts b/src/api/providers/__tests__/litellm.test.ts index bc20156d29..efe5c20201 100644 --- a/src/api/providers/__tests__/litellm.test.ts +++ b/src/api/providers/__tests__/litellm.test.ts @@ -77,11 +77,11 @@ describe("LiteLLMHandler", () => { it("returns correct model info when modelId is provided and found in getModels", async () => { const handler = new LiteLLMHandler(defaultMockOptions) const result = await handler.fetchModel() - expect(mockGetModels).toHaveBeenCalledWith( - "litellm", - defaultMockOptions.litellmApiKey, - defaultMockOptions.litellmBaseUrl, - ) + expect(mockGetModels).toHaveBeenCalledWith({ + provider: "litellm", + apiKey: defaultMockOptions.litellmApiKey, + baseUrl: defaultMockOptions.litellmBaseUrl, + }) expect(result).toEqual({ id: defaultMockOptions.litellmModelId, info: mockModelInfo }) }) diff --git a/src/api/providers/fetchers/modelCache.ts b/src/api/providers/fetchers/modelCache.ts index 175a6b2e15..5fad1f73c0 100644 --- a/src/api/providers/fetchers/modelCache.ts +++ b/src/api/providers/fetchers/modelCache.ts @@ -5,7 +5,7 @@ import NodeCache from "node-cache" import { ContextProxy } from "../../../core/config/ContextProxy" import { getCacheDirectoryPath } from "../../../shared/storagePathManager" -import { RouterName, ModelRecord } from "../../../shared/api" +import { RouterName, ModelRecord, GetModelsOptions } from "../../../shared/api" import { fileExistsAtPath } from "../../../utils/fs" import { getOpenRouterModels } from "./openrouter" @@ -30,18 +30,6 @@ async function readModels(router: RouterName): Promise return exists ? JSON.parse(await fs.readFile(filePath, "utf8")) : undefined } -/** - * Options for fetching models from different routers. - * This is a discriminated union type where the router property determines - * which other properties are required. - */ -export type GetModelsOptions = - | { router: "openrouter" } - | { router: "glama" } - | { router: "requesty"; apiKey?: string } - | { router: "unbound"; apiKey?: string } - | { router: "litellm"; apiKey: string; baseUrl: string } - /** * Get models from the cache or fetch them from the provider and cache them. * There are two caches: @@ -52,14 +40,14 @@ export type GetModelsOptions = * @returns The models from the cache or the fetched models. */ export const getModels = async (options: GetModelsOptions): Promise => { - const { router } = options - let models = memoryCache.get(router) + const { provider } = options + let models = memoryCache.get(provider) if (models) { return models } try { - switch (router) { + switch (provider) { case "openrouter": models = await getOpenRouterModels() break @@ -80,26 +68,26 @@ export const getModels = async (options: GetModelsOptions): Promise break default: // Ensures router is exhaustively checked if RouterName is a strict union - const exhaustiveCheck: never = router + const exhaustiveCheck: never = provider throw new Error(`Unknown router: ${exhaustiveCheck}`) } // Cache the fetched models (even if empty, to signify a successful fetch with no models) - memoryCache.set(router, models) - await writeModels(router, models).catch((err) => - console.error(`[getModels] Error writing ${router} models to file cache:`, err), + memoryCache.set(provider, models) + await writeModels(provider, models).catch((err) => + console.error(`[getModels] Error writing ${provider} models to file cache:`, err), ) try { - models = await readModels(router) + models = await readModels(provider) // console.log(`[getModels] read ${router} models from file cache`) } catch (error) { - console.error(`[getModels] error reading ${router} models from file cache`, error) + console.error(`[getModels] error reading ${provider} models from file cache`, error) } return models || {} } catch (error) { // Log the error and re-throw it so the caller can handle it (e.g., show a UI message). - console.error(`[getModels] Failed to fetch models for ${router}:`, error) + console.error(`[getModels] Failed to fetch models in modelCache for ${provider}:`, error) throw error // Re-throw the original error to be handled by the caller. } diff --git a/src/api/providers/litellm.ts b/src/api/providers/litellm.ts index be88ede5f6..5b2e938683 100644 --- a/src/api/providers/litellm.ts +++ b/src/api/providers/litellm.ts @@ -19,7 +19,7 @@ export class LiteLLMHandler extends RouterProvider implements SingleCompletionHa options, name: "litellm", baseURL: `${options.litellmBaseUrl || "http://localhost:4000"}`, - apiKey: options.litellmApiKey || "dummy-key", + apiKey: options.litellmApiKey || "sk-1234", modelId: options.litellmModelId, defaultModelId: litellmDefaultModelId, defaultModelInfo: litellmDefaultModelInfo, diff --git a/src/api/providers/openrouter.ts b/src/api/providers/openrouter.ts index 2d9c7f8b8a..8383a423b3 100644 --- a/src/api/providers/openrouter.ts +++ b/src/api/providers/openrouter.ts @@ -171,7 +171,7 @@ export class OpenRouterHandler extends BaseProvider implements SingleCompletionH public async fetchModel() { const [models, endpoints] = await Promise.all([ - getModels("openrouter"), + getModels({ provider: "openrouter" }), getModelEndpoints({ router: "openrouter", modelId: this.options.openRouterModelId, diff --git a/src/api/providers/requesty.ts b/src/api/providers/requesty.ts index fe8bba7e6e..130e5be10f 100644 --- a/src/api/providers/requesty.ts +++ b/src/api/providers/requesty.ts @@ -45,7 +45,7 @@ export class RequestyHandler extends BaseProvider implements SingleCompletionHan } public async fetchModel() { - this.models = await getModels("requesty") + this.models = await getModels({ provider: "requesty", apiKey: this.options.requestyApiKey }) return this.getModel() } diff --git a/src/api/providers/router-provider.ts b/src/api/providers/router-provider.ts index 548ce6cdf0..2a472eb9b4 100644 --- a/src/api/providers/router-provider.ts +++ b/src/api/providers/router-provider.ts @@ -1,6 +1,6 @@ import OpenAI from "openai" -import { ApiHandlerOptions, RouterName, ModelRecord, ModelInfo } from "../../shared/api" +import { ApiHandlerOptions, RouterName, ModelRecord, ModelInfo, GetModelsOptions } from "../../shared/api" import { BaseProvider } from "./base-provider" import { getModels } from "./fetchers/modelCache" @@ -51,7 +51,31 @@ export abstract class RouterProvider extends BaseProvider { } public async fetchModel() { - this.models = await getModels(this.name, this.apiKey, this.baseURL) + // Create the appropriate options based on router type + let options: GetModelsOptions + + switch (this.name) { + case "openrouter": + options = { provider: "openrouter" } + break + case "glama": + options = { provider: "glama" } + break + case "requesty": + options = { provider: "requesty", apiKey: this.apiKey } + break + case "unbound": + options = { provider: "unbound", apiKey: this.apiKey } + break + case "litellm": + options = { provider: "litellm", apiKey: this.apiKey, baseUrl: this.baseURL } + break + default: + const exhaustiveCheck: never = this.name + throw new Error(`Unknown provider: ${exhaustiveCheck}`) + } + + this.models = await getModels(options) return this.getModel() } diff --git a/src/core/webview/__tests__/webviewMessageHandler.test.ts b/src/core/webview/__tests__/webviewMessageHandler.test.ts index be457871c5..d619bab202 100644 --- a/src/core/webview/__tests__/webviewMessageHandler.test.ts +++ b/src/core/webview/__tests__/webviewMessageHandler.test.ts @@ -33,10 +33,11 @@ describe("webviewMessageHandler", () => { describe("requestRouterModels", () => { test("handles all successful model fetches correctly", async () => { // Mock all getModels calls to succeed with different data - ;(getModels as jest.Mock).mockImplementation((router) => { + ;(getModels as jest.Mock).mockImplementation((options) => { + const provider = options.provider return Promise.resolve({ - [`${router}-model-1`]: { name: `${router} Model 1` }, - [`${router}-model-2`]: { name: `${router} Model 2` }, + [`${provider}-model-1`]: { name: `${provider} Model 1` }, + [`${provider}-model-2`]: { name: `${provider} Model 2` }, }) }) @@ -75,14 +76,15 @@ describe("webviewMessageHandler", () => { test("handles some failed model fetches correctly", async () => { // Mock some getModels calls to succeed and others to fail - ;(getModels as jest.Mock).mockImplementation((router) => { - if (router === "openrouter" || router === "litellm") { + ;(getModels as jest.Mock).mockImplementation((options) => { + const provider = options.provider + if (provider === "openrouter" || provider === "litellm") { return Promise.resolve({ - [`${router}-model-1`]: { name: `${router} Model 1` }, + [`${provider}-model-1`]: { name: `${provider} Model 1` }, }) } - // For other routers, throw an error - return Promise.reject(new Error(`Failed to fetch ${router} models`)) + // For other providers, throw an error + return Promise.reject(new Error(`Failed to fetch ${provider} models`)) }) // Call the handler @@ -129,4 +131,46 @@ describe("webviewMessageHandler", () => { }) }) }) + + describe("requestProviderModels", () => { + test("when getModels succeeds, it posts a providerModelsResponse with models", async () => { + const mockLiteLLMModels = { "litellm-model-1": { name: "LiteLLM Model 1" } } + ;(getModels as jest.Mock).mockResolvedValueOnce(mockLiteLLMModels) + + await webviewMessageHandler(mockProvider as any, { + type: "requestProviderModels", + payload: { provider: "litellm", apiKey: "test-key", baseUrl: "test-url" }, + }) + + expect(mockProvider.postMessageToWebview).toHaveBeenCalledWith({ + type: "providerModelsResponse", + payload: { + provider: "litellm", + models: mockLiteLLMModels, + error: undefined, // Explicitly check error is undefined on success + }, + }) + expect(getModels).toHaveBeenCalledWith({ provider: "litellm", apiKey: "test-key", baseUrl: "test-url" }) + }) + + test("when getModels fails, it posts a providerModelsResponse with an error and empty models", async () => { + const errorMessage = "Failed to fetch LiteLLM models: No response from server." + ;(getModels as jest.Mock).mockRejectedValueOnce(new Error(errorMessage)) + + await webviewMessageHandler(mockProvider as any, { + type: "requestProviderModels", + payload: { provider: "litellm", apiKey: "test-key", baseUrl: "test-url" }, + }) + + expect(mockProvider.postMessageToWebview).toHaveBeenCalledWith({ + type: "providerModelsResponse", + payload: { + provider: "litellm", + models: {}, + error: errorMessage, + }, + }) + expect(getModels).toHaveBeenCalledWith({ provider: "litellm", apiKey: "test-key", baseUrl: "test-url" }) + }) + }) }) diff --git a/src/core/webview/webviewMessageHandler.ts b/src/core/webview/webviewMessageHandler.ts index 92fa8218b6..e2259b5866 100644 --- a/src/core/webview/webviewMessageHandler.ts +++ b/src/core/webview/webviewMessageHandler.ts @@ -6,7 +6,7 @@ import * as vscode from "vscode" import { ClineProvider } from "./ClineProvider" import { Language, ProviderSettings } from "../../schemas" import { changeLanguage, t } from "../../i18n" -import { RouterName, toRouterName } from "../../shared/api" +import { RouterName, toRouterName, ModelRecord } from "../../shared/api" import { supportPrompt } from "../../shared/support-prompt" import { checkoutDiffPayloadSchema, checkoutRestorePayloadSchema, WebviewMessage } from "../../shared/WebviewMessage" @@ -34,18 +34,12 @@ import { TelemetrySetting } from "../../shared/TelemetrySetting" import { getWorkspacePath } from "../../utils/path" import { Mode, defaultModeSlug } from "../../shared/modes" import { GlobalState } from "../../schemas" -import { getModels, flushModels } from "../../api/providers/fetchers/modelCache" +import { flushModels, getModels } from "../../api/providers/fetchers/modelCache" +import { GetModelsOptions } from "../../shared/api" import { generateSystemPrompt } from "./generateSystemPrompt" const ALLOWED_VSCODE_SETTINGS = new Set(["terminal.integrated.inheritEnv"]) -// Define a type for the payload of requestProviderModels for clarity -interface RequestProviderModelsPayload { - provider: RouterName // Should be 'litellm' or 'requesty' here - apiKey?: string - baseUrl?: string -} - export const webviewMessageHandler = async (provider: ClineProvider, message: WebviewMessage) => { // Utility functions provided for concise get/update of global state via contextProxy API. const getGlobalState = (key: K) => provider.contextProxy.getValue(key) @@ -284,42 +278,54 @@ export const webviewMessageHandler = async (provider: ClineProvider, message: We await flushModels(routerNameFlush) break case "requestProviderModels": { - const payload = message.payload as RequestProviderModelsPayload | undefined - if (!payload || !payload.provider) { + const optionsFromPayload = message.payload as any // Check payload structure first + + if ( + typeof optionsFromPayload !== "object" || + optionsFromPayload === null || + typeof optionsFromPayload.provider !== "string" || + !optionsFromPayload.provider + ) { + const providerNameForError = + typeof optionsFromPayload?.provider === "string" && optionsFromPayload.provider + ? (optionsFromPayload.provider as RouterName) + : ("unknown" as RouterName) + provider.postMessageToWebview({ type: "providerModelsResponse", payload: { - provider: payload?.provider || ("unknown" as RouterName), - error: "Invalid payload for requestProviderModels", + provider: providerNameForError, + error: "Invalid payload for requestProviderModels: payload must be an object with a valid 'provider' string property.", }, }) break } - const targetProvider = payload.provider as RouterName - let models = {} + const options = optionsFromPayload as GetModelsOptions // Now cast to GetModelsOptions + + let models: ModelRecord = {} let error: string | undefined try { - await flushModels(targetProvider) - - models = await getModels(targetProvider, payload.apiKey, payload.baseUrl) + await flushModels(options.provider) + models = await getModels(options) } catch (e: any) { - error = e.message || `Failed to fetch models for ${targetProvider}. Check console for details.` + error = + e.message || + `Failed to fetch models in webviewMessageHandler requestProviderModels for ${options.provider}. Check console for details.` models = {} } provider.postMessageToWebview({ type: "providerModelsResponse", - payload: { provider: targetProvider, models, error }, + payload: { provider: options.provider, models, error }, }) break } case "requestRouterModels": const { apiConfiguration } = await provider.getState() - // Handle each model fetch independently to avoid one failure affecting others - const routerModels = { + const routerModels: Partial> = { openrouter: {}, requesty: {}, glama: {}, @@ -327,53 +333,58 @@ export const webviewMessageHandler = async (provider: ClineProvider, message: We litellm: {}, } - // Helper function to safely fetch models - const safeGetModels = async (router: RouterName, apiKey?: string, baseUrl?: string) => { + const safeGetModels = async (options: GetModelsOptions): Promise => { try { - return await getModels(router, apiKey, baseUrl) + return await getModels(options) } catch (error) { - console.error(`Failed to fetch models for ${router}:`, error) - return {} // Return empty object on failure + console.error( + `Failed to fetch models in webviewMessageHandler requestRouterModels for ${options.provider}:`, + error, + ) + return {} } } - // Fetch all models in parallel but handle failures independently + const modelFetchPromises: Array<{ key: RouterName; options: GetModelsOptions }> = [ + { key: "openrouter", options: { provider: "openrouter" } }, + { key: "requesty", options: { provider: "requesty", apiKey: apiConfiguration.requestyApiKey } }, + { key: "glama", options: { provider: "glama" } }, + { key: "unbound", options: { provider: "unbound", apiKey: apiConfiguration.unboundApiKey } }, + ] + + const litellmApiKey = apiConfiguration.litellmApiKey + const litellmBaseUrl = apiConfiguration.litellmBaseUrl + + if (litellmApiKey && litellmBaseUrl) { + modelFetchPromises.push({ + key: "litellm", + options: { provider: "litellm", apiKey: litellmApiKey, baseUrl: litellmBaseUrl }, + }) + } + const results = await Promise.allSettled( - [ - { key: "openrouter", promise: safeGetModels("openrouter", apiConfiguration.openRouterApiKey) }, - { key: "requesty", promise: safeGetModels("requesty", apiConfiguration.requestyApiKey) }, - { key: "glama", promise: safeGetModels("glama", apiConfiguration.glamaApiKey) }, - { key: "unbound", promise: safeGetModels("unbound", apiConfiguration.unboundApiKey) }, - { - key: "litellm", - promise: safeGetModels( - "litellm", - apiConfiguration.litellmApiKey, - apiConfiguration.litellmBaseUrl, - ), - }, - ].map(async ({ key, promise }) => { + modelFetchPromises.map(async ({ key, options }) => { try { - const models = await promise + const models = await safeGetModels(options) return { key, models } } catch (error) { - console.error(`Error in router models fetch for ${key}:`, error) + console.error(`Outer catch: Error in router models fetch for ${key}:`, error) return { key, models: {} } } }), ) - // Process results and assign to routerModels results.forEach((result) => { if (result.status === "fulfilled") { - const key = result.value.key as keyof typeof routerModels - routerModels[key] = result.value.models + routerModels[result.value.key] = result.value.models + } else { + console.error("A model fetching promise was rejected:", result.reason) } }) provider.postMessageToWebview({ type: "routerModels", - routerModels, + routerModels: routerModels as Record, }) break case "requestOpenAiModels": diff --git a/src/shared/WebviewMessage.ts b/src/shared/WebviewMessage.ts index 1cc461cd47..d2c02e45d5 100644 --- a/src/shared/WebviewMessage.ts +++ b/src/shared/WebviewMessage.ts @@ -1,6 +1,6 @@ import { z } from "zod" -import { ProviderSettings, RouterName } from "./api" +import { ProviderSettings, GetModelsOptions } from "./api" import { Mode, PromptComponent, ModeConfig } from "./modes" export type ClineAskResponse = "yesButtonClicked" | "noButtonClicked" | "messageResponse" @@ -153,7 +153,7 @@ export interface WebviewMessage { slug?: string modeConfig?: ModeConfig timeout?: number - payload?: WebViewMessagePayload | RequestProviderModelsPayload + payload?: WebViewMessagePayload source?: "global" | "project" requestId?: string ids?: string[] @@ -179,11 +179,4 @@ export const checkoutRestorePayloadSchema = z.object({ export type CheckpointRestorePayload = z.infer -export type WebViewMessagePayload = CheckpointDiffPayload | CheckpointRestorePayload - -// Payload for requestProviderModels -export interface RequestProviderModelsPayload { - provider: RouterName - apiKey?: string - baseUrl?: string -} +export type WebViewMessagePayload = CheckpointDiffPayload | CheckpointRestorePayload | GetModelsOptions diff --git a/src/shared/api.ts b/src/shared/api.ts index dd8bd5bef4..6a846c29d3 100644 --- a/src/shared/api.ts +++ b/src/shared/api.ts @@ -1792,3 +1792,15 @@ export function toRouterName(value?: string): RouterName { export type ModelRecord = Record export type RouterModels = Record + +/** + * Options for fetching models from different providers. + * This is a discriminated union type where the provider property determines + * which other properties are required. + */ +export type GetModelsOptions = + | { provider: "openrouter" } + | { provider: "glama" } + | { provider: "requesty"; apiKey?: string } + | { provider: "unbound"; apiKey?: string } + | { provider: "litellm"; apiKey: string; baseUrl: string } diff --git a/webview-ui/src/components/settings/providers/LiteLLM.tsx b/webview-ui/src/components/settings/providers/LiteLLM.tsx index 413db6c812..758d35d86c 100644 --- a/webview-ui/src/components/settings/providers/LiteLLM.tsx +++ b/webview-ui/src/components/settings/providers/LiteLLM.tsx @@ -39,12 +39,18 @@ export const LiteLLM = ({ apiConfiguration, setApiConfigurationField, routerMode const handleRefreshModels = () => { setRefreshStatus("loading") setRefreshError(undefined) + + // Due to the button's disabled state logic, litellmApiKey and litellmBaseUrl are guaranteed to be non-empty strings here. + // We use non-null assertions (!) to reflect this guarantee for type safety. + const key = apiConfiguration.litellmApiKey! + const url = apiConfiguration.litellmBaseUrl! + const message: WebviewMessage = { type: "requestProviderModels", payload: { provider: "litellm", - apiKey: apiConfiguration.litellmApiKey, - baseUrl: apiConfiguration.litellmBaseUrl || "http://localhost:4000", + apiKey: key, + baseUrl: url, }, } vscode.postMessage(message) @@ -53,21 +59,30 @@ export const LiteLLM = ({ apiConfiguration, setApiConfigurationField, routerMode // Listen for model refresh responses using useEvent useEvent("message", (event: MessageEvent) => { const message = event.data - if (message.type === "providerModelsResponse" && message.payload && message.payload.provider === "litellm") { - if (message.payload.error) { - setRefreshStatus("error") - setRefreshError(message.payload.error) + if (message.type === "providerModelsResponse") { + if (message.payload && message.payload.provider === "litellm") { + if (message.payload.error) { + console.log("LiteLLM.tsx: Error found in payload:", message.payload.error) + setRefreshStatus("error") + setRefreshError(message.payload.error) + } else { + setRefreshStatus("success") + // Parent (ApiOptions.tsx) will handle updating the routerModels prop for ModelPicker + } } else { - setRefreshStatus("success") - // Parent (ApiOptions.tsx) will handle updating the routerModels prop for ModelPicker + console.log( + "LiteLLM.tsx: Received providerModelsResponse but not for litellm or payload missing. Provider:", + message.payload?.provider, + ) } } }) + console.log("apiconfig1212", apiConfiguration) return ( <> @@ -90,7 +105,9 @@ export const LiteLLM = ({ apiConfiguration, setApiConfigurationField, routerMode