From 2328695101bd14eefdbf11ba757d9e8856c8c057 Mon Sep 17 00:00:00 2001 From: slytechnical Date: Fri, 16 May 2025 16:05:16 -0500 Subject: [PATCH] Ability to refresh litellm models by refresh button. Provider-specific model fetching. Success vs fail user feedback on model fetch responses --- src/api/providers/fetchers/litellm.ts | 24 ++++- src/api/providers/fetchers/modelCache.ts | 102 ++++++++---------- src/core/webview/webviewMessageHandler.ts | 43 +++++++- src/shared/ExtensionMessage.ts | 11 +- src/shared/WebviewMessage.ts | 12 ++- .../src/components/settings/ApiOptions.tsx | 66 ++++++++++-- .../components/settings/providers/LiteLLM.tsx | 76 ++++++++++++- .../settings/providers/Requesty.tsx | 79 ++++++++++---- webview-ui/src/i18n/locales/en/settings.json | 5 +- 9 files changed, 322 insertions(+), 96 deletions(-) diff --git a/src/api/providers/fetchers/litellm.ts b/src/api/providers/fetchers/litellm.ts index 5bb37898fb..1c61ae3439 100644 --- a/src/api/providers/fetchers/litellm.ts +++ b/src/api/providers/fetchers/litellm.ts @@ -7,6 +7,7 @@ import { ModelRecord } from "../../../shared/api" * @param apiKey The API key for the LiteLLM server * @param baseUrl The base URL of the LiteLLM server * @returns A promise that resolves to a record of model IDs to model info + * @throws Will throw an error if the request fails or the response is not as expected. */ export async function getLiteLLMModels(apiKey: string, baseUrl: string): Promise { try { @@ -18,7 +19,8 @@ export async function getLiteLLMModels(apiKey: string, baseUrl: string): Promise headers["Authorization"] = `Bearer ${apiKey}` } - const response = await axios.get(`${baseUrl}/v1/model/info`, { headers }) + // Added timeout to prevent indefinite hanging + const response = await axios.get(`${baseUrl}/v1/model/info`, { headers, timeout: 15000 }) const models: ModelRecord = {} // Process the model info from the response @@ -44,11 +46,25 @@ export async function getLiteLLMModels(apiKey: string, baseUrl: string): Promise description: `${modelName} via LiteLLM proxy`, } } + } else { + // If response.data.data is not in the expected format, consider it an error. + console.error("Error fetching LiteLLM models: Unexpected response format", response.data) + throw new Error("Failed to fetch LiteLLM models: Unexpected response format.") } return models - } catch (error) { - console.error("Error fetching LiteLLM models:", error) - return {} + } catch (error: any) { + console.error("Error fetching LiteLLM models:", error.message ? error.message : error) + if (axios.isAxiosError(error) && error.response) { + throw new Error( + `Failed to fetch LiteLLM models: ${error.response.status} ${error.response.statusText}. Check base URL and API key.`, + ) + } else if (axios.isAxiosError(error) && error.request) { + throw new Error( + "Failed to fetch LiteLLM models: No response from server. Check LiteLLM server status and base URL.", + ) + } else { + throw new Error(`Failed to fetch LiteLLM models: ${error.message || "An unknown error occurred."}`) + } } } diff --git a/src/api/providers/fetchers/modelCache.ts b/src/api/providers/fetchers/modelCache.ts index 9ab4b851fc..9758f490a9 100644 --- a/src/api/providers/fetchers/modelCache.ts +++ b/src/api/providers/fetchers/modelCache.ts @@ -6,7 +6,6 @@ import NodeCache from "node-cache" import { ContextProxy } from "../../../core/config/ContextProxy" import { getCacheDirectoryPath } from "../../../shared/storagePathManager" import { RouterName, ModelRecord } from "../../../shared/api" -import { fileExistsAtPath } from "../../../utils/fs" import { getOpenRouterModels } from "./openrouter" import { getRequestyModels } from "./requesty" @@ -22,14 +21,6 @@ async function writeModels(router: RouterName, data: ModelRecord) { await fs.writeFile(path.join(cacheDir, filename), JSON.stringify(data)) } -async function readModels(router: RouterName): Promise { - const filename = `${router}_models.json` - const cacheDir = await getCacheDirectoryPath(ContextProxy.instance.globalStorageUri.fsPath) - const filePath = path.join(cacheDir, filename) - const exists = await fileExistsAtPath(filePath) - return exists ? JSON.parse(await fs.readFile(filePath, "utf8")) : undefined -} - /** * Get models from the cache or fetch them from the provider and cache them. * There are two caches: @@ -46,58 +37,59 @@ export const getModels = async ( apiKey: string | undefined = undefined, baseUrl: string | undefined = undefined, ): Promise => { - let models = memoryCache.get(router) - - if (models) { - // console.log(`[getModels] NodeCache hit for ${router} -> ${Object.keys(models).length}`) - return models + // If this call is meant for a refresh (indicated by apiKey/baseUrl for specific routers), + // the memory cache should have been flushed by the caller (e.g., webviewMessageHandler). + // Otherwise, for general calls, check memory cache first. + const modelsFromMemory = memoryCache.get(router) + if (modelsFromMemory) { + return modelsFromMemory } - switch (router) { - case "openrouter": - models = await getOpenRouterModels() - break - case "requesty": - // Requesty models endpoint requires an API key for per-user custom policies - models = await getRequestyModels(apiKey) - break - case "glama": - models = await getGlamaModels() - break - case "unbound": - models = await getUnboundModels() - break - case "litellm": - if (apiKey && baseUrl) { - models = await getLiteLLMModels(apiKey, baseUrl) - } else { - models = {} - } - break - } - - if (Object.keys(models).length > 0) { - // console.log(`[getModels] API fetch for ${router} -> ${Object.keys(models).length}`) - memoryCache.set(router, models) - - try { - await writeModels(router, models) - // console.log(`[getModels] wrote ${router} models to file cache`) - } catch (error) { - console.error(`[getModels] error writing ${router} models to file cache`, error) + let fetchedModels: ModelRecord + try { + switch (router) { + case "openrouter": + fetchedModels = await getOpenRouterModels() + break + case "requesty": + // Assuming getRequestyModels will throw if apiKey is needed and not provided or invalid. + fetchedModels = await getRequestyModels(apiKey) + break + case "glama": + fetchedModels = await getGlamaModels() + break + case "unbound": + fetchedModels = await getUnboundModels() + break + case "litellm": + // getLiteLLMModels now throws on error. + // It needs baseUrl. apiKey is optional for the protocol but might be needed by the server. + console.log("litellm1212", baseUrl, apiKey) + if (!baseUrl || !apiKey) { + // This case should ideally be handled by the caller if baseUrl is strictly required. + // However, for robustness, if called without baseUrl for litellm, it would fail in getLiteLLMModels or here. + throw new Error("Base URL and api key are required for LiteLLM models.") + } + fetchedModels = await getLiteLLMModels(apiKey || "", baseUrl) + break + default: + // Ensures router is exhaustively checked if RouterName is a strict union + const exhaustiveCheck: never = router + throw new Error(`Unknown router: ${exhaustiveCheck}`) } - return models - } - - try { - models = await readModels(router) - // console.log(`[getModels] read ${router} models from file cache`) + // Cache the fetched models (even if empty, to signify a successful fetch with no models) + memoryCache.set(router, fetchedModels) + await writeModels(router, fetchedModels).catch((err) => + console.error(`[getModels] Error writing ${router} models to file cache:`, err), + ) + return fetchedModels } catch (error) { - console.error(`[getModels] error reading ${router} models from file cache`, 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) - return models ?? {} + throw error // Re-throw the original error to be handled by the caller. + } } /** diff --git a/src/core/webview/webviewMessageHandler.ts b/src/core/webview/webviewMessageHandler.ts index 9c8b90ea8a..b31496bdcb 100644 --- a/src/core/webview/webviewMessageHandler.ts +++ b/src/core/webview/webviewMessageHandler.ts @@ -39,6 +39,13 @@ 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) @@ -273,9 +280,41 @@ export const webviewMessageHandler = async (provider: ClineProvider, message: We await provider.resetState() break case "flushRouterModels": - const routerName: RouterName = toRouterName(message.text) - await flushModels(routerName) + const routerNameFlush: RouterName = toRouterName(message.text) + await flushModels(routerNameFlush) break + case "requestProviderModels": { + const payload = message.payload as RequestProviderModelsPayload | undefined + if (!payload || !payload.provider) { + provider.postMessageToWebview({ + type: "providerModelsResponse", + payload: { + provider: payload?.provider || ("unknown" as RouterName), + error: "Invalid payload for requestProviderModels", + }, + }) + break + } + + const targetProvider = payload.provider as RouterName + let models = {} + let error: string | undefined + + try { + await flushModels(targetProvider) + + models = await getModels(targetProvider, payload.apiKey, payload.baseUrl) + } catch (e: any) { + error = e.message || `Failed to fetch models for ${targetProvider}. Check console for details.` + models = {} + } + + provider.postMessageToWebview({ + type: "providerModelsResponse", + payload: { provider: targetProvider, models, error }, + }) + break + } case "requestRouterModels": const { apiConfiguration } = await provider.getState() diff --git a/src/shared/ExtensionMessage.ts b/src/shared/ExtensionMessage.ts index 6330556024..e5ef6e5a06 100644 --- a/src/shared/ExtensionMessage.ts +++ b/src/shared/ExtensionMessage.ts @@ -15,7 +15,7 @@ import { } from "../schemas" import { McpServer } from "./mcp" import { Mode } from "./modes" -import { RouterModels } from "./api" +import { RouterModels, ModelRecord, RouterName } from "./api" export type { ProviderSettingsEntry, ToolProgressStatus } @@ -69,6 +69,7 @@ export interface ExtensionMessage { | "setHistoryPreviewCollapsed" | "commandExecutionStatus" | "vsCodeSetting" + | "providerModelsResponse" text?: string action?: | "chatButtonClicked" @@ -107,6 +108,7 @@ export interface ExtensionMessage { error?: string setting?: string value?: any + payload?: ProviderModelsResponsePayload } export type ExtensionState = Pick< @@ -289,3 +291,10 @@ export interface ClineApiReqInfo { } export type ClineApiReqCancelReason = "streaming_failed" | "user_cancelled" + +// Payload for providerModelsResponse +export interface ProviderModelsResponsePayload { + provider: RouterName + models?: ModelRecord + error?: string +} diff --git a/src/shared/WebviewMessage.ts b/src/shared/WebviewMessage.ts index 22fe5c7d3e..1cc461cd47 100644 --- a/src/shared/WebviewMessage.ts +++ b/src/shared/WebviewMessage.ts @@ -1,6 +1,6 @@ import { z } from "zod" -import { ProviderSettings } from "./api" +import { ProviderSettings, RouterName } from "./api" import { Mode, PromptComponent, ModeConfig } from "./modes" export type ClineAskResponse = "yesButtonClicked" | "noButtonClicked" | "messageResponse" @@ -130,6 +130,7 @@ export interface WebviewMessage { | "searchFiles" | "toggleApiConfigPin" | "setHistoryPreviewCollapsed" + | "requestProviderModels" text?: string disabled?: boolean askResponse?: ClineAskResponse @@ -152,7 +153,7 @@ export interface WebviewMessage { slug?: string modeConfig?: ModeConfig timeout?: number - payload?: WebViewMessagePayload + payload?: WebViewMessagePayload | RequestProviderModelsPayload source?: "global" | "project" requestId?: string ids?: string[] @@ -179,3 +180,10 @@ 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 +} diff --git a/webview-ui/src/components/settings/ApiOptions.tsx b/webview-ui/src/components/settings/ApiOptions.tsx index 1f45e92a88..e63ca7dd5b 100644 --- a/webview-ui/src/components/settings/ApiOptions.tsx +++ b/webview-ui/src/components/settings/ApiOptions.tsx @@ -11,6 +11,8 @@ import { glamaDefaultModelId, unboundDefaultModelId, litellmDefaultModelId, + RouterModels, + ModelRecord, } from "@roo/shared/api" import { vscode } from "@src/utils/vscode" @@ -63,6 +65,8 @@ export interface ApiOptionsProps { setErrorMessage: React.Dispatch> } +const emptyModelRecord: ModelRecord = {} + const ApiOptions = ({ uriScheme, apiConfiguration, @@ -123,7 +127,52 @@ const ApiOptions = ({ info: selectedModelInfo, } = useSelectedModel(apiConfiguration) - const { data: routerModels, refetch: refetchRouterModels } = useRouterModels() + const { data: initialRouterModels } = useRouterModels() + + const defaultRouterModels: RouterModels = useMemo( + () => ({ + openrouter: emptyModelRecord, + requesty: emptyModelRecord, + glama: emptyModelRecord, + unbound: emptyModelRecord, + litellm: emptyModelRecord, + }), + [], + ) + + const [currentRouterModels, setCurrentRouterModels] = useState( + initialRouterModels || defaultRouterModels, + ) + + useEffect(() => { + if (initialRouterModels) { + setCurrentRouterModels(initialRouterModels) + } else { + setCurrentRouterModels(defaultRouterModels) + } + }, [initialRouterModels, defaultRouterModels]) + + // Listen for specific provider model updates + useEffect(() => { + const handler = (event: MessageEvent) => { + const message = event.data + if (message.type === "providerModelsResponse" && message.payload) { + const { provider, models, error } = message.payload as { + provider: keyof RouterModels + models?: ModelRecord + error?: string + } + if (provider && models && !error) { + setCurrentRouterModels((prevModels) => ({ + ...prevModels, // prevModels is now guaranteed to be RouterModels + [provider]: models, + })) + } + } + } + window.addEventListener("message", handler) + return () => window.removeEventListener("message", handler) + }, []) // Update `apiModelId` whenever `selectedModelId` changes. useEffect(() => { @@ -175,10 +224,10 @@ const ApiOptions = ({ useEffect(() => { const apiValidationResult = - validateApiConfiguration(apiConfiguration) || validateModelId(apiConfiguration, routerModels) + validateApiConfiguration(apiConfiguration) || validateModelId(apiConfiguration, currentRouterModels) setErrorMessage(apiValidationResult) - }, [apiConfiguration, routerModels, setErrorMessage]) + }, [apiConfiguration, currentRouterModels, setErrorMessage]) const selectedProviderModels = useMemo( () => @@ -294,7 +343,7 @@ const ApiOptions = ({ )} @@ -314,7 +362,7 @@ const ApiOptions = ({ )} @@ -323,7 +371,7 @@ const ApiOptions = ({ )} @@ -394,7 +442,7 @@ const ApiOptions = ({ )} diff --git a/webview-ui/src/components/settings/providers/LiteLLM.tsx b/webview-ui/src/components/settings/providers/LiteLLM.tsx index 8f98d97e5d..aae723b663 100644 --- a/webview-ui/src/components/settings/providers/LiteLLM.tsx +++ b/webview-ui/src/components/settings/providers/LiteLLM.tsx @@ -1,21 +1,28 @@ -import { useCallback } from "react" +import { useCallback, useState, useEffect } from "react" import { VSCodeTextField } from "@vscode/webview-ui-toolkit/react" import { ProviderSettings, RouterModels, litellmDefaultModelId } from "@roo/shared/api" - +import { vscode } from "@src/utils/vscode" +import { Button } from "@src/components/ui" import { useAppTranslation } from "@src/i18n/TranslationContext" import { inputEventTransform } from "../transforms" import { ModelPicker } from "../ModelPicker" +import { WebviewMessage } from "@roo/shared/WebviewMessage" +import { ExtensionMessage, ProviderModelsResponsePayload } from "@roo/shared/ExtensionMessage" type LiteLLMProps = { apiConfiguration: ProviderSettings setApiConfigurationField: (field: keyof ProviderSettings, value: ProviderSettings[keyof ProviderSettings]) => void + // routerModels prop might need to be updated by parent if we want to show new models immediately. + // For now, this component will manage its own refresh feedback. routerModels?: RouterModels } export const LiteLLM = ({ apiConfiguration, setApiConfigurationField, routerModels }: LiteLLMProps) => { const { t } = useAppTranslation() + const [refreshStatus, setRefreshStatus] = useState<"idle" | "loading" | "success" | "error">("idle") + const [refreshError, setRefreshError] = useState() const handleInputChange = useCallback( ( @@ -28,6 +35,43 @@ export const LiteLLM = ({ apiConfiguration, setApiConfigurationField, routerMode [setApiConfigurationField], ) + const handleRefreshModels = () => { + setRefreshStatus("loading") + setRefreshError(undefined) + const message: WebviewMessage = { + type: "requestProviderModels", + payload: { + provider: "litellm", + apiKey: apiConfiguration.litellmApiKey, + baseUrl: apiConfiguration.litellmBaseUrl || "http://localhost:4000", + }, + } + vscode.postMessage(message) + } + + // Effect to listen for model refresh responses + useEffect(() => { + const handler = (event: MessageEvent) => { + const message = event.data + if ( + message.type === "providerModelsResponse" && + message.payload && + message.payload.provider === "litellm" + ) { + const payload = message.payload as ProviderModelsResponsePayload + if (payload.error) { + setRefreshStatus("error") + setRefreshError(payload.error) + } else { + setRefreshStatus("success") + // Parent (ApiOptions.tsx) will handle updating the routerModels prop for ModelPicker + } + } + } + window.addEventListener("message", handler) + return () => window.removeEventListener("message", handler) + }, []) + return ( <> + + {refreshStatus === "loading" && ( +
+ {t("settings:providers.refreshModels.loading")} +
+ )} + {refreshStatus === "success" && ( +
{t("settings:providers.refreshModels.success")}
+ )} + {refreshStatus === "error" && ( +
+ {refreshError || t("settings:providers.refreshModels.error")} +
+ )} + void routerModels?: RouterModels - refetchRouterModels: () => void } -export const Requesty = ({ - apiConfiguration, - setApiConfigurationField, - routerModels, - refetchRouterModels, -}: RequestyProps) => { +export const Requesty = ({ apiConfiguration, setApiConfigurationField, routerModels }: RequestyProps) => { const { t } = useAppTranslation() - - const [didRefetch, setDidRefetch] = useState() + const [refreshStatus, setRefreshStatus] = useState<"idle" | "loading" | "success" | "error">("idle") + const [refreshError, setRefreshError] = useState() const handleInputChange = useCallback( ( @@ -40,6 +35,40 @@ export const Requesty = ({ [setApiConfigurationField], ) + const handleRefreshModels = () => { + setRefreshStatus("loading") + setRefreshError(undefined) + const message: WebviewMessage = { + type: "requestProviderModels", + payload: { + provider: "requesty", + apiKey: apiConfiguration.requestyApiKey, + }, + } + vscode.postMessage(message) + } + + useEffect(() => { + const handler = (event: MessageEvent) => { + const message = event.data + if ( + message.type === "providerModelsResponse" && + message.payload && + message.payload.provider === "requesty" + ) { + const payload = message.payload as ProviderModelsResponsePayload + if (payload.error) { + setRefreshStatus("error") + setRefreshError(payload.error) + } else { + setRefreshStatus("success") + } + } + } + window.addEventListener("message", handler) + return () => window.removeEventListener("message", handler) + }, []) + return ( <> { - vscode.postMessage({ type: "flushRouterModels", text: "requesty" }) - refetchRouterModels() - setDidRefetch(true) - }}> + onClick={handleRefreshModels} + disabled={refreshStatus === "loading"} + className="w-full">
- + {refreshStatus === "loading" ? ( + + ) : ( + + )} {t("settings:providers.refreshModels.label")}
- {didRefetch && ( -
- {t("settings:providers.refreshModels.hint")} + {refreshStatus === "loading" && ( +
+ {t("settings:providers.refreshModels.loading")} +
+ )} + {refreshStatus === "success" && ( +
{t("settings:providers.refreshModels.success")}
+ )} + {refreshStatus === "error" && ( +
+ {refreshError || t("settings:providers.refreshModels.error")}
)} OpenRouter Transforms)",