Ability to refresh litellm models by refresh button. Provider-specific model fetching. Success vs fail user feedback on model fetch responses

This commit is contained in:
slytechnical 2025-05-16 16:05:16 -05:00
parent cc3189b6f5
commit 2328695101
9 changed files with 322 additions and 96 deletions

View file

@ -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<ModelRecord> {
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."}`)
}
}
}

View file

@ -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<ModelRecord | undefined> {
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<ModelRecord> => {
let models = memoryCache.get<ModelRecord>(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<ModelRecord>(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.
}
}
/**

View file

@ -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 = <K extends keyof GlobalState>(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()

View file

@ -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
}

View file

@ -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<typeof checkoutRestorePayloadSchema>
export type WebViewMessagePayload = CheckpointDiffPayload | CheckpointRestorePayload
// Payload for requestProviderModels
export interface RequestProviderModelsPayload {
provider: RouterName
apiKey?: string
baseUrl?: string
}

View file

@ -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<React.SetStateAction<string | undefined>>
}
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<RouterModels>(
initialRouterModels || defaultRouterModels,
)
useEffect(() => {
if (initialRouterModels) {
setCurrentRouterModels(initialRouterModels)
} else {
setCurrentRouterModels(defaultRouterModels)
}
}, [initialRouterModels, defaultRouterModels])
// Listen for specific provider model updates
useEffect(() => {
const handler = (event: MessageEvent<any>) => {
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 = ({
<OpenRouter
apiConfiguration={apiConfiguration}
setApiConfigurationField={setApiConfigurationField}
routerModels={routerModels}
routerModels={currentRouterModels}
selectedModelId={selectedModelId}
uriScheme={uriScheme}
fromWelcomeView={fromWelcomeView}
@ -305,8 +354,7 @@ const ApiOptions = ({
<Requesty
apiConfiguration={apiConfiguration}
setApiConfigurationField={setApiConfigurationField}
routerModels={routerModels}
refetchRouterModels={refetchRouterModels}
routerModels={currentRouterModels}
/>
)}
@ -314,7 +362,7 @@ const ApiOptions = ({
<Glama
apiConfiguration={apiConfiguration}
setApiConfigurationField={setApiConfigurationField}
routerModels={routerModels}
routerModels={currentRouterModels}
uriScheme={uriScheme}
/>
)}
@ -323,7 +371,7 @@ const ApiOptions = ({
<Unbound
apiConfiguration={apiConfiguration}
setApiConfigurationField={setApiConfigurationField}
routerModels={routerModels}
routerModels={currentRouterModels}
/>
)}
@ -394,7 +442,7 @@ const ApiOptions = ({
<LiteLLM
apiConfiguration={apiConfiguration}
setApiConfigurationField={setApiConfigurationField}
routerModels={routerModels}
routerModels={currentRouterModels}
/>
)}

View file

@ -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<string | undefined>()
const handleInputChange = useCallback(
<K extends keyof ProviderSettings, E>(
@ -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<ExtensionMessage>) => {
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 (
<>
<VSCodeTextField
@ -51,6 +95,34 @@ export const LiteLLM = ({ apiConfiguration, setApiConfigurationField, routerMode
{t("settings:providers.apiKeyStorageNotice")}
</div>
<Button
variant="outline"
onClick={handleRefreshModels}
disabled={refreshStatus === "loading"}
className="w-full">
<div className="flex items-center gap-2">
{refreshStatus === "loading" ? (
<span className="codicon codicon-loading codicon-modifier-spin" />
) : (
<span className="codicon codicon-refresh" />
)}
{t("settings:providers.refreshModels.label")}
</div>
</Button>
{refreshStatus === "loading" && (
<div className="text-sm text-vscode-descriptionForeground">
{t("settings:providers.refreshModels.loading")}
</div>
)}
{refreshStatus === "success" && (
<div className="text-sm text-vscode-foreground">{t("settings:providers.refreshModels.success")}</div>
)}
{refreshStatus === "error" && (
<div className="text-sm text-vscode-errorForeground">
{refreshError || t("settings:providers.refreshModels.error")}
</div>
)}
<ModelPicker
apiConfiguration={apiConfiguration}
defaultModelId={litellmDefaultModelId}

View file

@ -1,8 +1,7 @@
import { useCallback, useState } from "react"
import { useCallback, useState, useEffect } from "react"
import { VSCodeTextField } from "@vscode/webview-ui-toolkit/react"
import { ProviderSettings, RouterModels, requestyDefaultModelId } from "@roo/shared/api"
import { vscode } from "@src/utils/vscode"
import { useAppTranslation } from "@src/i18n/TranslationContext"
import { VSCodeButtonLink } from "@src/components/common/VSCodeButtonLink"
@ -11,23 +10,19 @@ import { Button } from "@src/components/ui"
import { inputEventTransform } from "../transforms"
import { ModelPicker } from "../ModelPicker"
import { RequestyBalanceDisplay } from "./RequestyBalanceDisplay"
import { WebviewMessage } from "@roo/shared/WebviewMessage"
import { ExtensionMessage, ProviderModelsResponsePayload } from "@roo/shared/ExtensionMessage"
type RequestyProps = {
apiConfiguration: ProviderSettings
setApiConfigurationField: (field: keyof ProviderSettings, value: ProviderSettings[keyof ProviderSettings]) => 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<boolean>()
const [refreshStatus, setRefreshStatus] = useState<"idle" | "loading" | "success" | "error">("idle")
const [refreshError, setRefreshError] = useState<string | undefined>()
const handleInputChange = useCallback(
<K extends keyof ProviderSettings, E>(
@ -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<ExtensionMessage>) => {
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 (
<>
<VSCodeTextField
@ -68,19 +97,29 @@ export const Requesty = ({
)}
<Button
variant="outline"
onClick={() => {
vscode.postMessage({ type: "flushRouterModels", text: "requesty" })
refetchRouterModels()
setDidRefetch(true)
}}>
onClick={handleRefreshModels}
disabled={refreshStatus === "loading"}
className="w-full">
<div className="flex items-center gap-2">
<span className="codicon codicon-refresh" />
{refreshStatus === "loading" ? (
<span className="codicon codicon-loading codicon-modifier-spin" />
) : (
<span className="codicon codicon-refresh" />
)}
{t("settings:providers.refreshModels.label")}
</div>
</Button>
{didRefetch && (
<div className="flex items-center text-vscode-errorForeground">
{t("settings:providers.refreshModels.hint")}
{refreshStatus === "loading" && (
<div className="text-sm text-vscode-descriptionForeground">
{t("settings:providers.refreshModels.loading")}
</div>
)}
{refreshStatus === "success" && (
<div className="text-sm text-vscode-foreground">{t("settings:providers.refreshModels.success")}</div>
)}
{refreshStatus === "error" && (
<div className="text-sm text-vscode-errorForeground">
{refreshError || t("settings:providers.refreshModels.error")}
</div>
)}
<ModelPicker

View file

@ -118,7 +118,10 @@
"requestyApiKey": "Requesty API Key",
"refreshModels": {
"label": "Refresh Models",
"hint": "Please reopen the settings to see the latest models."
"hint": "Please reopen the settings to see the latest models.",
"loading": "Refreshing models...",
"success": "Models refreshed successfully.",
"error": "Failed to refresh models. Please check your configuration and try again."
},
"getRequestyApiKey": "Get Requesty API Key",
"openRouterTransformsText": "Compress prompts and message chains to the context size (<a>OpenRouter Transforms</a>)",