diff --git a/apps/vscode-e2e/src/visual/__screenshots__/electron-chat-dark-sidebar.png b/apps/vscode-e2e/src/visual/__screenshots__/electron-chat-dark-sidebar.png index e8497070ef..3c03cbac03 100644 Binary files a/apps/vscode-e2e/src/visual/__screenshots__/electron-chat-dark-sidebar.png and b/apps/vscode-e2e/src/visual/__screenshots__/electron-chat-dark-sidebar.png differ diff --git a/packages/types/src/__tests__/organization-allow-list.test.ts b/packages/types/src/__tests__/organization-allow-list.test.ts new file mode 100644 index 0000000000..b71a901540 --- /dev/null +++ b/packages/types/src/__tests__/organization-allow-list.test.ts @@ -0,0 +1,81 @@ +import { describe, expect, it } from "vitest" + +import type { OrganizationAllowList } from "../cloud.js" +import { isModelAllowedForOrganization, isProviderAllowAll } from "../organization-allow-list.js" +import { providerIdentifiers } from "../provider-identifiers.js" + +describe("isProviderAllowAll", () => { + it("allows everything when the allow list is missing", () => { + expect(isProviderAllowAll(undefined, providerIdentifiers.anthropic)).toBe(true) + }) + + it("allows everything when the organization allows all", () => { + const allowList: OrganizationAllowList = { allowAll: true, providers: {} } + expect(isProviderAllowAll(allowList, providerIdentifiers.anthropic)).toBe(true) + }) + + it("rejects when the provider is missing but the organization is restricted", () => { + const allowList: OrganizationAllowList = { allowAll: false, providers: {} } + expect(isProviderAllowAll(allowList, undefined)).toBe(false) + }) + + it("rejects a provider without an explicit entry", () => { + const allowList: OrganizationAllowList = { allowAll: false, providers: {} } + expect(isProviderAllowAll(allowList, providerIdentifiers.anthropic)).toBe(false) + }) + + it("allows a provider narrowed to a model list (the list still constrains models)", () => { + const allowList: OrganizationAllowList = { + allowAll: false, + providers: { [providerIdentifiers.anthropic]: { allowAll: false, models: ["m"] } }, + } + // A non-allowAll provider must not expose the custom-model escape hatch. + expect(isProviderAllowAll(allowList, providerIdentifiers.anthropic)).toBe(false) + }) + + it("allows a provider explicitly marked allowAll", () => { + const allowList: OrganizationAllowList = { + allowAll: false, + providers: { [providerIdentifiers.openrouter]: { allowAll: true } }, + } + expect(isProviderAllowAll(allowList, providerIdentifiers.openrouter)).toBe(true) + }) +}) + +describe("isModelAllowedForOrganization", () => { + it("allows any model when the allow list is missing", () => { + expect(isModelAllowedForOrganization(undefined, providerIdentifiers.anthropic, "any")).toBe(true) + }) + + it("allows any model when the organization allows all", () => { + const allowList: OrganizationAllowList = { allowAll: true, providers: {} } + expect(isModelAllowedForOrganization(allowList, providerIdentifiers.anthropic, "any")).toBe(true) + }) + + it("rejects when the provider is missing but the organization is restricted", () => { + const allowList: OrganizationAllowList = { allowAll: false, providers: {} } + expect(isModelAllowedForOrganization(allowList, undefined, "any")).toBe(false) + }) + + it("rejects a model from a provider without an explicit entry", () => { + const allowList: OrganizationAllowList = { allowAll: false, providers: {} } + expect(isModelAllowedForOrganization(allowList, providerIdentifiers.anthropic, "m")).toBe(false) + }) + + it("allows a listed model and rejects an unlisted one", () => { + const allowList: OrganizationAllowList = { + allowAll: false, + providers: { [providerIdentifiers.anthropic]: { allowAll: false, models: ["allowed"] } }, + } + expect(isModelAllowedForOrganization(allowList, providerIdentifiers.anthropic, "allowed")).toBe(true) + expect(isModelAllowedForOrganization(allowList, providerIdentifiers.anthropic, "other")).toBe(false) + }) + + it("allows any model for a provider explicitly marked allowAll", () => { + const allowList: OrganizationAllowList = { + allowAll: false, + providers: { [providerIdentifiers.openrouter]: { allowAll: true } }, + } + expect(isModelAllowedForOrganization(allowList, providerIdentifiers.openrouter, "anything")).toBe(true) + }) +}) diff --git a/packages/types/src/index.ts b/packages/types/src/index.ts index 82588ae537..18fbb08e3c 100644 --- a/packages/types/src/index.ts +++ b/packages/types/src/index.ts @@ -18,6 +18,7 @@ export * from "./mcp.js" export * from "./message.js" export * from "./mode.js" export * from "./model.js" +export * from "./organization-allow-list.js" export * from "./provider-identifiers.js" export * from "./provider-settings.js" export * from "./task.js" diff --git a/packages/types/src/organization-allow-list.ts b/packages/types/src/organization-allow-list.ts new file mode 100644 index 0000000000..fa8ffe597d --- /dev/null +++ b/packages/types/src/organization-allow-list.ts @@ -0,0 +1,54 @@ +import type { OrganizationAllowList } from "./cloud.js" +import type { ProviderName } from "./provider-settings.js" + +/** + * Whether organization policy allows an arbitrary/custom model id for a provider. + * + * Only providers that are explicitly `allowAll` may use free-form ids; when a + * provider is narrowed to an allow-list, callers must not expose a custom-model + * escape hatch. This is the single source of truth shared by the settings + * model picker and the chat model selector. + */ +export const isProviderAllowAll = ( + allowList: OrganizationAllowList | undefined, + provider: ProviderName | undefined, +): boolean => { + if (!allowList || allowList.allowAll) { + return true + } + + if (!provider) { + return false + } + + return allowList.providers[provider]?.allowAll === true +} + +/** + * Whether organization policy allows selecting the given model id for a provider. + * + * `allowAll` (org-wide or per provider) permits any id; otherwise the id must be + * present in the provider's explicit `models` list. Intended as a defense-in-depth + * check that rejects a selection regardless of how it was produced (option click, + * keyboard navigation, or custom entry). + */ +export const isModelAllowedForOrganization = ( + allowList: OrganizationAllowList | undefined, + provider: ProviderName | undefined, + modelId: string, +): boolean => { + if (!allowList || allowList.allowAll) { + return true + } + + if (!provider) { + return false + } + + const providerConfig = allowList.providers[provider] + if (!providerConfig) { + return false + } + + return providerConfig.allowAll === true || (providerConfig.models?.includes(modelId) ?? false) +} diff --git a/packages/types/src/vscode-extension-host.ts b/packages/types/src/vscode-extension-host.ts index c0e8509105..eef84f326f 100644 --- a/packages/types/src/vscode-extension-host.ts +++ b/packages/types/src/vscode-extension-host.ts @@ -42,6 +42,7 @@ export interface ExtensionMessage { | "enhancedPrompt" | "commitSearchResults" | "listApiConfig" + | "apiConfigurationSaved" | typeof RouterModelsMessageType.routerModels | "zooGatewayCredentialsReady" | typeof OpenAiModelsMessageType.openAiModels @@ -464,6 +465,7 @@ export type EditQueuedMessagePayload = Pick { vi.mocked(axios.get).mockClear() }) + it("forwards cancellation to the HTTP request", async () => { + const controller = new AbortController() + vi.mocked(axios.get).mockImplementationOnce( + (_url, options) => + new Promise((_resolve, reject) => { + options?.signal?.addEventListener?.("abort", () => reject(new Error("cancelled"))) + }), + ) + const pending = getOpenAiModels("https://example.test/v1", "test-key", undefined, controller.signal) + expect(axios.get).toHaveBeenCalledWith( + "https://example.test/v1/models", + expect.objectContaining({ signal: controller.signal }), + ) + controller.abort() + expect(await pending).toEqual([]) + }) + it("should return empty array when baseUrl is not provided", async () => { const result = await getOpenAiModels(undefined, "test-key") expect(result).toEqual([]) diff --git a/src/api/providers/fetchers/__tests__/kimi-code.spec.ts b/src/api/providers/fetchers/__tests__/kimi-code.spec.ts index a96e760253..b8e25c4e43 100644 --- a/src/api/providers/fetchers/__tests__/kimi-code.spec.ts +++ b/src/api/providers/fetchers/__tests__/kimi-code.spec.ts @@ -1,3 +1,5 @@ +import { getEventListeners } from "events" + import { getKimiCodeModels, kimiCodeModelSchema, mapKimiCodeModel } from "../kimi-code" describe("Kimi Code model discovery", () => { @@ -83,6 +85,23 @@ describe("Kimi Code model discovery", () => { expect(models["model-b"].supportsReasoningEffort).toEqual(["low", "high", "max"]) }) + it.each(["success", "failure"])("cleans up caller listeners and timeout on %s", async (outcome) => { + vi.useFakeTimers() + const controller = new AbortController() + const transport = vi.spyOn(globalThis, "fetch") + const failure = new Error("network failed") + if (outcome === "success") transport.mockResolvedValue(new Response(JSON.stringify({ data: [] }))) + else transport.mockRejectedValue(failure) + + const result = getKimiCodeModels("token", { signal: controller.signal }) + if (outcome === "success") await expect(result).resolves.toEqual({}) + else await expect(result).rejects.toBe(failure) + expect(getEventListeners(controller.signal, "abort")).toHaveLength(0) + expect(vi.getTimerCount()).toBe(0) + controller.abort() + expect(transport.mock.calls[0][1]?.signal?.aborted).toBe(false) + }) + it("aborts model discovery after its deadline", async () => { vi.useFakeTimers() vi.spyOn(globalThis, "fetch").mockImplementation((_input, init) => { @@ -90,11 +109,13 @@ describe("Kimi Code model discovery", () => { init?.signal?.addEventListener("abort", () => reject(init.signal?.reason), { once: true }) }) }) - const result = expect(getKimiCodeModels("token")).rejects.toThrow("timed out") + const controller = new AbortController() + const result = expect(getKimiCodeModels("token", { signal: controller.signal })).rejects.toThrow("timed out") await vi.advanceTimersByTimeAsync(10_000) await result expect(vi.mocked(fetch).mock.calls[0][1]?.signal?.aborted).toBe(true) + expect(getEventListeners(controller.signal, "abort")).toHaveLength(0) expect(vi.getTimerCount()).toBe(0) }) diff --git a/src/api/providers/fetchers/__tests__/modelCache.authCancellation.spec.ts b/src/api/providers/fetchers/__tests__/modelCache.authCancellation.spec.ts new file mode 100644 index 0000000000..15c2df1b38 --- /dev/null +++ b/src/api/providers/fetchers/__tests__/modelCache.authCancellation.spec.ts @@ -0,0 +1,91 @@ +import { getEventListeners } from "events" +import axios from "axios" +import { providerIdentifiers } from "@roo-code/types" + +import { ModelRequestRegistry } from "../../../../core/webview/ModelRequestRegistry" +import { getModels, refreshModels } from "../modelCache" + +vi.mock("../../../../services/zoo-code-auth", () => ({ + getZooCodeBaseUrl: () => "https://example.test", + resolveZooGatewaySessionToken: (token?: string) => token, +})) +vi.mock("@roo-code/telemetry", () => ({ + TelemetryService: { instance: { captureEvent: vi.fn(), isTelemetryEnabled: () => false } }, +})) + +// Keep the registry, cache dispatch, and auth-scoped fetchers real; only HTTP is replaced. +describe.each([getModels, refreshModels])("auth-scoped cancellation via %s", (loadModels) => { + describe.each([providerIdentifiers.zooGateway, providerIdentifiers.kimiCode])("%s", (provider) => { + let registry: ModelRequestRegistry + let signals: AbortSignal[] + let events: string[] + + beforeEach(() => { + vi.useFakeTimers() + registry = new ModelRequestRegistry() + signals = [] + events = [] + vi.spyOn(console, "error").mockImplementation(() => {}) + const pendingRequest = (signal?: AbortSignal | null) => { + expect(signal).toBeDefined() + if (!signal) throw new Error("HTTP request has no cancellation signal") + signals.push(signal) + const index = signals.length + events.push(`start:${index}`) + return new Promise((_resolve, reject) => { + signal.addEventListener( + "abort", + () => { + events.push(`abort:${index}`) + reject(signal.reason) + }, + { once: true }, + ) + }) + } + vi.spyOn(axios, "get").mockImplementation((_url, config) => { + expect(config?.timeout).toBe(15_000) + return pendingRequest(config?.signal instanceof AbortSignal ? config.signal : undefined) + }) + vi.spyOn(globalThis, "fetch").mockImplementation((_url, init) => pendingRequest(init?.signal)) + }) + + afterEach(() => { + registry.dispose() + vi.restoreAllMocks() + vi.useRealTimers() + }) + + it.each(["cancel", "replace"])("aborts HTTP before restarting after %s", async (action) => { + const ownerSignals: AbortSignal[] = [] + const discover = async (signal: AbortSignal) => { + ownerSignals.push(signal) + await loadModels({ provider, apiKey: "synthetic-token", signal }) + } + const first = registry.run("router", discover) + expect(events).toEqual(["start:1"]) + if (action === "cancel") registry.cancel("router") + const second = registry.run("router", discover) + expect(events).toEqual(["start:1", "abort:1", "start:2"]) + expect(signals[0].aborted).toBe(true) + expect(signals[1].aborted).toBe(false) + await first + registry.cancel("router") + await second + expect(events).toEqual(["start:1", "abort:1", "start:2", "abort:2"]) + for (const signal of ownerSignals) expect(getEventListeners(signal, "abort")).toHaveLength(0) + expect(vi.getTimerCount()).toBe(0) + }) + + it("does not start HTTP for a pre-aborted caller", async () => { + const controller = new AbortController() + controller.abort() + const result = loadModels({ provider, apiKey: "synthetic-token", signal: controller.signal }) + if (loadModels === getModels) await expect(result).rejects.toMatchObject({ name: "AbortError" }) + else await expect(result).resolves.toEqual({}) + expect(axios.get).not.toHaveBeenCalled() + expect(fetch).not.toHaveBeenCalled() + expect(vi.getTimerCount()).toBe(0) + }) + }) +}) diff --git a/src/api/providers/fetchers/__tests__/modelCache.spec.ts b/src/api/providers/fetchers/__tests__/modelCache.spec.ts index 3b95f32234..9631f553a0 100644 --- a/src/api/providers/fetchers/__tests__/modelCache.spec.ts +++ b/src/api/providers/fetchers/__tests__/modelCache.spec.ts @@ -41,6 +41,8 @@ vi.mock("fs", () => ({ })) // Mock all the model fetchers +vi.mock("../poe") +vi.mock("../lmstudio") vi.mock("../litellm") vi.mock("../openrouter") vi.mock("../requesty") @@ -69,6 +71,8 @@ import * as fsSync from "fs" import NodeCache from "node-cache" import { TelemetryService } from "@roo-code/telemetry" import { getModels, getModelsFromCache } from "../modelCache" +import { getPoeModels } from "../poe" +import { getLMStudioModels } from "../lmstudio" import { getLiteLLMModels } from "../litellm" import { getOpenRouterModels } from "../openrouter" import { getRequestyModels } from "../requesty" @@ -1645,7 +1649,7 @@ it("releases the entry for a fetcher double that honors no cancellation at all", expect(mockGetOpenRouterModels).toHaveBeenCalledTimes(2) }) -it("ignores the caller signal on the auth-scoped bypass without entering the flight map", async () => { +it("forwards the caller signal on the auth-scoped bypass without entering the flight map", async () => { setupCancellationMocks() // The single-flight arms its per-flight fetch bound whenever it creates a flight; the // auth-scoped bypass must never touch that machinery. @@ -1663,14 +1667,10 @@ it("ignores the caller signal on the auth-scoped bypass without entering the fli // its own fetch, so the bypass never shares (or poisons) a flight with anything. const second = getModels({ provider: providerIdentifiers.zooGateway, apiKey: "token-a" }) expect(mockGetZooGatewayModels).toHaveBeenCalledTimes(2) - // The bypass carries no cancellation: the fetcher receives exactly its own options - // argument, so the caller's bound is never threaded to this path. - expect(mockGetZooGatewayModels.mock.calls[0]).toHaveLength(1) + // Each auth-scoped fetch receives only its own caller's cancellation signal. + expect(mockGetZooGatewayModels.mock.calls[0][1]).toEqual({ signal: controller.signal }) expect(mockGetZooGatewayModels.mock.calls[1]).toHaveLength(1) - // The caller's signal is ignored on this path: aborting changes nothing for a fetch the - // single-flight never owns, and the fetcher's own request bound remains the stop mechanism. - controller.abort() await expect(first).resolves.toEqual(cancelledModelsB) await expect(second).resolves.toEqual(cancelledModelsB) expect(boundSpy).not.toHaveBeenCalled() @@ -1678,3 +1678,82 @@ it("ignores the caller signal on the auth-scoped bypass without entering the fli boundSpy.mockRestore() } }) + +// These SDKs cannot stop discovery when their caller aborts. Keep the pending +// operation registered across selector unmount/remount and even after its bound. +it.each([providerIdentifiers.poe, providerIdentifiers.lmstudio])( + "reuses pending %s discovery across cancellation and releases it at the bound", + async (provider) => { + setupCancellationMocks() + const timeoutController = new AbortController() + const timeoutSpy = vi.spyOn(AbortSignal, "timeout").mockReturnValue(timeoutController.signal) + const fetcher = provider === providerIdentifiers.poe ? vi.mocked(getPoeModels) : vi.mocked(getLMStudioModels) + let resolveDiscovery!: (models: ModelRecord) => void + fetcher.mockReturnValueOnce( + new Promise((resolve) => { + resolveDiscovery = resolve + }), + ) + try { + const controller = new AbortController() + const first = getModels({ provider, signal: controller.signal }) + controller.abort() + await expect(first).rejects.toMatchObject({ name: "AbortError" }) + expect(getEventListeners(controller.signal, "abort")).toHaveLength(0) + + const remounted = getModels({ provider }) + expect(fetcher).toHaveBeenCalledTimes(1) + timeoutController.abort() + await expect(remounted).rejects.toMatchObject({ name: "AbortError" }) + // The deadline released the entry even though the non-cancellable SDK call is still + // pending, so a retry starts a fresh flight instead of joining the expired one. Give + // the fresh flight its own (unaborted) bound; the original call is left to settle. + timeoutSpy.mockReturnValue(new AbortController().signal) + fetcher.mockResolvedValueOnce(cancelledModelsB) + await expect(getModels({ provider })).resolves.toEqual(cancelledModelsB) + expect(fetcher).toHaveBeenCalledTimes(2) + expect(getEventListeners(timeoutController.signal, "abort")).toHaveLength(0) + + resolveDiscovery(cancelledModels) + await drainMicrotasks() + expect(getEventListeners(timeoutController.signal, "abort")).toHaveLength(0) + } finally { + resolveDiscovery(cancelledModels) + timeoutSpy.mockRestore() + } + }, +) + +it.each([providerIdentifiers.poe, providerIdentifiers.lmstudio])( + "delivers the same %s discovery to a remounted caller and releases listeners on settlement", + async (provider) => { + setupCancellationMocks() + const timeoutController = new AbortController() + const timeoutSpy = vi.spyOn(AbortSignal, "timeout").mockReturnValue(timeoutController.signal) + const fetcher = provider === providerIdentifiers.poe ? vi.mocked(getPoeModels) : vi.mocked(getLMStudioModels) + let resolveDiscovery!: (models: ModelRecord) => void + fetcher.mockReturnValueOnce( + new Promise((resolve) => { + resolveDiscovery = resolve + }), + ) + try { + const controller = new AbortController() + const first = getModels({ provider, signal: controller.signal }) + controller.abort() + await expect(first).rejects.toMatchObject({ name: "AbortError" }) + const remounted = getModels({ provider }) + expect(fetcher).toHaveBeenCalledTimes(1) + const args = fetcher.mock.calls[0] + const options = args[args.length - 1] + expect(options).toEqual({ signal: expect.objectContaining({ aborted: false }) }) + resolveDiscovery(cancelledModels) + await expect(remounted).resolves.toEqual(cancelledModels) + expect(getEventListeners(controller.signal, "abort")).toHaveLength(0) + expect(getEventListeners(timeoutController.signal, "abort")).toHaveLength(0) + } finally { + resolveDiscovery(cancelledModels) + timeoutSpy.mockRestore() + } + }, +) diff --git a/src/api/providers/fetchers/kimi-code.ts b/src/api/providers/fetchers/kimi-code.ts index 02bf9180f7..6e757874e8 100644 --- a/src/api/providers/fetchers/kimi-code.ts +++ b/src/api/providers/fetchers/kimi-code.ts @@ -1,5 +1,7 @@ import { z } from "zod" +import { throwIfAborted } from "../utils/abort-signal" + import { KIMI_CODE_BASE_URL, kimiCodeDefaultModelInfo, @@ -37,9 +39,12 @@ export function mapKimiCodeModel(model: z.infer): Mo } } -export async function getKimiCodeModels(apiKey?: string): Promise { +export async function getKimiCodeModels(apiKey?: string, opts?: { signal?: AbortSignal }): Promise { + throwIfAborted(opts?.signal) if (!apiKey) throw new Error("Kimi Code authentication is required to fetch models") const controller = new AbortController() + const onAbort = () => controller.abort(opts?.signal?.reason) + opts?.signal?.addEventListener("abort", onAbort, { once: true }) const timeout = setTimeout( () => controller.abort(new Error("Kimi Code models request timed out")), KIMI_CODE_MODELS_TIMEOUT_MS, @@ -58,5 +63,6 @@ export async function getKimiCodeModels(apiKey?: string): Promise { return Object.fromEntries(parsed.data.map((model) => [model.id, mapKimiCodeModel(model)])) } finally { clearTimeout(timeout) + opts?.signal?.removeEventListener("abort", onAbort) } } diff --git a/src/api/providers/fetchers/lmstudio.ts b/src/api/providers/fetchers/lmstudio.ts index 32387e2f08..469ba71847 100644 --- a/src/api/providers/fetchers/lmstudio.ts +++ b/src/api/providers/fetchers/lmstudio.ts @@ -74,8 +74,8 @@ export async function getLMStudioModels( await axios.get(`${baseUrl}/v1/models`, { signal: opts?.signal }) // The SDK's model-list calls expose no cancellation option, so an abort - // during them cannot reach the network; releasing the shared cache entry - // and stopping the waiters happens at the model-cache layer. + // during them cannot reach the network. The model cache stops cancelled waiters + // but retains this operation until settlement to prevent duplicate discovery. const client = new LMStudioClient({ baseUrl: lmsUrl }) // First, try to get all downloaded models diff --git a/src/api/providers/fetchers/modelCache.ts b/src/api/providers/fetchers/modelCache.ts index 6f0898c71b..b09ca858ec 100644 --- a/src/api/providers/fetchers/modelCache.ts +++ b/src/api/providers/fetchers/modelCache.ts @@ -47,9 +47,9 @@ const modelRecordSchema = z.record(z.string(), modelInfoSchema) const inFlightRefresh = new Map() // Upper bound for any fetch started through the single-flight, so a hung endpoint can never keep -// an in-flight entry pending indefinitely. The value is the maximum of the 5–15 s bounds the -// individual fetchers it subsumes used to apply, relaxing rather than tightening endpoints that -// already had a bound. +// a waiter pending indefinitely. Non-cancellable SDK operations retain their entry until +// settlement. The value is the maximum of the 5–15 s bounds the individual fetchers used to +// apply, relaxing rather than tightening endpoints that already had a bound. const MODEL_CATALOG_FETCH_TIMEOUT_MS = 15_000 /** @@ -59,9 +59,9 @@ const MODEL_CATALOG_FETCH_TIMEOUT_MS = 15_000 * - The internal AbortController is the only object that may cancel the flight's network * request; caller signals are only ever merged into a per-waiter abort view, so one waiter * aborting can never cancel the fetch out from under the others. - * - When the last waiter detaches while the fetch is still pending, the flight is aborted and - * its map entry removed synchronously, so a caller arriving immediately afterwards starts a - * fresh fetch instead of joining a doomed one. + * - When the last waiter detaches, cancellable flights are aborted and removed immediately. + * SDK calls without cancellation support remain registered until settlement so new waiters + * cannot start overlapping discovery. * - Settlement never stores data in the map: it only removes the entry it created, guarded by * flight identity so a late-settling stale flight can never evict a newer one. */ @@ -71,6 +71,7 @@ type FlightRecord = { timeoutSignal: AbortSignal waiters: number pending: boolean + cancelWhenUnused: boolean } // Cache keys (see getCacheKey) for which we've already reported an empty model response this @@ -250,7 +251,7 @@ async function readModels(cacheKey: string): Promise { * @param options - Provider options for fetching models * @param signal - Cancellation signal forwarded to the dispatched fetcher. The single-flight * (dedupedFetch) passes its internal controller's signal; the auth-scoped direct path passes - * none, so those fetchers keep their own bounds. + * the caller's signal while fetchers retain their own timeout bounds. * @returns Fresh models from the provider API */ async function fetchModelsFromProvider(options: GetModelsOptions, signal?: AbortSignal): Promise { @@ -305,10 +306,13 @@ async function fetchModelsFromProvider(options: GetModelsOptions, signal?: Abort models = await getMoonshotModels(options.baseUrl, options.apiKey, ...fetchOpts) break case providerIdentifiers.zooGateway: - models = await getZooGatewayModels({ zooSessionToken: options.apiKey, zooGatewayBaseUrl: options.baseUrl }) + models = await getZooGatewayModels( + { zooSessionToken: options.apiKey, zooGatewayBaseUrl: options.baseUrl }, + ...fetchOpts, + ) break case providerIdentifiers.kimiCode: - models = await getKimiCodeModels(options.apiKey) + models = await getKimiCodeModels(options.apiKey, ...fetchOpts) break default: { // Ensures router is exhaustively checked if RouterName is a strict union. @@ -356,10 +360,10 @@ export const getModels = async (options: GetModelsOptions): Promise // getModels(), and a fetch failure joined from refreshModels() still re-throws for // getModels() callers. try { - // The auth-scoped fetch bypasses the single-flight entirely, so options.signal is - // deliberately not forwarded there: there is no shared entry to release on abort, and - // these fetchers bound their own requests. - const sharedFetch = shouldSkipCache ? fetchModelsFromProvider(options) : dedupedFetch(cacheKey, options) + // Auth-scoped fetches belong to this caller, so cancellation goes directly to the fetcher. + const sharedFetch = shouldSkipCache + ? fetchModelsFromProvider(options, options.signal) + : dedupedFetch(cacheKey, options) const fetched = await sharedFetch const modelCount = Object.keys(fetched).length @@ -412,9 +416,7 @@ function dedupedFetch(cacheKey: string, options: GetModelsOptions): Promise controller.abort(timeoutSignal.reason) - timeoutSignal.addEventListener("abort", onTimeout, { once: true }) - const removeTimeoutListener = () => timeoutSignal.removeEventListener("abort", onTimeout) + let removeTimeoutListener = () => {} // Settlement and a last-waiter abort may happen in either order; both paths are idempotent, // and the identity guard makes the two interleavings equivalent. @@ -456,11 +458,28 @@ function dedupedFetch(cacheKey: string, options: GetModelsOptions): Promise { + controller.abort(timeoutSignal.reason) + if (record.waiters === 0 && record.pending) { + guardedDelete() + } + } + timeoutSignal.addEventListener("abort", onTimeout, { once: true }) + removeTimeoutListener = () => timeoutSignal.removeEventListener("abort", onTimeout) + // The settle reactions above can only run after this function's current synchronous run -- // including the set() below -- completes, since that's the earliest a promise reaction can // fire. So the entry is always registered before any settle handler can delete it, even if @@ -474,8 +493,8 @@ function dedupedFetch(cacheKey: string, options: GetModelsOptions): Promise { throwIfAborted(callerSignal) @@ -491,7 +510,7 @@ function joinFlight(cacheKey: string, record: FlightRecord, callerSignal?: Abort detached = true removeViewListener?.() record.waiters-- - if (record.waiters === 0 && record.pending) { + if (record.waiters === 0 && record.pending && (record.cancelWhenUnused || record.timeoutSignal.aborted)) { // Synchronous release: abort the shared fetch and drop the entry before any further // await point runs, so a late joiner never sees a doomed flight. record.controller.abort() @@ -557,10 +576,10 @@ export const refreshModels = async (options: GetModelsOptions): Promise> { +export async function getZooGatewayModels( + options?: ApiHandlerOptions, + opts?: { signal?: AbortSignal }, +): Promise> { + throwIfAborted(opts?.signal) const models: Record = {} const baseURL = options?.zooGatewayBaseUrl ?? `${getZooCodeBaseUrl()}/api/gateway/v1` @@ -37,6 +43,7 @@ export async function getZooGatewayModels(options?: ApiHandlerOptions): Promise< const response = await axios.get(`${baseURL}/models`, { headers, timeout: MODEL_DISCOVERY_TIMEOUT_MS, + signal: opts?.signal, }) const result = vercelAiGatewayModelsResponseSchema.safeParse(response.data) @@ -57,6 +64,7 @@ export async function getZooGatewayModels(options?: ApiHandlerOptions): Promise< models[id] = parseZooGatewayModel({ id, model }) } } catch (error) { + throwIfAborted(opts?.signal) // Log only safe fields; never serialize the full error object because it // includes request config/headers which carry the bearer session token. const err = error as { diff --git a/src/api/providers/openai.ts b/src/api/providers/openai.ts index 8ccf3335f0..e44fb1774e 100644 --- a/src/api/providers/openai.ts +++ b/src/api/providers/openai.ts @@ -573,8 +573,14 @@ export class OpenAiHandler extends BaseProvider implements SingleCompletionHandl } } -export async function getOpenAiModels(baseUrl?: string, apiKey?: string, openAiHeaders?: Record) { +export async function getOpenAiModels( + baseUrl?: string, + apiKey?: string, + openAiHeaders?: Record, + signal?: AbortSignal, +) { try { + signal?.throwIfAborted() if (!baseUrl) { return [] } @@ -586,7 +592,7 @@ export async function getOpenAiModels(baseUrl?: string, apiKey?: string, openAiH return [] } - const config: Record = {} + const config: Record = signal ? { signal } : {} const headers: Record = { ...DEFAULT_HEADERS, ...(openAiHeaders || {}), diff --git a/src/core/config/ProviderSettingsManager.ts b/src/core/config/ProviderSettingsManager.ts index 3fcc0e6e43..0aa93052a0 100644 --- a/src/core/config/ProviderSettingsManager.ts +++ b/src/core/config/ProviderSettingsManager.ts @@ -4,6 +4,7 @@ import deepEqual from "fast-deep-equal" import { type ProviderSettingsWithId, + type ProviderSettings, providerSettingsWithIdSchema, discriminatedProviderSettingsWithIdSchema, isSecretStateKey, @@ -52,6 +53,12 @@ export const providerProfilesSchema = z.object({ export type ProviderProfiles = z.infer +export class ProviderProfileConflictError extends Error { + constructor(name: string) { + super(`Provider profile '${name}' changed during update; rollback was skipped`) + } +} + export class ProviderSettingsManager { private static readonly SCOPE_PREFIX = "roo_cline_config_" private readonly defaultConfigId = this.generateId() @@ -374,34 +381,90 @@ export class ProviderSettingsManager { } } + private normalizeConfig(config: ProviderSettingsWithId): ProviderSettingsWithId { + const normalizedConfig = downgradeLegacyRooConfig(config as Record) + .config as ProviderSettingsWithId + + // For active providers, filter out settings from other providers. + // For retired providers, preserve full profile fields (including legacy + // provider-specific keys) to avoid data loss — passthrough() keeps + // unknown keys that strict parse() would strip. + return typeof normalizedConfig.apiProvider === "string" && isRetiredProvider(normalizedConfig.apiProvider) + ? providerSettingsWithIdSchema.passthrough().parse(normalizedConfig) + : discriminatedProviderSettingsWithIdSchema.parse(normalizedConfig) + } + + private assertConfigUnchanged(profiles: ProviderProfiles, name: string, expected?: ProviderSettingsWithId) { + // Match the JSON representation written to storage (undefined fields are omitted). + if ( + expected && + !deepEqual(profiles.apiConfigs[name], JSON.parse(JSON.stringify(this.normalizeConfig(expected)))) + ) { + throw new ProviderProfileConflictError(name) + } + } + /** * Save a config with the given name. * Preserves the ID from the input 'config' object if it exists, * otherwise generates a new one (for creation scenarios). + * A baseline applies only fields edited by the caller, rejecting overlapping edits. + * Renames update both names in one stored write under the same lock. */ - public async saveConfig(name: string, config: ProviderSettingsWithId): Promise { + public async saveConfig( + name: string, + config: ProviderSettingsWithId, + expected?: ProviderSettingsWithId, + baseline?: ProviderSettings, + renameFrom?: string, + ): Promise { try { return await this.lock(async () => { const providerProfiles = await this.load() + const sourceName = renameFrom ?? name + this.assertConfigUnchanged(providerProfiles, sourceName, expected) + if (renameFrom && (!providerProfiles.apiConfigs[renameFrom] || providerProfiles.apiConfigs[name])) { + throw new ProviderProfileConflictError(name) + } // Preserve the existing ID if this is an update to an existing config. - const existingId = providerProfiles.apiConfigs[name]?.id + const existingId = providerProfiles.apiConfigs[sourceName]?.id const id = config.id || existingId || this.generateId() - const normalizedConfig = downgradeLegacyRooConfig(config as Record) - .config as ProviderSettingsWithId - - // For active providers, filter out settings from other providers. - // For retired providers, preserve full profile fields (including legacy - // provider-specific keys) to avoid data loss — passthrough() keeps - // unknown keys that strict parse() would strip. - const filteredConfig = - typeof normalizedConfig.apiProvider === "string" && isRetiredProvider(normalizedConfig.apiProvider) - ? providerSettingsWithIdSchema.passthrough().parse(normalizedConfig) - : discriminatedProviderSettingsWithIdSchema.parse(normalizedConfig) + let filteredConfig = this.normalizeConfig(config) + if (baseline && providerProfiles.apiConfigs[sourceName]) { + const current = providerProfiles.apiConfigs[sourceName] + const original = this.normalizeConfig(baseline) + if ( + current.apiProvider !== original.apiProvider && + current.apiProvider !== filteredConfig.apiProvider + ) { + throw new ProviderProfileConflictError(sourceName) + } + const merged = { ...current } + for (const key of new Set([...Object.keys(original), ...Object.keys(filteredConfig)])) { + if (key === "id") continue + const field = key as keyof ProviderSettings + if (deepEqual(original[field], filteredConfig[field])) continue + if ( + !deepEqual(current[field], original[field]) && + !deepEqual(current[field], filteredConfig[field]) + ) { + throw new ProviderProfileConflictError(name) + } + Object.assign(merged, { [field]: filteredConfig[field] }) + } + filteredConfig = this.normalizeConfig(merged) + } providerProfiles.apiConfigs[name] = { ...filteredConfig, id } + if (renameFrom) { + delete providerProfiles.apiConfigs[renameFrom] + if (providerProfiles.currentApiConfigName === renameFrom) + providerProfiles.currentApiConfigName = name + } await this.store(providerProfiles) return id }) } catch (error) { + if (error instanceof ProviderProfileConflictError) throw error throw new Error(`Failed to save config: ${error}`) } } @@ -468,10 +531,11 @@ export class ProviderSettingsManager { /** * Delete a config by name. */ - public async deleteConfig(name: string) { + public async deleteConfig(name: string, expected?: ProviderSettingsWithId) { try { return await this.lock(async () => { const providerProfiles = await this.load() + this.assertConfigUnchanged(providerProfiles, name, expected) if (!providerProfiles.apiConfigs[name]) { throw new Error(`Config '${name}' not found`) @@ -485,6 +549,7 @@ export class ProviderSettingsManager { await this.store(providerProfiles) }) } catch (error) { + if (error instanceof ProviderProfileConflictError) throw error throw new Error(`Failed to delete config: ${error}`) } } @@ -506,7 +571,7 @@ export class ProviderSettingsManager { /** * Set the API config for a specific mode. */ - public async setModeConfig(mode: Mode, configId: string) { + public async setModeConfig(mode: Mode, configId: string | undefined) { try { return await this.lock(async () => { const providerProfiles = await this.load() @@ -515,7 +580,11 @@ export class ProviderSettingsManager { providerProfiles.modeApiConfigs = {} } // Assign the chosen config ID to this mode - providerProfiles.modeApiConfigs[mode] = configId + if (configId === undefined) { + delete providerProfiles.modeApiConfigs[mode] + } else { + providerProfiles.modeApiConfigs[mode] = configId + } await this.store(providerProfiles) }) } catch (error) { diff --git a/src/core/config/__tests__/ProviderSettingsManager.spec.ts b/src/core/config/__tests__/ProviderSettingsManager.spec.ts index 13e1aeeb2d..57eb13de2a 100644 --- a/src/core/config/__tests__/ProviderSettingsManager.spec.ts +++ b/src/core/config/__tests__/ProviderSettingsManager.spec.ts @@ -12,7 +12,12 @@ import { import { clearAllMocks } from "../../../test-utils/reset" import { makeExtensionContext } from "../../../test-utils/vscode" -import { ProviderSettingsManager, ProviderProfiles, SyncCloudProfilesResult } from "../ProviderSettingsManager" +import { + ProviderProfileConflictError, + ProviderSettingsManager, + ProviderProfiles, + SyncCloudProfilesResult, +} from "../ProviderSettingsManager" // `export()` builds an API handler per profile to read model capabilities. Mock // buildApiHandler with the real @roo-code/types model definitions so the token-field @@ -73,6 +78,147 @@ describe("ProviderSettingsManager", () => { providerSettingsManager = new ProviderSettingsManager(mockContext) }) + describe("concurrent profile edits", () => { + const original = { + apiProvider: providerIdentifiers.openrouter, + openRouterApiKey: "old-key", + openRouterModelId: "old-model", + } + beforeEach(async () => { + const storage = new Map() + mockSecrets.get.mockImplementation(async (key: string) => storage.get(key)) + mockSecrets.store.mockImplementation(async (key: string, value: string) => { + storage.set(key, value) + }) + await providerSettingsManager.saveConfig("profile", original) + }) + + it.each(["save", "delete"] as const)("allows %s with a matching expected profile", async (operation) => { + const expected = await providerSettingsManager.getProfile({ name: "profile" }) + if (operation === "save") { + await providerSettingsManager.saveConfig("profile", { ...original, openRouterModelId: "new" }, expected) + expect(await providerSettingsManager.getProfile({ name: "profile" })).toMatchObject({ + openRouterModelId: "new", + }) + } else { + await providerSettingsManager.deleteConfig("profile", expected) + expect(await providerSettingsManager.hasConfig("profile")).toBe(false) + } + }) + + it.each(["save", "delete"] as const)( + "rejects %s without writing when the expected profile is stale", + async (operation) => { + const expected = await providerSettingsManager.getProfile({ name: "profile" }) + await providerSettingsManager.saveConfig("profile", { ...original, openRouterApiKey: "new-key" }) + const stored = await providerSettingsManager.export() + mockSecrets.store.mockClear() + const mutation = + operation === "save" + ? providerSettingsManager.saveConfig("profile", original, expected) + : providerSettingsManager.deleteConfig("profile", expected) + await expect(mutation).rejects.toBeInstanceOf(ProviderProfileConflictError) + expect(mockSecrets.store).not.toHaveBeenCalled() + expect(await providerSettingsManager.export()).toEqual(stored) + }, + ) + + it.each([true, false])( + "preserves concurrent credential and model edits (model first: %s)", + async (modelFirst) => { + const model = { ...original, openRouterModelId: "new-model" } + const settings = { ...original, openRouterApiKey: "new-key" } + for (const edit of modelFirst ? [model, settings] : [settings, model]) { + await providerSettingsManager.saveConfig("profile", edit, undefined, original) + } + expect(await providerSettingsManager.getProfile({ name: "profile" })).toMatchObject({ + openRouterApiKey: "new-key", + openRouterModelId: "new-model", + }) + }, + ) + + it("rejects conflicting edits and provider switches without writes", async () => { + await providerSettingsManager.saveConfig("profile", { ...original, openRouterModelId: "new-model" }) + mockSecrets.store.mockClear() + await expect( + providerSettingsManager.saveConfig( + "profile", + { ...original, openRouterModelId: "other-model" }, + undefined, + original, + ), + ).rejects.toBeInstanceOf(ProviderProfileConflictError) + expect(mockSecrets.store).not.toHaveBeenCalled() + await providerSettingsManager.saveConfig("profile", { apiProvider: providerIdentifiers.anthropic }) + mockSecrets.store.mockClear() + await expect( + providerSettingsManager.saveConfig( + "profile", + { ...original, openRouterModelId: "other-model" }, + undefined, + original, + ), + ).rejects.toBeInstanceOf(ProviderProfileConflictError) + expect(mockSecrets.store).not.toHaveBeenCalled() + }) + + it("accepts identical concurrent provider changes", async () => { + const updated = { apiProvider: providerIdentifiers.anthropic, apiKey: "new-provider-key" } + await providerSettingsManager.saveConfig("profile", updated, undefined, original) + await providerSettingsManager.saveConfig("profile", updated, undefined, original) + expect(await providerSettingsManager.getProfile({ name: "profile" })).toMatchObject(updated) + }) + + it("renames atomically while preserving a concurrent model edit and profile ID", async () => { + const { id } = await providerSettingsManager.getProfile({ name: "profile" }) + await providerSettingsManager.saveConfig("profile", { ...original, openRouterModelId: "new-model" }) + await providerSettingsManager.saveConfig( + "renamed", + { ...original, openRouterApiKey: "new-key" }, + undefined, + original, + "profile", + ) + expect(await providerSettingsManager.hasConfig("profile")).toBe(false) + expect(await providerSettingsManager.getProfile({ name: "renamed" })).toMatchObject({ + id, + openRouterModelId: "new-model", + openRouterApiKey: "new-key", + }) + await expect( + providerSettingsManager.saveConfig("default", original, undefined, original, "renamed"), + ).rejects.toBeInstanceOf(ProviderProfileConflictError) + expect(await providerSettingsManager.hasConfig("renamed")).toBe(true) + }) + }) + + describe("setModeConfig", () => { + it("persists a string mapping and removes only the unset mode after reload", async () => { + const storage = new Map() + mockSecrets.get.mockImplementation(async (key: string) => storage.get(key)) + mockSecrets.store.mockImplementation(async (key: string, value: string) => { + storage.set(key, value) + }) + + await providerSettingsManager.setModeConfig("code", "code-profile") + await providerSettingsManager.setModeConfig("ask", "ask-profile") + const reloaded = new ProviderSettingsManager(mockContext) + expect(await reloaded.getModeConfigId("code")).toBe("code-profile") + const storedModes = (await reloaded.export()).modeApiConfigs + expect(storedModes).toMatchObject({ code: "code-profile", ask: "ask-profile" }) + + await reloaded.setModeConfig("code", undefined) + const afterUnset = new ProviderSettingsManager(mockContext) + expect(await afterUnset.getModeConfigId("code")).toBeUndefined() + const remainingModes = (await afterUnset.export()).modeApiConfigs + expect(remainingModes).not.toHaveProperty("code") + expect(remainingModes).toEqual( + Object.fromEntries(Object.entries(storedModes ?? {}).filter(([mode]) => mode !== "code")), + ) + }) + }) + describe("initialize", () => { it("should not write to storage when secrets.get returns null", async () => { // Mock readConfig to return null diff --git a/src/core/webview/ClineProvider.ts b/src/core/webview/ClineProvider.ts index 60a92eebaa..d9c902d74c 100644 --- a/src/core/webview/ClineProvider.ts +++ b/src/core/webview/ClineProvider.ts @@ -110,11 +110,12 @@ import { buildApiHandler } from "../../api" import { forceFullModelDetailsLoad, hasLoadedFullDetails } from "../../api/providers/fetchers/lmstudio" import { ContextProxy } from "../config/ContextProxy" -import { ProviderSettingsManager } from "../config/ProviderSettingsManager" +import { ProviderProfileConflictError, ProviderSettingsManager } from "../config/ProviderSettingsManager" import { CustomModesManager } from "../config/CustomModesManager" import { PendingActionSettlementError, Task } from "../task/Task" import { webviewMessageHandler } from "./webviewMessageHandler" +import { ModelRequestRegistry } from "./ModelRequestRegistry" import type { WebviewFocusTracker } from "./WebviewFocusTracker" import type { ClineMessage, TodoItem } from "@roo-code/types" import { @@ -251,6 +252,7 @@ export class ClineProvider public readonly taskHistoryStore: TaskHistoryStore private taskHistoryStoreInitialized = false public static readonly PENDING_OPERATION_TIMEOUT_MS = 30000 // 30 seconds + public readonly modelRequests = new ModelRequestRegistry() private providerProfileMutationQueue = Promise.resolve() private historyTaskCreationQueue = Promise.resolve() @@ -260,15 +262,28 @@ export class ClineProvider private enqueueProviderProfileMutation(fn: (signal: AbortSignal) => Promise): Promise { const controller = new AbortController() - // Run fn after either outcome so a rejected mutation never poisons the queue. - const run = this.providerProfileMutationQueue.then( - () => fn(controller.signal), - () => fn(controller.signal), - ) - const callerResult = this.withProviderProfileMutationTimeout(run, () => { - controller.abort() - this.log("Provider profile mutation timed out; aborting in-flight mutation") + // The queue tail always resolves. Skip callers whose queue wait expired before + // allowing their mutation to perform any work. + const run = this.providerProfileMutationQueue.then(() => { + controller.signal.throwIfAborted() + return fn(controller.signal) }) + // Allow a separate wait window for an earlier operation and compensation; + // waiting must not consume this mutation's execution timeout. A stalled queue + // still returns an error to the caller without permitting concurrent writes. + const callerResult = this.withProviderProfileMutationTimeout( + this.providerProfileMutationQueue, + () => { + controller.abort() + this.log("Provider profile mutation queue wait timed out; skipping queued mutation") + }, + ClineProvider.PENDING_OPERATION_TIMEOUT_MS * 2, + ).then(() => + this.withProviderProfileMutationTimeout(run, () => { + controller.abort() + this.log("Provider profile mutation timed out; aborting in-flight mutation") + }), + ) void run.then( () => { @@ -287,22 +302,26 @@ export class ClineProvider }, ) - // Advance from the timeout-bounded result. Each fn checks its AbortSignal before - // writing state, so advancing the queue on timeout cannot produce stale overwrites. - this.providerProfileMutationQueue = callerResult.then( + // Keep ownership until all writes and compensation settle, even if the caller + // has timed out. Otherwise an older rollback can overwrite a later mutation. + this.providerProfileMutationQueue = run.then( () => undefined, () => undefined, ) return callerResult } - private withProviderProfileMutationTimeout(operation: Promise, onTimeout: () => void): Promise { + private withProviderProfileMutationTimeout( + operation: Promise, + onTimeout: () => void, + timeoutMs = ClineProvider.PENDING_OPERATION_TIMEOUT_MS, + ): Promise { let timeoutId: ReturnType | undefined const timeout = new Promise((_, reject) => { timeoutId = setTimeout(() => { onTimeout() reject(new Error("Provider profile mutation timed out")) - }, ClineProvider.PENDING_OPERATION_TIMEOUT_MS) + }, timeoutMs) }) return Promise.race([operation, timeout]).finally(() => { @@ -840,6 +859,7 @@ export class ClineProvider } this._disposed = true + this.modelRequests.dispose() this._postStateToWebviewThrottled.cancel() this.log("Disposing ClineProvider...") @@ -1876,57 +1896,118 @@ export class ClineProvider name: string, providerSettings: ProviderSettings, activate: boolean = true, + baseline?: ProviderSettings, + renameFrom?: string, ): Promise { try { return await this.enqueueProviderProfileMutation(async (signal) => { - // TODO: Do we need to be calling `activateProfile`? It's not - // clear to me what the source of truth should be; in some cases - // we rely on the `ContextProxy`'s data store and in other cases - // we rely on the `ProviderSettingsManager`'s data store. It might - // be simpler to unify these two. - const id = await this.providerSettingsManager.saveConfig(name, providerSettings) - - if (signal.aborted) return id - - if (activate) { - const { mode } = await this.getState() - - // These promises do the following: - // 1. Adds or updates the list of provider profiles. - // 2. Sets the current provider profile. - // 3. Sets the current mode's provider profile. - // 4. Copies the provider settings to the context. - // - // Note: 1, 2, and 4 can be done in one `ContextProxy` call: - // this.contextProxy.setValues({ ...providerSettings, listApiConfigMeta: ..., currentApiConfigName: ... }) - // We should probably switch to that and verify that it works. - // I left the original implementation in just to be safe. - await Promise.all([ - this.updateGlobalState("listApiConfigMeta", await this.providerSettingsManager.listConfig()), - this.updateGlobalState("currentApiConfigName", name), - this.providerSettingsManager.setModeConfig(mode, id), - this.contextProxy.setProviderSettings(providerSettings), - ]) - - // Change the provider for the current task. - // TODO: We should rename `buildApiHandler` for clarity (e.g. `getProviderClient`). - this.updateTaskApiHandlerIfNeeded(providerSettings, { forceRebuild: true }) - - // Keep the current task's sticky provider profile in sync with the newly-activated profile. - await this.persistStickyProviderProfileToCurrentTask(name) - } else { - await this.updateGlobalState("listApiConfigMeta", await this.providerSettingsManager.listConfig()) - } + signal.throwIfAborted() + // Snapshot and compensate inside the queue so a failed save cannot + // overwrite a later successful mutation. + const sourceName = renameFrom ?? name + const previousProfile = (await this.providerSettingsManager.hasConfig(sourceName)) + ? await this.providerSettingsManager.getProfile({ name: sourceName }) + : undefined + const previousSettings = this.contextProxy.getProviderSettings() + const previousName = this.contextProxy.getValue("currentApiConfigName") + const previousMeta = this.contextProxy.getValue("listApiConfigMeta") + const task = activate ? this.getCurrentTask() : undefined + const previousTaskSettings = task?.apiConfiguration + const mode = activate ? (await this.getState()).mode : undefined + const previousModeConfigId = mode ? await this.providerSettingsManager.getModeConfigId(mode) : undefined + let savedId: string | undefined - await this.postStateToWebview() - return id + try { + signal.throwIfAborted() + const id = await this.providerSettingsManager.saveConfig( + name, + providerSettings, + previousProfile, + baseline, + renameFrom, + ) + savedId = id + // Activate and compensate the merged stored profile, including concurrent edits. + providerSettings = await this.providerSettingsManager.getProfile({ name }) + const { organizationAllowList } = await this.getState() + if ( + !ProfileValidator.isProfileAllowed( + providerSettings, + organizationAllowList ?? ORGANIZATION_ALLOW_ALL, + ) + ) { + throw new OrganizationAllowListViolationError( + t("common:errors.violated_organization_allowlist"), + ) + } + signal.throwIfAborted() + const listApiConfig = await this.providerSettingsManager.listConfig() + await this.updateGlobalState("listApiConfigMeta", listApiConfig) + if (activate && mode) { + // Sequential writes must settle before compensation starts. + await this.updateGlobalState("currentApiConfigName", name) + await this.providerSettingsManager.setModeConfig(mode, id) + await this.contextProxy.setProviderSettings(providerSettings) + this.updateTaskApiHandlerIfNeeded(providerSettings, { forceRebuild: true }) + } + signal.throwIfAborted() + await this.postStateToWebview() + signal.throwIfAborted() + if (activate) await this.persistStickyProviderProfileToCurrentTask(name) + return id + } catch (error) { + if (savedId) { + // Attempt every compensation even if an individual store is unavailable. + let profileConflict = false + const restore = async (write: () => unknown | Promise) => { + try { + await write() + } catch (rollbackError) { + this.log(`Failed to roll back provider profile state: ${String(rollbackError)}`) + if (rollbackError instanceof ProviderProfileConflictError) profileConflict = true + } + } + await restore(() => + previousProfile + ? this.providerSettingsManager.saveConfig( + sourceName, + previousProfile, + { + ...providerSettings, + id: savedId, + }, + undefined, + renameFrom ? name : undefined, + ) + : this.providerSettingsManager.deleteConfig(name, { ...providerSettings, id: savedId }), + ) + if (activate && mode) { + await restore(() => this.providerSettingsManager.setModeConfig(mode, previousModeConfigId)) + await restore(() => this.contextProxy.setProviderSettings(previousSettings)) + await restore(() => this.updateGlobalState("currentApiConfigName", previousName)) + if (task && previousTaskSettings) { + await restore(() => task.updateApiConfiguration(previousTaskSettings)) + } + } + // A direct settings save may have published newer profile metadata. + if (!profileConflict) { + await restore(() => this.updateGlobalState("listApiConfigMeta", previousMeta)) + } + await restore(() => this.postStateToWebview()) + } + throw error + } }) } catch (error) { this.log( `Error create new api configuration: ${JSON.stringify(error, Object.getOwnPropertyNames(error), 2)}`, ) - vscode.window.showErrorMessage(t("common:errors.create_api_config")) + vscode.window.showErrorMessage( + error instanceof OrganizationAllowListViolationError + ? error.message + : t("common:errors.create_api_config"), + ) return undefined } } diff --git a/src/core/webview/ModelRequestRegistry.ts b/src/core/webview/ModelRequestRegistry.ts new file mode 100644 index 0000000000..ba4b8cbd3a --- /dev/null +++ b/src/core/webview/ModelRequestRegistry.ts @@ -0,0 +1,29 @@ +/** Model discovery work owned by one webview. */ +export class ModelRequestRegistry { + private readonly requests = new Map() + private disposed = false + + async run(id: string, fetch: (signal: AbortSignal) => Promise): Promise { + if (this.disposed) return + this.cancel(id) + const controller = new AbortController() + this.requests.set(id, controller) + try { + await fetch(controller.signal) + } catch (error) { + if (!controller.signal.aborted) throw error + } finally { + if (this.requests.get(id) === controller) this.requests.delete(id) + } + } + + cancel(id: string): void { + this.requests.get(id)?.abort() + this.requests.delete(id) + } + + dispose(): void { + this.disposed = true + for (const id of this.requests.keys()) this.cancel(id) + } +} diff --git a/src/core/webview/__tests__/ClineProvider.apiHandlerRebuild.spec.ts b/src/core/webview/__tests__/ClineProvider.apiHandlerRebuild.spec.ts index 4561d688fe..bfdd6f019e 100644 --- a/src/core/webview/__tests__/ClineProvider.apiHandlerRebuild.spec.ts +++ b/src/core/webview/__tests__/ClineProvider.apiHandlerRebuild.spec.ts @@ -1,3 +1,5 @@ +import { t } from "../../../i18n" +import { ProviderSettingsManager } from "../../config/ProviderSettingsManager" import { WebviewFocusTracker } from "../WebviewFocusTracker" // npx vitest core/webview/__tests__/ClineProvider.apiHandlerRebuild.spec.ts @@ -5,7 +7,7 @@ import * as vscode from "vscode" import { makeCompositeDisposable } from "../../../test-utils/vscode" import { TelemetryService } from "@roo-code/telemetry" -import { getModelId, RooCodeEventName } from "@roo-code/types" +import { getModelId, RooCodeEventName, type ProviderSettings } from "@roo-code/types" import { ContextProxy } from "../../config/ContextProxy" import type { Mode } from "../../../shared/modes" @@ -241,9 +243,18 @@ describe("ClineProvider - API Handler Rebuild Guard", () => { new WebviewFocusTracker(), ) + let storedSettings: ProviderSettings = { + apiProvider: providerIdentifiers.openrouter, + openRouterModelId: "openai/gpt-4", + } // Mock providerSettingsManager ;(provider as any).providerSettingsManager = { - saveConfig: vi.fn().mockResolvedValue("test-id"), + saveConfig: vi.fn().mockImplementation(async (_name: string, settings: ProviderSettings) => { + storedSettings = settings + return "test-id" + }), + hasConfig: vi.fn().mockResolvedValue(true), + deleteConfig: vi.fn(), listConfig: vi.fn().mockResolvedValue([ { name: "test-config", @@ -260,12 +271,9 @@ describe("ClineProvider - API Handler Rebuild Guard", () => { apiProvider: providerIdentifiers.openrouter, openRouterModelId: "openai/gpt-4", }), - getProfile: vi.fn().mockResolvedValue({ - name: "test-config", - id: "test-id", - apiProvider: providerIdentifiers.openrouter, - openRouterModelId: "openai/gpt-4", - }), + getProfile: vi + .fn() + .mockImplementation(async () => ({ name: "test-config", id: "test-id", ...storedSettings })), } // Get the buildApiHandler mock @@ -426,6 +434,287 @@ describe("ClineProvider - API Handler Rebuild Guard", () => { // Should not call buildApiHandler when there's no task expect(buildApiHandlerMock).not.toHaveBeenCalled() }) + + test("rolls back in-memory provider state when a failure occurs after mutation", async () => { + // Seed a known previously-active state. + await provider["contextProxy"].setProviderSettings({ + apiProvider: providerIdentifiers.openrouter, + openRouterModelId: "openai/gpt-4", + }) + await provider["updateGlobalState"]("currentApiConfigName", "previous-config") + + // Fail the first state broadcast — which happens only after the + // in-memory state has already been mutated — to simulate a partial + // failure. The rollback's own broadcast should still succeed. + const postSpy = vi + .spyOn(provider, "postStateToWebview") + .mockRejectedValueOnce(new Error("broadcast failed")) + + const result = await provider.upsertProviderProfile( + "new-config", + { + apiProvider: providerIdentifiers.anthropic, + apiModelId: "claude-3-5-sonnet-20241022", + }, + true, + ) + + // The half-applied profile must not survive: the previous settings and + // profile name are restored so the store and in-memory context stay in sync. + expect(result).toBeUndefined() + expect(provider["contextProxy"].getProviderSettings()).toMatchObject({ + apiProvider: providerIdentifiers.openrouter, + openRouterModelId: "openai/gpt-4", + }) + expect(provider["contextProxy"].getValue("currentApiConfigName")).toBe("previous-config") + + postSpy.mockRestore() + }) + }) + + describe("persisted profile compensation", () => { + const oldSettings = { apiProvider: providerIdentifiers.openrouter, openRouterModelId: "openai/gpt-4" } + const newSettings = { apiProvider: providerIdentifiers.openrouter, openRouterModelId: "new-model" } + beforeEach(async () => { + Object.defineProperty(provider, "providerSettingsManager", { + value: new ProviderSettingsManager(mockContext), + configurable: true, + }) + await provider.contextProxy.setValue("mode", "code") + await provider.providerSettingsManager.saveConfig("existing", oldSettings) + await provider.providerSettingsManager.setModeConfig("code", undefined) + await provider.contextProxy.setProviderSettings(oldSettings) + await provider.contextProxy.setValue("currentApiConfigName", "existing") + await provider.contextProxy.setValue( + "listApiConfigMeta", + await provider.providerSettingsManager.listConfig(), + ) + }) + + test.each([true, false])("publishes merged settings and model edits (model first: %s)", async (modelFirst) => { + const settings = { ...oldSettings, openRouterApiKey: "updated-key" } + const edits = modelFirst ? [newSettings, settings] : [settings, newSettings] + const task = new Task({ ...defaultTaskOptions, apiConfiguration: oldSettings }) + await provider.addClineToStack(task) + const results = await Promise.all( + edits.map((edit) => provider.upsertProviderProfile("existing", edit, true, oldSettings)), + ) + expect(results.every(Boolean)).toBe(true) + const expected = { ...newSettings, openRouterApiKey: "updated-key" } + expect(await provider.providerSettingsManager.getProfile({ name: "existing" })).toMatchObject(expected) + expect(provider.contextProxy.getProviderSettings()).toMatchObject(expected) + expect(task.apiConfiguration).toMatchObject(expected) + }) + + test("rolls back a policy rejection and reports its specific error for an inactive save", async () => { + vi.spyOn(provider, "getState").mockResolvedValue({ + ...(await provider.getState()), + organizationAllowList: { + allowAll: false, + providers: { + [providerIdentifiers.openrouter]: { allowAll: false, models: [oldSettings.openRouterModelId] }, + }, + }, + }) + expect(await provider.upsertProviderProfile("existing", newSettings, false, oldSettings)).toBeUndefined() + expect(await provider.providerSettingsManager.getProfile({ name: "existing" })).toMatchObject(oldSettings) + expect(vscode.window.showErrorMessage).toHaveBeenCalledWith( + t("common:errors.violated_organization_allowlist"), + ) + }) + + test("compensates a rename and its activation before processing another save", async () => { + const before = await provider.providerSettingsManager.getProfile({ name: "existing" }) + vi.spyOn(provider, "postStateToWebview").mockRejectedValueOnce(new Error("broadcast failed")) + expect( + await provider.upsertProviderProfile("renamed", newSettings, true, oldSettings, "existing"), + ).toBeUndefined() + expect(await provider.providerSettingsManager.hasConfig("renamed")).toBe(false) + expect(await provider.providerSettingsManager.getProfile({ name: "existing" })).toEqual(before) + expect(provider.contextProxy.getValue("currentApiConfigName")).toBe("existing") + expect(provider.contextProxy.getProviderSettings()).toMatchObject(oldSettings) + expect(await provider.upsertProviderProfile("existing", newSettings, true, oldSettings)).toBeTruthy() + }) + + test.each(["list", "activation", "broadcast"] as const)( + "restores existing and removes new profiles after %s failure", + async (phase) => { + const manager = provider.providerSettingsManager + const previous = await manager.getProfile({ name: "existing" }) + const previousMeta = provider.contextProxy.getValue("listApiConfigMeta") + const task = new Task({ ...defaultTaskOptions, apiConfiguration: oldSettings }) + await provider.addClineToStack(task) + for (const name of ["existing", "new-profile"]) { + if (phase === "list") + vi.spyOn(manager, "listConfig").mockRejectedValueOnce(new Error("list failed")) + if (phase === "activation") + vi.spyOn(provider.contextProxy, "setProviderSettings").mockRejectedValueOnce( + new Error("activation failed"), + ) + if (phase === "broadcast") + vi.spyOn(provider, "postStateToWebview").mockRejectedValueOnce(new Error("broadcast failed")) + expect(await provider.upsertProviderProfile(name, newSettings)).toBeUndefined() + expect(await manager.getProfile({ name: "existing" })).toEqual(previous) + expect(await manager.hasConfig("new-profile")).toBe(false) + expect(await manager.getModeConfigId("code")).toBeUndefined() + expect(provider.contextProxy.getProviderSettings()).toMatchObject(oldSettings) + expect(provider.contextProxy.getValue("currentApiConfigName")).toBe("existing") + expect(provider.contextProxy.getValue("listApiConfigMeta")).toEqual(previousMeta) + expect(task.apiConfiguration).toEqual(oldSettings) + expect(task.setTaskApiConfigName).not.toHaveBeenCalled() + } + }, + ) + + test.each(["existing", "new-profile"])( + "preserves a concurrent settings save for %s when upsert fails", + async (name) => { + const manager = provider.providerSettingsManager + const finalSettings = { ...oldSettings, openRouterModelId: "concurrent-model" } + const task = new Task({ ...defaultTaskOptions, apiConfiguration: oldSettings }) + await provider.addClineToStack(task) + let failBroadcast!: (reason: Error) => void + const broadcast = vi.spyOn(provider, "postStateToWebview").mockImplementationOnce( + () => + new Promise((_resolve, reject) => { + failBroadcast = reject + }), + ) + const log = vi.spyOn(provider, "log") + const pending = provider.upsertProviderProfile(name, newSettings) + await vi.waitFor(() => expect(broadcast).toHaveBeenCalledTimes(1)) + const upserted = await manager.getProfile({ name }) + + // Non-webview direct manager writers can bypass the provider queue. + await manager.saveConfig(name, finalSettings) + const metadata = await manager.listConfig() + await provider.contextProxy.setValue("listApiConfigMeta", metadata) + failBroadcast(new Error("broadcast failed")) + expect(await pending).toBeUndefined() + + const reloaded = new ProviderSettingsManager(mockContext) + expect(await reloaded.getProfile({ name })).toMatchObject({ ...finalSettings, id: upserted.id }) + expect(provider.contextProxy.getValue("listApiConfigMeta")).toEqual(metadata) + expect(await reloaded.getModeConfigId("code")).toBeUndefined() + expect(provider.contextProxy.getProviderSettings()).toMatchObject(oldSettings) + expect(provider.contextProxy.getValue("currentApiConfigName")).toBe("existing") + expect(task.apiConfiguration).toEqual(oldSettings) + expect(broadcast).toHaveBeenCalledTimes(2) + expect(log).toHaveBeenCalledWith(expect.stringContaining("rollback was skipped")) + }, + ) + + test("holds later model updates until a timed-out save finishes compensation", async () => { + vi.useFakeTimers() + try { + const manager = provider.providerSettingsManager + const save = manager.saveConfig.bind(manager) + let releaseBroadcast!: () => void + let releaseRollback!: () => void + const broadcast = vi.spyOn(provider, "postStateToWebview").mockImplementationOnce( + () => + new Promise((resolve) => { + releaseBroadcast = resolve + }), + ) + const saves = vi + .spyOn(manager, "saveConfig") + .mockImplementationOnce(save) + .mockImplementationOnce(async (name, settings, expected) => { + await new Promise((resolve) => { + releaseRollback = resolve + }) + return save(name, settings, expected) + }) + const first = provider.upsertProviderProfile("existing", newSettings) + await vi.advanceTimersByTimeAsync(0) + expect(broadcast).toHaveBeenCalledTimes(1) + await vi.advanceTimersByTimeAsync(ClineProvider.PENDING_OPERATION_TIMEOUT_MS) + expect(await first).toBeUndefined() + + const finalSettings = { ...oldSettings, openRouterModelId: "final-model" } + let secondSettled = false + const second = provider.upsertProviderProfile("existing", finalSettings).then((result) => { + secondSettled = true + return result + }) + await vi.advanceTimersByTimeAsync(0) + expect(saves).toHaveBeenCalledTimes(1) + releaseBroadcast() + await vi.advanceTimersByTimeAsync(0) + expect(saves).toHaveBeenCalledTimes(2) + expect(saves.mock.calls[1][1]).toMatchObject(oldSettings) + // Waiting for compensation must not consume the next save's execution window. + await vi.advanceTimersByTimeAsync(ClineProvider.PENDING_OPERATION_TIMEOUT_MS + 1) + expect(secondSettled).toBe(false) + expect(saves).toHaveBeenCalledTimes(2) + releaseRollback() + expect(await second).toBeTruthy() + expect(saves).toHaveBeenCalledTimes(3) + const reloaded = new ProviderSettingsManager(mockContext) + const profile = await reloaded.getProfile({ name: "existing" }) + expect(profile).toMatchObject(finalSettings) + expect(await reloaded.getModeConfigId("code")).toBe(profile.id) + expect(provider.contextProxy.getProviderSettings()).toMatchObject(finalSettings) + } finally { + vi.useRealTimers() + } + }) + + test("keeps a later update queued until persisted rollback completes", async () => { + const manager = provider.providerSettingsManager + const save = manager.saveConfig.bind(manager) + let releaseRollback!: () => void + vi.spyOn(manager, "saveConfig") + .mockImplementationOnce(save) + .mockImplementationOnce(async (name, settings, expected) => { + await new Promise((resolve) => { + releaseRollback = resolve + }) + return save(name, settings, expected) + }) + vi.spyOn(provider, "postStateToWebview").mockRejectedValueOnce(new Error("broadcast failed")) + const first = provider.upsertProviderProfile("existing", newSettings) + await vi.waitFor(() => expect(manager.saveConfig).toHaveBeenCalledTimes(2)) + const finalSettings = { ...oldSettings, openRouterModelId: "final-model" } + const hasConfig = vi.spyOn(manager, "hasConfig") + const second = provider.upsertProviderProfile("existing", finalSettings) + await new Promise((resolve) => setTimeout(resolve, 0)) + expect(hasConfig).not.toHaveBeenCalled() + expect(manager.saveConfig).toHaveBeenCalledTimes(2) + releaseRollback() + expect(await first).toBeUndefined() + expect(await second).toBeTruthy() + expect(await manager.getProfile({ name: "existing" })).toMatchObject(finalSettings) + expect(provider.contextProxy.getProviderSettings()).toMatchObject(finalSettings) + }) + + test("snapshots queued updates after the prior mutation and finishes rollback before the next", async () => { + let release!: () => void + const broadcast = vi + .spyOn(provider, "postStateToWebview") + .mockImplementationOnce( + () => + new Promise((resolve) => { + release = resolve + }), + ) + .mockRejectedValueOnce(new Error("second broadcast failed")) + const first = provider.upsertProviderProfile("existing", newSettings) + await vi.waitFor(() => expect(broadcast).toHaveBeenCalledTimes(1)) + const second = provider.upsertProviderProfile("existing", { + ...newSettings, + openRouterModelId: "failed-model", + }) + release() + expect(await first).toBeTruthy() + expect(await second).toBeUndefined() + expect(await provider.providerSettingsManager.getProfile({ name: "existing" })).toMatchObject(newSettings) + expect(provider.contextProxy.getProviderSettings()).toMatchObject(newSettings) + expect(await provider.providerSettingsManager.getModeConfigId("code")).toBe( + (await provider.providerSettingsManager.getProfile({ name: "existing" })).id, + ) + }) }) describe("activateProviderProfile", () => { @@ -490,7 +779,7 @@ describe("ClineProvider - API Handler Rebuild Guard", () => { expect(setValueSpy).toHaveBeenCalledWith("currentApiConfigName", "second-profile") }) - test("timed-out mutations abort before writing state and advance the queue", async () => { + test("timed-out mutations abort before writing state and release the queue after settling", async () => { vi.useFakeTimers() const logSpy = vi.spyOn(provider, "log") const setValueSpy = vi.spyOn(provider.contextProxy, "setValue") @@ -522,9 +811,9 @@ describe("ClineProvider - API Handler Rebuild Guard", () => { await vi.advanceTimersByTimeAsync(ClineProvider.PENDING_OPERATION_TIMEOUT_MS) await firstResult - // Queue advanced immediately on timeout — second enqueues now. + // The caller timed out, but the underlying activation still owns the queue. const second = provider.activateProviderProfile({ name: "second-profile" }) - // activateProfile not yet called for second (it runs in the next microtask). + await vi.advanceTimersByTimeAsync(0) expect(provider["providerSettingsManager"].activateProfile).toHaveBeenCalledTimes(1) // Resolve the first activation's inner promise so its in-flight mock can return. @@ -541,6 +830,44 @@ describe("ClineProvider - API Handler Rebuild Guard", () => { } }) + test("bounds queue waits and skips expired mutations when the blocked write eventually settles", async () => { + vi.useFakeTimers() + try { + let release!: () => void + const first = provider["enqueueProviderProfileMutation"]( + () => + new Promise((resolve) => { + release = resolve + }), + ) + const firstResult = expect(first).rejects.toThrow("Provider profile mutation timed out") + await vi.advanceTimersByTimeAsync(ClineProvider.PENDING_OPERATION_TIMEOUT_MS) + await firstResult + + const expiredWrite = vi.fn(async () => {}) + const second = provider["enqueueProviderProfileMutation"](expiredWrite) + const onRejected = vi.fn() + void second.catch(onRejected) + await vi.advanceTimersByTimeAsync(ClineProvider.PENDING_OPERATION_TIMEOUT_MS * 2) + expect(onRejected).toHaveBeenCalledWith( + expect.objectContaining({ message: "Provider profile mutation timed out" }), + ) + expect(expiredWrite).not.toHaveBeenCalled() + + const nextWrite = vi.fn(async () => {}) + const third = provider["enqueueProviderProfileMutation"](nextWrite) + await vi.advanceTimersByTimeAsync(0) + expect(nextWrite).not.toHaveBeenCalled() + release() + await third + expect(expiredWrite).not.toHaveBeenCalled() + expect(nextWrite).toHaveBeenCalledTimes(1) + expect(vi.getTimerCount()).toBe(0) + } finally { + vi.useRealTimers() + } + }) + test("mode switch preserves its default task when queued behind a profile mutation", async () => { let releaseProfileActivation!: () => void const profileActivation = new Promise((resolve) => { diff --git a/src/core/webview/__tests__/ClineProvider.spec.ts b/src/core/webview/__tests__/ClineProvider.spec.ts index 7c85a6b372..7b6082a6a4 100644 --- a/src/core/webview/__tests__/ClineProvider.spec.ts +++ b/src/core/webview/__tests__/ClineProvider.spec.ts @@ -2630,8 +2630,13 @@ describe("ClineProvider", () => { .mockResolvedValue([ { name: "test-config", id: "test-id", apiProvider: providerIdentifiers.anthropic }, ]), + hasConfig: vi.fn().mockResolvedValue(false), + getProfile: vi + .fn() + .mockResolvedValue({ id: "test-id", apiProvider: providerIdentifiers.anthropic, apiKey: "test-key" }), saveConfig: vi.fn().mockResolvedValue("test-id"), setModeConfig: vi.fn(), + getModeConfigId: vi.fn().mockResolvedValue(undefined), } as any // Update API configuration @@ -3366,7 +3371,16 @@ describe("ClineProvider", () => { const messageHandler = (mockWebviewView.webview.onDidReceiveMessage as any).mock.calls[0][0] ;(provider as any).providerSettingsManager = { + hasConfig: vi.fn().mockResolvedValue(false), + getProfile: vi.fn().mockResolvedValue({ + id: "test-id", + apiProvider: providerIdentifiers.anthropic, + apiKey: "test-key", + }), + deleteConfig: vi.fn(), setModeConfig: vi.fn().mockRejectedValue(new Error("Failed to update mode config")), + saveConfig: vi.fn().mockResolvedValue("test-id"), + getModeConfigId: vi.fn().mockResolvedValue(undefined), listConfig: vi .fn() .mockResolvedValue([ @@ -3399,8 +3413,16 @@ describe("ClineProvider", () => { const messageHandler = (mockWebviewView.webview.onDidReceiveMessage as any).mock.calls[0][0] ;(provider as any).providerSettingsManager = { + hasConfig: vi.fn().mockResolvedValue(false), + getProfile: vi.fn().mockResolvedValue({ + id: "test-id", + apiProvider: providerIdentifiers.anthropic, + apiKey: "test-key", + }), + deleteConfig: vi.fn(), setModeConfig: vi.fn(), - saveConfig: vi.fn().mockResolvedValue(undefined), + saveConfig: vi.fn().mockResolvedValue("test-id"), + getModeConfigId: vi.fn().mockResolvedValue(undefined), listConfig: vi .fn() .mockResolvedValue([ @@ -3421,7 +3443,13 @@ describe("ClineProvider", () => { }) // Verify config was saved - expect(provider.providerSettingsManager.saveConfig).toHaveBeenCalledWith("test-config", testApiConfig) + expect(provider.providerSettingsManager.saveConfig).toHaveBeenCalledWith( + "test-config", + testApiConfig, + undefined, + undefined, + undefined, + ) // Verify state updates expect(mockContext.globalState.update).toHaveBeenCalledWith("listApiConfigMeta", [ @@ -3444,8 +3472,16 @@ describe("ClineProvider", () => { throw new Error("API handler error") }) ;(provider as any).providerSettingsManager = { + hasConfig: vi.fn().mockResolvedValue(false), + getProfile: vi.fn().mockResolvedValue({ + id: "test-id", + apiProvider: providerIdentifiers.anthropic, + apiKey: "test-key", + }), + deleteConfig: vi.fn(), setModeConfig: vi.fn(), - saveConfig: vi.fn().mockResolvedValue(undefined), + saveConfig: vi.fn().mockResolvedValue("test-id"), + getModeConfigId: vi.fn().mockResolvedValue(undefined), listConfig: vi .fn() .mockResolvedValue([ @@ -3475,11 +3511,14 @@ describe("ClineProvider", () => { ) expect(vscode.window.showErrorMessage).toHaveBeenCalledWith("errors.create_api_config") - // Verify state was still updated + // The initial metadata write is followed by restoration of the snapshot. + expect(provider.providerSettingsManager.deleteConfig).toHaveBeenCalledWith( + "test-config", + expect.objectContaining({ ...testApiConfig, id: "test-id" }), + ) expect(mockContext.globalState.update).toHaveBeenCalledWith("listApiConfigMeta", [ { name: "test-config", id: "test-id", apiProvider: providerIdentifiers.anthropic }, ]) - expect(mockContext.globalState.update).toHaveBeenCalledWith("currentApiConfigName", "test-config") }) test("handles successful saveApiConfiguration", async () => { @@ -3487,8 +3526,14 @@ describe("ClineProvider", () => { const messageHandler = (mockWebviewView.webview.onDidReceiveMessage as any).mock.calls[0][0] ;(provider as any).providerSettingsManager = { + hasConfig: vi.fn().mockResolvedValue(false), + getProfile: vi.fn().mockResolvedValue({ + id: "test-id", + apiProvider: providerIdentifiers.anthropic, + apiKey: "test-key", + }), setModeConfig: vi.fn(), - saveConfig: vi.fn().mockResolvedValue(undefined), + saveConfig: vi.fn().mockResolvedValue("test-id"), listConfig: vi .fn() .mockResolvedValue([ @@ -3509,7 +3554,13 @@ describe("ClineProvider", () => { }) // Verify config was saved - expect(provider.providerSettingsManager.saveConfig).toHaveBeenCalledWith("test-config", testApiConfig) + expect(provider.providerSettingsManager.saveConfig).toHaveBeenCalledWith( + "test-config", + testApiConfig, + undefined, + undefined, + undefined, + ) // Verify state updates expect(mockContext.globalState.update).toHaveBeenCalledWith("listApiConfigMeta", [ diff --git a/src/core/webview/__tests__/ModelRequestRegistry.spec.ts b/src/core/webview/__tests__/ModelRequestRegistry.spec.ts new file mode 100644 index 0000000000..c820831210 --- /dev/null +++ b/src/core/webview/__tests__/ModelRequestRegistry.spec.ts @@ -0,0 +1,76 @@ +import { ModelRequestRegistry } from "../ModelRequestRegistry" + +it("aborts all pending requests on disposal and does not start new work", async () => { + const registry = new ModelRequestRegistry() + const signals: AbortSignal[] = [] + const fetch = vi.fn( + (signal: AbortSignal) => + new Promise((resolve) => { + signals.push(signal) + signal.addEventListener("abort", () => resolve(), { once: true }) + }), + ) + const first = registry.run("first", fetch) + const second = registry.run("second", fetch) + registry.dispose() + await Promise.all([first, second, registry.run("third", fetch)]) + expect(signals.every((signal) => signal.aborted)).toBe(true) + expect(fetch).toHaveBeenCalledTimes(2) +}) + +it("aborts a replaced request and keeps its replacement registered after late completion", async () => { + const registry = new ModelRequestRegistry() + let finishFirst!: () => void + let firstSignal!: AbortSignal + let secondSignal!: AbortSignal + const first = registry.run("same-id", (signal) => { + firstSignal = signal + return new Promise((resolve) => { + finishFirst = resolve + }) + }) + const second = registry.run("same-id", (signal) => { + expect(firstSignal.aborted).toBe(true) + secondSignal = signal + return new Promise((resolve) => { + signal.addEventListener("abort", () => resolve(), { once: true }) + }) + }) + finishFirst() + await first + expect(secondSignal.aborted).toBe(false) + registry.cancel("same-id") + expect(secondSignal.aborted).toBe(true) + await second +}) + +it("swallows abort-triggered rejections on replacement and explicit cancellation", async () => { + const registry = new ModelRequestRegistry() + const signals: AbortSignal[] = [] + const fetch = (signal: AbortSignal) => + new Promise((_resolve, reject) => { + signals.push(signal) + signal.addEventListener("abort", () => reject(signal.reason), { once: true }) + }) + const first = registry.run("same-id", fetch) + const firstResolved = expect(first).resolves.toBeUndefined() + const second = registry.run("same-id", fetch) + const secondResolved = expect(second).resolves.toBeUndefined() + + await firstResolved + expect(signals[0].aborted).toBe(true) + expect(signals[1].aborted).toBe(false) + registry.cancel("same-id") + expect(signals[1].aborted).toBe(true) + await secondResolved +}) + +it("propagates errors from requests that were not aborted", async () => { + const registry = new ModelRequestRegistry() + const failure = new Error("model discovery failed") + await expect( + registry.run("failed", async () => { + throw failure + }), + ).rejects.toBe(failure) +}) diff --git a/src/core/webview/__tests__/webviewMessageHandler.spec.ts b/src/core/webview/__tests__/webviewMessageHandler.spec.ts index dcd70f92f1..ed8b43ac4e 100644 --- a/src/core/webview/__tests__/webviewMessageHandler.spec.ts +++ b/src/core/webview/__tests__/webviewMessageHandler.spec.ts @@ -1,3 +1,14 @@ +import { ModelRequestRegistry } from "../ModelRequestRegistry" +import { getOpenAiModels } from "../../../api/providers/openai" +import { getVsCodeLmModels } from "../../../api/providers/vscode-lm" +vi.mock("../../../api/providers/openai", async (importOriginal) => ({ + ...(await importOriginal()), + getOpenAiModels: vi.fn(), +})) +vi.mock("../../../api/providers/vscode-lm", async (importOriginal) => ({ + ...(await importOriginal()), + getVsCodeLmModels: vi.fn(), +})) // npx vitest core/webview/__tests__/webviewMessageHandler.spec.ts import type { Mock } from "vitest" @@ -97,6 +108,7 @@ const mockFetchOpenAiCodexRateLimitInfo = vi.mocked(fetchOpenAiCodexRateLimitInf // Mock ClineProvider const mockClineProvider = { + modelRequests: new ModelRequestRegistry(), getState: vi.fn(), postMessageToWebview: vi.fn(), customModesManager: { @@ -169,6 +181,192 @@ describe("webviewMessageHandler - theme fixture probes", () => { }) }) +describe("webviewMessageHandler - organization allowlist enforcement", () => { + // Provider narrowed to a single model; anything else is disallowed. + const restrictiveAllowList = { + allowAll: false, + providers: { + [providerIdentifiers.anthropic]: { allowAll: false, models: ["claude-sonnet-4-20250514"] }, + }, + } + + const setUpsertSpy = (spy: ReturnType) => { + ;(mockClineProvider as unknown as { upsertProviderProfile: ReturnType }).upsertProviderProfile = + spy + } + + beforeEach(() => { + vi.clearAllMocks() + mockClineProvider.getState = vi.fn().mockResolvedValue({ organizationAllowList: restrictiveAllowList }) + }) + + it("rejects a disallowed model on upsertApiConfiguration without persisting", async () => { + const upsertProviderProfile = vi.fn().mockResolvedValue("id") + setUpsertSpy(upsertProviderProfile) + + await webviewMessageHandler(mockClineProvider, { + type: "upsertApiConfiguration", + text: "profile", + apiConfiguration: { apiProvider: providerIdentifiers.anthropic, apiModelId: "claude-opus-4-20250514" }, + }) + + expect(upsertProviderProfile).not.toHaveBeenCalled() + expect(vscode.window.showErrorMessage).toHaveBeenCalled() + }) + + it("rejects a disallowed save without persisting or updating state", async () => { + const saveConfig = vi.fn() + const listConfig = vi.fn() + Object.assign(mockClineProvider, { providerSettingsManager: { saveConfig, listConfig } }) + await webviewMessageHandler(mockClineProvider, { + type: "saveApiConfiguration", + text: "profile", + apiConfiguration: { apiProvider: providerIdentifiers.anthropic, apiModelId: "claude-opus-4-20250514" }, + }) + expect(saveConfig).not.toHaveBeenCalled() + expect(listConfig).not.toHaveBeenCalled() + expect(mockClineProvider.contextProxy.setValue).not.toHaveBeenCalled() + expect(mockClineProvider.postStateToWebview).not.toHaveBeenCalled() + expect(vscode.window.showErrorMessage).toHaveBeenCalledWith(t("common:errors.violated_organization_allowlist")) + }) + + it("rejects a disallowed rename before saving, deleting, or activating a profile", async () => { + const saveConfig = vi.fn() + const deleteConfig = vi.fn() + const activateProviderProfile = vi.fn() + Object.assign(mockClineProvider, { + providerSettingsManager: { + getProfile: vi.fn().mockResolvedValue({ id: "profile-id" }), + saveConfig, + deleteConfig, + }, + activateProviderProfile, + }) + + await webviewMessageHandler(mockClineProvider, { + type: "renameApiConfiguration", + values: { oldName: "old", newName: "new" }, + apiConfiguration: { apiProvider: providerIdentifiers.anthropic, apiModelId: "claude-opus-4-20250514" }, + }) + + expect(saveConfig).not.toHaveBeenCalled() + expect(deleteConfig).not.toHaveBeenCalled() + expect(activateProviderProfile).not.toHaveBeenCalled() + expect(vscode.window.showErrorMessage).toHaveBeenCalledWith(t("common:errors.violated_organization_allowlist")) + }) + + it.each([restrictiveAllowList, undefined])( + "renames an allowed profile and preserves its ID (%j)", + async (organizationAllowList) => { + mockClineProvider.getState = vi.fn().mockResolvedValue({ organizationAllowList }) + const saveConfig = vi.fn() + const deleteConfig = vi.fn() + const activateProviderProfile = vi.fn() + Object.assign(mockClineProvider, { + providerSettingsManager: { + getProfile: vi.fn().mockResolvedValue({ id: "profile-id" }), + saveConfig, + deleteConfig, + }, + activateProviderProfile, + }) + const apiConfiguration = { + apiProvider: providerIdentifiers.anthropic, + apiModelId: "claude-sonnet-4-20250514", + } + + await webviewMessageHandler(mockClineProvider, { + type: "renameApiConfiguration", + values: { oldName: "old", newName: "new" }, + apiConfiguration, + }) + + expect(mockClineProvider.upsertProviderProfile).toHaveBeenCalledWith( + "new", + apiConfiguration, + true, + undefined, + "old", + ) + expect(saveConfig).not.toHaveBeenCalled() + expect(deleteConfig).not.toHaveBeenCalled() + expect(vscode.window.showErrorMessage).not.toHaveBeenCalled() + }, + ) + + it.each([{ text: "profile" }, { apiConfiguration: {} }])( + "rejects incomplete save requests with an acknowledgement", + async (fields) => { + await webviewMessageHandler(mockClineProvider, { + type: "upsertApiConfiguration", + requestId: "invalid-save", + ...fields, + }) + expect(mockClineProvider.postMessageToWebview).toHaveBeenCalledWith({ + type: "apiConfigurationSaved", + requestId: "invalid-save", + success: false, + }) + }, + ) + + it.each(["success", "conflict", "error"])("acknowledges an API configuration save after %s", async (outcome) => { + const save = + outcome === "error" + ? vi.fn().mockRejectedValue(new Error("save failed")) + : vi.fn().mockResolvedValue(outcome === "success" ? "id" : undefined) + setUpsertSpy(save) + await webviewMessageHandler(mockClineProvider, { + type: "upsertApiConfiguration", + text: "profile", + requestId: "save-request", + apiConfiguration: { apiProvider: providerIdentifiers.anthropic, apiModelId: "claude-sonnet-4-20250514" }, + }) + expect(mockClineProvider.postMessageToWebview).toHaveBeenCalledWith({ + type: "apiConfigurationSaved", + requestId: "save-request", + success: outcome === "success", + }) + }) + + it.each(["success", "conflict", "invalid", "unchanged"])("acknowledges profile rename %s", async (outcome) => { + setUpsertSpy(vi.fn().mockResolvedValue(outcome === "success" ? "id" : undefined)) + await webviewMessageHandler(mockClineProvider, { + type: "renameApiConfiguration", + requestId: "rename-request", + values: { oldName: "old", newName: outcome === "unchanged" ? "old" : "new" }, + apiConfiguration: + outcome === "invalid" + ? undefined + : { apiProvider: providerIdentifiers.anthropic, apiModelId: "claude-sonnet-4-20250514" }, + }) + expect(mockClineProvider.postMessageToWebview).toHaveBeenCalledWith({ + type: "apiConfigurationSaved", + requestId: "rename-request", + success: outcome === "success" || outcome === "unchanged", + }) + }) + + it("allows an allow-listed model on upsertApiConfiguration", async () => { + const upsertProviderProfile = vi.fn().mockResolvedValue("id") + setUpsertSpy(upsertProviderProfile) + + await webviewMessageHandler(mockClineProvider, { + type: "upsertApiConfiguration", + text: "profile", + apiConfiguration: { apiProvider: providerIdentifiers.anthropic, apiModelId: "claude-sonnet-4-20250514" }, + }) + + expect(upsertProviderProfile).toHaveBeenCalledWith( + "profile", + expect.objectContaining({ apiModelId: "claude-sonnet-4-20250514" }), + true, + undefined, + ) + expect(vscode.window.showErrorMessage).not.toHaveBeenCalled() + }) +}) + import { t } from "../../../i18n" vi.mock("vscode", () => { @@ -273,6 +471,92 @@ describe("webviewMessageHandler - requestLmStudioModels", () => { }) }) + it.each(["cancel", "dispose"])( + "detaches a preview on %s while retaining discovery for the next waiter", + async (action) => { + let finish!: (models: ModelRecord) => void + mockGetLMStudioModels.mockImplementationOnce( + () => + new Promise((resolve) => { + finish = resolve + }), + ) + const values = { baseUrl: "http://localhost:4567" } + const first = webviewMessageHandler(mockClineProvider, { + type: "requestLmStudioModels", + requestId: "first-preview", + values, + }) + await vi.waitFor(() => expect(mockGetLMStudioModels).toHaveBeenCalledTimes(1)) + if (action === "cancel") { + await webviewMessageHandler(mockClineProvider, { + type: "cancelModelRequest", + requestId: "first-preview", + }) + } else { + mockClineProvider.modelRequests.dispose() + Object.defineProperty(mockClineProvider, "modelRequests", { + value: new ModelRequestRegistry(), + configurable: true, + }) + } + await first + const second = webviewMessageHandler(mockClineProvider, { + type: "requestLmStudioModels", + requestId: "second-preview", + values, + }) + await Promise.resolve() + expect(mockGetLMStudioModels).toHaveBeenCalledTimes(1) + finish({}) + await second + expect(mockClineProvider.postMessageToWebview).not.toHaveBeenCalledWith( + expect.objectContaining({ requestId: "first-preview" }), + ) + expect(mockClineProvider.postMessageToWebview).toHaveBeenCalledWith( + expect.objectContaining({ requestId: "second-preview" }), + ) + }, + ) + + it("bounds preview waiters without starting overlapping SDK discovery after timeout", async () => { + const timeout = new AbortController() + const timeoutSignal = vi.spyOn(AbortSignal, "timeout").mockReturnValue(timeout.signal) + let finish!: (models: ModelRecord) => void + mockGetLMStudioModels.mockImplementationOnce( + () => + new Promise((resolve) => { + finish = resolve + }), + ) + try { + const values = { baseUrl: "http://localhost:5432" } + const first = webviewMessageHandler(mockClineProvider, { + type: "requestLmStudioModels", + requestId: "timeout-preview", + values, + }) + await vi.waitFor(() => expect(mockGetLMStudioModels).toHaveBeenCalledTimes(1)) + timeout.abort() + await first + expect(mockClineProvider.postMessageToWebview).toHaveBeenCalledWith( + expect.objectContaining({ requestId: "timeout-preview", lmStudioModels: {} }), + ) + timeoutSignal.mockReturnValue(new AbortController().signal) + const second = webviewMessageHandler(mockClineProvider, { + type: "requestLmStudioModels", + requestId: "retry-preview", + values, + }) + await Promise.resolve() + expect(mockGetLMStudioModels).toHaveBeenCalledTimes(1) + finish({}) + await second + } finally { + timeoutSignal.mockRestore() + } + }) + it("successfully fetches models from LMStudio", async () => { const mockModels: ModelRecord = { "model-1": { @@ -293,28 +577,69 @@ describe("webviewMessageHandler - requestLmStudioModels", () => { await webviewMessageHandler(mockClineProvider, { type: "requestLmStudioModels", + requestId: "models-request", }) expect(mockGetModels).toHaveBeenCalledWith({ provider: providerIdentifiers.lmstudio, + signal: expect.any(AbortSignal), baseUrl: "http://localhost:1234", }) expect(mockClineProvider.postMessageToWebview).toHaveBeenCalledWith({ type: "lmStudioModels", + requestId: "models-request", lmStudioModels: mockModels, }) }) + it.each([ + { source: "persisted settings", values: undefined }, + { source: "request payload", values: { baseUrl: "http://127.0.0.1:4321" } }, + ])("completes empty and failed fetches using $source with the request ID", async ({ values }) => { + const fetchModels = values ? mockGetLMStudioModels : mockGetModels + fetchModels.mockResolvedValueOnce({}).mockRejectedValueOnce(new Error("LM Studio unavailable")) + + for (const requestId of ["empty-models", "failed-models"]) { + vi.mocked(mockClineProvider.postMessageToWebview).mockClear() + await webviewMessageHandler(mockClineProvider, { type: "requestLmStudioModels", requestId, values }) + + expect(mockClineProvider.postMessageToWebview).toHaveBeenCalledExactlyOnceWith({ + type: "lmStudioModels", + lmStudioModels: {}, + requestId, + }) + } + }) + + it("does not send a completion after a model request is cancelled", async () => { + mockGetLMStudioModels.mockImplementation(async () => { + await webviewMessageHandler(mockClineProvider, { + type: "cancelModelRequest", + requestId: "cancelled-models", + }) + throw new Error("Request cancelled") + }) + + await webviewMessageHandler(mockClineProvider, { + type: "requestLmStudioModels", + requestId: "cancelled-models", + values: { baseUrl: "http://127.0.0.1:4321" }, + }) + + expect(mockClineProvider.postMessageToWebview).not.toHaveBeenCalled() + }) + it("prefers the request payload base URL over persisted settings", async () => { mockGetLMStudioModels.mockResolvedValue({}) await webviewMessageHandler(mockClineProvider, { type: "requestLmStudioModels", + requestId: "models-request", values: { baseUrl: "http://127.0.0.1:4321" }, }) - expect(mockGetLMStudioModels).toHaveBeenCalledWith("http://127.0.0.1:4321") + expect(mockGetLMStudioModels).toHaveBeenCalledWith("http://127.0.0.1:4321", { signal: expect.any(AbortSignal) }) expect(mockGetModels).not.toHaveBeenCalled() }) @@ -323,10 +648,11 @@ describe("webviewMessageHandler - requestLmStudioModels", () => { await webviewMessageHandler(mockClineProvider, { type: "requestLmStudioModels", + requestId: "models-request", values: { baseUrl: "" }, }) - expect(mockGetLMStudioModels).toHaveBeenCalledWith("") + expect(mockGetLMStudioModels).toHaveBeenCalledWith("http://localhost:1234", { signal: expect.any(AbortSignal) }) expect(mockGetModels).not.toHaveBeenCalled() }) }) @@ -396,15 +722,18 @@ describe("webviewMessageHandler - requestOllamaModels", () => { await webviewMessageHandler(mockClineProvider, { type: "requestOllamaModels", + requestId: "models-request", }) expect(mockGetModels).toHaveBeenCalledWith({ provider: providerIdentifiers.ollama, + signal: expect.any(AbortSignal), baseUrl: "http://localhost:1234", }) expect(mockClineProvider.postMessageToWebview).toHaveBeenCalledWith({ type: "ollamaModels", + requestId: "models-request", ollamaModels: mockModels, }) }) @@ -414,10 +743,12 @@ describe("webviewMessageHandler - requestOllamaModels", () => { await webviewMessageHandler(mockClineProvider, { type: "requestOllamaModels", + requestId: "models-request", }) expect(mockClineProvider.postMessageToWebview).toHaveBeenCalledWith({ type: "ollamaModels", + requestId: "models-request", ollamaModels: {}, }) }) @@ -427,10 +758,12 @@ describe("webviewMessageHandler - requestOllamaModels", () => { await webviewMessageHandler(mockClineProvider, { type: "requestOllamaModels", + requestId: "models-request", }) expect(mockClineProvider.postMessageToWebview).toHaveBeenCalledWith({ type: "ollamaModels", + requestId: "models-request", ollamaModels: {}, error: "Connection refused", }) @@ -445,6 +778,7 @@ describe("webviewMessageHandler - requestOllamaModels", () => { await webviewMessageHandler(mockClineProvider, { type: "requestOllamaModels", + requestId: "models-request", values: { baseUrl: "https://ollama.example.com" }, }) @@ -454,6 +788,7 @@ describe("webviewMessageHandler - requestOllamaModels", () => { ) expect(mockClineProvider.postMessageToWebview).toHaveBeenCalledWith({ type: "ollamaModels", + requestId: "models-request", ollamaModels: {}, error: "Cache write failed", }) @@ -473,6 +808,7 @@ describe("webviewMessageHandler - requestOllamaModels", () => { await webviewMessageHandler(mockClineProvider, { type: "requestOllamaModels", + requestId: "models-request", values: { baseUrl: "https://ollama.example.com", apiKey: "secret-key", @@ -483,6 +819,7 @@ describe("webviewMessageHandler - requestOllamaModels", () => { expect(mockFlushModels).toHaveBeenCalledWith( { provider: providerIdentifiers.ollama, + signal: expect.any(AbortSignal), baseUrl: "https://ollama.example.com", apiKey: "secret-key", }, @@ -490,12 +827,14 @@ describe("webviewMessageHandler - requestOllamaModels", () => { ) expect(mockGetModels).toHaveBeenCalledWith({ provider: providerIdentifiers.ollama, + signal: expect.any(AbortSignal), baseUrl: "https://ollama.example.com", apiKey: "secret-key", }) expect(mockClineProvider.postMessageToWebview).toHaveBeenCalledWith({ type: "ollamaModels", + requestId: "models-request", ollamaModels: mockModels, }) }) @@ -2307,3 +2646,178 @@ describe("webviewMessageHandler - telemetrySetting", () => { expect(TelemetryService.instance.updateTelemetryState).not.toHaveBeenCalled() }) }) + +describe("model request correlation and cancellation", () => { + beforeEach(() => { + vi.clearAllMocks() + mockClineProvider.getState = vi.fn().mockResolvedValue({ apiConfiguration: {} }) + mockFlushModels.mockReset().mockResolvedValue(undefined) + mockGetModels.mockReset().mockResolvedValue({}) + }) + + it("echoes the OpenAI request ID", async () => { + vi.mocked(getOpenAiModels).mockResolvedValue(["model"]) + await webviewMessageHandler(mockClineProvider, { + type: "requestOpenAiModels", + requestId: "openai-request", + values: { baseUrl: "https://example.test/v1", apiKey: "test-key" }, + }) + expect(getOpenAiModels).toHaveBeenCalledWith( + "https://example.test/v1", + "test-key", + undefined, + expect.any(AbortSignal), + ) + expect(mockClineProvider.postMessageToWebview).toHaveBeenCalledWith({ + type: "openAiModels", + requestId: "openai-request", + openAiModels: ["model"], + }) + }) + + it("echoes the VS Code LM request ID", async () => { + vi.mocked(getVsCodeLmModels).mockResolvedValue([]) + await webviewMessageHandler(mockClineProvider, { type: "requestVsCodeLmModels", requestId: "vscode-request" }) + expect(mockClineProvider.postMessageToWebview).toHaveBeenCalledWith({ + type: "vsCodeLmModels", + requestId: "vscode-request", + vsCodeLmModels: [], + }) + }) + + it.each(["requestOllamaModels", "requestLmStudioModels"] as const)( + "cancels %s during refresh before reading models", + async (type) => { + let finish!: () => void + mockFlushModels.mockImplementationOnce( + () => + new Promise((resolve) => { + finish = resolve + }), + ) + const pending = webviewMessageHandler(mockClineProvider, { type, requestId: "cancel-refresh" }) + await vi.waitFor(() => expect(mockFlushModels).toHaveBeenCalled()) + const signal = mockFlushModels.mock.calls[0][0].signal + await webviewMessageHandler(mockClineProvider, { type: "cancelModelRequest", requestId: "cancel-refresh" }) + expect(signal?.aborted).toBe(true) + finish() + await pending + expect(mockGetModels).not.toHaveBeenCalled() + expect(mockClineProvider.postMessageToWebview).not.toHaveBeenCalled() + }, + ) + + it.each(["requestOllamaModels", "requestLmStudioModels", "requestRouterModels"] as const)( + "cancels %s during model reads without publishing results", + async (type) => { + let finish!: (models: ModelRecord) => void + mockGetModels.mockImplementationOnce( + () => + new Promise((resolve) => { + finish = resolve + }), + ) + const pending = webviewMessageHandler(mockClineProvider, { + type, + requestId: "cancel-read", + values: { provider: providerIdentifiers.openrouter }, + }) + await vi.waitFor(() => expect(mockGetModels).toHaveBeenCalled()) + const signal = mockGetModels.mock.calls[0][0].signal + await webviewMessageHandler(mockClineProvider, { type: "cancelModelRequest", requestId: "cancel-read" }) + expect(signal?.aborted).toBe(true) + finish({}) + await pending + expect(mockClineProvider.postMessageToWebview).not.toHaveBeenCalled() + }, + ) + + it.each(["cancel", "dispose"])( + "shares pending VS Code discovery after %s and refreshes after settlement", + async (action) => { + let finish!: (models: Awaited>) => void + vi.mocked(getVsCodeLmModels).mockImplementationOnce( + () => + new Promise((resolve) => { + finish = resolve + }), + ) + const pending = webviewMessageHandler(mockClineProvider, { + type: "requestVsCodeLmModels", + requestId: "cancel-vscode", + }) + if (action === "cancel") { + await webviewMessageHandler(mockClineProvider, { + type: "cancelModelRequest", + requestId: "cancel-vscode", + }) + } else { + mockClineProvider.modelRequests.dispose() + Object.defineProperty(mockClineProvider, "modelRequests", { + value: new ModelRequestRegistry(), + configurable: true, + }) + } + await pending + expect(mockClineProvider.postMessageToWebview).not.toHaveBeenCalled() + const restarted = webviewMessageHandler(mockClineProvider, { + type: "requestVsCodeLmModels", + requestId: "restarted-vscode", + }) + const sibling = webviewMessageHandler(mockClineProvider, { + type: "requestVsCodeLmModels", + requestId: "sibling-vscode", + }) + await webviewMessageHandler(mockClineProvider, { type: "cancelModelRequest", requestId: "sibling-vscode" }) + await sibling + expect(getVsCodeLmModels).toHaveBeenCalledTimes(1) + finish([]) + await restarted + expect(mockClineProvider.postMessageToWebview).toHaveBeenCalledExactlyOnceWith({ + type: "vsCodeLmModels", + requestId: "restarted-vscode", + vsCodeLmModels: [], + }) + await webviewMessageHandler(mockClineProvider, { type: "requestVsCodeLmModels", requestId: "fresh-vscode" }) + expect(getVsCodeLmModels).toHaveBeenCalledTimes(2) + }, + ) + + it("allows VS Code discovery to retry after a shared failure", async () => { + vi.mocked(getVsCodeLmModels).mockRejectedValueOnce(new Error("discovery failed")) + const first = webviewMessageHandler(mockClineProvider, { + type: "requestVsCodeLmModels", + requestId: "failed-first", + }) + const second = webviewMessageHandler(mockClineProvider, { + type: "requestVsCodeLmModels", + requestId: "failed-second", + }) + await expect(first).rejects.toThrow("discovery failed") + await expect(second).rejects.toThrow("discovery failed") + expect(getVsCodeLmModels).toHaveBeenCalledTimes(1) + expect(mockClineProvider.postMessageToWebview).not.toHaveBeenCalled() + await webviewMessageHandler(mockClineProvider, { type: "requestVsCodeLmModels", requestId: "retry" }) + expect(getVsCodeLmModels).toHaveBeenCalledTimes(2) + }) + + it("aborts OpenAI HTTP work and suppresses its late response", async () => { + let finish!: (models: string[]) => void + vi.mocked(getOpenAiModels).mockImplementationOnce( + () => + new Promise((resolve) => { + finish = resolve + }), + ) + const pending = webviewMessageHandler(mockClineProvider, { + type: "requestOpenAiModels", + requestId: "cancel-openai", + values: { baseUrl: "https://example.test/v1", apiKey: "test-key" }, + }) + await webviewMessageHandler(mockClineProvider, { type: "cancelModelRequest", requestId: "cancel-openai" }) + expect(vi.mocked(getOpenAiModels).mock.calls[0][3]?.aborted).toBe(true) + finish(["stale"]) + await pending + expect(mockClineProvider.postMessageToWebview).not.toHaveBeenCalled() + }) +}) diff --git a/src/core/webview/webviewMessageHandler.ts b/src/core/webview/webviewMessageHandler.ts index 193540455b..20a49ecaa5 100644 --- a/src/core/webview/webviewMessageHandler.ts +++ b/src/core/webview/webviewMessageHandler.ts @@ -1,3 +1,4 @@ +import { mergeAbortSignalAndTimeout, resolveModelWithAbort } from "../../api/providers/utils/abort-signal" import { safeWriteJson } from "../../utils/safeWriteJson" import * as path from "path" import * as os from "os" @@ -30,6 +31,7 @@ import { RouterModelsMessageType, VsCodeLmModelsMessageType, isTelemetryOptedIn, + ORGANIZATION_ALLOW_ALL, } from "@roo-code/types" import { customToolRegistry } from "@roo-code/core" import { CloudService } from "@roo-code/cloud" @@ -43,6 +45,8 @@ import { ClineProvider } from "./ClineProvider" import { findOriginalContent } from "./stripOriginalContent" import { handleCheckpointRestoreOperation } from "./checkpointRestoreHandler" import { generateErrorDiagnostics } from "./diagnosticsHandler" +import { ProfileValidator } from "../../shared/ProfileValidator" +import { OrganizationAllowListViolationError } from "../../utils/errors" import { handleRequestSkills, handleCreateSkill, @@ -119,11 +123,71 @@ import { handleCheckoutBranch, } from "./worktree" +// Native discovery has no cancellation token. Keep it shared across request IDs +// and provider instances until it settles; cancellation only detaches a waiter. +let pendingVsCodeLmModels: ReturnType | undefined +const discoverVsCodeLmModels = () => { + pendingVsCodeLmModels ??= getVsCodeLmModels().finally(() => { + pendingVsCodeLmModels = undefined + }) + return pendingVsCodeLmModels +} + +const pendingLmStudioPreviews = new Map>() +const discoverLmStudioPreview = (baseUrl: string) => { + const key = baseUrl || "http://localhost:1234" + let pending = pendingLmStudioPreviews.get(key) + if (!pending) { + // Bound the cancellable HTTP probe. SDK discovery must remain registered until + // it actually settles, even when every timed-out waiter has detached. + pending = getLMStudioModels(key, { signal: AbortSignal.timeout(15_000) }).finally(() => + pendingLmStudioPreviews.delete(key), + ) + pendingLmStudioPreviews.set(key, pending) + } + return pending +} + export const webviewMessageHandler = async ( provider: ClineProvider, message: WebviewMessage, marketplaceManager?: MarketplaceManager, ) => { + if (message.type === "cancelModelRequest") { + if (message.requestId) provider.modelRequests.cancel(message.requestId) + return + } + if ( + message.requestId && + [ + "requestOllamaModels", + "requestLmStudioModels", + "requestOpenAiModels", + "requestVsCodeLmModels", + "requestRouterModels", + ].includes(message.type) + ) { + return provider.modelRequests.run(message.requestId, (signal) => + handleWebviewMessage(provider, message, marketplaceManager, signal), + ) + } + return handleWebviewMessage(provider, message, marketplaceManager) +} + +const handleWebviewMessage = async ( + provider: ClineProvider, + message: WebviewMessage, + marketplaceManager?: MarketplaceManager, + modelRequestSignal?: AbortSignal, +) => { + const requestOptions = (options: GetModelsOptions): GetModelsOptions => { + modelRequestSignal?.throwIfAborted() + return modelRequestSignal ? { ...options, signal: modelRequestSignal } : options + } + const getRequestModels = (options: GetModelsOptions) => getModels(requestOptions(options)) + const flushRequestModels = (options: GetModelsOptions, refresh: boolean) => + flushModels(requestOptions(options), refresh) + // Utility functions provided for concise get/update of global state via contextProxy API. const getGlobalState = (key: K) => provider.contextProxy.getValue(key) const updateGlobalState = async (key: K, value: GlobalState[K]) => @@ -1139,8 +1203,9 @@ export const webviewMessageHandler = async ( const safeGetModels = async (options: GetModelsOptions): Promise => { try { - return await getModels(options) + return await getRequestModels(options) } catch (error) { + modelRequestSignal?.throwIfAborted() console.error( `Failed to fetch models in webviewMessageHandler requestRouterModels for ${options.provider}:`, error, @@ -1195,7 +1260,7 @@ export const webviewMessageHandler = async ( // If explicit credentials are provided in message.values (from Refresh Models button), // flush the cache first to ensure we fetch fresh data with the new credentials if (message?.values?.litellmApiKey || message?.values?.litellmBaseUrl) { - await flushModels( + await flushRequestModels( { provider: providerIdentifiers.litellm, apiKey: litellmApiKey, baseUrl: litellmBaseUrl }, true, ) @@ -1213,7 +1278,7 @@ export const webviewMessageHandler = async ( if (poeApiKey) { if (message?.values?.poeApiKey || message?.values?.poeBaseUrl) { - await flushModels( + await flushRequestModels( { provider: providerIdentifiers.poe, apiKey: poeApiKey, baseUrl: poeBaseUrl }, true, ) @@ -1231,7 +1296,7 @@ export const webviewMessageHandler = async ( if (deepSeekApiKey) { if (message?.values?.deepSeekApiKey || message?.values?.deepSeekBaseUrl) { - await flushModels( + await flushRequestModels( { provider: providerIdentifiers.deepseek, apiKey: deepSeekApiKey, baseUrl: deepSeekBaseUrl }, true, ) @@ -1253,7 +1318,7 @@ export const webviewMessageHandler = async ( if (moonshotApiKey) { if (message?.values?.moonshotApiKey || message?.values?.moonshotBaseUrl) { - await flushModels( + await flushRequestModels( { provider: providerIdentifiers.moonshot, apiKey: moonshotApiKey, baseUrl: moonshotBaseUrl }, true, ) @@ -1278,7 +1343,7 @@ export const webviewMessageHandler = async ( // Refresh the cache when a new key is explicitly provided (e.g. the Refresh Models button). if (message?.values?.opencodeGoApiKey) { - await flushModels({ provider: providerIdentifiers.opencodeGo, apiKey: opencodeGoApiKey }, true) + await flushRequestModels({ provider: providerIdentifiers.opencodeGo, apiKey: opencodeGoApiKey }, true) } candidates.push({ @@ -1295,7 +1360,7 @@ export const webviewMessageHandler = async ( // Refresh the cache when a new key is explicitly provided (e.g. the Refresh Models button). if (message?.values?.kenariApiKey) { - await flushModels({ provider: providerIdentifiers.kenari, apiKey: kenariApiKey }, true) + await flushRequestModels({ provider: providerIdentifiers.kenari, apiKey: kenariApiKey }, true) } candidates.push({ @@ -1308,7 +1373,7 @@ export const webviewMessageHandler = async ( // same key-scoped options for refresh and retrieval. const nanoGptApiKey = message?.values?.nanoGptApiKey ?? apiConfiguration.nanoGptApiKey if (message?.values?.nanoGptApiKey !== undefined) { - await flushModels({ provider: providerIdentifiers.nanogpt, apiKey: nanoGptApiKey }, true) + await flushRequestModels({ provider: providerIdentifiers.nanogpt, apiKey: nanoGptApiKey }, true) } candidates.push({ @@ -1340,7 +1405,7 @@ export const webviewMessageHandler = async ( // If refresh flag is set and we have a specific provider, flush its cache first if (shouldRefresh && providerFilter && modelFetchPromises.length > 0) { const targetCandidate = modelFetchPromises[0] - await flushModels(targetCandidate.options, true) + await flushRequestModels(targetCandidate.options, true) } const results = await Promise.allSettled( @@ -1350,6 +1415,7 @@ export const webviewMessageHandler = async ( }), ) + modelRequestSignal?.throwIfAborted() results.forEach((result, index) => { const routerName = modelFetchPromises[index].key @@ -1373,8 +1439,10 @@ export const webviewMessageHandler = async ( } }) + modelRequestSignal?.throwIfAborted() await provider.postMessageToWebview({ type: RouterModelsMessageType.routerModels, + requestId: message.requestId, routerModels, values: providerFilter ? { provider: requestedProvider } : undefined, }) @@ -1399,67 +1467,84 @@ export const webviewMessageHandler = async ( // Refresh the cache before reading the models. Keep this error // separate from the read below so diagnostics identify which // cache operation failed. - await flushModels(ollamaOptions, true) + await flushRequestModels(ollamaOptions, true) } catch (error) { + modelRequestSignal?.throwIfAborted() const errorMsg = error instanceof Error ? error.message : String(error) provider.log(`[requestOllamaModels] Failed to refresh model cache for ${logBaseUrl}: ${errorMsg}`) + modelRequestSignal?.throwIfAborted() await provider.postMessageToWebview({ type: OllamaModelsMessageType.ollamaModels, ollamaModels: {}, error: errorMsg, + requestId: message.requestId, }) break } try { - const ollamaModels = await getModels(ollamaOptions) + const ollamaModels = await getRequestModels(ollamaOptions) // Always post a response so the webview refresh status can // transition out of "loading" — even when no models are found. - await provider.postMessageToWebview({ type: OllamaModelsMessageType.ollamaModels, ollamaModels }) + modelRequestSignal?.throwIfAborted() + await provider.postMessageToWebview({ + type: OllamaModelsMessageType.ollamaModels, + ollamaModels, + requestId: message.requestId, + }) } catch (error) { + modelRequestSignal?.throwIfAborted() const errorMsg = error instanceof Error ? error.message : String(error) provider.log(`[requestOllamaModels] Failed to read models for ${logBaseUrl}: ${errorMsg}`) + modelRequestSignal?.throwIfAborted() await provider.postMessageToWebview({ type: OllamaModelsMessageType.ollamaModels, ollamaModels: {}, error: errorMsg, + requestId: message.requestId, }) } break } case LmStudioModelsMessageType.requestLmStudioModels: { // Specific handler for LM Studio models only. - const { apiConfiguration: lmStudioApiConfig } = await provider.getState() + let lmStudioModels: ModelRecord = {} try { + const { apiConfiguration: lmStudioApiConfig } = await provider.getState() const requestedBaseUrl = message.values?.baseUrl const hasPreviewBaseUrl = typeof requestedBaseUrl === "string" - let lmStudioModels: ModelRecord if (hasPreviewBaseUrl) { - lmStudioModels = await getLMStudioModels(requestedBaseUrl) + modelRequestSignal?.throwIfAborted() + lmStudioModels = await resolveModelWithAbort( + () => discoverLmStudioPreview(requestedBaseUrl), + mergeAbortSignalAndTimeout(modelRequestSignal, 15_000), + "LM Studio", + ) } else { const lmStudioOptions = { provider: providerIdentifiers.lmstudio, baseUrl: lmStudioApiConfig.lmStudioBaseUrl, } // Flush cache and refresh to ensure fresh models. - await flushModels(lmStudioOptions, true) - lmStudioModels = await getModels(lmStudioOptions) - } - - if (Object.keys(lmStudioModels).length > 0) { - await provider.postMessageToWebview({ - type: LmStudioModelsMessageType.lmStudioModels, - lmStudioModels: lmStudioModels, - }) + await flushRequestModels(lmStudioOptions, true) + lmStudioModels = await getRequestModels(lmStudioOptions) } } catch (error) { + modelRequestSignal?.throwIfAborted() // Silently fail - user hasn't configured LM Studio yet. console.debug("LM Studio models fetch failed:", error) } + modelRequestSignal?.throwIfAborted() + await provider.postMessageToWebview({ + type: LmStudioModelsMessageType.lmStudioModels, + lmStudioModels, + requestId: message.requestId, + }) break } case "requestRooModels": { + modelRequestSignal?.throwIfAborted() await provider.postMessageToWebview({ type: RouterModelsMessageType.singleRouterModelFetchResponse, success: false, @@ -1474,17 +1559,30 @@ export const webviewMessageHandler = async ( message?.values?.baseUrl, message?.values?.apiKey, message?.values?.openAiHeaders, + modelRequestSignal, ) - await provider.postMessageToWebview({ type: OpenAiModelsMessageType.openAiModels, openAiModels }) + modelRequestSignal?.throwIfAborted() + await provider.postMessageToWebview({ + type: OpenAiModelsMessageType.openAiModels, + openAiModels, + requestId: message.requestId, + }) } break - case VsCodeLmModelsMessageType.requestVsCodeLmModels: - const vsCodeLmModels = await getVsCodeLmModels() + case VsCodeLmModelsMessageType.requestVsCodeLmModels: { + modelRequestSignal?.throwIfAborted() + const vsCodeLmModels = await resolveModelWithAbort(discoverVsCodeLmModels, modelRequestSignal, "VS Code LM") // TODO: Cache like we do for OpenRouter, etc? - await provider.postMessageToWebview({ type: VsCodeLmModelsMessageType.vsCodeLmModels, vsCodeLmModels }) + modelRequestSignal?.throwIfAborted() + await provider.postMessageToWebview({ + type: VsCodeLmModelsMessageType.vsCodeLmModels, + vsCodeLmModels, + requestId: message.requestId, + }) break + } case "openImage": await openImage(message.text!, { values: message.values }) break @@ -2291,50 +2389,145 @@ export const webviewMessageHandler = async ( case "saveApiConfiguration": if (message.text && message.apiConfiguration) { try { - await provider.providerSettingsManager.saveConfig(message.text, message.apiConfiguration) - const listApiConfig = await provider.providerSettingsManager.listConfig() - await updateGlobalState("listApiConfigMeta", listApiConfig) + // Enforce the organization allowlist at the persistence + // boundary, not only in the webview, so a crafted or stale + // profile update cannot bypass the policy. + const { organizationAllowList } = await provider.getState() + if ( + !ProfileValidator.isProfileAllowed( + message.apiConfiguration, + organizationAllowList ?? ORGANIZATION_ALLOW_ALL, + ) + ) { + throw new OrganizationAllowListViolationError( + t("common:errors.violated_organization_allowlist"), + ) + } + await provider.upsertProviderProfile( + message.text, + message.apiConfiguration, + false, + message.apiConfigurationBaseline, + ) } catch (error) { provider.log( `Error save api configuration: ${JSON.stringify(error, Object.getOwnPropertyNames(error), 2)}`, ) - vscode.window.showErrorMessage(t("common:errors.save_api_config")) + vscode.window.showErrorMessage( + error instanceof OrganizationAllowListViolationError + ? error.message + : t("common:errors.save_api_config"), + ) } } break case "upsertApiConfiguration": if (message.text && message.apiConfiguration) { - await provider.upsertProviderProfile(message.text, message.apiConfiguration) + let success = false + try { + // Enforce the organization allowlist at the persistence + // boundary, not only in the webview, so a crafted or stale + // profile update cannot bypass the policy. + const { organizationAllowList } = await provider.getState() + if ( + !ProfileValidator.isProfileAllowed( + message.apiConfiguration, + organizationAllowList ?? ORGANIZATION_ALLOW_ALL, + ) + ) { + throw new OrganizationAllowListViolationError( + t("common:errors.violated_organization_allowlist"), + ) + } + const savedId = await provider.upsertProviderProfile( + message.text, + message.apiConfiguration, + true, + message.apiConfigurationBaseline, + ) + success = savedId !== undefined + } catch (error) { + provider.log( + `Error upsert api configuration: ${JSON.stringify(error, Object.getOwnPropertyNames(error), 2)}`, + ) + vscode.window.showErrorMessage( + error instanceof OrganizationAllowListViolationError + ? error.message + : t("common:errors.save_api_config"), + ) + } finally { + if (message.requestId) { + await provider.postMessageToWebview({ + type: "apiConfigurationSaved", + requestId: message.requestId, + success, + }) + } + } + } else if (message.requestId) { + await provider.postMessageToWebview({ + type: "apiConfigurationSaved", + requestId: message.requestId, + success: false, + }) } break case "renameApiConfiguration": if (message.values && message.apiConfiguration) { + let success = false try { const { oldName, newName } = message.values if (oldName === newName) { + success = true break } - // Load the old configuration to get its ID. - const { id } = await provider.providerSettingsManager.getProfile({ name: oldName }) - - // Create a new configuration with the new name and old ID. - await provider.providerSettingsManager.saveConfig(newName, { ...message.apiConfiguration, id }) - - // Delete the old configuration. - await provider.providerSettingsManager.deleteConfig(oldName) + const { organizationAllowList } = await provider.getState() + if ( + !ProfileValidator.isProfileAllowed( + message.apiConfiguration, + organizationAllowList ?? ORGANIZATION_ALLOW_ALL, + ) + ) { + throw new OrganizationAllowListViolationError( + t("common:errors.violated_organization_allowlist"), + ) + } - // Re-activate to update the global settings related to the - // currently activated provider profile. - await provider.activateProviderProfile({ name: newName }) + const savedId = await provider.upsertProviderProfile( + newName, + message.apiConfiguration, + true, + message.apiConfigurationBaseline, + oldName, + ) + success = savedId !== undefined } catch (error) { provider.log( `Error rename api configuration: ${JSON.stringify(error, Object.getOwnPropertyNames(error), 2)}`, ) - vscode.window.showErrorMessage(t("common:errors.rename_api_config")) + vscode.window.showErrorMessage( + error instanceof OrganizationAllowListViolationError + ? error.message + : t("common:errors.rename_api_config"), + ) + } finally { + if (message.requestId) { + await provider.postMessageToWebview({ + type: "apiConfigurationSaved", + requestId: message.requestId, + success, + }) + } } + } else if (message.requestId) { + await provider.postMessageToWebview({ + type: "apiConfigurationSaved", + requestId: message.requestId, + success: false, + }) } break case "loadApiConfiguration": diff --git a/src/eslint-suppressions.json b/src/eslint-suppressions.json index e4b15aa27e..c8d7ffe57a 100644 --- a/src/eslint-suppressions.json +++ b/src/eslint-suppressions.json @@ -401,7 +401,7 @@ }, "api/providers/openai.ts": { "@typescript-eslint/no-explicit-any": { - "count": 3 + "count": 2 } }, "api/providers/openrouter.ts": { diff --git a/src/shared/ProfileValidator.ts b/src/shared/ProfileValidator.ts index 923582bd06..b002c0e5ee 100644 --- a/src/shared/ProfileValidator.ts +++ b/src/shared/ProfileValidator.ts @@ -69,9 +69,17 @@ export class ProfileValidator { return profile.litellmModelId case providerIdentifiers.lmstudio: return profile.lmStudioModelId - case providerIdentifiers.vscodeLm: - // We probably need something more flexible for this one, if we need to really support it here. - return profile.vsCodeLmModelSelector?.id + case providerIdentifiers.vscodeLm: { + // The webview's model selector validates and stores the compound `vendor/family` + // identity when no explicit `id` is set, so the validator must extract the same + // value; otherwise a restrictive allowlist entry keyed on `vendor/family` would + // never match the stored selection. + const selector = profile.vsCodeLmModelSelector + return ( + selector?.id ?? + (selector?.vendor && selector?.family ? `${selector.vendor}/${selector.family}` : undefined) + ) + } case providerIdentifiers.openrouter: return profile.openRouterModelId case providerIdentifiers.ollama: diff --git a/src/shared/__tests__/ProfileValidator.spec.ts b/src/shared/__tests__/ProfileValidator.spec.ts index 865fa5cf51..996ee2ed3f 100644 --- a/src/shared/__tests__/ProfileValidator.spec.ts +++ b/src/shared/__tests__/ProfileValidator.spec.ts @@ -296,6 +296,38 @@ describe("ProfileValidator", () => { expect(ProfileValidator.isProfileAllowed(profile, allowList)).toBe(true) }) + it("should fall back to the vendor/family compound id for vscode-lm when id is absent", () => { + const allowList: OrganizationAllowList = { + allowAll: false, + providers: { + [providerIdentifiers.vscodeLm]: { allowAll: false, models: ["copilot/gpt-4"] }, + }, + } + const profile: ProviderSettings = { + apiProvider: providerIdentifiers.vscodeLm, + vsCodeLmModelSelector: { vendor: "copilot", family: "gpt-4" }, + } + + expect(ProfileValidator.isProfileAllowed(profile, allowList)).toBe(true) + }) + + it.each([undefined, {}, { vendor: "copilot" }, { family: "gpt-4" }])( + "rejects incomplete VS Code LM selector %j", + (vsCodeLmModelSelector) => { + expect( + ProfileValidator.isProfileAllowed( + { apiProvider: providerIdentifiers.vscodeLm, vsCodeLmModelSelector }, + { + allowAll: false, + providers: { + [providerIdentifiers.vscodeLm]: { allowAll: false, models: ["copilot/gpt-4"] }, + }, + }, + ), + ).toBe(false) + }, + ) + it("should extract lmStudioModelId for lmstudio provider", () => { const allowList: OrganizationAllowList = { allowAll: false, diff --git a/webview-ui/src/components/chat/ChatModelSelector.tsx b/webview-ui/src/components/chat/ChatModelSelector.tsx new file mode 100644 index 0000000000..6296143971 --- /dev/null +++ b/webview-ui/src/components/chat/ChatModelSelector.tsx @@ -0,0 +1,242 @@ +import { useMemo, useState, useCallback, useRef, type KeyboardEvent } from "react" + +import { cn } from "@/lib/utils" +import { useRooPortal } from "@/components/ui/hooks/useRooPortal" +import { Popover, PopoverContent, PopoverTrigger, StandardTooltip } from "@/components/ui" +import { useAppTranslation } from "@/i18n/TranslationContext" +import { vscode } from "@/utils/vscode" +import { useExtensionState } from "@/context/ExtensionStateContext" +import { useSelectedModel } from "@/components/ui/hooks/useSelectedModel" +import { isModelAllowedForOrganization, isProviderAllowAll } from "@roo-code/types" + +import { filterModels } from "@/components/settings/utils/organizationFilters" +import { useChatModelSelector } from "./hooks/useChatModelSelector" + +interface ChatModelSelectorProps { + disabled?: boolean + title: string + triggerClassName?: string +} + +export const ChatModelSelector = ({ disabled = false, title, triggerClassName = "" }: ChatModelSelectorProps) => { + const { t } = useAppTranslation() + const { apiConfiguration, organizationAllowList, currentApiConfigName } = useExtensionState() + const { id: selectedModelId } = useSelectedModel(apiConfiguration) + const { provider, models, modelIdKey, defaultModelId, isLoading, valueTransform, displayTransform } = + useChatModelSelector() + + const [open, setOpen] = useState(false) + const [searchValue, setSearchValue] = useState("") + const searchInputRef = useRef(null) + const portalContainer = useRooPortal("roo-portal") + + // Filter deprecated models but always keep the currently selected one visible. + const modelIds = useMemo(() => { + const filteredModels = filterModels(models, provider, organizationAllowList) + const available = Object.entries(filteredModels ?? {}) + .filter(([modelId, modelInfo]) => { + if (modelId === selectedModelId) return true + return !modelInfo.deprecated + }) + .map(([modelId]) => modelId) + .sort((a, b) => a.localeCompare(b)) + return available + }, [models, provider, organizationAllowList, selectedModelId]) + + // Resolve the display value (custom transform for compound config values like VSCode LM). + // The trigger label must reflect the persisted selection, never the transient search + // text — otherwise typing in the popover search box rewrites the trigger/highlight. + const displayValue = useMemo(() => { + if (displayTransform && modelIdKey) { + const storedValue = apiConfiguration?.[modelIdKey] + return storedValue ? displayTransform(storedValue) : undefined + } + return selectedModelId || undefined + }, [displayTransform, modelIdKey, apiConfiguration, selectedModelId]) + + const filteredModelIds = useMemo(() => { + if (!searchValue) return modelIds + const q = searchValue.toLowerCase() + return modelIds.filter((id) => id.toLowerCase().includes(q)) + }, [modelIds, searchValue]) + + // Org policy permits an arbitrary/custom model id only when the provider is + // explicitly `allowAll`; the options list is already filtered by `filterModels`. + const canUseCustomModel = useMemo( + () => isProviderAllowAll(organizationAllowList, provider), + [organizationAllowList, provider], + ) + + // Defense-in-depth: reject a selection the org policy disallows, regardless + // of how it was produced (option click, keyboard, or custom entry). + const isModelAllowed = useCallback( + (modelId: string): boolean => isModelAllowedForOrganization(organizationAllowList, provider, modelId), + [organizationAllowList, provider], + ) + + const onSelect = useCallback( + (modelId: string) => { + if (!modelId || !modelIdKey || !apiConfiguration || !currentApiConfigName) { + return + } + + if (!isModelAllowed(modelId)) { + return + } + + setOpen(false) + setSearchValue("") + + // Transform the model id for storage if needed (e.g. VSCode LM selector object). + const valueToStore = valueTransform ? valueTransform(modelId) : modelId + + // Persist the change to the current API configuration profile. The + // backend will save the profile, activate it and broadcast the + // updated apiConfiguration back to the webview. + vscode.postMessage({ + type: "upsertApiConfiguration", + text: currentApiConfigName, + apiConfigurationBaseline: apiConfiguration, + apiConfiguration: { + ...apiConfiguration, + [modelIdKey]: valueToStore, + }, + }) + }, + [modelIdKey, apiConfiguration, valueTransform, currentApiConfigName, isModelAllowed], + ) + + const onClearSearch = useCallback(() => { + setSearchValue("") + searchInputRef.current?.focus() + }, []) + + const onKeyDown = (event: KeyboardEvent) => { + if (event.nativeEvent.isComposing) return + + const options = Array.from(event.currentTarget.querySelectorAll("[data-model-option]")) + const currentIndex = options.findIndex((option) => option === event.target) + const isSearchInput = event.target === searchInputRef.current + if (!isSearchInput && currentIndex === -1) return + + if (isSearchInput && event.key === "Enter") { + event.preventDefault() + options[0]?.click() + return + } + + let nextIndex: number + switch (event.key) { + case "ArrowDown": + nextIndex = currentIndex + 1 + break + case "ArrowUp": + nextIndex = currentIndex === -1 ? options.length - 1 : currentIndex - 1 + break + case "Home": + case "End": + if (isSearchInput) return + nextIndex = event.key === "Home" ? 0 : options.length - 1 + break + default: + return + } + event.preventDefault() + if (nextIndex < 0 || nextIndex >= options.length) { + searchInputRef.current?.focus() + } else { + options[nextIndex]?.focus() + } + } + + return ( + + + + + {isLoading ? "…" : displayValue || defaultModelId || t("chat:selectModel")} + + + + +
+
+ setSearchValue(e.target.value)} + placeholder={t("chat:searchModel")} + className="w-full h-8 px-2 py-1 text-xs bg-vscode-input-background text-vscode-input-foreground border border-vscode-input-border rounded focus:outline-0" + autoFocus + /> + {searchValue.length > 0 && ( +
+
+ )} +
+ + {filteredModelIds.length === 0 && !isLoading && ( +
{t("chat:modelListEmpty")}
+ )} + + {modelIds.length > 0 && ( +
+ {filteredModelIds.map((modelId) => ( + + ))} +
+ )} + + {searchValue && canUseCustomModel && !modelIds.includes(searchValue) && ( + + )} +
+
+
+ ) +} diff --git a/webview-ui/src/components/chat/ChatTextArea.tsx b/webview-ui/src/components/chat/ChatTextArea.tsx index a78f3f0f8e..4fea41b578 100644 --- a/webview-ui/src/components/chat/ChatTextArea.tsx +++ b/webview-ui/src/components/chat/ChatTextArea.tsx @@ -27,6 +27,7 @@ import { StandardTooltip } from "@src/components/ui" import Thumbnails from "../common/Thumbnails" import { ModeSelector } from "./ModeSelector" import { ApiConfigSelector } from "./ApiConfigSelector" +import { ChatModelSelector } from "./ChatModelSelector" import { AutoApproveDropdown } from "./AutoApproveDropdown" import { MAX_IMAGES_PER_MESSAGE } from "./constants" import ContextMenu from "./ContextMenu" @@ -1319,6 +1320,11 @@ export const ChatTextArea = forwardRef( lockApiConfigAcrossModes={!!lockApiConfigAcrossModes} onToggleLockApiConfig={handleToggleLockApiConfig} /> +
diff --git a/webview-ui/src/components/chat/__tests__/ChatModelSelector.spec.tsx b/webview-ui/src/components/chat/__tests__/ChatModelSelector.spec.tsx new file mode 100644 index 0000000000..82e2648137 --- /dev/null +++ b/webview-ui/src/components/chat/__tests__/ChatModelSelector.spec.tsx @@ -0,0 +1,500 @@ +import type { ReactNode } from "react" +import userEvent from "@testing-library/user-event" +import { render, screen, fireEvent } from "@/utils/test-utils" +import { isModelAllowedForOrganization, providerIdentifiers } from "@roo-code/types" +import { vscode } from "@/utils/vscode" + +import { ChatModelSelector } from "../ChatModelSelector" +import { useChatModelSelector } from "../hooks/useChatModelSelector" + +vi.mock("@roo-code/types", async (importOriginal) => { + const actual = await importOriginal() + return { ...actual, isModelAllowedForOrganization: vi.fn(actual.isModelAllowedForOrganization) } +}) + +// Mock the dependencies +vi.mock("@/utils/vscode", () => ({ + vscode: { + postMessage: vi.fn(), + }, +})) + +vi.mock("@/i18n/TranslationContext", () => ({ + useAppTranslation: () => ({ + t: (key: string) => key, + }), +})) + +vi.mock("@/components/ui/hooks/useRooPortal", () => ({ + useRooPortal: () => document.body, +})) + +// Mock the ExtensionStateContext (configurable per test) +const { mockUseExtensionState } = vi.hoisted(() => ({ + mockUseExtensionState: vi.fn(), +})) + +vi.mock("@/context/ExtensionStateContext", () => ({ + useExtensionState: (...args: unknown[]) => mockUseExtensionState(...args), +})) + +// Mock useSelectedModel +vi.mock("@/components/ui/hooks/useSelectedModel", () => ({ + useSelectedModel: () => ({ id: "claude-opus-4-20250514", info: undefined }), +})) + +// Mock useChatModelSelector hook +vi.mock("../hooks/useChatModelSelector", () => ({ + useChatModelSelector: vi.fn(), +})) + +// Exercise the real popover's focus management and Escape handling. +vi.mock("@/components/ui", async () => ({ + ...(await import("@/components/ui/popover")), + StandardTooltip: ({ children }: { children: ReactNode }) => <>{children}, +})) + +const mockUseChatModelSelector = useChatModelSelector as ReturnType + +const anthropicModels = { + "claude-opus-4-20250514": { + maxTokens: 32000, + contextWindow: 200000, + supportsPromptCache: true, + supportsImages: true, + }, + "claude-sonnet-4-20250514": { + maxTokens: 32000, + contextWindow: 200000, + supportsPromptCache: true, + supportsImages: true, + }, +} + +// Build the hook's return value for a mocked render. A factory keeps the +// default shape in one place so tests can override a single field without +// invoking the mock itself (which is not callable under TypeScript). +const createSelectorState = (overrides: Partial> = {}) => ({ + provider: providerIdentifiers.anthropic, + models: anthropicModels, + modelIdKey: "apiModelId", + defaultModelId: "claude-sonnet-4-20250514", + isLoading: false, + ...overrides, +}) + +describe("ChatModelSelector", () => { + const defaultProps = { + title: "Select model", + } + + const focusDescriptor = Object.getOwnPropertyDescriptor(HTMLElement.prototype, "focus")! + + beforeAll(() => { + // The global FAST compatibility mock makes focus a no-op. Borrow JSDOM's + // native implementation from a fresh realm for keyboard interaction tests. + const iframe = document.createElement("iframe") + document.body.appendChild(iframe) + const nativeFocus = iframe.contentDocument!.createElement("button").focus + Object.defineProperty(HTMLElement.prototype, "focus", { + configurable: true, + writable: true, + value: nativeFocus, + }) + iframe.remove() + }) + + afterAll(() => { + Object.defineProperty(HTMLElement.prototype, "focus", focusDescriptor) + }) + + beforeEach(() => { + vi.clearAllMocks() + mockUseExtensionState.mockReturnValue({ + apiConfiguration: { + apiProvider: providerIdentifiers.anthropic, + apiModelId: "claude-opus-4-20250514", + }, + organizationAllowList: { allowAll: true, providers: {} }, + currentApiConfigName: "default", + }) + mockUseChatModelSelector.mockReturnValue(createSelectorState()) + }) + + test("renders the trigger with the current model id", () => { + render() + + const trigger = screen.getByTestId("chat-model-selector-trigger") + expect(trigger).toBeInTheDocument() + expect(trigger).toHaveTextContent("claude-opus-4-20250514") + }) + + test("disables the trigger when disabled prop is true", () => { + render() + + const trigger = screen.getByTestId("chat-model-selector-trigger") + expect(trigger).toBeDisabled() + }) + + test("renders model list when opened", () => { + render() + + const trigger = screen.getByTestId("chat-model-selector-trigger") + fireEvent.click(trigger) + + // Use testids instead of text queries because the selected model id also + // appears in the trigger label. + expect(screen.getByTestId("chat-model-option-claude-opus-4-20250514")).toBeInTheDocument() + expect(screen.getByTestId("chat-model-option-claude-sonnet-4-20250514")).toBeInTheDocument() + }) + + test("hides deprecated models while retaining the persisted deprecated selection", () => { + mockUseChatModelSelector.mockReturnValue( + createSelectorState({ + models: { + ...anthropicModels, + "claude-opus-4-20250514": { ...anthropicModels["claude-opus-4-20250514"], deprecated: true }, + "retired-model": { ...anthropicModels["claude-sonnet-4-20250514"], deprecated: true }, + }, + }), + ) + render() + fireEvent.click(screen.getByTestId("chat-model-selector-trigger")) + + expect(screen.queryByTestId("chat-model-option-retired-model")).not.toBeInTheDocument() + expect(screen.getByTestId("chat-model-option-claude-opus-4-20250514")).toBeInTheDocument() + expect(screen.getByTestId("chat-model-option-claude-sonnet-4-20250514")).toBeInTheDocument() + }) + + test.each(["ArrowDown", "ArrowUp", "Enter"])("ignores %s while composing text", async (key) => { + const user = userEvent.setup() + render() + await user.click(screen.getByTestId("chat-model-selector-trigger")) + const search = screen.getByRole("textbox") + expect(search).toHaveFocus() + + fireEvent.keyDown(search, { key, isComposing: true }) + + expect(search).toHaveFocus() + expect(vscode.postMessage).not.toHaveBeenCalled() + // Normal navigation still works after composition ends. + await user.keyboard("{ArrowDown}{Enter}") + expect(vscode.postMessage).toHaveBeenCalledWith(expect.objectContaining({ type: "upsertApiConfiguration" })) + }) + + test("shows the empty catalog message once loading finishes", async () => { + const user = userEvent.setup() + mockUseChatModelSelector.mockReturnValue({ + provider: providerIdentifiers.anthropic, + models: {}, + modelIdKey: "apiModelId", + defaultModelId: "", + isLoading: false, + }) + render() + await user.click(screen.getByTestId("chat-model-selector-trigger")) + expect(screen.getByText("chat:modelListEmpty")).toBeInTheDocument() + }) + + test.each([false, true])("handles an unmatched search while isLoading is %s", async (isLoading) => { + const user = userEvent.setup() + mockUseChatModelSelector.mockReturnValue(createSelectorState({ isLoading })) + render() + await user.click(screen.getByTestId("chat-model-selector-trigger")) + await user.type(screen.getByRole("textbox"), "custom-model") + + expect(screen.queryByTestId("chat-model-option-claude-opus-4-20250514")).not.toBeInTheDocument() + expect(screen.queryByTestId("chat-model-option-claude-sonnet-4-20250514")).not.toBeInTheDocument() + expect(screen.queryByText("chat:modelListEmpty") !== null).toBe(!isLoading) + expect(screen.getByTestId("chat-model-use-custom")).toBeInTheDocument() + + await user.click(screen.getByRole("button", { name: "chat:clearSearch" })) + expect(screen.queryByText("chat:modelListEmpty")).not.toBeInTheDocument() + expect(screen.getByTestId("chat-model-option-claude-opus-4-20250514")).toBeInTheDocument() + expect(screen.getByTestId("chat-model-option-claude-sonnet-4-20250514")).toBeInTheDocument() + }) + + test("posts the transformed selection for compound configurations", async () => { + const user = userEvent.setup() + mockUseExtensionState.mockReturnValue({ + apiConfiguration: { apiProvider: providerIdentifiers.vscodeLm }, + organizationAllowList: { allowAll: true, providers: {} }, + currentApiConfigName: "default", + }) + mockUseChatModelSelector.mockReturnValue({ + provider: providerIdentifiers.vscodeLm, + models: { "copilot/gpt-4o": { maxTokens: 1, contextWindow: 1 } }, + modelIdKey: "vsCodeLmModelSelector", + defaultModelId: "", + isLoading: false, + valueTransform: (modelId: string) => { + const [vendor, family] = modelId.split("/") + return { vendor, family } + }, + }) + render() + await user.click(screen.getByTestId("chat-model-selector-trigger")) + await user.click(screen.getByTestId("chat-model-option-copilot/gpt-4o")) + expect(vscode.postMessage).toHaveBeenCalledWith({ + type: "upsertApiConfiguration", + apiConfigurationBaseline: mockUseExtensionState().apiConfiguration, + text: "default", + apiConfiguration: { + apiProvider: providerIdentifiers.vscodeLm, + vsCodeLmModelSelector: { vendor: "copilot", family: "gpt-4o" }, + }, + }) + }) + + test("selecting a model posts upsertApiConfiguration message", () => { + render() + + const trigger = screen.getByTestId("chat-model-selector-trigger") + fireEvent.click(trigger) + + const option = screen.getByTestId("chat-model-option-claude-sonnet-4-20250514") + fireEvent.click(option) + + expect(vscode.postMessage).toHaveBeenCalledWith({ + type: "upsertApiConfiguration", + apiConfigurationBaseline: mockUseExtensionState().apiConfiguration, + text: "default", + apiConfiguration: { + apiProvider: providerIdentifiers.anthropic, + apiModelId: "claude-sonnet-4-20250514", + }, + }) + }) + + test("supports custom model via search", () => { + render() + + const trigger = screen.getByTestId("chat-model-selector-trigger") + fireEvent.click(trigger) + + // Type in the search box. + const searchInput = screen.getByPlaceholderText("chat:searchModel") + fireEvent.change(searchInput, { target: { value: "claude-3-5-haiku" } }) + + const customOption = screen.getByTestId("chat-model-use-custom") + fireEvent.click(customOption) + + expect(vscode.postMessage).toHaveBeenCalledWith({ + type: "upsertApiConfiguration", + apiConfigurationBaseline: mockUseExtensionState().apiConfiguration, + text: "default", + apiConfiguration: { + apiProvider: providerIdentifiers.anthropic, + apiModelId: "claude-3-5-haiku", + }, + }) + }) + + test("clears the search query via the clear button", () => { + render() + + const trigger = screen.getByTestId("chat-model-selector-trigger") + fireEvent.click(trigger) + + const searchInput = screen.getByPlaceholderText("chat:searchModel") as HTMLInputElement + fireEvent.change(searchInput, { target: { value: "claude-3" } }) + expect(searchInput.value).toBe("claude-3") + + const clearButton = screen.getByRole("button", { name: "chat:clearSearch" }) + fireEvent.click(clearButton) + + expect((screen.getByPlaceholderText("chat:searchModel") as HTMLInputElement).value).toBe("") + }) + + test("does not post a message when modelIdKey becomes unset", () => { + const { rerender } = render() + fireEvent.click(screen.getByTestId("chat-model-selector-trigger")) + mockUseChatModelSelector.mockReturnValue({ + provider: providerIdentifiers.anthropic, + models: anthropicModels, + modelIdKey: undefined, + defaultModelId: undefined, + isLoading: false, + }) + rerender() + + expect(screen.getByTestId("chat-model-selector-trigger")).toBeDisabled() + fireEvent.click(screen.getByTestId("chat-model-option-claude-opus-4-20250514")) + expect(vscode.postMessage).not.toHaveBeenCalled() + }) + + test("does not post a message when apiConfiguration is missing and modelIdKey is set", () => { + mockUseExtensionState.mockReturnValue({ + apiConfiguration: undefined, + organizationAllowList: { allowAll: true, providers: {} }, + currentApiConfigName: "default", + }) + render() + const trigger = screen.getByTestId("chat-model-selector-trigger") + expect(trigger).toBeEnabled() + fireEvent.click(trigger) + fireEvent.click(screen.getByTestId("chat-model-option-claude-opus-4-20250514")) + expect(vscode.postMessage).not.toHaveBeenCalled() + }) + + test("opens, navigates options, and selects a model using the keyboard", async () => { + const user = userEvent.setup() + render() + await user.tab() + await user.keyboard("{Enter}") + expect(screen.getByRole("textbox")).toHaveFocus() + await user.keyboard("{ArrowDown}") + expect(screen.getByTestId("chat-model-option-claude-opus-4-20250514")).toHaveFocus() + await user.keyboard("{End}") + expect(screen.getByRole("button", { name: "claude-sonnet-4-20250514" })).toHaveFocus() + await user.keyboard("{Home}{ArrowUp}") + expect(screen.getByRole("textbox")).toHaveFocus() + await user.keyboard("{ArrowUp}{Enter}") + expect(vscode.postMessage).toHaveBeenCalledWith( + expect.objectContaining({ + apiConfiguration: expect.objectContaining({ apiModelId: "claude-sonnet-4-20250514" }), + }), + ) + expect(screen.queryByRole("textbox")).not.toBeInTheDocument() + expect(screen.getByTestId("chat-model-selector-trigger")).toHaveFocus() + }) + + test("selects the first filtered model with Enter from search", async () => { + const user = userEvent.setup() + render() + await user.click(screen.getByTestId("chat-model-selector-trigger")) + await user.keyboard("sonnet{Enter}") + expect(vscode.postMessage).toHaveBeenCalledWith( + expect.objectContaining({ + apiConfiguration: expect.objectContaining({ apiModelId: "claude-sonnet-4-20250514" }), + }), + ) + }) + + test.each(["{Enter}", "{ArrowDown} "])("selects a custom model with %s", async (keys) => { + const user = userEvent.setup() + render() + await user.click(screen.getByTestId("chat-model-selector-trigger")) + await user.keyboard(`custom-model${keys}`) + expect(vscode.postMessage).toHaveBeenCalledWith( + expect.objectContaining({ + apiConfiguration: expect.objectContaining({ apiModelId: "custom-model" }), + }), + ) + }) + + test("clears search by keyboard and dismisses the popover with Escape", async () => { + const user = userEvent.setup() + render() + await user.click(screen.getByTestId("chat-model-selector-trigger")) + await user.keyboard("sonnet") + await user.tab() + expect(screen.getByRole("button", { name: "chat:clearSearch" })).toHaveFocus() + await user.keyboard(" ") + expect(screen.getByRole("textbox")).toHaveValue("") + expect(screen.getByRole("textbox")).toHaveFocus() + await user.keyboard("{Escape}") + expect(screen.queryByRole("textbox")).not.toBeInTheDocument() + expect(screen.getByTestId("chat-model-selector-trigger")).toHaveFocus() + expect(vscode.postMessage).not.toHaveBeenCalled() + }) + + test("renders the display-transformed value for compound configs (e.g. VSCode LM)", () => { + mockUseExtensionState.mockReturnValue({ + apiConfiguration: { + apiProvider: providerIdentifiers.vscodeLm, + vsCodeLmModelSelector: { vendor: "copilot", family: "gpt-4o" }, + }, + organizationAllowList: { allowAll: true, providers: {} }, + currentApiConfigName: "default", + }) + mockUseChatModelSelector.mockReturnValue({ + provider: providerIdentifiers.vscodeLm, + models: { "copilot/gpt-4o": { maxTokens: 1, contextWindow: 1 } }, + modelIdKey: "vsCodeLmModelSelector", + defaultModelId: undefined, + isLoading: false, + displayTransform: (value: unknown) => { + const selector = value as { vendor?: string; family?: string } + return selector.vendor && selector.family ? `${selector.vendor}/${selector.family}` : "" + }, + }) + + render() + + // The trigger shows the display-transformed value instead of the raw model id. + const trigger = screen.getByTestId("chat-model-selector-trigger") + expect(trigger).toHaveTextContent("copilot/gpt-4o") + }) + + test("disables the trigger when currentApiConfigName is missing", () => { + mockUseExtensionState.mockReturnValue({ + apiConfiguration: { + apiProvider: providerIdentifiers.anthropic, + apiModelId: "claude-opus-4-20250514", + }, + organizationAllowList: { allowAll: true, providers: {} }, + currentApiConfigName: undefined, + }) + + render() + + // Without a resolved profile name the upsert would be dropped, so the + // control must not be interactive — and must not look interactive. + expect(screen.getByTestId("chat-model-selector-trigger")).toBeDisabled() + }) + + test("hides the custom-model escape hatch when the provider is restricted", () => { + mockUseExtensionState.mockReturnValue({ + apiConfiguration: { + apiProvider: providerIdentifiers.anthropic, + apiModelId: "claude-opus-4-20250514", + }, + organizationAllowList: { + allowAll: false, + providers: { + [providerIdentifiers.anthropic]: { allowAll: false, models: ["claude-sonnet-4-20250514"] }, + }, + }, + currentApiConfigName: "default", + }) + + render() + fireEvent.click(screen.getByTestId("chat-model-selector-trigger")) + fireEvent.change(screen.getByPlaceholderText("chat:searchModel"), { target: { value: "gpt-999" } }) + + // The provider is narrowed to an allow-list, so free-form ids are not offered. + expect(screen.queryByTestId("chat-model-use-custom")).not.toBeInTheDocument() + }) + + test("rejects a rendered selection when the policy check denies it", () => { + render() + fireEvent.click(screen.getByTestId("chat-model-selector-trigger")) + vi.mocked(isModelAllowedForOrganization).mockReturnValueOnce(false) + fireEvent.click(screen.getByTestId("chat-model-option-claude-sonnet-4-20250514")) + expect(vscode.postMessage).not.toHaveBeenCalledWith(expect.objectContaining({ type: "upsertApiConfiguration" })) + }) + + test("filters out models that the organization policy disallows", () => { + mockUseExtensionState.mockReturnValue({ + apiConfiguration: { + apiProvider: providerIdentifiers.anthropic, + apiModelId: "claude-opus-4-20250514", + }, + organizationAllowList: { + allowAll: false, + providers: { + [providerIdentifiers.anthropic]: { allowAll: false, models: ["claude-sonnet-4-20250514"] }, + }, + }, + currentApiConfigName: "default", + }) + + render() + fireEvent.click(screen.getByTestId("chat-model-selector-trigger")) + + // Only the allow-listed model is offered; the disallowed Opus entry is absent. + expect(screen.getByTestId("chat-model-option-claude-sonnet-4-20250514")).toBeInTheDocument() + expect(screen.queryByTestId("chat-model-option-claude-opus-4-20250514")).not.toBeInTheDocument() + }) +}) diff --git a/webview-ui/src/components/chat/__tests__/ChatTextArea.visual.tsx b/webview-ui/src/components/chat/__tests__/ChatTextArea.visual.tsx index 9c447fe011..8b54af07b6 100644 --- a/webview-ui/src/components/chat/__tests__/ChatTextArea.visual.tsx +++ b/webview-ui/src/components/chat/__tests__/ChatTextArea.visual.tsx @@ -14,14 +14,10 @@ for (const theme of visualThemes) { await expect(editor).toBeVisible() await expect(story).toHaveScreenshot(`chat-composer-resting-${theme.name}.png`) - await page.evaluate(() => (document.activeElement as HTMLElement | null)?.blur()) - for ( - let index = 0; - index < 10 && !(await editor.evaluate((element) => element === document.activeElement)); - index++ - ) { - await page.keyboard.press("Tab") - } + // Re-enter the editor by keyboard without depending on the composer's tab-stop count. + await editor.focus() + await page.keyboard.press("Tab") + await page.keyboard.press("Shift+Tab") await expect(editor).toBeFocused() await expect(story).toHaveScreenshot(`chat-composer-focus-${theme.name}.png`) await expectBoundedLayout(page, story, { diff --git a/webview-ui/src/components/chat/__tests__/__screenshots__/chat-composer-focus-dark.png b/webview-ui/src/components/chat/__tests__/__screenshots__/chat-composer-focus-dark.png index 4de566f736..6c90b6f610 100644 Binary files a/webview-ui/src/components/chat/__tests__/__screenshots__/chat-composer-focus-dark.png and b/webview-ui/src/components/chat/__tests__/__screenshots__/chat-composer-focus-dark.png differ diff --git a/webview-ui/src/components/chat/__tests__/__screenshots__/chat-composer-focus-high-contrast-light.png b/webview-ui/src/components/chat/__tests__/__screenshots__/chat-composer-focus-high-contrast-light.png index f5c759fad8..8e9c01c7fe 100644 Binary files a/webview-ui/src/components/chat/__tests__/__screenshots__/chat-composer-focus-high-contrast-light.png and b/webview-ui/src/components/chat/__tests__/__screenshots__/chat-composer-focus-high-contrast-light.png differ diff --git a/webview-ui/src/components/chat/__tests__/__screenshots__/chat-composer-focus-high-contrast.png b/webview-ui/src/components/chat/__tests__/__screenshots__/chat-composer-focus-high-contrast.png index 465ffc0263..8d232cb2ab 100644 Binary files a/webview-ui/src/components/chat/__tests__/__screenshots__/chat-composer-focus-high-contrast.png and b/webview-ui/src/components/chat/__tests__/__screenshots__/chat-composer-focus-high-contrast.png differ diff --git a/webview-ui/src/components/chat/__tests__/__screenshots__/chat-composer-focus-light.png b/webview-ui/src/components/chat/__tests__/__screenshots__/chat-composer-focus-light.png index 1f238a4619..b6976d6cff 100644 Binary files a/webview-ui/src/components/chat/__tests__/__screenshots__/chat-composer-focus-light.png and b/webview-ui/src/components/chat/__tests__/__screenshots__/chat-composer-focus-light.png differ diff --git a/webview-ui/src/components/chat/__tests__/__screenshots__/chat-composer-resting-dark.png b/webview-ui/src/components/chat/__tests__/__screenshots__/chat-composer-resting-dark.png index fad9aa30b9..6c90b6f610 100644 Binary files a/webview-ui/src/components/chat/__tests__/__screenshots__/chat-composer-resting-dark.png and b/webview-ui/src/components/chat/__tests__/__screenshots__/chat-composer-resting-dark.png differ diff --git a/webview-ui/src/components/chat/__tests__/__screenshots__/chat-composer-resting-high-contrast-light.png b/webview-ui/src/components/chat/__tests__/__screenshots__/chat-composer-resting-high-contrast-light.png index f5c759fad8..8e9c01c7fe 100644 Binary files a/webview-ui/src/components/chat/__tests__/__screenshots__/chat-composer-resting-high-contrast-light.png and b/webview-ui/src/components/chat/__tests__/__screenshots__/chat-composer-resting-high-contrast-light.png differ diff --git a/webview-ui/src/components/chat/__tests__/__screenshots__/chat-composer-resting-high-contrast.png b/webview-ui/src/components/chat/__tests__/__screenshots__/chat-composer-resting-high-contrast.png index 465ffc0263..8d232cb2ab 100644 Binary files a/webview-ui/src/components/chat/__tests__/__screenshots__/chat-composer-resting-high-contrast.png and b/webview-ui/src/components/chat/__tests__/__screenshots__/chat-composer-resting-high-contrast.png differ diff --git a/webview-ui/src/components/chat/__tests__/__screenshots__/chat-composer-resting-light.png b/webview-ui/src/components/chat/__tests__/__screenshots__/chat-composer-resting-light.png index 1f238a4619..b6976d6cff 100644 Binary files a/webview-ui/src/components/chat/__tests__/__screenshots__/chat-composer-resting-light.png and b/webview-ui/src/components/chat/__tests__/__screenshots__/chat-composer-resting-light.png differ diff --git a/webview-ui/src/components/chat/hooks/__tests__/useChatModelSelector.spec.tsx b/webview-ui/src/components/chat/hooks/__tests__/useChatModelSelector.spec.tsx new file mode 100644 index 0000000000..45c798a4c1 --- /dev/null +++ b/webview-ui/src/components/chat/hooks/__tests__/useChatModelSelector.spec.tsx @@ -0,0 +1,621 @@ +import React from "react" +import { act, renderHook, waitFor } from "@testing-library/react" +import { QueryClient, QueryClientProvider } from "@tanstack/react-query" + +import { + anthropicDefaultModelId, + openRouterDefaultModelId, + zooGatewayDefaultModelId, + opencodeGoDefaultModelId, + kenariDefaultModelId, + nanoGptDefaultModelId, + requestyDefaultModelId, + unboundDefaultModelId, + vercelAiGatewayDefaultModelId, + kimiCodeDefaultModelId, + mainlandZAiDefaultModelId, + providerIdentifiers, + retiredProviderIdentifiers, +} from "@roo-code/types" + +import { useChatModelSelector } from "../useChatModelSelector" +import { useExtensionState } from "@/context/ExtensionStateContext" +import { useRouterModels } from "@/components/ui/hooks/useRouterModels" +import { vscode } from "@/utils/vscode" + +vi.mock("@/context/ExtensionStateContext", () => ({ + useExtensionState: vi.fn(), +})) + +vi.mock("@/components/ui/hooks/useRouterModels", () => ({ + useRouterModels: vi.fn(), +})) + +vi.mock("@/utils/vscode", () => ({ + vscode: { + postMessage: vi.fn(), + }, +})) + +const mockUseExtensionState = useExtensionState as ReturnType +const mockUseRouterModels = useRouterModels as ReturnType +const mockPostMessage = vscode.postMessage as ReturnType + +const wrapper = ({ children }: { children: React.ReactNode }) => { + const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } }) + return {children} +} + +const emitMessage = (message: unknown) => { + window.dispatchEvent(new MessageEvent("message", { data: message })) +} + +describe("useChatModelSelector", () => { + beforeEach(() => { + vi.clearAllMocks() + mockUseRouterModels.mockReturnValue({ + data: undefined, + isLoading: false, + isError: false, + }) + mockUseExtensionState.mockReturnValue({ + apiConfiguration: { apiProvider: providerIdentifiers.anthropic, apiModelId: "claude-opus-4-20250514" }, + routerModels: undefined, + }) + }) + + describe("request cleanup", () => { + it.each([ + providerIdentifiers.ollama, + providerIdentifiers.lmstudio, + providerIdentifiers.openai, + providerIdentifiers.vscodeLm, + ])("cancels %s on provider switch and rapid remount", (apiProvider) => { + const config = { apiProvider, openAiBaseUrl: "https://example.test/v1", openAiApiKey: "test-key" } + mockUseExtensionState.mockReturnValue({ apiConfiguration: config }) + const first = renderHook(() => useChatModelSelector(), { wrapper }) + const request = mockPostMessage.mock.calls.at(-1)![0] + mockUseExtensionState.mockReturnValue({ apiConfiguration: { apiProvider: providerIdentifiers.anthropic } }) + first.rerender() + expect(mockPostMessage).toHaveBeenLastCalledWith({ + type: "cancelModelRequest", + requestId: request.requestId, + }) + mockUseExtensionState.mockReturnValue({ apiConfiguration: config }) + first.rerender() + const restarted = mockPostMessage.mock.calls.at(-1)![0] + expect(restarted.requestId).not.toBe(request.requestId) + first.unmount() + expect(mockPostMessage).toHaveBeenLastCalledWith({ + type: "cancelModelRequest", + requestId: restarted.requestId, + }) + const second = renderHook(() => useChatModelSelector(), { wrapper }) + const remounted = mockPostMessage.mock.calls.at(-1)![0] + expect(remounted.requestId).not.toBe(restarted.requestId) + second.unmount() + expect(mockPostMessage).toHaveBeenLastCalledWith({ + type: "cancelModelRequest", + requestId: remounted.requestId, + }) + }) + + it("cancels the replayed StrictMode effect and the final request on unmount", () => { + mockUseExtensionState.mockReturnValue({ apiConfiguration: { apiProvider: providerIdentifiers.ollama } }) + const { unmount } = renderHook(() => useChatModelSelector(), { + wrapper: ({ children }) => {children}, + }) + const messages = mockPostMessage.mock.calls.map(([message]) => message) + expect(messages.map((message) => message.type)).toEqual([ + "requestOllamaModels", + "cancelModelRequest", + "requestOllamaModels", + ]) + expect(messages[1].requestId).toBe(messages[0].requestId) + expect(messages[2].requestId).not.toBe(messages[0].requestId) + unmount() + expect(mockPostMessage).toHaveBeenLastCalledWith({ + type: "cancelModelRequest", + requestId: messages[2].requestId, + }) + }) + }) + + it("defaults an unset provider to OpenRouter and configures its catalog query", () => { + mockUseExtensionState.mockReturnValue({ apiConfiguration: {} }) + const { result } = renderHook(() => useChatModelSelector(), { wrapper }) + expect(result.current.provider).toBe(providerIdentifiers.openrouter) + expect(result.current.modelIdKey).toBe("openRouterModelId") + expect(result.current.defaultModelId).toBe(openRouterDefaultModelId) + expect(mockUseRouterModels).toHaveBeenCalledWith({ provider: providerIdentifiers.openrouter, enabled: true }) + }) + + describe("static providers", () => { + it("returns static models for anthropic with apiModelId key", () => { + const { result } = renderHook(() => useChatModelSelector(), { wrapper }) + + expect(result.current.provider).toBe(providerIdentifiers.anthropic) + expect(result.current.modelIdKey).toBe("apiModelId") + expect(result.current.models).not.toBeNull() + expect(Object.keys(result.current.models!)).toContain("claude-opus-4-20250514") + expect(result.current.defaultModelId).toBe(anthropicDefaultModelId) + }) + + it("returns the mainland coding models and default for Z.AI china_coding", () => { + mockUseExtensionState.mockReturnValue({ + apiConfiguration: { apiProvider: providerIdentifiers.zai, zaiApiLine: "china_coding" }, + }) + const { result } = renderHook(() => useChatModelSelector(), { wrapper }) + expect(result.current.defaultModelId).toBe(mainlandZAiDefaultModelId) + expect(Object.keys(result.current.models!).sort()).toEqual([ + "glm-4.5", + "glm-4.5-air", + "glm-4.5-airx", + "glm-4.5-flash", + "glm-4.5-x", + "glm-4.5v", + "glm-4.6", + "glm-4.6v", + "glm-4.6v-flash", + "glm-4.6v-flashx", + "glm-4.7", + "glm-4.7-flash", + "glm-4.7-flashx", + "glm-5", + "glm-5-turbo", + "glm-5.1", + "glm-5.2", + "glm-5.3", + "glm-5.3-flash", + ]) + }) + + it("does not request message-based models for static providers", () => { + renderHook(() => useChatModelSelector(), { wrapper }) + + expect(mockPostMessage).not.toHaveBeenCalled() + }) + }) + + describe("dynamic router providers", () => { + it.each([ + [ + providerIdentifiers.zooGateway, + "zooGatewayModelId", + "anthropic/claude-sonnet-4", + zooGatewayDefaultModelId, + ], + [providerIdentifiers.opencodeGo, "opencodeGoModelId", "glm-5.2", opencodeGoDefaultModelId], + [providerIdentifiers.kenari, "kenariModelId", "glm-5-2", kenariDefaultModelId], + [providerIdentifiers.nanogpt, "nanoGptModelId", "openai/gpt-5.6-sol", nanoGptDefaultModelId], + [providerIdentifiers.requesty, "requestyModelId", "openai/gpt-5.1", requestyDefaultModelId], + [providerIdentifiers.unbound, "unboundModelId", "openai/gpt-4o", unboundDefaultModelId], + [ + providerIdentifiers.vercelAiGateway, + "vercelAiGatewayModelId", + "openai/gpt-4o-mini", + vercelAiGatewayDefaultModelId, + ], + [providerIdentifiers.kimiCode, "apiModelId", "kimi-k2", kimiCodeDefaultModelId], + ] as const)( + "reads %s models from the react-query routerModels", + (provider, modelIdKey, modelId, defaultModelId) => { + mockUseRouterModels.mockReturnValue({ + data: { [provider]: { [modelId]: { maxTokens: 1, contextWindow: 1 } } }, + isLoading: false, + isError: false, + }) + mockUseExtensionState.mockReturnValue({ + apiConfiguration: { apiProvider: provider, [modelIdKey]: modelId }, + routerModels: undefined, + }) + + const { result } = renderHook(() => useChatModelSelector(), { wrapper }) + + expect(result.current.modelIdKey).toBe(modelIdKey) + expect(result.current.defaultModelId).toBe(defaultModelId) + expect(Object.keys(result.current.models!)).toEqual([modelId]) + }, + ) + }) + + describe("openai (OpenAI compatible)", () => { + it("requests openAi models on mount when baseUrl and apiKey are set", () => { + mockUseExtensionState.mockReturnValue({ + apiConfiguration: { + apiProvider: providerIdentifiers.openai, + openAiBaseUrl: "https://api.example.com/v1", + openAiApiKey: "test-key", + openAiHeaders: {}, + }, + routerModels: undefined, + }) + + renderHook(() => useChatModelSelector(), { wrapper }) + + expect(mockPostMessage).toHaveBeenCalledWith( + expect.objectContaining({ + type: "requestOpenAiModels", + values: { + baseUrl: "https://api.example.com/v1", + apiKey: "test-key", + customHeaders: {}, + openAiHeaders: {}, + }, + }), + ) + }) + + it.each(["openAiBaseUrl", "openAiApiKey"] as const)( + "clears loaded models and suppresses requests when %s is removed", + (missingField) => { + const config = { + apiProvider: providerIdentifiers.openai, + openAiBaseUrl: "https://api.example.com/v1", + openAiApiKey: "test-key", + } + mockUseExtensionState.mockReturnValue({ apiConfiguration: config }) + const { result, rerender } = renderHook(() => useChatModelSelector(), { wrapper }) + const requestId = mockPostMessage.mock.calls[0][0].requestId + act(() => emitMessage({ type: "openAiModels", openAiModels: ["gpt-4o"], requestId })) + expect(Object.keys(result.current.models!)).toEqual(["gpt-4o"]) + + mockPostMessage.mockClear() + mockUseExtensionState.mockReturnValue({ apiConfiguration: { ...config, [missingField]: undefined } }) + rerender() + + expect(mockPostMessage.mock.calls.map(([message]) => message)).toEqual([ + { type: "cancelModelRequest", requestId }, + ]) + expect(result.current.models).toBeNull() + expect(result.current.isLoading).toBe(false) + act(() => emitMessage({ type: "openAiModels", openAiModels: ["stale-model"], requestId })) + expect(result.current.models).toBeNull() + expect(result.current.isLoading).toBe(false) + }, + ) + + it("only requests again when serialized header values change", () => { + const configure = (headers?: Record) => + mockUseExtensionState.mockReturnValue({ + apiConfiguration: { + apiProvider: providerIdentifiers.openai, + openAiBaseUrl: "https://api.example.com/v1", + openAiApiKey: "test-key", + openAiHeaders: headers, + }, + }) + configure() + const { rerender } = renderHook(() => useChatModelSelector(), { wrapper }) + configure({}) + rerender() + expect(mockPostMessage).toHaveBeenCalledTimes(1) + configure({ "X-Test": "first" }) + rerender() + expect(mockPostMessage).toHaveBeenCalledTimes(3) + configure({ "X-Test": "first" }) + rerender() + expect(mockPostMessage).toHaveBeenCalledTimes(3) + configure({ "X-Test": "second" }) + rerender() + expect(mockPostMessage).toHaveBeenCalledTimes(5) + expect(mockPostMessage).toHaveBeenLastCalledWith( + expect.objectContaining({ + values: expect.objectContaining({ openAiHeaders: { "X-Test": "second" } }), + }), + ) + }) + + it("uses models delivered through the openAiModels message", async () => { + mockUseExtensionState.mockReturnValue({ + apiConfiguration: { + apiProvider: providerIdentifiers.openai, + openAiBaseUrl: "https://api.example.com/v1", + openAiApiKey: "test-key", + openAiHeaders: {}, + }, + routerModels: undefined, + }) + + const { result } = renderHook(() => useChatModelSelector(), { wrapper }) + + act(() => { + emitMessage({ type: "openAiModels", openAiModels: ["gpt-4o", "gpt-4o-mini"] }) + }) + + await waitFor(() => { + expect(result.current.models).not.toBeNull() + }) + + expect(Object.keys(result.current.models!)).toEqual(expect.arrayContaining(["gpt-4o", "gpt-4o-mini"])) + expect(result.current.modelIdKey).toBe("openAiModelId") + }) + + it("stops loading when the OpenAI models response is empty", async () => { + mockUseExtensionState.mockReturnValue({ + apiConfiguration: { + apiProvider: providerIdentifiers.openai, + openAiBaseUrl: "https://api.example.com/v1", + openAiApiKey: "test-key", + openAiHeaders: {}, + }, + routerModels: undefined, + }) + + const { result } = renderHook(() => useChatModelSelector(), { wrapper }) + + expect(result.current.isLoading).toBe(true) + + const requestId = mockPostMessage.mock.calls[0][0].requestId as string + act(() => { + emitMessage({ type: "openAiModels", openAiModels: [], requestId }) + }) + + // An empty (but successful) response must end the spinner; otherwise the + // trigger shows "…" forever for an endpoint with no models. + await waitFor(() => { + expect(result.current.isLoading).toBe(false) + }) + expect(result.current.models).toBeNull() + }) + + it("ignores model responses whose requestId does not match the latest request", async () => { + mockUseExtensionState.mockReturnValue({ + apiConfiguration: { + apiProvider: providerIdentifiers.openai, + openAiBaseUrl: "https://api.example.com/v1", + openAiApiKey: "test-key", + openAiHeaders: {}, + }, + routerModels: undefined, + }) + + const { result } = renderHook(() => useChatModelSelector(), { wrapper }) + + const requestId = mockPostMessage.mock.calls[0][0].requestId as string + expect(typeof requestId).toBe("string") + + // A response for a different request is stale and must be discarded. + act(() => { + emitMessage({ type: "openAiModels", openAiModels: ["stale-model"], requestId: "different-request" }) + }) + expect(result.current.models).toBeNull() + + // The matching response is applied. + act(() => { + emitMessage({ type: "openAiModels", openAiModels: ["fresh-model"], requestId }) + }) + await waitFor(() => { + expect(result.current.models).not.toBeNull() + }) + expect(Object.keys(result.current.models!)).toEqual(["fresh-model"]) + }) + }) + + describe("router providers", () => { + it("reads openrouter models from the react-query routerModels", () => { + mockUseRouterModels.mockReturnValue({ + data: { [providerIdentifiers.openrouter]: { "openai/gpt-4o": { maxTokens: 1, contextWindow: 1 } } }, + isLoading: false, + isError: false, + }) + mockUseExtensionState.mockReturnValue({ + apiConfiguration: { apiProvider: providerIdentifiers.openrouter, openRouterModelId: "openai/gpt-4o" }, + routerModels: undefined, + }) + + const { result } = renderHook(() => useChatModelSelector(), { wrapper }) + + expect(result.current.modelIdKey).toBe("openRouterModelId") + expect(Object.keys(result.current.models!)).toEqual(["openai/gpt-4o"]) + }) + + it.each([providerIdentifiers.poe, providerIdentifiers.litellm])( + "does not query %s on mount or remount when models come from shared state", + (provider) => { + mockUseExtensionState.mockReturnValue({ + apiConfiguration: { apiProvider: provider }, + routerModels: { [provider]: {} }, + }) + // A disabled query's status must not hide the shared state's empty result. + mockUseRouterModels.mockReturnValue({ isLoading: true }) + const first = renderHook(() => useChatModelSelector(), { wrapper }) + expect(first.result.current.isLoading).toBe(false) + first.unmount() + const second = renderHook(() => useChatModelSelector(), { wrapper }) + expect(second.result.current.models).toEqual({}) + second.unmount() + for (const [options] of mockUseRouterModels.mock.calls) { + expect(options).toEqual({ provider, enabled: false }) + } + expect(mockPostMessage).not.toHaveBeenCalled() + }, + ) + + it("reads poe models from the backend-broadcast state routerModels", () => { + mockUseExtensionState.mockReturnValue({ + apiConfiguration: { apiProvider: providerIdentifiers.poe, poeApiKey: "test" }, + routerModels: { [providerIdentifiers.poe]: { "claude-sonnet": { maxTokens: 1, contextWindow: 1 } } }, + }) + + const { result } = renderHook(() => useChatModelSelector(), { wrapper }) + + expect(result.current.modelIdKey).toBe("apiModelId") + expect(Object.keys(result.current.models!)).toEqual(["claude-sonnet"]) + }) + + it("reads litellm models from the backend-broadcast state routerModels", () => { + mockUseExtensionState.mockReturnValue({ + apiConfiguration: { + apiProvider: providerIdentifiers.litellm, + litellmApiKey: "test", + litellmBaseUrl: "http://x", + }, + routerModels: { [providerIdentifiers.litellm]: { "gpt-4o": { maxTokens: 1, contextWindow: 1 } } }, + }) + + const { result } = renderHook(() => useChatModelSelector(), { wrapper }) + + expect(result.current.modelIdKey).toBe("litellmModelId") + expect(Object.keys(result.current.models!)).toEqual(["gpt-4o"]) + }) + }) + + describe("ollama / lmstudio / vscode-lm", () => { + it("requests ollama models on mount", () => { + mockUseExtensionState.mockReturnValue({ + apiConfiguration: { apiProvider: providerIdentifiers.ollama, ollamaModelId: "llama3" }, + routerModels: undefined, + }) + + renderHook(() => useChatModelSelector(), { wrapper }) + + expect(mockPostMessage).toHaveBeenCalledWith(expect.objectContaining({ type: "requestOllamaModels" })) + }) + + it("uses models delivered through the ollamaModels message", async () => { + mockUseExtensionState.mockReturnValue({ + apiConfiguration: { apiProvider: providerIdentifiers.ollama, ollamaModelId: "llama3" }, + routerModels: undefined, + }) + + const { result } = renderHook(() => useChatModelSelector(), { wrapper }) + + act(() => { + emitMessage({ type: "ollamaModels", ollamaModels: { llama3: { maxTokens: 1, contextWindow: 1 } } }) + }) + + await waitFor(() => { + expect(result.current.models).not.toBeNull() + }) + expect(result.current.modelIdKey).toBe("ollamaModelId") + expect(Object.keys(result.current.models!)).toEqual(["llama3"]) + }) + + it("requests lmstudio models on mount", () => { + mockUseExtensionState.mockReturnValue({ + apiConfiguration: { apiProvider: providerIdentifiers.lmstudio, lmStudioModelId: "local-model" }, + routerModels: undefined, + }) + + renderHook(() => useChatModelSelector(), { wrapper }) + + expect(mockPostMessage).toHaveBeenCalledWith(expect.objectContaining({ type: "requestLmStudioModels" })) + }) + + it("uses models delivered through the lmStudioModels message", async () => { + mockUseExtensionState.mockReturnValue({ + apiConfiguration: { apiProvider: providerIdentifiers.lmstudio, lmStudioModelId: "local-model" }, + routerModels: undefined, + }) + + const { result } = renderHook(() => useChatModelSelector(), { wrapper }) + + act(() => { + emitMessage({ + type: "lmStudioModels", + lmStudioModels: { "local-model": { maxTokens: 1, contextWindow: 1 } }, + }) + }) + + await waitFor(() => { + expect(result.current.models).not.toBeNull() + }) + expect(result.current.modelIdKey).toBe("lmStudioModelId") + expect(Object.keys(result.current.models!)).toEqual(["local-model"]) + }) + + it("requests vsCodeLm models on mount and builds model records", async () => { + mockUseExtensionState.mockReturnValue({ + apiConfiguration: { apiProvider: providerIdentifiers.vscodeLm }, + routerModels: undefined, + }) + + const { result } = renderHook(() => useChatModelSelector(), { wrapper }) + + expect(mockPostMessage).toHaveBeenCalledWith(expect.objectContaining({ type: "requestVsCodeLmModels" })) + + act(() => { + emitMessage({ + type: "vsCodeLmModels", + vsCodeLmModels: [{ vendor: "copilot", family: "gpt-4o" }], + }) + }) + + await waitFor(() => { + expect(result.current.models).not.toBeNull() + }) + expect(Object.keys(result.current.models!)).toEqual(["copilot/gpt-4o"]) + expect(result.current.modelIdKey).toBe("vsCodeLmModelSelector") + }) + + it("transforms vsCodeLm model ids to vendor/family and formats display values", async () => { + mockUseExtensionState.mockReturnValue({ + apiConfiguration: { + apiProvider: providerIdentifiers.vscodeLm, + vsCodeLmModelSelector: { vendor: "copilot", family: "gpt-4o" }, + }, + routerModels: undefined, + }) + + const { result } = renderHook(() => useChatModelSelector(), { wrapper }) + + act(() => { + emitMessage({ + type: "vsCodeLmModels", + vsCodeLmModels: [{ vendor: "copilot", family: "gpt-4o" }], + }) + }) + + await waitFor(() => { + expect(result.current.models).not.toBeNull() + }) + + // valueTransform: "copilot/gpt-4o" -> { vendor: "copilot", family: "gpt-4o" } + expect(result.current.valueTransform!("copilot/gpt-4o")).toEqual({ vendor: "copilot", family: "gpt-4o" }) + // displayTransform: { vendor, family } -> "copilot/gpt-4o" + expect(result.current.displayTransform!({ vendor: "copilot", family: "gpt-4o" })).toBe("copilot/gpt-4o") + // displayTransform returns "" for missing values + expect(result.current.displayTransform!(undefined)).toBe("") + expect(result.current.displayTransform!({ vendor: "copilot" })).toBe("") + }) + }) + + it.each([ + [providerIdentifiers.openai, { type: "openAiModels", openAiModels: ["old-model"] }], + [providerIdentifiers.ollama, { type: "ollamaModels", ollamaModels: { "old-model": { contextWindow: 1 } } }], + [ + providerIdentifiers.lmstudio, + { type: "lmStudioModels", lmStudioModels: { "old-model": { contextWindow: 1 } } }, + ], + [ + providerIdentifiers.vscodeLm, + { type: "vsCodeLmModels", vsCodeLmModels: [{ vendor: "old", family: "model" }] }, + ], + ])("discards cached %s models when switching providers", (provider, message) => { + mockUseExtensionState.mockReturnValue({ apiConfiguration: { apiProvider: provider } }) + const { result, rerender } = renderHook(() => useChatModelSelector(), { wrapper }) + act(() => emitMessage(message)) + expect(Object.keys(result.current.models!)).toHaveLength(1) + mockUseExtensionState.mockReturnValue({ apiConfiguration: { apiProvider: providerIdentifiers.anthropic } }) + rerender() + mockUseExtensionState.mockReturnValue({ apiConfiguration: { apiProvider: provider } }) + rerender() + expect(Object.keys(result.current.models ?? {})).toHaveLength(0) + }) + + describe("retired providers", () => { + it("returns empty data for retired providers", () => { + mockUseExtensionState.mockReturnValue({ + apiConfiguration: { apiProvider: retiredProviderIdentifiers.groq as unknown as "openrouter" }, + routerModels: undefined, + }) + + const { result } = renderHook(() => useChatModelSelector(), { wrapper }) + + expect(result.current.provider).toBeUndefined() + expect(result.current.models).toBeNull() + expect(result.current.modelIdKey).toBeUndefined() + }) + }) +}) diff --git a/webview-ui/src/components/chat/hooks/useChatModelSelector.ts b/webview-ui/src/components/chat/hooks/useChatModelSelector.ts new file mode 100644 index 0000000000..327518c6a7 --- /dev/null +++ b/webview-ui/src/components/chat/hooks/useChatModelSelector.ts @@ -0,0 +1,329 @@ +import { useCallback, useEffect, useMemo, useRef, useState } from "react" +import { useEvent } from "react-use" + +import { + type ModelInfo, + type ModelRecord, + type ExtensionMessage, + type LanguageModelChatSelector, + type ProviderName, + type ProviderSettings, + getProviderDefaultModelId, + isRetiredProvider, + openAiModelInfoSaneDefaults, + providerIdentifiers, +} from "@roo-code/types" + +import { useRouterModels } from "@/components/ui/hooks/useRouterModels" +import { getStaticModelsForProvider } from "@/components/settings/utils/providerModelConfig" +import { MODELS_BY_PROVIDER } from "@/components/settings/constants" +import { useExtensionState } from "@/context/ExtensionStateContext" +import { vscode } from "@/utils/vscode" + +type ModelIdKey = keyof Pick< + ProviderSettings, + | "openRouterModelId" + | "requestyModelId" + | "unboundModelId" + | "openAiModelId" + | "litellmModelId" + | "vercelAiGatewayModelId" + | "zooGatewayModelId" + | "opencodeGoModelId" + | "kenariModelId" + | "nanoGptModelId" + | "apiModelId" + | "ollamaModelId" + | "lmStudioModelId" + | "vsCodeLmModelSelector" +> + +// Router providers: model list comes from `useRouterModels` (or the backend +// broadcast cache for litellm/poe). The map holds the config field each one +// stores its selection under, so the switch below never re-declares them. +// Mirrors the router set used by `getProviderDefaultModelId`. +const ROUTER_PROVIDERS: Partial> = { + [providerIdentifiers.openrouter]: { modelIdKey: "openRouterModelId" }, + [providerIdentifiers.requesty]: { modelIdKey: "requestyModelId" }, + [providerIdentifiers.unbound]: { modelIdKey: "unboundModelId" }, + [providerIdentifiers.litellm]: { modelIdKey: "litellmModelId", fromState: true }, + [providerIdentifiers.vercelAiGateway]: { modelIdKey: "vercelAiGatewayModelId" }, + [providerIdentifiers.zooGateway]: { modelIdKey: "zooGatewayModelId" }, + [providerIdentifiers.opencodeGo]: { modelIdKey: "opencodeGoModelId" }, + [providerIdentifiers.kenari]: { modelIdKey: "kenariModelId" }, + [providerIdentifiers.nanogpt]: { modelIdKey: "nanoGptModelId" }, + [providerIdentifiers.kimiCode]: { modelIdKey: "apiModelId" }, + [providerIdentifiers.poe]: { modelIdKey: "apiModelId", fromState: true }, +} + +// Providers whose model list is delivered through a dedicated message event +// (mirrors the settings page provider components). +const MESSAGE_BASED_PROVIDERS = new Set([ + providerIdentifiers.openai, + providerIdentifiers.ollama, + providerIdentifiers.lmstudio, + providerIdentifiers.vscodeLm, +]) + +export interface ChatModelSelectorData { + /** The provider key used to determine model source (undefined for retired providers). */ + provider: ProviderName | undefined + /** Record of available models for the current provider (null until loaded or when N/A). */ + models: Record | null + /** The configuration field key that stores the selected model for this provider. */ + modelIdKey: ModelIdKey | undefined + /** The default model id for this provider. */ + defaultModelId: string + /** Whether the model list is still loading. */ + isLoading: boolean + /** Transform a selected model id into the stored configuration value (e.g. VSCode LM selector object). */ + valueTransform?: (modelId: string) => unknown + /** Transform the stored configuration value back to a display string (e.g. VSCode LM selector). */ + displayTransform?: (value: unknown) => string +} + +/** + * Resolves the model list, storage key and defaults for the currently active + * provider so the chat input bar can render a compact model picker. + * + * The data sources mirror the settings page (`ApiOptions` and the provider + * components) so the chat selector shows the exact same models: + * - Router providers (see `ROUTER_PROVIDERS`): react-query `useRouterModels` + * request, except litellm/poe which read the backend-broadcast cache. + * - openai (OpenAI compatible), ollama, lmstudio, vscode-lm: request on mount + * and listen for the corresponding `*Models` message event. + * - Static providers: `getStaticModelsForProvider`. + * + * Default model ids come from the shared `getProviderDefaultModelId` helper so + * this hook never re-declares the provider matrix. + */ +export const useChatModelSelector = (): ChatModelSelectorData => { + const { apiConfiguration, routerModels: stateRouterModels } = useExtensionState() + + const provider = (apiConfiguration?.apiProvider || providerIdentifiers.openrouter) as ProviderName + const activeProvider = isRetiredProvider(provider) ? undefined : provider + const routerConfig = activeProvider ? ROUTER_PROVIDERS[activeProvider] : undefined + + const routerModels = useRouterModels({ + provider: activeProvider, + enabled: !!routerConfig && !routerConfig.fromState, + }) + + // Message-based providers: request the models on mount and keep the + // latest list delivered by the backend. Each request carries a unique + // `requestId` and a response is only applied when it matches the most + // recent request, so stale/out-of-order responses (for a previous provider + // or an earlier base URL) cannot clobber the current list. + const [openAiModels, setOpenAiModels] = useState([]) + const [ollamaModels, setOllamaModels] = useState({}) + const [lmStudioModels, setLmStudioModels] = useState({}) + const [vsCodeLmModels, setVsCodeLmModels] = useState([]) + // Track request completion explicitly. `openAiModels.length === 0` is a valid + // successful result (an endpoint with no models), so it alone must not keep + // the picker in a permanent loading state. + const [openAiModelsReceived, setOpenAiModelsReceived] = useState(false) + // Only one message-based provider is active at a time, so a single latest + // request id is sufficient to reject stale responses. + const latestRequestId = useRef(undefined) + + const onMessage = useCallback((event: MessageEvent) => { + const message: ExtensionMessage = event.data + // Responses without a requestId are broadcast/legacy payloads and are + // accepted; responses with a mismatching id are stale and ignored. + const isStale = (): boolean => message.requestId !== undefined && message.requestId !== latestRequestId.current + switch (message.type) { + case "openAiModels": + if (isStale()) break + setOpenAiModels(message.openAiModels ?? []) + setOpenAiModelsReceived(true) + break + case "ollamaModels": + if (isStale()) break + setOllamaModels(message.ollamaModels ?? {}) + break + case "lmStudioModels": + if (isStale()) break + setLmStudioModels(message.lmStudioModels ?? {}) + break + case "vsCodeLmModels": + if (isStale()) break + setVsCodeLmModels(message.vsCodeLmModels ?? []) + break + } + }, []) + useEvent("message", onMessage) + + useEffect(() => { + setOpenAiModels([]) + setOllamaModels({}) + setLmStudioModels({}) + setVsCodeLmModels([]) + setOpenAiModelsReceived(false) + latestRequestId.current = undefined + }, [activeProvider]) + + const serializedOpenAiHeaders = JSON.stringify(apiConfiguration?.openAiHeaders ?? {}) + + // Request models on mount when a message-based provider is active + // (mirrors Ollama.tsx / LMStudio.tsx / OpenAICompatible.tsx behaviors). + useEffect(() => { + if (!activeProvider || !MESSAGE_BASED_PROVIDERS.has(activeProvider)) { + return + } + + // The effect reruns when the OpenAI request inputs change (base URL, API key, headers) + // without changing the provider. Drop the previous list first so credentials that were + // removed or replaced don't leave stale models selectable when no replacement request is + // sent. Request-ID creation and cleanup are unchanged, preserving cancellation behavior. + if (activeProvider === providerIdentifiers.openai) { + setOpenAiModels([]) + setOpenAiModelsReceived(false) + } + + const requestId = `${activeProvider}-${Date.now()}-${Math.random().toString(36).slice(2)}` + latestRequestId.current = requestId + + switch (activeProvider) { + case providerIdentifiers.openai: + if (apiConfiguration?.openAiBaseUrl && apiConfiguration?.openAiApiKey) { + vscode.postMessage({ + type: "requestOpenAiModels", + requestId, + values: { + baseUrl: apiConfiguration.openAiBaseUrl, + apiKey: apiConfiguration.openAiApiKey, + customHeaders: {}, + openAiHeaders: JSON.parse(serializedOpenAiHeaders), + }, + }) + } + break + case providerIdentifiers.ollama: + vscode.postMessage({ type: "requestOllamaModels", requestId }) + break + case providerIdentifiers.lmstudio: + vscode.postMessage({ type: "requestLmStudioModels", requestId }) + break + case providerIdentifiers.vscodeLm: + vscode.postMessage({ type: "requestVsCodeLmModels", requestId }) + break + } + return () => { + latestRequestId.current = undefined + vscode.postMessage({ type: "cancelModelRequest", requestId }) + } + }, [activeProvider, apiConfiguration?.openAiBaseUrl, apiConfiguration?.openAiApiKey, serializedOpenAiHeaders]) + + // Resolve the model list + storage key + transforms for the active provider. + // The default id comes from the shared provider registry. + return useMemo(() => { + if (!activeProvider) { + return { + provider: undefined, + models: null, + modelIdKey: undefined, + defaultModelId: "", + isLoading: false, + } + } + + const defaultModelId = getProviderDefaultModelId(activeProvider, { + isChina: activeProvider === providerIdentifiers.zai && apiConfiguration?.zaiApiLine === "china_coding", + }) + + if (routerConfig) { + // RouterModels is keyed by dynamic/local providers only, so a narrow + // structural cast is needed to index it by the active provider name. + const source = (routerConfig.fromState ? stateRouterModels : routerModels.data) as + | Partial> + | undefined + return { + provider: activeProvider, + models: source?.[activeProvider] ?? null, + modelIdKey: routerConfig.modelIdKey, + defaultModelId, + isLoading: routerConfig.fromState ? false : routerModels.isLoading, + } + } + + let models: Record | null = null + let modelIdKey: ModelIdKey | undefined = undefined + let valueTransform: ((modelId: string) => unknown) | undefined + let displayTransform: ((value: unknown) => string) | undefined + let isLoading = false + + switch (activeProvider) { + case providerIdentifiers.openai: + // OpenAI Compatible: the list is fetched from the baseUrl via + // `requestOpenAiModels` and delivered through `openAiModels`. + models = + Object.keys(openAiModels).length > 0 + ? Object.fromEntries(openAiModels.map((item) => [item, openAiModelInfoSaneDefaults])) + : null + modelIdKey = "openAiModelId" + isLoading = + !!apiConfiguration?.openAiBaseUrl && !!apiConfiguration?.openAiApiKey && !openAiModelsReceived + break + case providerIdentifiers.ollama: + models = Object.keys(ollamaModels).length > 0 ? ollamaModels : null + modelIdKey = "ollamaModelId" + break + case providerIdentifiers.lmstudio: + models = Object.keys(lmStudioModels).length > 0 ? lmStudioModels : null + modelIdKey = "lmStudioModelId" + break + case providerIdentifiers.vscodeLm: + models = vsCodeLmModels.reduce( + (acc, model) => { + const modelId = `${model.vendor}/${model.family}` + acc[modelId] = { + maxTokens: 0, + contextWindow: 0, + supportsPromptCache: false, + description: `${model.vendor} - ${model.family}`, + } + return acc + }, + {} as Record, + ) + modelIdKey = "vsCodeLmModelSelector" + valueTransform = (modelId) => { + const [vendor, family] = modelId.split("/") + return { vendor, family } + } + displayTransform = (value) => { + if (!value) return "" + const selector = value as { vendor?: string; family?: string } + return selector.vendor && selector.family ? `${selector.vendor}/${selector.family}` : "" + } + break + default: + // Static models providers (anthropic, bedrock, gemini, etc.). + models = MODELS_BY_PROVIDER[activeProvider] + ? getStaticModelsForProvider(activeProvider, undefined, apiConfiguration) + : null + modelIdKey = "apiModelId" + } + + return { + provider: activeProvider, + models, + modelIdKey, + defaultModelId, + isLoading, + valueTransform, + displayTransform, + } + }, [ + activeProvider, + routerConfig, + apiConfiguration, + routerModels, + stateRouterModels, + openAiModels, + openAiModelsReceived, + ollamaModels, + lmStudioModels, + vsCodeLmModels, + ]) +} diff --git a/webview-ui/src/components/settings/ApiConfigManager.tsx b/webview-ui/src/components/settings/ApiConfigManager.tsx index d57cd004a9..ac9d038cc0 100644 --- a/webview-ui/src/components/settings/ApiConfigManager.tsx +++ b/webview-ui/src/components/settings/ApiConfigManager.tsx @@ -17,6 +17,7 @@ import { } from "@/components/ui" interface ApiConfigManagerProps { + disabled?: boolean currentApiConfigName?: string listApiConfigMeta?: ProviderSettingsEntry[] organizationAllowList?: OrganizationAllowList @@ -27,6 +28,7 @@ interface ApiConfigManagerProps { } const ApiConfigManager = ({ + disabled = false, currentApiConfigName = "", listApiConfigMeta = [], organizationAllowList, @@ -119,16 +121,18 @@ const ApiConfigManager = ({ }, [currentApiConfigName]) const handleSelectConfig = (configName: string) => { - if (!configName) return + if (disabled || !configName) return onSelectConfig(configName) } const handleAdd = () => { + if (disabled) return resetCreateState() setIsCreating(true) } const handleStartRename = () => { + if (disabled) return setIsRenaming(true) setInputValue(currentApiConfigName || "") setError(null) @@ -139,6 +143,7 @@ const ApiConfigManager = ({ } const handleSave = () => { + if (disabled) return const trimmedValue = inputValue.trim() const error = validateName(trimmedValue, false) @@ -159,6 +164,7 @@ const ApiConfigManager = ({ } const handleNewProfileSave = () => { + if (disabled) return const trimmedValue = newProfileName.trim() const error = validateName(trimmedValue, true) @@ -172,7 +178,7 @@ const ApiConfigManager = ({ } const handleDelete = () => { - if (!currentApiConfigName || !listApiConfigMeta || listApiConfigMeta.length <= 1) return + if (disabled || !currentApiConfigName || !listApiConfigMeta || listApiConfigMeta.length <= 1) return // Let the extension handle both deletion and selection. onDeleteConfig(currentApiConfigName) @@ -188,6 +194,7 @@ const ApiConfigManager = ({
{ @@ -209,7 +216,7 @@ const ApiConfigManager = ({ @@ -269,6 +282,7 @@ const ApiConfigManager = ({ @@ -313,6 +327,7 @@ const ApiConfigManager = ({ {t("settings:providers.newProfile")} { @@ -343,7 +358,7 @@ const ApiConfigManager = ({ @@ -773,6 +832,7 @@ const SettingsView = forwardRef(({ onDone, t
@@ -783,20 +843,14 @@ const SettingsView = forwardRef(({ onDone, t onDeleteConfig={(configName: string) => vscode.postMessage({ type: "deleteApiConfiguration", text: configName }) } - onRenameConfig={(oldName: string, newName: string) => { - vscode.postMessage({ + onRenameConfig={(oldName: string, newName: string) => + postApiConfiguration({ type: "renameApiConfiguration", values: { oldName, newName }, - apiConfiguration, }) - prevApiConfigName.current = newName - }} + } onUpsertConfig={(configName: string) => - vscode.postMessage({ - type: "upsertApiConfiguration", - text: configName, - apiConfiguration, - }) + postApiConfiguration({ type: "upsertApiConfiguration", text: configName }) } /> { const getRenameForm = () => screen.getByTestId("rename-form") const getDialogContent = () => screen.getByTestId("dialog-content") + it("disables profile actions while a save is pending", () => { + render() + for (const id of ["add-profile-button", "rename-profile-button", "delete-profile-button"]) { + expect(screen.getByTestId(id)).toBeDisabled() + fireEvent.click(screen.getByTestId(id)) + } + fireEvent.change(screen.getByTestId("select-component"), { target: { value: "Another Config" } }) + expect(mockOnSelectConfig).not.toHaveBeenCalled() + expect(mockOnDeleteConfig).not.toHaveBeenCalled() + expect(mockOnRenameConfig).not.toHaveBeenCalled() + expect(mockOnUpsertConfig).not.toHaveBeenCalled() + }) + + it.each(["rename", "create"])( + "retains a %s draft while disabled and submits after the pending save finishes", + (action) => { + const view = render() + fireEvent.click(screen.getByTestId(action === "rename" ? "rename-profile-button" : "add-profile-button")) + const input = + action === "rename" + ? screen.getByPlaceholderText("settings:providers.enterNewName") + : screen.getByTestId("new-profile-input") + fireEvent.input(input, { target: { value: "New Profile" } }) + view.rerender() + const submit = screen.getByTestId(action === "rename" ? "save-rename-button" : "create-profile-button") + expect(submit).toBeDisabled() + fireEvent.keyDown(input, { key: "Enter" }) + expect(mockOnRenameConfig).not.toHaveBeenCalled() + expect(mockOnUpsertConfig).not.toHaveBeenCalled() + expect(input).toHaveValue("New Profile") + view.rerender() + fireEvent.click(submit) + expect(action === "rename" ? mockOnRenameConfig : mockOnUpsertConfig).toHaveBeenCalledTimes(1) + if (action === "rename") expect(mockOnRenameConfig).toHaveBeenCalledWith("Default Config", "New Profile") + else expect(mockOnUpsertConfig).toHaveBeenCalledWith("New Profile") + }, + ) + it("opens new profile dialog when clicking add button", () => { render() diff --git a/webview-ui/src/components/settings/__tests__/SettingsView.change-detection.spec.tsx b/webview-ui/src/components/settings/__tests__/SettingsView.change-detection.spec.tsx index 6704aa1cf0..83c09575a5 100644 --- a/webview-ui/src/components/settings/__tests__/SettingsView.change-detection.spec.tsx +++ b/webview-ui/src/components/settings/__tests__/SettingsView.change-detection.spec.tsx @@ -508,6 +508,8 @@ describe("SettingsView - Change Detection Fix", () => { fireEvent.click(screen.getByTestId("save-button")) expect(mockPostMessage).toHaveBeenCalledWith({ type: "upsertApiConfiguration", + requestId: expect.any(String), + apiConfigurationBaseline: expect.any(Object), text: "default", apiConfiguration: expect.objectContaining({ reasoningEffort: "high" }), }) @@ -560,6 +562,8 @@ describe("SettingsView - Change Detection Fix", () => { fireEvent.click(screen.getByTestId("save-button")) expect(mockPostMessage).toHaveBeenCalledWith({ type: "upsertApiConfiguration", + requestId: expect.any(String), + apiConfigurationBaseline: expect.any(Object), text: "default", apiConfiguration: expect.objectContaining({ apiProvider: providerIdentifiers.baseten, @@ -592,12 +596,23 @@ describe("SettingsView - Change Detection Fix", () => { }) expect(screen.getByTestId("provider-value")).toHaveTextContent("deepseek") + const saved = mockPostMessage.mock.calls + .map(([message]) => message) + .find((message) => message.type === "upsertApiConfiguration")! + fireEvent( + window, + new MessageEvent("message", { + data: { type: "apiConfigurationSaved", requestId: saved.requestId, success: true }, + }), + ) mockPostMessage.mockClear() fireEvent.click(screen.getByTestId("save-button")) expect(mockPostMessage).toHaveBeenCalledWith({ type: "upsertApiConfiguration", + requestId: expect.any(String), + apiConfigurationBaseline: expect.any(Object), text: "default", apiConfiguration: expect.objectContaining({ apiProvider: providerIdentifiers.deepseek, diff --git a/webview-ui/src/components/settings/__tests__/SettingsView.spec.tsx b/webview-ui/src/components/settings/__tests__/SettingsView.spec.tsx index 97b390f6f4..70e9a23ddd 100644 --- a/webview-ui/src/components/settings/__tests__/SettingsView.spec.tsx +++ b/webview-ui/src/components/settings/__tests__/SettingsView.spec.tsx @@ -447,6 +447,17 @@ describe("SettingsView - Sound Settings", () => { }), ) + const saved = vi + .mocked(vscode.postMessage) + .mock.calls.map(([message]) => message) + .find((message) => message.type === "upsertApiConfiguration")! + fireEvent( + window, + new MessageEvent("message", { + data: { type: "apiConfigurationSaved", requestId: saved.requestId, success: true }, + }), + ) + // Reset clears the override; it is persisted as null (not undefined). fireEvent.click(within(getSettingsContent()).getByTestId("chat-font-size-reset")) fireEvent.click(screen.getByTestId("save-button")) diff --git a/webview-ui/src/components/settings/__tests__/SettingsView.unsaved-changes.spec.tsx b/webview-ui/src/components/settings/__tests__/SettingsView.unsaved-changes.spec.tsx index 575da5a833..ddb51da776 100644 --- a/webview-ui/src/components/settings/__tests__/SettingsView.unsaved-changes.spec.tsx +++ b/webview-ui/src/components/settings/__tests__/SettingsView.unsaved-changes.spec.tsx @@ -1,5 +1,5 @@ import { providerIdentifiers, openAiModelInfoSaneDefaults, type ProviderSettings } from "@roo-code/types" -import { screen, fireEvent, waitFor } from "@testing-library/react" +import { act, screen, fireEvent, waitFor } from "@testing-library/react" import { renderWithExtensionState } from "@/utils/test-utils" import { vi, describe, it, expect, beforeEach } from "vitest" @@ -249,6 +249,7 @@ vi.mock("../SettingsSearch", () => ({ import { useExtensionState } from "@src/context/ExtensionStateContext" import ApiOptions from "../ApiOptions" +import ApiConfigManager from "../ApiConfigManager" describe("SettingsView - Unsaved Changes Detection", () => { let queryClient: QueryClient @@ -320,6 +321,9 @@ describe("SettingsView - Unsaved Changes Detection", () => { beforeEach(() => { vi.clearAllMocks() + vi.mocked(ApiConfigManager).mockImplementation(() => ( +
ApiConfigManager
+ )) // Reset the ApiOptions mock to its default implementation vi.mocked(ApiOptions).mockImplementation(() => { // Don't do anything with props, just render a div @@ -674,7 +678,9 @@ describe("SettingsView - Unsaved Changes Detection", () => { expect(postMessage).toHaveBeenCalledWith({ type: "upsertApiConfiguration", + requestId: expect.any(String), text: "default", + apiConfigurationBaseline: liveApiConfiguration, apiConfiguration: { apiProvider: providerIdentifiers.nanogpt, nanoGptApiKey: "unsaved-key", @@ -682,6 +688,124 @@ describe("SettingsView - Unsaved Changes Detection", () => { nanoGptRoutingPreference: "tools", }, }) + const firstSave = postMessage.mock.calls + .map(([message]) => message) + .find((message) => message.type === "upsertApiConfiguration")! + fireEvent( + window, + new MessageEvent("message", { + data: { type: "apiConfigurationSaved", requestId: "unrelated", success: true }, + }), + ) + expect(screen.getByTestId("save-button")).toBeDisabled() + fireEvent( + window, + new MessageEvent("message", { + data: { type: "apiConfigurationSaved", requestId: firstSave.requestId, success: false }, + }), + ) + expect(screen.getByTestId("save-button")).toBeEnabled() + postMessage.mockClear() + fireEvent.click(screen.getByTestId("save-button")) + const retry = postMessage.mock.calls + .map(([message]) => message) + .find((message) => message.type === "upsertApiConfiguration")! + expect(retry.apiConfigurationBaseline).toEqual(liveApiConfiguration) + expect(retry.apiConfiguration).toEqual(firstSave.apiConfiguration) + fireEvent( + window, + new MessageEvent("message", { + data: { type: "apiConfigurationSaved", requestId: retry.requestId, success: true }, + }), + ) + postMessage.mockClear() + fireEvent.change(screen.getByTestId("cached-nanogpt-model"), { target: { value: "openai/original" } }) + fireEvent.click(screen.getByTestId("save-button")) + expect(postMessage).toHaveBeenCalledWith( + expect.objectContaining({ + type: "upsertApiConfiguration", + apiConfigurationBaseline: expect.objectContaining({ nanoGptModelId: "openai/next" }), + apiConfiguration: expect.objectContaining({ nanoGptModelId: "openai/original" }), + }), + ) + }) + + it.each(["upsert", "rename"])("advances the baseline after a confirmed profile-manager %s", (operation) => { + vi.mocked(ApiOptions).mockImplementation(({ apiConfiguration, setApiConfigurationField }) => ( + setApiConfigurationField("modelTemperature", Number(event.target.value))} + /> + )) + vi.mocked(ApiConfigManager).mockImplementation(({ onRenameConfig, onUpsertConfig }) => ( + + )) + renderWithExtensionState(, { queryClient }) + fireEvent.change(screen.getByTestId("profile-temperature"), { target: { value: "0.5" } }) + fireEvent.click(screen.getByTestId("persist-profile")) + const first = postMessage.mock.calls + .map(([message]) => message) + .find( + (message) => + message.type === (operation === "rename" ? "renameApiConfiguration" : "upsertApiConfiguration"), + )! + expect(first.requestId).toEqual(expect.any(String)) + fireEvent( + window, + new MessageEvent("message", { + data: { type: "apiConfigurationSaved", requestId: first.requestId, success: true }, + }), + ) + postMessage.mockClear() + fireEvent.change(screen.getByTestId("profile-temperature"), { target: { value: "0.7" } }) + fireEvent.click(screen.getByTestId("save-button")) + expect(postMessage).toHaveBeenCalledWith( + expect.objectContaining({ + type: "upsertApiConfiguration", + apiConfigurationBaseline: expect.objectContaining({ modelTemperature: 0.5 }), + apiConfiguration: expect.objectContaining({ modelTemperature: 0.7 }), + }), + ) + }) + + it("re-enables Save with the original baseline when its acknowledgement is lost", () => { + vi.mocked(ApiOptions).mockImplementation(({ setApiConfigurationField }) => ( + + )) + const view = renderWithExtensionState(, { queryClient }) + fireEvent.click(screen.getByTestId("edit-before-timeout")) + vi.useFakeTimers() + try { + fireEvent.click(screen.getByTestId("save-button")) + const first = postMessage.mock.calls + .map(([message]) => message) + .find((message) => message.type === "upsertApiConfiguration")! + expect(screen.getByTestId("save-button")).toBeDisabled() + expect(vi.mocked(ApiConfigManager).mock.lastCall?.[0].disabled).toBe(true) + act(() => vi.advanceTimersByTime(120_000)) + expect(vi.mocked(ApiConfigManager).mock.lastCall?.[0].disabled).toBe(false) + expect(screen.getByTestId("save-button")).toBeEnabled() + postMessage.mockClear() + fireEvent.click(screen.getByTestId("save-button")) + const retry = postMessage.mock.calls + .map(([message]) => message) + .find((message) => message.type === "upsertApiConfiguration")! + expect(retry.requestId).not.toBe(first.requestId) + expect(retry.apiConfigurationBaseline).toEqual(first.apiConfigurationBaseline) + expect(retry.apiConfiguration).toEqual(first.apiConfiguration) + } finally { + view.unmount() + vi.useRealTimers() + } }) it("keeps OpenAI-compatible reasoning edits cached until Save despite a live state refresh", async () => { @@ -714,7 +838,7 @@ describe("SettingsView - Unsaved Changes Detection", () => { vi.mocked(useExtensionState, { partial: true }).mockReturnValue({ ...defaultExtensionState, - apiConfiguration: { ...configuration }, + apiConfiguration: { ...configuration, openAiModelId: "concurrent-model" }, }) view.rerender() @@ -726,7 +850,9 @@ describe("SettingsView - Unsaved Changes Detection", () => { fireEvent.click(screen.getByTestId("save-button")) expect(postMessage).toHaveBeenCalledWith({ type: "upsertApiConfiguration", + requestId: expect.any(String), text: "default", + apiConfigurationBaseline: configuration, apiConfiguration: { ...configuration, enableReasoningEffort: true, diff --git a/webview-ui/src/components/ui/hooks/__tests__/useRouterModels.spec.tsx b/webview-ui/src/components/ui/hooks/__tests__/useRouterModels.spec.tsx new file mode 100644 index 0000000000..574ecefbc8 --- /dev/null +++ b/webview-ui/src/components/ui/hooks/__tests__/useRouterModels.spec.tsx @@ -0,0 +1,111 @@ +import { providerIdentifiers } from "@roo-code/types" +import React from "react" +import { renderHook } from "@testing-library/react" +import { QueryClient, QueryClientProvider } from "@tanstack/react-query" +import { fetchRouterModels, useRouterModels } from "../useRouterModels" +import { vscode } from "@/utils/vscode" + +vi.mock("@/utils/vscode", () => ({ vscode: { postMessage: vi.fn() } })) + +describe("router request cleanup", () => { + beforeEach(() => { + vi.useFakeTimers() + vi.clearAllMocks() + }) + afterEach(() => { + vi.useRealTimers() + vi.restoreAllMocks() + }) + + it("rejects an already-aborted signal without starting any work", async () => { + const add = vi.spyOn(window, "addEventListener") + const timeout = vi.spyOn(globalThis, "setTimeout") + const controller = new AbortController() + const reason = new Error("Model request already cancelled") + controller.abort(reason) + + await expect(fetchRouterModels(providerIdentifiers.openrouter, controller.signal)).rejects.toBe(reason) + + expect(add).not.toHaveBeenCalled() + expect(timeout).not.toHaveBeenCalled() + expect(vi.getTimerCount()).toBe(0) + expect(vscode.postMessage).not.toHaveBeenCalled() + }) + + it("removes its listener and timer and cancels backend work on abort", async () => { + const add = vi.spyOn(window, "addEventListener") + const remove = vi.spyOn(window, "removeEventListener") + const controller = new AbortController() + const pending = fetchRouterModels(providerIdentifiers.openrouter, controller.signal) + const rejected = expect(pending).rejects.toBeDefined() + const request = vi.mocked(vscode.postMessage).mock.calls[0][0] + const listener = add.mock.calls.find(([type]) => type === "message")![1] + expect(vi.getTimerCount()).toBe(1) + controller.abort() + await rejected + expect(remove).toHaveBeenCalledWith("message", listener) + expect(vi.getTimerCount()).toBe(0) + expect(vscode.postMessage).toHaveBeenLastCalledWith({ + type: "cancelModelRequest", + requestId: request.requestId, + }) + }) + + it("times out, cancels the matching request, and releases listeners and timers", async () => { + const add = vi.spyOn(window, "addEventListener") + const remove = vi.spyOn(window, "removeEventListener") + const controller = new AbortController() + const addAbort = vi.spyOn(controller.signal, "addEventListener") + const removeAbort = vi.spyOn(controller.signal, "removeEventListener") + const pending = fetchRouterModels(providerIdentifiers.openrouter, controller.signal) + const rejected = expect(pending).rejects.toThrow("Router models request timed out") + const request = vi.mocked(vscode.postMessage).mock.calls[0][0] + const listener = add.mock.calls.find(([type]) => type === "message")![1] + await vi.advanceTimersByTimeAsync(10_000) + await rejected + expect(remove).toHaveBeenCalledWith("message", listener) + expect(removeAbort).toHaveBeenCalledWith("abort", addAbort.mock.calls[0][1]) + expect(vi.getTimerCount()).toBe(0) + expect(vscode.postMessage).toHaveBeenLastCalledWith({ + type: "cancelModelRequest", + requestId: request.requestId, + }) + controller.abort() + expect(vscode.postMessage).toHaveBeenCalledTimes(2) + }) + + it("aborts a pending query when its last observer unmounts", () => { + const client = new QueryClient({ defaultOptions: { queries: { retry: false, gcTime: Infinity } } }) + const remove = vi.spyOn(window, "removeEventListener") + const { unmount } = renderHook(() => useRouterModels({ provider: providerIdentifiers.openrouter }), { + wrapper: ({ children }) => {children}, + }) + const request = vi.mocked(vscode.postMessage).mock.calls[0][0] + unmount() + expect(vscode.postMessage).toHaveBeenLastCalledWith({ + type: "cancelModelRequest", + requestId: request.requestId, + }) + expect(remove).toHaveBeenCalledWith("message", expect.any(Function)) + expect(vi.getTimerCount()).toBe(0) + client.clear() + }) + + it("cleans up after a matching response and ignores another request's response", async () => { + const controller = new AbortController() + const pending = fetchRouterModels(providerIdentifiers.openrouter, controller.signal) + const request = vi.mocked(vscode.postMessage).mock.calls[0][0] + const response = { + type: "routerModels", + values: { provider: providerIdentifiers.openrouter }, + routerModels: { openrouter: {} }, + } + window.dispatchEvent(new MessageEvent("message", { data: { ...response, requestId: "stale" } })) + expect(vi.getTimerCount()).toBe(1) + window.dispatchEvent(new MessageEvent("message", { data: { ...response, requestId: request.requestId } })) + expect(await pending).toEqual(response.routerModels) + expect(vi.getTimerCount()).toBe(0) + controller.abort() + expect(vscode.postMessage).toHaveBeenCalledTimes(1) + }) +}) diff --git a/webview-ui/src/components/ui/hooks/useRouterModels.ts b/webview-ui/src/components/ui/hooks/useRouterModels.ts index a7d6b36725..9205f952ca 100644 --- a/webview-ui/src/components/ui/hooks/useRouterModels.ts +++ b/webview-ui/src/components/ui/hooks/useRouterModels.ts @@ -14,9 +14,21 @@ type UseRouterModelsOptions = { enabled?: boolean // gate fetching entirely } -export const fetchRouterModels = async (provider?: string) => +export const fetchRouterModels = async (provider?: string, signal?: AbortSignal) => new Promise((resolve, reject) => { + const requestId = crypto.randomUUID() + if (signal?.aborted) { + reject(signal.reason) + return + } + const abort = () => { + cleanup() + vscode.postMessage({ type: "cancelModelRequest", requestId }) + reject(signal?.reason ?? new Error("Router models request aborted")) + } const cleanup = () => { + clearTimeout(timeout) + signal?.removeEventListener("abort", abort) if (typeof window !== "undefined") { window.removeEventListener("message", handler) } @@ -24,6 +36,7 @@ export const fetchRouterModels = async (provider?: string) => const timeout = setTimeout(() => { cleanup() + vscode.postMessage({ type: "cancelModelRequest", requestId }) reject(new Error("Router models request timed out")) }, 10000) @@ -34,7 +47,7 @@ export const fetchRouterModels = async (provider?: string) => const msgProvider = message?.values?.provider as string | undefined // Verify response matches request - if (provider !== msgProvider) { + if (provider !== msgProvider || (message.requestId !== undefined && message.requestId !== requestId)) { // Not our response; ignore and wait for the matching one return } @@ -50,11 +63,12 @@ export const fetchRouterModels = async (provider?: string) => } } + signal?.addEventListener("abort", abort, { once: true }) window.addEventListener("message", handler) if (provider) { - vscode.postMessage({ type: RouterModelsMessageType.requestRouterModels, values: { provider } }) + vscode.postMessage({ type: RouterModelsMessageType.requestRouterModels, requestId, values: { provider } }) } else { - vscode.postMessage({ type: RouterModelsMessageType.requestRouterModels }) + vscode.postMessage({ type: RouterModelsMessageType.requestRouterModels, requestId }) } }) @@ -62,7 +76,7 @@ export const useRouterModels = (opts: UseRouterModelsOptions = {}) => { const provider = opts.provider || undefined return useQuery({ queryKey: [RouterModelsMessageType.routerModels, provider || allRouterModelsProvider], - queryFn: () => fetchRouterModels(provider), + queryFn: ({ signal }) => fetchRouterModels(provider, signal), enabled: opts.enabled !== false, }) } diff --git a/webview-ui/src/i18n/locales/ca/chat.json b/webview-ui/src/i18n/locales/ca/chat.json index 7779a0a5cb..5763a76ba3 100644 --- a/webview-ui/src/i18n/locales/ca/chat.json +++ b/webview-ui/src/i18n/locales/ca/chat.json @@ -1,5 +1,9 @@ { "greeting": "Benvingut a Zoo Code!", + "clearSearch": "Neteja la cerca", + "searchModel": "Cerca models...", + "modelListEmpty": "Aquest proveïdor no té models disponibles", + "useCustomModel": "Utilitza el model personalitzat: {{modelId}}", "task": { "title": "Tasca", "expand": "Expandir tasca", diff --git a/webview-ui/src/i18n/locales/de/chat.json b/webview-ui/src/i18n/locales/de/chat.json index 45f9d34c3b..84dac08695 100644 --- a/webview-ui/src/i18n/locales/de/chat.json +++ b/webview-ui/src/i18n/locales/de/chat.json @@ -1,5 +1,9 @@ { "greeting": "Willkommen bei Zoo Code!", + "clearSearch": "Suche löschen", + "searchModel": "Modelle durchsuchen...", + "modelListEmpty": "Für diesen Anbieter sind keine Modelle verfügbar", + "useCustomModel": "Benutzerdefiniertes Modell verwenden: {{modelId}}", "task": { "title": "Aufgabe", "expand": "Aufgabe erweitern", diff --git a/webview-ui/src/i18n/locales/en/chat.json b/webview-ui/src/i18n/locales/en/chat.json index 460e94413e..269699d634 100644 --- a/webview-ui/src/i18n/locales/en/chat.json +++ b/webview-ui/src/i18n/locales/en/chat.json @@ -1,5 +1,9 @@ { "greeting": "Welcome to Zoo Code!", + "clearSearch": "Clear search", + "searchModel": "Search models...", + "modelListEmpty": "No models available for this provider", + "useCustomModel": "Use custom model: {{modelId}}", "task": { "title": "Task", "expand": "Expand task", diff --git a/webview-ui/src/i18n/locales/es/chat.json b/webview-ui/src/i18n/locales/es/chat.json index 9f2727493f..2970e3f283 100644 --- a/webview-ui/src/i18n/locales/es/chat.json +++ b/webview-ui/src/i18n/locales/es/chat.json @@ -1,5 +1,9 @@ { "greeting": "¡Bienvenido a Zoo Code!", + "clearSearch": "Limpiar búsqueda", + "searchModel": "Buscar modelos...", + "modelListEmpty": "No hay modelos disponibles para este proveedor", + "useCustomModel": "Usar modelo personalizado: {{modelId}}", "task": { "title": "Tarea", "expand": "Expandir tarea", diff --git a/webview-ui/src/i18n/locales/fr/chat.json b/webview-ui/src/i18n/locales/fr/chat.json index 2b4cf8c557..1e40b80c4c 100644 --- a/webview-ui/src/i18n/locales/fr/chat.json +++ b/webview-ui/src/i18n/locales/fr/chat.json @@ -1,5 +1,9 @@ { "greeting": "Bienvenue sur Zoo Code !", + "clearSearch": "Effacer la recherche", + "searchModel": "Rechercher des modèles...", + "modelListEmpty": "Aucun modèle disponible pour ce fournisseur", + "useCustomModel": "Utiliser un modèle personnalisé : {{modelId}}", "task": { "title": "Tâche", "expand": "Développer la tâche", diff --git a/webview-ui/src/i18n/locales/hi/chat.json b/webview-ui/src/i18n/locales/hi/chat.json index ebbd679434..dd9046c762 100644 --- a/webview-ui/src/i18n/locales/hi/chat.json +++ b/webview-ui/src/i18n/locales/hi/chat.json @@ -1,5 +1,9 @@ { "greeting": "Zoo Code में आपका स्वागत है!", + "clearSearch": "खोज साफ़ करें", + "searchModel": "मॉडल खोजें...", + "modelListEmpty": "इस प्रदाता के लिए कोई मॉडल उपलब्ध नहीं है", + "useCustomModel": "कस्टम मॉडल का उपयोग करें: {{modelId}}", "task": { "title": "कार्य", "expand": "कार्य विस्तृत करें", diff --git a/webview-ui/src/i18n/locales/id/chat.json b/webview-ui/src/i18n/locales/id/chat.json index d097894162..60307938fd 100644 --- a/webview-ui/src/i18n/locales/id/chat.json +++ b/webview-ui/src/i18n/locales/id/chat.json @@ -1,5 +1,9 @@ { "greeting": "Selamat datang di Zoo Code!", + "clearSearch": "Hapus pencarian", + "searchModel": "Cari model...", + "modelListEmpty": "Tidak ada model yang tersedia untuk penyedia ini", + "useCustomModel": "Gunakan model kustom: {{modelId}}", "task": { "title": "Tugas", "expand": "Perluas tugas", diff --git a/webview-ui/src/i18n/locales/it/chat.json b/webview-ui/src/i18n/locales/it/chat.json index e99c76e702..26f84e574d 100644 --- a/webview-ui/src/i18n/locales/it/chat.json +++ b/webview-ui/src/i18n/locales/it/chat.json @@ -1,5 +1,9 @@ { "greeting": "Benvenuto in Zoo Code!", + "clearSearch": "Cancella ricerca", + "searchModel": "Cerca modelli...", + "modelListEmpty": "Nessun modello disponibile per questo provider", + "useCustomModel": "Usa modello personalizzato: {{modelId}}", "task": { "title": "Attività", "expand": "Espandi attività", diff --git a/webview-ui/src/i18n/locales/ja/chat.json b/webview-ui/src/i18n/locales/ja/chat.json index 5ffd5c873c..dfe66eca53 100644 --- a/webview-ui/src/i18n/locales/ja/chat.json +++ b/webview-ui/src/i18n/locales/ja/chat.json @@ -1,5 +1,9 @@ { "greeting": "Zoo Code へようこそ!", + "clearSearch": "検索をクリア", + "searchModel": "モデルを検索...", + "modelListEmpty": "このプロバイダーで利用可能なモデルはありません", + "useCustomModel": "カスタムモデルを使用: {{modelId}}", "task": { "title": "タスク", "expand": "タスクを展開", diff --git a/webview-ui/src/i18n/locales/ko/chat.json b/webview-ui/src/i18n/locales/ko/chat.json index 313d3b4722..207bcfdc78 100644 --- a/webview-ui/src/i18n/locales/ko/chat.json +++ b/webview-ui/src/i18n/locales/ko/chat.json @@ -1,5 +1,9 @@ { "greeting": "Zoo Code에 오신 것을 환영합니다!", + "clearSearch": "검색 지우기", + "searchModel": "모델 검색...", + "modelListEmpty": "이 공급자에 사용 가능한 모델이 없습니다", + "useCustomModel": "사용자 지정 모델 사용: {{modelId}}", "task": { "title": "작업", "expand": "작업 펼치기", diff --git a/webview-ui/src/i18n/locales/nl/chat.json b/webview-ui/src/i18n/locales/nl/chat.json index 1e091996c6..7838130444 100644 --- a/webview-ui/src/i18n/locales/nl/chat.json +++ b/webview-ui/src/i18n/locales/nl/chat.json @@ -1,5 +1,9 @@ { "greeting": "Welkom bij Zoo Code!", + "clearSearch": "Zoekopdracht wissen", + "searchModel": "Modellen zoeken...", + "modelListEmpty": "Geen modellen beschikbaar voor deze provider", + "useCustomModel": "Aangepast model gebruiken: {{modelId}}", "task": { "title": "Taak", "expand": "Taak uitvouwen", diff --git a/webview-ui/src/i18n/locales/pl/chat.json b/webview-ui/src/i18n/locales/pl/chat.json index 76ff83065b..47f95a87b0 100644 --- a/webview-ui/src/i18n/locales/pl/chat.json +++ b/webview-ui/src/i18n/locales/pl/chat.json @@ -1,5 +1,9 @@ { "greeting": "Witamy w Zoo Code!", + "clearSearch": "Wyczyść wyszukiwanie", + "searchModel": "Szukaj modeli...", + "modelListEmpty": "Brak dostępnych modeli dla tego dostawcy", + "useCustomModel": "Użyj niestandardowego modelu: {{modelId}}", "task": { "title": "Zadanie", "expand": "Rozwiń zadanie", diff --git a/webview-ui/src/i18n/locales/pt-BR/chat.json b/webview-ui/src/i18n/locales/pt-BR/chat.json index 5884e3424b..a7658fb4f1 100644 --- a/webview-ui/src/i18n/locales/pt-BR/chat.json +++ b/webview-ui/src/i18n/locales/pt-BR/chat.json @@ -1,5 +1,9 @@ { "greeting": "Bem-vindo ao Zoo Code!", + "clearSearch": "Limpar pesquisa", + "searchModel": "Pesquisar modelos...", + "modelListEmpty": "Nenhum modelo disponível para este provedor", + "useCustomModel": "Usar modelo personalizado: {{modelId}}", "task": { "title": "Tarefa", "expand": "Expandir tarefa", diff --git a/webview-ui/src/i18n/locales/ru/chat.json b/webview-ui/src/i18n/locales/ru/chat.json index 8de32b56fe..0a75e44a25 100644 --- a/webview-ui/src/i18n/locales/ru/chat.json +++ b/webview-ui/src/i18n/locales/ru/chat.json @@ -1,5 +1,9 @@ { "greeting": "Добро пожаловать в Zoo Code!", + "clearSearch": "Очистить поиск", + "searchModel": "Поиск моделей...", + "modelListEmpty": "Нет доступных моделей для этого провайдера", + "useCustomModel": "Использовать пользовательскую модель: {{modelId}}", "task": { "title": "Задача", "expand": "Развернуть задачу", diff --git a/webview-ui/src/i18n/locales/tr/chat.json b/webview-ui/src/i18n/locales/tr/chat.json index 0cb4576ee0..5c65497ebf 100644 --- a/webview-ui/src/i18n/locales/tr/chat.json +++ b/webview-ui/src/i18n/locales/tr/chat.json @@ -1,5 +1,9 @@ { "greeting": "Zoo Code'a hoş geldin!", + "clearSearch": "Aramayı temizle", + "searchModel": "Modelleri ara...", + "modelListEmpty": "Bu sağlayıcı için kullanılabilir model yok", + "useCustomModel": "Özel model kullan: {{modelId}}", "task": { "title": "Görev", "expand": "Görevi genişlet", diff --git a/webview-ui/src/i18n/locales/vi/chat.json b/webview-ui/src/i18n/locales/vi/chat.json index 5a3965a5c1..cd6a339bda 100644 --- a/webview-ui/src/i18n/locales/vi/chat.json +++ b/webview-ui/src/i18n/locales/vi/chat.json @@ -1,5 +1,9 @@ { "greeting": "Chào mừng đến với Zoo Code!", + "clearSearch": "Xóa tìm kiếm", + "searchModel": "Tìm kiếm mô hình...", + "modelListEmpty": "Không có mô hình nào khả dụng cho nhà cung cấp này", + "useCustomModel": "Sử dụng mô hình tùy chỉnh: {{modelId}}", "task": { "title": "Nhiệm vụ", "expand": "Mở rộng nhiệm vụ", diff --git a/webview-ui/src/i18n/locales/zh-CN/chat.json b/webview-ui/src/i18n/locales/zh-CN/chat.json index 982220c2ea..285c95f17d 100644 --- a/webview-ui/src/i18n/locales/zh-CN/chat.json +++ b/webview-ui/src/i18n/locales/zh-CN/chat.json @@ -1,5 +1,9 @@ { "greeting": "欢迎使用 Zoo Code!", + "clearSearch": "清除搜索", + "searchModel": "搜索模型...", + "modelListEmpty": "此提供商没有可用的模型", + "useCustomModel": "使用自定义模型:{{modelId}}", "task": { "title": "任务", "expand": "展开任务", diff --git a/webview-ui/src/i18n/locales/zh-TW/chat.json b/webview-ui/src/i18n/locales/zh-TW/chat.json index 65fbc49176..ffcea90547 100644 --- a/webview-ui/src/i18n/locales/zh-TW/chat.json +++ b/webview-ui/src/i18n/locales/zh-TW/chat.json @@ -1,5 +1,9 @@ { "greeting": "歡迎使用 Zoo Code!", + "clearSearch": "清除搜尋", + "searchModel": "搜尋模型...", + "modelListEmpty": "此提供者沒有可用的模型", + "useCustomModel": "使用自訂模型:{{modelId}}", "task": { "title": "工作", "expand": "展開工作",