Switched GetModelsOptions to a discriminated union based on api provider. Adjusted related code and tests to match

This commit is contained in:
slytechnical 2025-05-19 15:11:21 -05:00
parent bd1f93a755
commit d4918c369d
12 changed files with 210 additions and 128 deletions

View file

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

View file

@ -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<ModelRecord | undefined>
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<ModelRecord> => {
const { router } = options
let models = memoryCache.get<ModelRecord>(router)
const { provider } = options
let models = memoryCache.get<ModelRecord>(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<ModelRecord>
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.
}

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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 = <K extends keyof GlobalState>(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<Record<RouterName, ModelRecord>> = {
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<ModelRecord> => {
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<RouterName, ModelRecord>,
})
break
case "requestOpenAiModels":

View file

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

View file

@ -1792,3 +1792,15 @@ export function toRouterName(value?: string): RouterName {
export type ModelRecord = Record<string, ModelInfo>
export type RouterModels = Record<RouterName, ModelRecord>
/**
* 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 }

View file

@ -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<ExtensionMessage>) => {
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 (
<>
<VSCodeTextField
value={apiConfiguration?.litellmBaseUrl || "http://localhost:4000"}
value={apiConfiguration?.litellmBaseUrl || ""}
onInput={handleInputChange("litellmBaseUrl")}
placeholder="http://localhost:4000"
className="w-full">
@ -90,7 +105,9 @@ export const LiteLLM = ({ apiConfiguration, setApiConfigurationField, routerMode
<Button
variant="outline"
onClick={handleRefreshModels}
disabled={refreshStatus === "loading"}
disabled={
refreshStatus === "loading" || !apiConfiguration.litellmApiKey || !apiConfiguration.litellmBaseUrl
}
className="w-full">
<div className="flex items-center gap-2">
{refreshStatus === "loading" ? (

View file

@ -1,5 +1,6 @@
import { useCallback, useState, useEffect } from "react"
import { useCallback, useState } from "react"
import { VSCodeTextField } from "@vscode/webview-ui-toolkit/react"
import { useEvent } from "react-use"
import { ProviderSettings, RouterModels, requestyDefaultModelId } from "@roo/shared/api"
import { vscode } from "@src/utils/vscode"
@ -48,26 +49,18 @@ export const Requesty = ({ apiConfiguration, setApiConfigurationField, routerMod
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")
}
useEvent("message", (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 (
<>