mirror of
https://github.com/RooVetGit/Roo-Code.git
synced 2026-09-06 08:18:39 +00:00
Switched GetModelsOptions to a discriminated union based on api provider. Adjusted related code and tests to match
This commit is contained in:
parent
bd1f93a755
commit
d4918c369d
12 changed files with 210 additions and 128 deletions
|
|
@ -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 })
|
||||
})
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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" })
|
||||
})
|
||||
})
|
||||
})
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 }
|
||||
|
|
|
|||
|
|
@ -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" ? (
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
<>
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue