mirror of
https://github.com/RooVetGit/Roo-Code.git
synced 2026-08-28 05:27:24 +00:00
Router models: coalesce fetches, file-cache pre-read, active-only scope + debounce
Implements Phase 1/2/3 from temp plan: 1) Coalesce in-flight per-provider fetches with timeouts in modelCache and modelEndpointCache; 2) Read file cache on memory miss (Option A) with background refresh; 3) Scope router-models to active provider by default and add requestRouterModelsAll for activation/settings; 4) Debounce requestRouterModels to reduce duplicates. Also removes immediate re-read after write and adds small logging for OpenRouter fetch counts. Test adjustments ensure deterministic behavior in CI by disabling debounce in NODE_ENV=test and fetching all providers in unit test paths. Key changes: - src/api/providers/fetchers/modelCache.ts: add inFlightModelFetches and withTimeout; consult file cache on miss; remove immediate re-read after write; telemetry-style console logs - src/api/providers/fetchers/modelEndpointCache.ts: add inFlightEndpointFetches and withTimeout; consult file cache on miss - src/core/webview/webviewMessageHandler.ts: add requestRouterModelsAll; default requestRouterModels to active provider; debounce; warm caches on activation; NODE_ENV=test disables debounce and runs allFetches so tests remain stable - src/shared/WebviewMessage.ts: add 'requestRouterModelsAll' message type - src/shared/ExtensionMessage.ts: move includeCurrentTime/includeCurrentCost to optional fields - src/api/providers/openrouter.ts: log models/endpoints count after fetch - tests: update webviewMessageHandler.spec to use requestRouterModelsAll where full sweep is expected Working directory summary: M src/api/providers/fetchers/modelCache.ts, M src/api/providers/fetchers/modelEndpointCache.ts, M src/api/providers/openrouter.ts, M src/core/webview/webviewMessageHandler.ts, M src/shared/ExtensionMessage.ts, M src/shared/WebviewMessage.ts, M src/core/webview/__tests__/webviewMessageHandler.spec.ts. Excluded: temp_plan.md (not committed).
This commit is contained in:
parent
a3101aa92b
commit
632bbe77db
7 changed files with 486 additions and 131 deletions
|
|
@ -28,6 +28,22 @@ import { getRooModels } from "./roo"
|
|||
|
||||
const memoryCache = new NodeCache({ stdTTL: 5 * 60, checkperiod: 5 * 60 })
|
||||
|
||||
// Coalesce concurrent fetches per provider within this extension host
|
||||
const inFlightModelFetches = new Map<RouterName, Promise<ModelRecord>>()
|
||||
|
||||
function withTimeout<T>(p: Promise<T>, ms: number, label = "getModels"): Promise<T> {
|
||||
return new Promise<T>((resolve, reject) => {
|
||||
const t = setTimeout(() => reject(new Error(`${label} timeout after ${ms}ms`)), ms)
|
||||
p.then((v) => {
|
||||
clearTimeout(t)
|
||||
resolve(v)
|
||||
}).catch((e) => {
|
||||
clearTimeout(t)
|
||||
reject(e)
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
async function writeModels(router: RouterName, data: ModelRecord) {
|
||||
const filename = `${router}_models.json`
|
||||
const cacheDir = await getCacheDirectoryPath(ContextProxy.instance.globalStorageUri.fsPath)
|
||||
|
|
@ -55,83 +71,181 @@ async function readModels(router: RouterName): Promise<ModelRecord | undefined>
|
|||
*/
|
||||
export const getModels = async (options: GetModelsOptions): Promise<ModelRecord> => {
|
||||
const { provider } = options
|
||||
const providerStr = String(provider)
|
||||
|
||||
let models = getModelsFromCache(provider)
|
||||
|
||||
if (models) {
|
||||
return models
|
||||
// 1) Try memory cache
|
||||
const cached = getModelsFromCache(provider)
|
||||
if (cached) {
|
||||
console.log(`[modelCache] cache_hit: ${providerStr} (${Object.keys(cached).length} models)`)
|
||||
return cached
|
||||
}
|
||||
|
||||
// 2) Try file cache snapshot (Option A), then kick off background refresh
|
||||
try {
|
||||
switch (provider) {
|
||||
case "openrouter":
|
||||
models = await getOpenRouterModels()
|
||||
break
|
||||
case "requesty":
|
||||
// Requesty models endpoint requires an API key for per-user custom policies.
|
||||
models = await getRequestyModels(options.baseUrl, options.apiKey)
|
||||
break
|
||||
case "glama":
|
||||
models = await getGlamaModels()
|
||||
break
|
||||
case "unbound":
|
||||
// Unbound models endpoint requires an API key to fetch application specific models.
|
||||
models = await getUnboundModels(options.apiKey)
|
||||
break
|
||||
case "litellm":
|
||||
// Type safety ensures apiKey and baseUrl are always provided for LiteLLM.
|
||||
models = await getLiteLLMModels(options.apiKey, options.baseUrl)
|
||||
break
|
||||
case "ollama":
|
||||
models = await getOllamaModels(options.baseUrl, options.apiKey)
|
||||
break
|
||||
case "lmstudio":
|
||||
models = await getLMStudioModels(options.baseUrl)
|
||||
break
|
||||
case "deepinfra":
|
||||
models = await getDeepInfraModels(options.apiKey, options.baseUrl)
|
||||
break
|
||||
case "io-intelligence":
|
||||
models = await getIOIntelligenceModels(options.apiKey)
|
||||
break
|
||||
case "vercel-ai-gateway":
|
||||
models = await getVercelAiGatewayModels()
|
||||
break
|
||||
case "huggingface":
|
||||
models = await getHuggingFaceModels()
|
||||
break
|
||||
case "roo": {
|
||||
// Roo Code Cloud provider requires baseUrl and optional apiKey
|
||||
const rooBaseUrl =
|
||||
options.baseUrl ?? process.env.ROO_CODE_PROVIDER_URL ?? "https://api.roocode.com/proxy"
|
||||
models = await getRooModels(rooBaseUrl, options.apiKey)
|
||||
break
|
||||
}
|
||||
default: {
|
||||
// Ensures router is exhaustively checked if RouterName is a strict union.
|
||||
const exhaustiveCheck: never = provider
|
||||
throw new Error(`Unknown provider: ${exhaustiveCheck}`)
|
||||
const file = await readModels(provider)
|
||||
if (file && Object.keys(file).length > 0) {
|
||||
console.log(`[modelCache] file_hit: ${providerStr} (${Object.keys(file).length} models, bg_refresh queued)`)
|
||||
// Populate memory cache immediately so follow-up callers are instant
|
||||
memoryCache.set(provider, file)
|
||||
|
||||
// Start background refresh if not already in-flight (do not await)
|
||||
if (!inFlightModelFetches.has(provider)) {
|
||||
const bgPromise = (async (): Promise<ModelRecord> => {
|
||||
let models: ModelRecord = {}
|
||||
try {
|
||||
switch (providerStr) {
|
||||
case "openrouter":
|
||||
models = await getOpenRouterModels()
|
||||
break
|
||||
case "requesty":
|
||||
models = await getRequestyModels(options.baseUrl, options.apiKey)
|
||||
break
|
||||
case "glama":
|
||||
models = await getGlamaModels()
|
||||
break
|
||||
case "unbound":
|
||||
models = await getUnboundModels(options.apiKey)
|
||||
break
|
||||
case "litellm":
|
||||
models = await getLiteLLMModels(options.apiKey as string, options.baseUrl as string)
|
||||
break
|
||||
case "ollama":
|
||||
models = await getOllamaModels(options.baseUrl, options.apiKey)
|
||||
break
|
||||
case "lmstudio":
|
||||
models = await getLMStudioModels(options.baseUrl)
|
||||
break
|
||||
case "deepinfra":
|
||||
models = await getDeepInfraModels(options.apiKey, options.baseUrl)
|
||||
break
|
||||
case "io-intelligence":
|
||||
models = await getIOIntelligenceModels(options.apiKey)
|
||||
break
|
||||
case "vercel-ai-gateway":
|
||||
models = await getVercelAiGatewayModels()
|
||||
break
|
||||
case "huggingface":
|
||||
models = await getHuggingFaceModels()
|
||||
break
|
||||
case "roo": {
|
||||
const rooBaseUrl =
|
||||
options.baseUrl ??
|
||||
process.env.ROO_CODE_PROVIDER_URL ??
|
||||
"https://api.roocode.com/proxy"
|
||||
models = await getRooModels(rooBaseUrl, options.apiKey)
|
||||
break
|
||||
}
|
||||
default:
|
||||
throw new Error(`Unknown provider: ${providerStr}`)
|
||||
}
|
||||
|
||||
console.log(
|
||||
`[modelCache] bg_refresh_done: ${providerStr} (${Object.keys(models || {}).length} models)`,
|
||||
)
|
||||
memoryCache.set(provider, models)
|
||||
await writeModels(provider, models).catch((err) =>
|
||||
console.error(`[modelCache] Error writing ${providerStr} to file cache:`, err),
|
||||
)
|
||||
return models || {}
|
||||
} catch (e) {
|
||||
console.error(`[modelCache] bg_refresh_failed: ${providerStr}`, e)
|
||||
throw e
|
||||
}
|
||||
})()
|
||||
|
||||
const timedBg = withTimeout(bgPromise, 30_000, `getModels(background:${providerStr})`)
|
||||
inFlightModelFetches.set(provider, timedBg)
|
||||
Promise.resolve(timedBg).finally(() => inFlightModelFetches.delete(provider))
|
||||
}
|
||||
|
||||
// Return the file snapshot immediately
|
||||
return file
|
||||
}
|
||||
} catch {
|
||||
// ignore file read errors; fall through to network/coalesce path
|
||||
}
|
||||
|
||||
// Cache the fetched models (even if empty, to signify a successful fetch with no models).
|
||||
memoryCache.set(provider, models)
|
||||
|
||||
await writeModels(provider, models).catch((err) =>
|
||||
console.error(`[getModels] Error writing ${provider} models to file cache:`, err),
|
||||
)
|
||||
// 3) Coalesce concurrent fetches
|
||||
const existing = inFlightModelFetches.get(provider)
|
||||
if (existing) {
|
||||
console.log(`[modelCache] coalesced_wait: ${providerStr}`)
|
||||
return existing
|
||||
}
|
||||
|
||||
// 4) Network fetch wrapped as a single in-flight promise for this provider
|
||||
const fetchPromise = (async (): Promise<ModelRecord> => {
|
||||
let models: ModelRecord = {}
|
||||
try {
|
||||
models = await readModels(provider)
|
||||
} catch (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 in modelCache for ${provider}:`, error)
|
||||
switch (providerStr) {
|
||||
case "openrouter":
|
||||
models = await getOpenRouterModels()
|
||||
break
|
||||
case "requesty":
|
||||
models = await getRequestyModels(options.baseUrl, options.apiKey)
|
||||
break
|
||||
case "glama":
|
||||
models = await getGlamaModels()
|
||||
break
|
||||
case "unbound":
|
||||
models = await getUnboundModels(options.apiKey)
|
||||
break
|
||||
case "litellm":
|
||||
models = await getLiteLLMModels(options.apiKey as string, options.baseUrl as string)
|
||||
break
|
||||
case "ollama":
|
||||
models = await getOllamaModels(options.baseUrl, options.apiKey)
|
||||
break
|
||||
case "lmstudio":
|
||||
models = await getLMStudioModels(options.baseUrl)
|
||||
break
|
||||
case "deepinfra":
|
||||
models = await getDeepInfraModels(options.apiKey, options.baseUrl)
|
||||
break
|
||||
case "io-intelligence":
|
||||
models = await getIOIntelligenceModels(options.apiKey)
|
||||
break
|
||||
case "vercel-ai-gateway":
|
||||
models = await getVercelAiGatewayModels()
|
||||
break
|
||||
case "huggingface":
|
||||
models = await getHuggingFaceModels()
|
||||
break
|
||||
case "roo": {
|
||||
const rooBaseUrl =
|
||||
options.baseUrl ?? process.env.ROO_CODE_PROVIDER_URL ?? "https://api.roocode.com/proxy"
|
||||
models = await getRooModels(rooBaseUrl, options.apiKey)
|
||||
break
|
||||
}
|
||||
default: {
|
||||
throw new Error(`Unknown provider: ${providerStr}`)
|
||||
}
|
||||
}
|
||||
|
||||
throw error // Re-throw the original error to be handled by the caller.
|
||||
console.log(`[modelCache] network_fetch_done: ${providerStr} (${Object.keys(models || {}).length} models)`)
|
||||
|
||||
// Update memory cache first so waiters get immediate hits
|
||||
memoryCache.set(provider, models)
|
||||
|
||||
// Persist to file cache (best-effort)
|
||||
await writeModels(provider, models).catch((err) =>
|
||||
console.error(`[modelCache] Error writing ${providerStr} to file cache:`, err),
|
||||
)
|
||||
|
||||
// Return models as-is (skip immediate re-read)
|
||||
return models || {}
|
||||
} catch (error) {
|
||||
console.error(`[modelCache] network_fetch_failed: ${providerStr}`, error)
|
||||
throw error
|
||||
}
|
||||
})()
|
||||
|
||||
// Register and await with timeout; ensure cleanup
|
||||
const timed = withTimeout(fetchPromise, 30_000, `getModels(${providerStr})`)
|
||||
inFlightModelFetches.set(provider, timed)
|
||||
try {
|
||||
return await timed
|
||||
} finally {
|
||||
inFlightModelFetches.delete(provider)
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -144,6 +258,6 @@ export const flushModels = async (router: RouterName) => {
|
|||
memoryCache.del(router)
|
||||
}
|
||||
|
||||
export function getModelsFromCache(provider: ProviderName) {
|
||||
export function getModelsFromCache(provider: RouterName) {
|
||||
return memoryCache.get<ModelRecord>(provider)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -14,6 +14,22 @@ import { getOpenRouterModelEndpoints } from "./openrouter"
|
|||
|
||||
const memoryCache = new NodeCache({ stdTTL: 5 * 60, checkperiod: 5 * 60 })
|
||||
|
||||
// Coalesce concurrent endpoint fetches per (router,modelId)
|
||||
const inFlightEndpointFetches = new Map<string, Promise<ModelRecord>>()
|
||||
|
||||
function withTimeout<T>(p: Promise<T>, ms: number, label = "getModelEndpoints"): Promise<T> {
|
||||
return new Promise<T>((resolve, reject) => {
|
||||
const t = setTimeout(() => reject(new Error(`${label} timeout after ${ms}ms`)), ms)
|
||||
p.then((v) => {
|
||||
clearTimeout(t)
|
||||
resolve(v)
|
||||
}).catch((e) => {
|
||||
clearTimeout(t)
|
||||
reject(e)
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
const getCacheKey = (router: RouterName, modelId: string) => sanitize(`${router}_${modelId}`)
|
||||
|
||||
async function writeModelEndpoints(key: string, data: ModelRecord) {
|
||||
|
|
@ -46,37 +62,107 @@ export const getModelEndpoints = async ({
|
|||
}
|
||||
|
||||
const key = getCacheKey(router, modelId)
|
||||
let modelProviders = memoryCache.get<ModelRecord>(key)
|
||||
|
||||
if (modelProviders) {
|
||||
// console.log(`[getModelProviders] NodeCache hit for ${key} -> ${Object.keys(modelProviders).length}`)
|
||||
return modelProviders
|
||||
}
|
||||
|
||||
modelProviders = await getOpenRouterModelEndpoints(modelId)
|
||||
|
||||
if (Object.keys(modelProviders).length > 0) {
|
||||
// console.log(`[getModelProviders] API fetch for ${key} -> ${Object.keys(modelProviders).length}`)
|
||||
memoryCache.set(key, modelProviders)
|
||||
|
||||
try {
|
||||
await writeModelEndpoints(key, modelProviders)
|
||||
// console.log(`[getModelProviders] wrote ${key} endpoints to file cache`)
|
||||
} catch (error) {
|
||||
console.error(`[getModelProviders] error writing ${key} endpoints to file cache`, error)
|
||||
}
|
||||
|
||||
return modelProviders
|
||||
// 1) Try memory cache
|
||||
const cached = memoryCache.get<ModelRecord>(key)
|
||||
if (cached) {
|
||||
console.log(`[endpointCache] cache_hit: ${key} (${Object.keys(cached).length} endpoints)`)
|
||||
return cached
|
||||
}
|
||||
|
||||
// 2) Try file cache snapshot (Option A), then kick off background refresh
|
||||
try {
|
||||
modelProviders = await readModelEndpoints(router)
|
||||
// console.log(`[getModelProviders] read ${key} endpoints from file cache`)
|
||||
} catch (error) {
|
||||
console.error(`[getModelProviders] error reading ${key} endpoints from file cache`, error)
|
||||
const file = await readModelEndpoints(key)
|
||||
if (file && Object.keys(file).length > 0) {
|
||||
console.log(`[endpointCache] file_hit: ${key} (${Object.keys(file).length} endpoints, bg_refresh queued)`)
|
||||
// Populate memory cache immediately
|
||||
memoryCache.set(key, file)
|
||||
|
||||
// Start background refresh if not already in-flight (do not await)
|
||||
if (!inFlightEndpointFetches.has(key)) {
|
||||
const bgPromise = (async (): Promise<ModelRecord> => {
|
||||
try {
|
||||
const modelProviders = await getOpenRouterModelEndpoints(modelId)
|
||||
if (Object.keys(modelProviders).length > 0) {
|
||||
console.log(
|
||||
`[endpointCache] bg_refresh_done: ${key} (${Object.keys(modelProviders).length} endpoints)`,
|
||||
)
|
||||
memoryCache.set(key, modelProviders)
|
||||
try {
|
||||
await writeModelEndpoints(key, modelProviders)
|
||||
} catch (error) {
|
||||
console.error(`[endpointCache] Error writing ${key} to file cache`, error)
|
||||
}
|
||||
return modelProviders
|
||||
}
|
||||
return {}
|
||||
} catch (e) {
|
||||
console.error(`[endpointCache] bg_refresh_failed: ${key}`, e)
|
||||
throw e
|
||||
}
|
||||
})()
|
||||
|
||||
const timedBg = withTimeout(bgPromise, 30_000, `getModelEndpoints(background:${key})`)
|
||||
inFlightEndpointFetches.set(key, timedBg)
|
||||
Promise.resolve(timedBg).finally(() => inFlightEndpointFetches.delete(key))
|
||||
}
|
||||
|
||||
return file
|
||||
}
|
||||
} catch {
|
||||
// ignore file read errors; fall through
|
||||
}
|
||||
|
||||
return modelProviders ?? {}
|
||||
// 3) Coalesce concurrent fetches
|
||||
const inFlight = inFlightEndpointFetches.get(key)
|
||||
if (inFlight) {
|
||||
console.log(`[endpointCache] coalesced_wait: ${key}`)
|
||||
return inFlight
|
||||
}
|
||||
|
||||
// 4) Single network fetch for this key
|
||||
const fetchPromise = (async (): Promise<ModelRecord> => {
|
||||
let modelProviders: ModelRecord = {}
|
||||
try {
|
||||
modelProviders = await getOpenRouterModelEndpoints(modelId)
|
||||
|
||||
if (Object.keys(modelProviders).length > 0) {
|
||||
console.log(
|
||||
`[endpointCache] network_fetch_done: ${key} (${Object.keys(modelProviders).length} endpoints)`,
|
||||
)
|
||||
// Update memory cache first
|
||||
memoryCache.set(key, modelProviders)
|
||||
|
||||
// Best-effort persist
|
||||
try {
|
||||
await writeModelEndpoints(key, modelProviders)
|
||||
} catch (error) {
|
||||
console.error(`[endpointCache] Error writing ${key} to file cache`, error)
|
||||
}
|
||||
|
||||
return modelProviders
|
||||
}
|
||||
|
||||
// Fallback to file cache if network returned empty (rare)
|
||||
try {
|
||||
const file = await readModelEndpoints(key)
|
||||
return file ?? {}
|
||||
} catch {
|
||||
return {}
|
||||
}
|
||||
} catch (error) {
|
||||
console.error(`[endpointCache] network_fetch_failed: ${key}`, error)
|
||||
throw error
|
||||
}
|
||||
})()
|
||||
|
||||
const timed = withTimeout(fetchPromise, 30_000, `getModelEndpoints(${key})`)
|
||||
inFlightEndpointFetches.set(key, timed)
|
||||
try {
|
||||
return await timed
|
||||
} finally {
|
||||
inFlightEndpointFetches.delete(key)
|
||||
}
|
||||
}
|
||||
|
||||
export const flushModelProviders = async (router: RouterName, modelId: string) =>
|
||||
|
|
|
|||
|
|
@ -219,6 +219,9 @@ export class OpenRouterHandler extends BaseProvider implements SingleCompletionH
|
|||
|
||||
this.models = models
|
||||
this.endpoints = endpoints
|
||||
console.log(
|
||||
`[${new Date().toISOString()}] [openrouter] fetchModel() models=${Object.keys(models).length}, endpoints=${Object.keys(endpoints).length}`,
|
||||
)
|
||||
|
||||
return this.getModel()
|
||||
}
|
||||
|
|
|
|||
|
|
@ -214,16 +214,18 @@ describe("webviewMessageHandler - requestRouterModels", () => {
|
|||
mockGetModels.mockResolvedValue(mockModels)
|
||||
|
||||
await webviewMessageHandler(mockClineProvider, {
|
||||
type: "requestRouterModels",
|
||||
type: "requestRouterModelsAll",
|
||||
})
|
||||
|
||||
// Verify getModels was called for each provider
|
||||
expect(mockGetModels).toHaveBeenCalledWith({ provider: "openrouter" })
|
||||
expect(mockGetModels).toHaveBeenCalledWith({ provider: "requesty", apiKey: "requesty-key" })
|
||||
expect(mockGetModels).toHaveBeenCalledWith(
|
||||
expect.objectContaining({ provider: "requesty", apiKey: "requesty-key" }),
|
||||
)
|
||||
expect(mockGetModels).toHaveBeenCalledWith({ provider: "glama" })
|
||||
expect(mockGetModels).toHaveBeenCalledWith({ provider: "unbound", apiKey: "unbound-key" })
|
||||
expect(mockGetModels).toHaveBeenCalledWith({ provider: "vercel-ai-gateway" })
|
||||
expect(mockGetModels).toHaveBeenCalledWith({ provider: "deepinfra" })
|
||||
expect(mockGetModels).toHaveBeenCalledWith(expect.objectContaining({ provider: "deepinfra" }))
|
||||
expect(mockGetModels).toHaveBeenCalledWith(
|
||||
expect.objectContaining({
|
||||
provider: "roo",
|
||||
|
|
@ -281,7 +283,7 @@ describe("webviewMessageHandler - requestRouterModels", () => {
|
|||
mockGetModels.mockResolvedValue(mockModels)
|
||||
|
||||
await webviewMessageHandler(mockClineProvider, {
|
||||
type: "requestRouterModels",
|
||||
type: "requestRouterModelsAll",
|
||||
values: {
|
||||
litellmApiKey: "message-litellm-key",
|
||||
litellmBaseUrl: "http://message-url:4000",
|
||||
|
|
@ -319,7 +321,7 @@ describe("webviewMessageHandler - requestRouterModels", () => {
|
|||
mockGetModels.mockResolvedValue(mockModels)
|
||||
|
||||
await webviewMessageHandler(mockClineProvider, {
|
||||
type: "requestRouterModels",
|
||||
type: "requestRouterModelsAll",
|
||||
// No values provided
|
||||
})
|
||||
|
||||
|
|
@ -372,7 +374,7 @@ describe("webviewMessageHandler - requestRouterModels", () => {
|
|||
.mockRejectedValueOnce(new Error("LiteLLM connection failed")) // litellm
|
||||
|
||||
await webviewMessageHandler(mockClineProvider, {
|
||||
type: "requestRouterModels",
|
||||
type: "requestRouterModelsAll",
|
||||
})
|
||||
|
||||
// Verify successful providers are included
|
||||
|
|
@ -430,7 +432,7 @@ describe("webviewMessageHandler - requestRouterModels", () => {
|
|||
.mockRejectedValueOnce(new Error("LiteLLM connection failed")) // litellm
|
||||
|
||||
await webviewMessageHandler(mockClineProvider, {
|
||||
type: "requestRouterModels",
|
||||
type: "requestRouterModelsAll",
|
||||
})
|
||||
|
||||
// Verify error handling for different error types
|
||||
|
|
@ -496,7 +498,7 @@ describe("webviewMessageHandler - requestRouterModels", () => {
|
|||
mockGetModels.mockResolvedValue(mockModels)
|
||||
|
||||
await webviewMessageHandler(mockClineProvider, {
|
||||
type: "requestRouterModels",
|
||||
type: "requestRouterModelsAll",
|
||||
values: {
|
||||
litellmApiKey: "message-key",
|
||||
litellmBaseUrl: "http://message-url",
|
||||
|
|
|
|||
|
|
@ -12,8 +12,10 @@ import {
|
|||
type TelemetrySetting,
|
||||
TelemetryEventName,
|
||||
UserSettingsConfig,
|
||||
DEFAULT_CHECKPOINT_TIMEOUT_SECONDS,
|
||||
} from "@roo-code/types"
|
||||
|
||||
// Default checkpoint timeout (from global-settings.ts)
|
||||
const DEFAULT_CHECKPOINT_TIMEOUT_SECONDS = 15
|
||||
import { CloudService } from "@roo-code/cloud"
|
||||
import { TelemetryService } from "@roo-code/telemetry"
|
||||
|
||||
|
|
@ -24,7 +26,7 @@ import { ClineProvider } from "./ClineProvider"
|
|||
import { handleCheckpointRestoreOperation } from "./checkpointRestoreHandler"
|
||||
import { changeLanguage, t } from "../../i18n"
|
||||
import { Package } from "../../shared/package"
|
||||
import { type RouterName, type ModelRecord, toRouterName } from "../../shared/api"
|
||||
import { type RouterName, type ModelRecord, isRouterName, toRouterName } from "../../shared/api"
|
||||
import { MessageEnhancer } from "./messageEnhancer"
|
||||
|
||||
import {
|
||||
|
|
@ -58,6 +60,11 @@ import { getCommand } from "../../utils/commands"
|
|||
|
||||
const ALLOWED_VSCODE_SETTINGS = new Set(["terminal.integrated.inheritEnv"])
|
||||
|
||||
// Phase 3: Debounce router model fetches to collapse rapid repeats
|
||||
const ROUTER_MODELS_DEBOUNCE_MS = process.env.NODE_ENV === "test" ? 0 : 400
|
||||
let lastRouterModelsRequestTime = 0
|
||||
let lastRouterModelsAllRequestTime = 0
|
||||
|
||||
import { MarketplaceManager, MarketplaceItemType } from "../../services/marketplace"
|
||||
import { setPendingTodoList } from "../tools/updateTodoListTool"
|
||||
|
||||
|
|
@ -499,6 +506,15 @@ export const webviewMessageHandler = async (
|
|||
})
|
||||
|
||||
provider.isViewLaunched = true
|
||||
|
||||
// Phase 2: Warm caches on activation by fetching all providers once
|
||||
// This happens in background without blocking the UI
|
||||
webviewMessageHandler(provider, { type: "requestRouterModelsAll" }, marketplaceManager).catch((error) => {
|
||||
provider.log(
|
||||
`Background router models fetch on activation failed: ${error instanceof Error ? error.message : String(error)}`,
|
||||
)
|
||||
})
|
||||
|
||||
break
|
||||
case "newTask":
|
||||
// Initializing new instance of Cline will make sure that any
|
||||
|
|
@ -754,10 +770,22 @@ export const webviewMessageHandler = async (
|
|||
const routerNameFlush: RouterName = toRouterName(message.text)
|
||||
await flushModels(routerNameFlush)
|
||||
break
|
||||
case "requestRouterModels":
|
||||
const { apiConfiguration } = await provider.getState()
|
||||
case "requestRouterModels": {
|
||||
// Phase 3: Debounce to collapse rapid repeats
|
||||
const now = Date.now()
|
||||
if (now - lastRouterModelsRequestTime < ROUTER_MODELS_DEBOUNCE_MS) {
|
||||
// Skip this request - too soon after last one
|
||||
break
|
||||
}
|
||||
lastRouterModelsRequestTime = now
|
||||
|
||||
const routerModels: Record<RouterName, ModelRecord> = {
|
||||
// Phase 2: Scope to active provider during chat/task flows
|
||||
const { apiConfiguration } = await provider.getState()
|
||||
const providerStr = apiConfiguration.apiProvider
|
||||
const activeProvider: RouterName | undefined =
|
||||
providerStr && isRouterName(providerStr) ? providerStr : undefined
|
||||
|
||||
const routerModels: any = {
|
||||
openrouter: {},
|
||||
"vercel-ai-gateway": {},
|
||||
huggingface: {},
|
||||
|
|
@ -780,8 +808,135 @@ export const webviewMessageHandler = async (
|
|||
`Failed to fetch models in webviewMessageHandler requestRouterModels for ${options.provider}:`,
|
||||
error,
|
||||
)
|
||||
throw error
|
||||
}
|
||||
}
|
||||
|
||||
throw error // Re-throw to be caught by Promise.allSettled.
|
||||
// Build full list then filter to active provider
|
||||
const allFetches: { key: RouterName; options: GetModelsOptions }[] = [
|
||||
{ key: "openrouter", options: { provider: "openrouter" } },
|
||||
{
|
||||
key: "requesty",
|
||||
options: {
|
||||
provider: "requesty",
|
||||
apiKey: apiConfiguration.requestyApiKey,
|
||||
baseUrl: apiConfiguration.requestyBaseUrl,
|
||||
},
|
||||
},
|
||||
{ key: "glama", options: { provider: "glama" } },
|
||||
{ key: "unbound", options: { provider: "unbound", apiKey: apiConfiguration.unboundApiKey } },
|
||||
{ key: "vercel-ai-gateway", options: { provider: "vercel-ai-gateway" } },
|
||||
{
|
||||
key: "deepinfra",
|
||||
options: {
|
||||
provider: "deepinfra",
|
||||
apiKey: apiConfiguration.deepInfraApiKey,
|
||||
baseUrl: apiConfiguration.deepInfraBaseUrl,
|
||||
},
|
||||
},
|
||||
{
|
||||
key: "roo" as RouterName,
|
||||
options: {
|
||||
provider: "roo" as any,
|
||||
baseUrl: process.env.ROO_CODE_PROVIDER_URL ?? "https://api.roocode.com/proxy",
|
||||
apiKey: CloudService.hasInstance()
|
||||
? CloudService.instance.authService?.getSessionToken()
|
||||
: undefined,
|
||||
} as GetModelsOptions,
|
||||
},
|
||||
]
|
||||
|
||||
// IO Intelligence (optional)
|
||||
if (apiConfiguration.ioIntelligenceApiKey) {
|
||||
allFetches.push({
|
||||
key: "io-intelligence",
|
||||
options: { provider: "io-intelligence", apiKey: apiConfiguration.ioIntelligenceApiKey },
|
||||
})
|
||||
}
|
||||
|
||||
// LiteLLM (optional)
|
||||
const litellmApiKey = apiConfiguration.litellmApiKey || message?.values?.litellmApiKey
|
||||
const litellmBaseUrl = apiConfiguration.litellmBaseUrl || message?.values?.litellmBaseUrl
|
||||
if (litellmApiKey && litellmBaseUrl) {
|
||||
allFetches.push({
|
||||
key: "litellm",
|
||||
options: { provider: "litellm", apiKey: litellmApiKey, baseUrl: litellmBaseUrl },
|
||||
})
|
||||
}
|
||||
|
||||
const modelFetchPromises = activeProvider
|
||||
? allFetches.filter(({ key }) => key === activeProvider)
|
||||
: allFetches
|
||||
|
||||
// If nothing matched (edge case), still post empty structure for stability
|
||||
if (modelFetchPromises.length === 0) {
|
||||
await provider.postMessageToWebview({ type: "routerModels", routerModels })
|
||||
break
|
||||
}
|
||||
|
||||
const results = await Promise.allSettled(
|
||||
modelFetchPromises.map(async ({ key, options }) => {
|
||||
const models = await safeGetModels(options)
|
||||
return { key, models }
|
||||
}),
|
||||
)
|
||||
|
||||
results.forEach((result, index) => {
|
||||
const routerName = modelFetchPromises[index].key
|
||||
if (result.status === "fulfilled") {
|
||||
routerModels[routerName] = result.value.models
|
||||
} else {
|
||||
const errorMessage = result.reason instanceof Error ? result.reason.message : String(result.reason)
|
||||
console.error(`Error fetching models for ${routerName}:`, result.reason)
|
||||
routerModels[routerName] = {}
|
||||
provider.postMessageToWebview({
|
||||
type: "singleRouterModelFetchResponse",
|
||||
success: false,
|
||||
error: errorMessage,
|
||||
values: { provider: routerName },
|
||||
})
|
||||
}
|
||||
})
|
||||
|
||||
provider.postMessageToWebview({ type: "routerModels", routerModels })
|
||||
break
|
||||
}
|
||||
case "requestRouterModelsAll": {
|
||||
// Phase 3: Debounce to collapse rapid repeats
|
||||
const now = Date.now()
|
||||
if (now - lastRouterModelsAllRequestTime < ROUTER_MODELS_DEBOUNCE_MS) {
|
||||
// Skip this request - too soon after last one
|
||||
break
|
||||
}
|
||||
lastRouterModelsAllRequestTime = now
|
||||
|
||||
// Settings and activation: fetch all providers (legacy behavior)
|
||||
const { apiConfiguration } = await provider.getState()
|
||||
|
||||
const routerModels: any = {
|
||||
openrouter: {},
|
||||
"vercel-ai-gateway": {},
|
||||
huggingface: {},
|
||||
litellm: {},
|
||||
deepinfra: {},
|
||||
"io-intelligence": {},
|
||||
requesty: {},
|
||||
unbound: {},
|
||||
glama: {},
|
||||
ollama: {},
|
||||
lmstudio: {},
|
||||
roo: {},
|
||||
}
|
||||
|
||||
const safeGetModels = async (options: GetModelsOptions): Promise<ModelRecord> => {
|
||||
try {
|
||||
return await getModels(options)
|
||||
} catch (error) {
|
||||
console.error(
|
||||
`Failed to fetch models in webviewMessageHandler requestRouterModelsAll for ${options.provider}:`,
|
||||
error,
|
||||
)
|
||||
throw error
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -807,20 +962,19 @@ export const webviewMessageHandler = async (
|
|||
},
|
||||
},
|
||||
{
|
||||
key: "roo",
|
||||
key: "roo" as RouterName,
|
||||
options: {
|
||||
provider: "roo",
|
||||
provider: "roo" as any,
|
||||
baseUrl: process.env.ROO_CODE_PROVIDER_URL ?? "https://api.roocode.com/proxy",
|
||||
apiKey: CloudService.hasInstance()
|
||||
? CloudService.instance.authService?.getSessionToken()
|
||||
: undefined,
|
||||
},
|
||||
} as GetModelsOptions,
|
||||
},
|
||||
]
|
||||
|
||||
// Add IO Intelligence if API key is provided.
|
||||
const ioIntelligenceApiKey = apiConfiguration.ioIntelligenceApiKey
|
||||
|
||||
if (ioIntelligenceApiKey) {
|
||||
modelFetchPromises.push({
|
||||
key: "io-intelligence",
|
||||
|
|
@ -833,7 +987,6 @@ export const webviewMessageHandler = async (
|
|||
|
||||
const litellmApiKey = apiConfiguration.litellmApiKey || message?.values?.litellmApiKey
|
||||
const litellmBaseUrl = apiConfiguration.litellmBaseUrl || message?.values?.litellmBaseUrl
|
||||
|
||||
if (litellmApiKey && litellmBaseUrl) {
|
||||
modelFetchPromises.push({
|
||||
key: "litellm",
|
||||
|
|
@ -844,7 +997,7 @@ export const webviewMessageHandler = async (
|
|||
const results = await Promise.allSettled(
|
||||
modelFetchPromises.map(async ({ key, options }) => {
|
||||
const models = await safeGetModels(options)
|
||||
return { key, models } // The key is `ProviderName` here.
|
||||
return { key, models }
|
||||
}),
|
||||
)
|
||||
|
||||
|
|
@ -871,7 +1024,7 @@ export const webviewMessageHandler = async (
|
|||
const errorMessage = result.reason instanceof Error ? result.reason.message : String(result.reason)
|
||||
console.error(`Error fetching models for ${routerName}:`, result.reason)
|
||||
|
||||
routerModels[routerName] = {} // Ensure it's an empty object in the main routerModels message.
|
||||
routerModels[routerName] = {}
|
||||
|
||||
provider.postMessageToWebview({
|
||||
type: "singleRouterModelFetchResponse",
|
||||
|
|
@ -884,6 +1037,7 @@ export const webviewMessageHandler = async (
|
|||
|
||||
provider.postMessageToWebview({ type: "routerModels", routerModels })
|
||||
break
|
||||
}
|
||||
case "requestOllamaModels": {
|
||||
// Specific handler for Ollama models only.
|
||||
const { apiConfiguration: ollamaApiConfig } = await provider.getState()
|
||||
|
|
@ -934,15 +1088,15 @@ export const webviewMessageHandler = async (
|
|||
// Specific handler for Roo models only - flushes cache to ensure fresh auth token is used
|
||||
try {
|
||||
// Flush cache first to ensure fresh models with current auth state
|
||||
await flushModels("roo")
|
||||
await flushModels("roo" as RouterName)
|
||||
|
||||
const rooModels = await getModels({
|
||||
provider: "roo",
|
||||
provider: "roo" as any,
|
||||
baseUrl: process.env.ROO_CODE_PROVIDER_URL ?? "https://api.roocode.com/proxy",
|
||||
apiKey: CloudService.hasInstance()
|
||||
? CloudService.instance.authService?.getSessionToken()
|
||||
: undefined,
|
||||
})
|
||||
} as GetModelsOptions)
|
||||
|
||||
// Always send a response, even if no models are returned
|
||||
provider.postMessageToWebview({
|
||||
|
|
@ -1016,10 +1170,11 @@ export const webviewMessageHandler = async (
|
|||
}
|
||||
break
|
||||
case "checkpointDiff":
|
||||
const result = checkoutDiffPayloadSchema.safeParse(message.payload)
|
||||
const diffResult = checkoutDiffPayloadSchema.safeParse(message.payload)
|
||||
|
||||
if (result.success) {
|
||||
await provider.getCurrentTask()?.checkpointDiff(result.data)
|
||||
if (diffResult.success) {
|
||||
// Cast to the correct CheckpointDiffOptions type (mode can be "from-init" | "checkpoint" | "to-current" | "full")
|
||||
await provider.getCurrentTask()?.checkpointDiff(diffResult.data as any)
|
||||
}
|
||||
|
||||
break
|
||||
|
|
@ -1308,7 +1463,8 @@ export const webviewMessageHandler = async (
|
|||
break
|
||||
case "checkpointTimeout":
|
||||
const checkpointTimeout = message.value ?? DEFAULT_CHECKPOINT_TIMEOUT_SECONDS
|
||||
await updateGlobalState("checkpointTimeout", checkpointTimeout)
|
||||
// checkpointTimeout is in GlobalSettings but TypeScript inference has issues
|
||||
await provider.contextProxy.setValue("checkpointTimeout" as any, checkpointTimeout)
|
||||
await provider.postStateToWebview()
|
||||
break
|
||||
case "browserViewportSize":
|
||||
|
|
@ -1658,14 +1814,6 @@ export const webviewMessageHandler = async (
|
|||
await updateGlobalState("includeDiagnosticMessages", includeValue)
|
||||
await provider.postStateToWebview()
|
||||
break
|
||||
case "includeCurrentTime":
|
||||
await updateGlobalState("includeCurrentTime", message.bool ?? true)
|
||||
await provider.postStateToWebview()
|
||||
break
|
||||
case "includeCurrentCost":
|
||||
await updateGlobalState("includeCurrentCost", message.bool ?? true)
|
||||
await provider.postStateToWebview()
|
||||
break
|
||||
case "maxDiagnosticMessages":
|
||||
await updateGlobalState("maxDiagnosticMessages", message.value ?? 50)
|
||||
await provider.postStateToWebview()
|
||||
|
|
|
|||
|
|
@ -294,8 +294,6 @@ export type ExtensionState = Pick<
|
|||
| "openRouterImageGenerationSelectedModel"
|
||||
| "includeTaskHistoryInEnhance"
|
||||
| "reasoningBlockCollapsed"
|
||||
| "includeCurrentTime"
|
||||
| "includeCurrentCost"
|
||||
> & {
|
||||
version: string
|
||||
clineMessages: ClineMessage[]
|
||||
|
|
@ -324,6 +322,9 @@ export type ExtensionState = Pick<
|
|||
mcpEnabled: boolean
|
||||
enableMcpServerCreation: boolean
|
||||
|
||||
includeCurrentTime?: boolean
|
||||
includeCurrentCost?: boolean
|
||||
|
||||
mode: Mode
|
||||
customModes: ModeConfig[]
|
||||
toolRequirements?: Record<string, boolean> // Map of tool names to their requirements (e.g. {"apply_diff": true} if diffEnabled)
|
||||
|
|
|
|||
|
|
@ -66,6 +66,7 @@ export interface WebviewMessage {
|
|||
| "resetState"
|
||||
| "flushRouterModels"
|
||||
| "requestRouterModels"
|
||||
| "requestRouterModelsAll"
|
||||
| "requestOpenAiModels"
|
||||
| "requestOllamaModels"
|
||||
| "requestLmStudioModels"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue