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/src/core/webview/ClineProvider.ts b/src/core/webview/ClineProvider.ts index 718d6430c1..2745c93270 100644 --- a/src/core/webview/ClineProvider.ts +++ b/src/core/webview/ClineProvider.ts @@ -185,6 +185,25 @@ type GetStateOptions = { includeTaskHistory?: boolean } +/** + * Internal options for {@link ClineProvider.upsertProviderProfile}. + */ +type UpsertProviderProfileOptions = { + /** + * Internal-only bypass of the organization model allow-list, used exclusively + * for Zoo Gateway credential synchronization (token refresh) and sign-out + * writes. These are auth writes, not model selections: a restrictive + * allow-list may omit `zoo-gateway` entirely, or list the provider without the + * active `zooGatewayModelId`, and must not reject the credential write and + * leave stale credentials behind in the active profile. + * + * This flag must never be set from a webview-originated code path: the webview + * is not a trusted boundary, so every user-driven profile write keeps the + * allow-list enforcement. + */ + bypassAllowList?: boolean +} + export class ClineProvider extends EventEmitter implements vscode.WebviewViewProvider, TelemetryPropertiesProvider, TaskProviderLike @@ -1835,7 +1854,52 @@ export class ClineProvider name: string, providerSettings: ProviderSettings, activate: boolean = true, + options: UpsertProviderProfileOptions = {}, ): Promise { + // Enforce the organization model allow-list before persisting or + // activating a profile. The webview is not a trusted boundary, so the + // model selector's client-side gating cannot be the only check. + // Task creation validates too, but rejecting here prevents an + // unauthorized profile from being written or activated at all. + // + // `bypassAllowList` is reserved for internal Zoo Gateway credential + // writes (token refresh / sign-out), which are auth writes rather than + // model selections; see `UpsertProviderProfileOptions`. + if (!options.bypassAllowList) { + let organizationAllowList = ORGANIZATION_ALLOW_ALL + + if (CloudService.hasInstance()) { + // The webview is not a trusted boundary, so a profile write must not + // fail open. Allow-all is only legitimate when no cloud instance exists + // (positively no organization policy). If a cloud instance exists but its + // policy cannot be read, reject the write rather than persisting a + // possibly disallowed model during a transient cloud/allow-list failure. + try { + organizationAllowList = await CloudService.instance.getAllowList() + } catch (error) { + this.log( + `[upsertProviderProfile] Blocked profile "${name}": organization allow-list unavailable: ${ + error instanceof Error ? error.message : String(error) + }`, + ) + void vscode.window.showErrorMessage(t("common:errors.violated_organization_allowlist")) + return undefined + } + } + + if (!ProfileValidator.isProfileAllowed(providerSettings, organizationAllowList)) { + this.log( + `[upsertProviderProfile] Blocked profile "${name}": model is not allowed by the organization allow-list`, + ) + // Surface the rejection to the user here rather than in the webview + // handler: direct callers (OAuth callbacks, sign-out) reach this + // method without the handler and would otherwise fail silently. No + // webview context is required, so non-webview callers are safe. + void vscode.window.showErrorMessage(t("common:errors.violated_organization_allowlist")) + return undefined + } + } + try { return await this.enqueueProviderProfileMutation(async (signal) => { // TODO: Do we need to be calling `activateProfile`? It's not @@ -1843,12 +1907,53 @@ export class ClineProvider // 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. + // Snapshot the pre-write state so a failure *after* `saveConfig` + // succeeds (during the activation writes below) can be rolled back + // instead of leaving the profile secret and the active/mode state + // divergent — the profile would carry the new model while the mode + // or current profile still pointed at the old one. + const priorCurrentApiConfigName = this.contextProxy.getValue("currentApiConfigName") + const priorProviderSettings = this.contextProxy.getProviderSettings() + + // `getProfile` wraps both "not found" and transient read failures in the + // same error, so it cannot decide whether the profile pre-existed. Probe + // existence explicitly: a destructive rollback (`deleteConfig`) must only + // run when absence is confirmed, never on a swallowed read error that would + // otherwise delete an existing profile and its secrets. + let profileExisted: boolean | undefined + try { + profileExisted = await this.providerSettingsManager.hasConfig(name) + } catch { + // Existence is unknown; rollback will restore/no-op rather than delete. + profileExisted = undefined + } + + let priorProfile: Awaited> | undefined + if (profileExisted !== false) { + try { + priorProfile = await this.providerSettingsManager.getProfile({ name }) + } catch (error) { + // The profile exists (or existence is unknown) but could not be read. + // When existence was confirmed, propagate so the write aborts *before* + // `saveConfig` rather than proceeding with an unknown prior state. + if (profileExisted === true) { + throw error + } + } + } + const id = await this.providerSettingsManager.saveConfig(name, providerSettings) if (signal.aborted) return id if (activate) { const { mode } = await this.getState() + let priorModeConfigId: string | undefined + try { + priorModeConfigId = await this.providerSettingsManager.getModeConfigId(mode) + } catch { + // No prior mapping (or unavailable); nothing to restore for the mode. + } // These promises do the following: // 1. Adds or updates the list of provider profiles. @@ -1860,19 +1965,55 @@ export class ClineProvider // 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) + try { + 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) + } catch (error) { + // Compensating rollback: restore the profile secret, the active + // profile name, the mode mapping and the in-memory provider + // settings to their pre-write values so a partial activation + // cannot report success while leaving inconsistent state. + try { + if (priorProfile) { + await this.providerSettingsManager.saveConfig(name, priorProfile) + } else if (profileExisted === false) { + // Absence was confirmed before the write; remove the new profile. + // A swallowed read error (unknown existence) must never reach here. + await this.providerSettingsManager.deleteConfig(name) + } + // A pre-existing mode mapping is restored; a newly created one + // has no delete API, so it is left as a best-effort remainder. + if (priorModeConfigId) { + await this.providerSettingsManager.setModeConfig(mode, priorModeConfigId) + } + await this.contextProxy.setValue("currentApiConfigName", priorCurrentApiConfigName) + await this.contextProxy.setValues({ + listApiConfigMeta: await this.providerSettingsManager.listConfig(), + }) + await this.contextProxy.setProviderSettings(priorProviderSettings) + } catch (rollbackError) { + this.log( + `[upsertProviderProfile] rollback failed for "${name}": ${ + rollbackError instanceof Error ? rollbackError.message : String(rollbackError) + }`, + ) + } + throw error + } } else { await this.updateGlobalState("listApiConfigMeta", await this.providerSettingsManager.listConfig()) } @@ -2124,7 +2265,13 @@ export class ClineProvider } // Activate only if zoo-gateway was the active provider (shouldn't happen if // no profiles exist, but defensive). - await this.upsertProviderProfile("Zoo Gateway", newConfiguration, isZooGatewayActive) + // + // `bypassAllowList`: internal auth credential write. `ProfileValidator` + // cannot map zoo-gateway to a model id, so a restrictive organization + // allow-list would reject this write and leave no credentials persisted. + await this.upsertProviderProfile("Zoo Gateway", newConfiguration, isZooGatewayActive, { + bypassAllowList: true, + }) } else { // Update every existing zoo-gateway profile with the new token and the // derived base URL so that environment-specific routing stays consistent. @@ -2139,7 +2286,8 @@ export class ClineProvider if (isActiveProfile) { // Use upsertProviderProfile with activate: true so the in-memory handler // picks up the new token immediately for the current task. - await this.upsertProviderProfile(entry.name, updated, true) + // `bypassAllowList`: internal auth credential write (see above). + await this.upsertProviderProfile(entry.name, updated, true, { bypassAllowList: true }) } else { // Non-active profiles just need the token saved to disk. await this.providerSettingsManager.saveConfig(entry.name, updated) diff --git a/src/core/webview/__tests__/ClineProvider.apiHandlerRebuild.spec.ts b/src/core/webview/__tests__/ClineProvider.apiHandlerRebuild.spec.ts index 99d254cb9a..70433ef63f 100644 --- a/src/core/webview/__tests__/ClineProvider.apiHandlerRebuild.spec.ts +++ b/src/core/webview/__tests__/ClineProvider.apiHandlerRebuild.spec.ts @@ -3,7 +3,7 @@ import * as vscode from "vscode" import { TelemetryService } from "@roo-code/telemetry" -import { getModelId, RooCodeEventName } from "@roo-code/types" +import { getModelId, ORGANIZATION_ALLOW_ALL, RooCodeEventName } from "@roo-code/types" import { ContextProxy } from "../../config/ContextProxy" import type { Mode } from "../../../shared/modes" @@ -11,6 +11,13 @@ import { Task, TaskOptions } from "../../task/Task" import { ClineProvider } from "../ClineProvider" import { providerIdentifiers } from "@roo-code/types/provider-identifiers" +// Partial mock: other provider modules (e.g. `deepseek.ts`) import +// `OpenAiHandler` from here, so the real exports must remain available. +vi.mock("../../../api/providers/openai", async (importOriginal) => ({ + ...(await importOriginal()), + getOpenAiModels: vi.fn(), +})) + // Mock setup vi.mock("fs/promises", () => ({ mkdir: vi.fn().mockResolvedValue(undefined), @@ -130,12 +137,20 @@ vi.mock("../../task/Task", () => ({ }), })) +// Hoisted so the `@roo-code/cloud` mock factory (which is hoisted above the +// imports) can expose the same `getAllowList`/`hasInstance` spies the tests configure. +const { mockGetAllowList, mockHasInstance } = vi.hoisted(() => ({ + mockGetAllowList: vi.fn(), + mockHasInstance: vi.fn(), +})) + vi.mock("@roo-code/cloud", () => ({ CloudService: { - hasInstance: vi.fn().mockReturnValue(true), + hasInstance: mockHasInstance, get instance() { return { isAuthenticated: vi.fn().mockReturnValue(false), + getAllowList: mockGetAllowList, } }, }, @@ -153,6 +168,10 @@ describe("ClineProvider - API Handler Rebuild Guard", () => { beforeEach(async () => { vi.clearAllMocks() + // Default to allow-all; individual tests override with a restrictive list. + mockGetAllowList.mockResolvedValue(ORGANIZATION_ALLOW_ALL) + // A cloud instance exists by default, which is the fail-closed case. + mockHasInstance.mockReturnValue(true) if (!TelemetryService.hasInstance()) { TelemetryService.createInstance([]) @@ -257,6 +276,9 @@ describe("ClineProvider - API Handler Rebuild Guard", () => { apiProvider: providerIdentifiers.openrouter, openRouterModelId: "openai/gpt-4", }), + // Default to "does not exist yet"; tests that simulate an existing + // profile must override this so the prior-profile snapshot is taken. + hasConfig: vi.fn().mockResolvedValue(false), } // Get the buildApiHandler mock @@ -417,6 +439,242 @@ describe("ClineProvider - API Handler Rebuild Guard", () => { // Should not call buildApiHandler when there's no task expect(buildApiHandlerMock).not.toHaveBeenCalled() }) + + test("persists and activates an allowed profile under a restrictive allow-list", async () => { + mockGetAllowList.mockResolvedValue({ + allowAll: false, + providers: { + [providerIdentifiers.openrouter]: { allowAll: false, models: ["openai/gpt-4"] }, + }, + }) + + const result = await provider.upsertProviderProfile("test-config", { + apiProvider: providerIdentifiers.openrouter, + openRouterModelId: "openai/gpt-4", + }) + + expect(result).toBe("test-id") + expect(provider["providerSettingsManager"].saveConfig).toHaveBeenCalledWith( + "test-config", + expect.objectContaining({ openRouterModelId: "openai/gpt-4" }), + ) + expect(mockContext.globalState.update).toHaveBeenCalledWith("currentApiConfigName", "test-config") + expect(vscode.window.showErrorMessage).not.toHaveBeenCalled() + }) + + test("rejects a disallowed profile: not persisted, not activated, user notified", async () => { + mockGetAllowList.mockResolvedValue({ + allowAll: false, + providers: { + [providerIdentifiers.openrouter]: { allowAll: false, models: ["openai/gpt-4"] }, + }, + }) + const saveConfig = provider["providerSettingsManager"].saveConfig + + const result = await provider.upsertProviderProfile("blocked-config", { + apiProvider: providerIdentifiers.openrouter, + openRouterModelId: "openai/forbidden", + }) + + // No id means the write was rejected before any persistence/activation. + expect(result).toBeUndefined() + expect(saveConfig).not.toHaveBeenCalled() + expect(mockContext.globalState.update).not.toHaveBeenCalledWith("currentApiConfigName", "blocked-config") + // The notification lives in `upsertProviderProfile` so direct callers + // (OAuth callbacks, sign-out) do not fail silently. + expect(vscode.window.showErrorMessage).toHaveBeenCalledWith("errors.violated_organization_allowlist") + }) + + test("rejects the write when a cloud instance exists but the allow-list is unavailable", async () => { + mockGetAllowList.mockRejectedValue(new Error("cloud unavailable")) + + const saveConfig = provider["providerSettingsManager"].saveConfig + const result = await provider.upsertProviderProfile("test-config", { + apiProvider: providerIdentifiers.openrouter, + openRouterModelId: "openai/gpt-4", + }) + + // Fail closed: a transient allow-list failure with a live cloud instance + // must not persist a possibly disallowed model. + expect(result).toBeUndefined() + expect(saveConfig).not.toHaveBeenCalled() + expect(vscode.window.showErrorMessage).toHaveBeenCalledWith("errors.violated_organization_allowlist") + }) + + test("allows the write when no cloud instance exists (positively no organization policy)", async () => { + mockHasInstance.mockReturnValue(false) + // Even a failing policy read must not matter: allow-all is legitimate + // when no organization policy positively applies. + mockGetAllowList.mockRejectedValue(new Error("should not be consulted for the gate")) + + const result = await provider.upsertProviderProfile("test-config", { + apiProvider: providerIdentifiers.openrouter, + openRouterModelId: "openai/gpt-4", + }) + + expect(result).toBe("test-id") + expect(provider["providerSettingsManager"].saveConfig).toHaveBeenCalled() + expect(vscode.window.showErrorMessage).not.toHaveBeenCalled() + }) + + test("bypassAllowList writes Zoo Gateway credentials even when the allow-list forbids the provider", async () => { + // A restrictive allow-list that does not mention zoo-gateway at all. This + // is exactly the case that used to strand a refreshed/cleared token, and + // it is why the credential write opts out of the model allow-list. + mockGetAllowList.mockResolvedValue({ + allowAll: false, + providers: { + [providerIdentifiers.openrouter]: { allowAll: false, models: ["openai/gpt-4"] }, + }, + }) + + const result = await provider.upsertProviderProfile( + "Zoo Gateway", + { apiProvider: providerIdentifiers.zooGateway, zooSessionToken: "zoo_ext_token" }, + false, + { bypassAllowList: true }, + ) + + // The write succeeded even though the allow-list forbids zoo-gateway: + // without the bypass this exact profile is rejected (see the next test). + expect(result).toBe("test-id") + expect(provider["providerSettingsManager"].saveConfig).toHaveBeenCalledWith( + "Zoo Gateway", + expect.objectContaining({ zooSessionToken: "zoo_ext_token" }), + ) + expect(vscode.window.showErrorMessage).not.toHaveBeenCalled() + }) + + test("accepts a Zoo Gateway profile whose model is on the allow-list", async () => { + // `ProfileValidator.getModelIdFromProfile` maps zoo-gateway to + // `zooGatewayModelId`, so a listed model passes without the bypass. + mockGetAllowList.mockResolvedValue({ + allowAll: false, + providers: { + [providerIdentifiers.zooGateway]: { + allowAll: false, + models: ["anthropic/claude-sonnet-4"], + }, + }, + }) + + const result = await provider.upsertProviderProfile("Zoo Gateway", { + apiProvider: providerIdentifiers.zooGateway, + zooSessionToken: "zoo_ext_token", + zooGatewayModelId: "anthropic/claude-sonnet-4", + }) + + expect(result).toBe("test-id") + expect(provider["providerSettingsManager"].saveConfig).toHaveBeenCalledWith( + "Zoo Gateway", + expect.objectContaining({ zooGatewayModelId: "anthropic/claude-sonnet-4" }), + ) + expect(vscode.window.showErrorMessage).not.toHaveBeenCalled() + }) + + test("rejects a Zoo Gateway profile whose model is not on the allow-list", async () => { + mockGetAllowList.mockResolvedValue({ + allowAll: false, + providers: { + [providerIdentifiers.zooGateway]: { + allowAll: false, + models: ["anthropic/claude-sonnet-4"], + }, + }, + }) + + const result = await provider.upsertProviderProfile("Zoo Gateway", { + apiProvider: providerIdentifiers.zooGateway, + zooSessionToken: "zoo_ext_token", + zooGatewayModelId: "anthropic/claude-opus-4", + }) + + expect(result).toBeUndefined() + expect(vscode.window.showErrorMessage).toHaveBeenCalledWith("errors.violated_organization_allowlist") + }) + + test("the Zoo Gateway bypass does not leak to other providers", async () => { + mockGetAllowList.mockResolvedValue({ + allowAll: false, + providers: { + [providerIdentifiers.openrouter]: { allowAll: false, models: ["openai/gpt-4"] }, + }, + }) + + // Same restrictive list, but without the internal bypass the write is + // still rejected — proving the escape hatch is opt-in per call. + const result = await provider.upsertProviderProfile("blocked-config", { + apiProvider: providerIdentifiers.zooGateway, + zooSessionToken: "zoo_ext_token", + }) + + expect(result).toBeUndefined() + expect(vscode.window.showErrorMessage).toHaveBeenCalledWith("errors.violated_organization_allowlist") + }) + + test("rolls back the profile and active state when an activation write fails after saveConfig", async () => { + // Seed the previously active profile name so the rollback restores a known value. + await provider.contextProxy.setValue("currentApiConfigName", "test-config") + const priorProfile = { + name: "new-config", + id: "prior-id", + apiProvider: providerIdentifiers.openrouter, + openRouterModelId: "openai/gpt-4", + } + // The profile pre-exists, so existence is confirmed and its prior value is read. + provider["providerSettingsManager"].hasConfig = vi.fn().mockResolvedValue(true) + provider["providerSettingsManager"].getProfile = vi.fn().mockResolvedValue(priorProfile) + provider["providerSettingsManager"].getModeConfigId = vi.fn().mockResolvedValue(undefined) + // Fail an activation write that runs *after* saveConfig succeeded. + provider["providerSettingsManager"].setModeConfig = vi + .fn() + .mockRejectedValue(new Error("mode write failed")) + const saveConfig = provider["providerSettingsManager"].saveConfig + + const result = await provider.upsertProviderProfile("new-config", { + apiProvider: providerIdentifiers.openrouter, + openRouterModelId: "openai/gpt-4-turbo", + }) + + // The partial write is reported as a failure, never a success. + expect(result).toBeUndefined() + // The new model was written first, then the prior profile was restored so + // the secret cannot carry the new model while the active state is stale. + expect(saveConfig).toHaveBeenNthCalledWith( + 1, + "new-config", + expect.objectContaining({ openRouterModelId: "openai/gpt-4-turbo" }), + ) + expect(saveConfig).toHaveBeenNthCalledWith(2, "new-config", priorProfile) + // The previously active profile name ("test-config") is restored. + expect(mockContext.globalState.update).toHaveBeenCalledWith("currentApiConfigName", "test-config") + // No activation success state leaked to the webview. + expect(vscode.window.showErrorMessage).toHaveBeenCalledWith("errors.create_api_config") + }) + + test("aborts without deleting an existing profile when the prior profile cannot be read", async () => { + // The profile exists, but reading it fails (e.g. a transient secrets error). + // A swallowed failure must NOT be treated as "absent": rollback must never + // delete the still-existing profile and its secrets. + provider["providerSettingsManager"].hasConfig = vi.fn().mockResolvedValue(true) + provider["providerSettingsManager"].getProfile = vi + .fn() + .mockRejectedValue(new Error("transient secrets read failure")) + const saveConfig = provider["providerSettingsManager"].saveConfig + const deleteConfig = vi.fn().mockResolvedValue(undefined) + provider["providerSettingsManager"].deleteConfig = deleteConfig + + const result = await provider.upsertProviderProfile("test-config", { + apiProvider: providerIdentifiers.openrouter, + openRouterModelId: "openai/gpt-4-turbo", + }) + + // The write aborts *before* saveConfig, so the existing profile is untouched. + expect(result).toBeUndefined() + expect(saveConfig).not.toHaveBeenCalled() + expect(deleteConfig).not.toHaveBeenCalled() + expect(vscode.window.showErrorMessage).toHaveBeenCalledWith("errors.create_api_config") + }) }) describe("activateProviderProfile", () => { diff --git a/src/core/webview/__tests__/ClineProvider.spec.ts b/src/core/webview/__tests__/ClineProvider.spec.ts index 29a8ed53f6..d2fcf8540e 100644 --- a/src/core/webview/__tests__/ClineProvider.spec.ts +++ b/src/core/webview/__tests__/ClineProvider.spec.ts @@ -401,6 +401,9 @@ vi.mock("@roo-code/cloud", () => ({ login: vi.fn().mockResolvedValue(undefined), logout: vi.fn().mockResolvedValue(undefined), off: vi.fn(), + // A cloud instance is present, so the fail-closed allow-list guard reads + // this. Default to allow-all so profile writes behave as before. + getAllowList: vi.fn().mockReturnValue({ allowAll: true, providers: {} }), } }, }, @@ -5197,6 +5200,7 @@ describe("ClineProvider - Comprehensive Edit/Delete Edge Cases", () => { zooGatewayBaseUrl: "https://www.zoocode.dev/api/gateway/v1", }), false, + { bypassAllowList: true }, ) }) @@ -5241,6 +5245,7 @@ describe("ClineProvider - Comprehensive Edit/Delete Edge Cases", () => { zooGatewayBaseUrl: "https://www.zoocode.dev/api/gateway/v1", }), true, + { bypassAllowList: true }, ) expect(saveConfig).toHaveBeenCalledWith( "Backup Zoo", diff --git a/src/core/webview/__tests__/ClineProvider.sticky-mode.spec.ts b/src/core/webview/__tests__/ClineProvider.sticky-mode.spec.ts index 55e26a9667..b70ae6aa6f 100644 --- a/src/core/webview/__tests__/ClineProvider.sticky-mode.spec.ts +++ b/src/core/webview/__tests__/ClineProvider.sticky-mode.spec.ts @@ -118,6 +118,9 @@ vi.mock("@roo-code/cloud", () => ({ get instance() { return { isAuthenticated: vi.fn().mockReturnValue(false), + // A cloud instance is present, so the fail-closed allow-list guard reads + // this. Default to allow-all so profile writes behave as before. + getAllowList: vi.fn().mockReturnValue({ allowAll: true, providers: {} }), } }, }, diff --git a/src/core/webview/__tests__/webviewMessageHandler.spec.ts b/src/core/webview/__tests__/webviewMessageHandler.spec.ts index 4c2a301965..508a940c77 100644 --- a/src/core/webview/__tests__/webviewMessageHandler.spec.ts +++ b/src/core/webview/__tests__/webviewMessageHandler.spec.ts @@ -10,6 +10,12 @@ vi.mock("../../../services/zoo-code-auth", () => ({ vi.mock("../../../api/providers/fetchers/lmstudio", () => ({ getLMStudioModels: vi.fn(), })) +// Partial mock: other provider modules import `OpenAiHandler` from here, so the +// real exports must remain available. +vi.mock("../../../api/providers/openai", async (importOriginal) => ({ + ...(await importOriginal()), + getOpenAiModels: vi.fn(), +})) vi.mock("../../../integrations/theme/getTheme", () => ({ getTheme: vi.fn().mockResolvedValue({}), @@ -75,6 +81,7 @@ import { webviewMessageHandler } from "../webviewMessageHandler" import type { ClineProvider } from "../ClineProvider" import { flushModels, getModels } from "../../../api/providers/fetchers/modelCache" import { getLMStudioModels } from "../../../api/providers/fetchers/lmstudio" +import { getOpenAiModels } from "../../../api/providers/openai" import { getCommands } from "../../../services/command/commands" import { ensureDcgInstalled } from "../../../services/destructive-command-guard" import { @@ -90,6 +97,7 @@ const { fetchOpenAiCodexRateLimitInfo } = await import("../../../integrations/op const mockGetModels = getModels as Mock const mockFlushModels = flushModels as Mock const mockGetLMStudioModels = getLMStudioModels as Mock +const mockGetOpenAiModels = getOpenAiModels as Mock const mockGetCommands = vi.mocked(getCommands) const mockGetAccessToken = vi.mocked(openAiCodexOAuthManager.getAccessToken) const mockGetAccountId = vi.mocked(openAiCodexOAuthManager.getAccountId) @@ -125,6 +133,23 @@ const mockClineProvider = { cwd: "/mock/workspace", } as unknown as ClineProvider +// Structural overrides for the auth/profile tests below. Declaring the shape keeps the +// collaborator swaps type-checked (no `as any`) — the same pattern as `providerForLaunch` +// used by the telemetry tests. +type MockProviderOverrides = { + contextProxy: { + getProviderSettings: ReturnType + getValues: ReturnType + } + providerSettingsManager: { + listConfig: ReturnType + getProfile: ReturnType + saveConfig: ReturnType + } + upsertProviderProfile: ReturnType + getState: (typeof mockClineProvider)["getState"] +} + describe("webviewMessageHandler - theme fixture probes", () => { const originalProbeSetting = process.env.ROO_CODE_THEME_FIXTURE_PROBE const themeFixture = { @@ -501,6 +526,35 @@ describe("webviewMessageHandler - requestOllamaModels", () => { }) }) +describe("webviewMessageHandler - requestOpenAiModels", () => { + beforeEach(() => { + vi.clearAllMocks() + mockGetOpenAiModels.mockReset() + }) + + it("echoes the caller's requestId on the openAiModels response", async () => { + mockGetOpenAiModels.mockResolvedValue(["gpt-4o", "gpt-4o-mini"]) + + await webviewMessageHandler(mockClineProvider, { + type: "requestOpenAiModels", + requestId: "req-123", + values: { + baseUrl: "https://api.example.com/v1", + apiKey: "test-api-key-not-real", + openAiHeaders: {}, + }, + }) + + // The request id must round-trip so the webview can correlate the + // response with the request it issued and drop stale replies. + expect(mockClineProvider.postMessageToWebview).toHaveBeenCalledWith({ + type: "openAiModels", + openAiModels: ["gpt-4o", "gpt-4o-mini"], + requestId: "req-123", + }) + }) +}) + describe("webviewMessageHandler - requestRouterModels", () => { beforeEach(() => { vi.clearAllMocks() @@ -1796,13 +1850,14 @@ describe("zooCodeSignOut", () => { const { disconnectZooCode } = await import("../../../services/zoo-code-auth") const upsertProviderProfile = vi.fn().mockResolvedValue(undefined) const saveConfig = vi.fn().mockResolvedValue(undefined) + const authProvider = mockClineProvider as unknown as MockProviderOverrides - ;(mockClineProvider as any).contextProxy = { + authProvider.contextProxy = { ...mockClineProvider.contextProxy, getProviderSettings: vi.fn().mockReturnValue({ apiProvider: providerIdentifiers.zooGateway }), getValues: vi.fn().mockReturnValue({ currentApiConfigName: "Zoo Gateway" }), } - ;(mockClineProvider as any).providerSettingsManager = { + authProvider.providerSettingsManager = { listConfig: vi.fn().mockResolvedValue([ { name: "Zoo Gateway", apiProvider: providerIdentifiers.zooGateway }, { name: "Backup Zoo", apiProvider: providerIdentifiers.zooGateway }, @@ -1820,7 +1875,7 @@ describe("zooCodeSignOut", () => { }), saveConfig, } - ;(mockClineProvider as any).upsertProviderProfile = upsertProviderProfile + authProvider.upsertProviderProfile = upsertProviderProfile await webviewMessageHandler(mockClineProvider, { type: "zooCodeSignOut" }) @@ -1829,6 +1884,9 @@ describe("zooCodeSignOut", () => { "Zoo Gateway", expect.not.objectContaining({ zooSessionToken: expect.anything() }), true, + // Internal auth write: bypasses the model allow-list (zoo-gateway has + // no model-id mapping). + { bypassAllowList: true }, ) expect(saveConfig).toHaveBeenCalledWith( "Backup Zoo", @@ -1837,15 +1895,51 @@ describe("zooCodeSignOut", () => { expect(mockClineProvider.postStateToWebview).toHaveBeenCalled() }) + it("reports a failure instead of clearing silently when the active profile write is rejected", async () => { + // `upsertProviderProfile` returns `undefined` when the write is rejected + // (e.g. the model-allow-list guard or a disk error). Sign-out must not + // log a successful cleanup in that case. + const upsertProviderProfile = vi.fn().mockResolvedValue(undefined) + const authProvider = mockClineProvider as unknown as MockProviderOverrides + + authProvider.contextProxy = { + ...mockClineProvider.contextProxy, + getProviderSettings: vi.fn().mockReturnValue({ apiProvider: providerIdentifiers.zooGateway }), + getValues: vi.fn().mockReturnValue({ currentApiConfigName: "Zoo Gateway" }), + } + authProvider.providerSettingsManager = { + listConfig: vi + .fn() + .mockResolvedValue([{ name: "Zoo Gateway", apiProvider: providerIdentifiers.zooGateway }]), + getProfile: vi.fn().mockResolvedValue({ + apiProvider: providerIdentifiers.zooGateway, + zooSessionToken: "token-active", + }), + saveConfig: vi.fn(), + } + authProvider.upsertProviderProfile = upsertProviderProfile + + await webviewMessageHandler(mockClineProvider, { type: "zooCodeSignOut" }) + + // The rejected write must be surfaced, not reported as a successful cleanup. + expect(mockClineProvider.log).toHaveBeenCalledWith( + expect.stringContaining('[zooCodeSignOut] Failed to clear profile token for "Zoo Gateway"'), + ) + expect(mockClineProvider.log).not.toHaveBeenCalledWith( + expect.stringContaining('[zooCodeSignOut] Cleared zooSessionToken from "Zoo Gateway"'), + ) + }) + it("still clears the in-memory handler when the active profile token is already empty on disk", async () => { const upsertProviderProfile = vi.fn().mockResolvedValue(undefined) + const authProvider = mockClineProvider as unknown as MockProviderOverrides - ;(mockClineProvider as any).contextProxy = { + authProvider.contextProxy = { ...mockClineProvider.contextProxy, getProviderSettings: vi.fn().mockReturnValue({ apiProvider: providerIdentifiers.zooGateway }), getValues: vi.fn().mockReturnValue({ currentApiConfigName: "Zoo Gateway" }), } - ;(mockClineProvider as any).providerSettingsManager = { + authProvider.providerSettingsManager = { listConfig: vi .fn() .mockResolvedValue([{ name: "Zoo Gateway", apiProvider: providerIdentifiers.zooGateway }]), @@ -1855,7 +1949,7 @@ describe("zooCodeSignOut", () => { }), saveConfig: vi.fn(), } - ;(mockClineProvider as any).upsertProviderProfile = upsertProviderProfile + authProvider.upsertProviderProfile = upsertProviderProfile await webviewMessageHandler(mockClineProvider, { type: "zooCodeSignOut" }) @@ -1863,10 +1957,42 @@ describe("zooCodeSignOut", () => { "Zoo Gateway", expect.not.objectContaining({ zooSessionToken: expect.anything() }), true, + { bypassAllowList: true }, ) }) }) +describe("webviewMessageHandler - upsertApiConfiguration allow-list", () => { + beforeEach(() => { + vi.clearAllMocks() + }) + + it("delegates to upsertProviderProfile without pre-validating (single enforcement point)", async () => { + const upsertProviderProfile = vi.fn().mockResolvedValue("profile-id") + const detailProvider = mockClineProvider as unknown as MockProviderOverrides + detailProvider.upsertProviderProfile = upsertProviderProfile + const getState = vi.fn().mockResolvedValue({ + apiConfiguration: {}, + organizationAllowList: { allowAll: false, providers: {} }, + }) + detailProvider.getState = getState + + const apiConfiguration = { apiProvider: providerIdentifiers.anthropic, apiModelId: "not-allowed" } + + await webviewMessageHandler(mockClineProvider, { + type: "upsertApiConfiguration", + text: "test-config", + apiConfiguration, + }) + + // The handler forwards the write and relies on `upsertProviderProfile` for + // enforcement; it does not call `getState()` or reject on its own. + expect(upsertProviderProfile).toHaveBeenCalledWith("test-config", apiConfiguration) + expect(getState).not.toHaveBeenCalled() + expect(vscode.window.showErrorMessage).not.toHaveBeenCalled() + }) +}) + describe("webviewMessageHandler - kimiCodeSignIn", () => { beforeEach(() => { vi.clearAllMocks() diff --git a/src/core/webview/webviewMessageHandler.ts b/src/core/webview/webviewMessageHandler.ts index 4ba94d454c..cbd5268d24 100644 --- a/src/core/webview/webviewMessageHandler.ts +++ b/src/core/webview/webviewMessageHandler.ts @@ -1380,6 +1380,11 @@ export const webviewMessageHandler = async ( break } case OllamaModelsMessageType.requestOllamaModels: { + // Echo the caller's request id so the webview can correlate the + // response with the request it issued and drop stale replies that + // arrive after a provider/profile switch. + const requestId = message.requestId + // Specific handler for Ollama models only. const { apiConfiguration: ollamaApiConfig } = await provider.getState() // Prefer the baseUrl/apiKey from the message values (which reflect @@ -1406,6 +1411,7 @@ export const webviewMessageHandler = async ( type: OllamaModelsMessageType.ollamaModels, ollamaModels: {}, error: errorMsg, + requestId, }) break } @@ -1415,7 +1421,11 @@ export const webviewMessageHandler = async ( // 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 }) + await provider.postMessageToWebview({ + type: OllamaModelsMessageType.ollamaModels, + ollamaModels, + requestId, + }) } catch (error) { const errorMsg = error instanceof Error ? error.message : String(error) provider.log(`[requestOllamaModels] Failed to read models for ${logBaseUrl}: ${errorMsg}`) @@ -1423,11 +1433,15 @@ export const webviewMessageHandler = async ( type: OllamaModelsMessageType.ollamaModels, ollamaModels: {}, error: errorMsg, + requestId, }) } break } case LmStudioModelsMessageType.requestLmStudioModels: { + // Echo the caller's request id (see `requestOllamaModels` above). + const requestId = message.requestId + // Specific handler for LM Studio models only. const { apiConfiguration: lmStudioApiConfig } = await provider.getState() try { @@ -1450,6 +1464,7 @@ export const webviewMessageHandler = async ( await provider.postMessageToWebview({ type: LmStudioModelsMessageType.lmStudioModels, lmStudioModels: lmStudioModels, + requestId, }) } } catch (error) { @@ -1475,14 +1490,26 @@ export const webviewMessageHandler = async ( message?.values?.openAiHeaders, ) - await provider.postMessageToWebview({ type: OpenAiModelsMessageType.openAiModels, openAiModels }) + // Echo the caller's request id so the webview can correlate the + // response with the request it issued and drop stale replies + // that arrive after a provider/profile switch. + await provider.postMessageToWebview({ + type: OpenAiModelsMessageType.openAiModels, + openAiModels, + requestId: message.requestId, + }) } break case VsCodeLmModelsMessageType.requestVsCodeLmModels: const vsCodeLmModels = await getVsCodeLmModels() // TODO: Cache like we do for OpenRouter, etc? - await provider.postMessageToWebview({ type: VsCodeLmModelsMessageType.vsCodeLmModels, vsCodeLmModels }) + // Echo the caller's request id so stale replies can be discarded. + await provider.postMessageToWebview({ + type: VsCodeLmModelsMessageType.vsCodeLmModels, + vsCodeLmModels, + requestId: message.requestId, + }) break case "openImage": await openImage(message.text!, { values: message.values }) @@ -2280,6 +2307,13 @@ export const webviewMessageHandler = async ( break case "upsertApiConfiguration": if (message.text && message.apiConfiguration) { + // Allow-list enforcement lives solely in + // `ClineProvider.upsertProviderProfile`: it is not a trusted + // boundary either way (the webview can forge a model id), and + // keeping a single enforcement point means direct callers + // (OAuth callbacks, sign-out) get the same user notification + // instead of failing silently. Validating here as well would + // duplicate the work (including an extra `getState()`). await provider.upsertProviderProfile(message.text, message.apiConfiguration) } break @@ -2989,7 +3023,27 @@ export const webviewMessageHandler = async ( const isThisProfileActive = isZooGatewayActive && currentApiConfigName === entry.name if (isThisProfileActive) { - await provider.upsertProviderProfile(entry.name, cleanedProfile, true) + // `bypassAllowList`: clearing a Zoo Gateway session token is an + // internal auth write, not a model selection. `ProfileValidator` + // cannot map zoo-gateway to a model id, so a restrictive + // allow-list would otherwise reject the cleanup and leave the + // stale token in the active handler. + const writeResult = await provider.upsertProviderProfile( + entry.name, + cleanedProfile, + true, + { bypassAllowList: true }, + ) + + // The write can still fail for other reasons (disk error, + // disabled profile enforcement). Never report a successful + // token cleanup in that case: surface it instead. The + // allow-list check is bypassed on this call, so a neutral + // message is used rather than an allow-list violation. + if (writeResult === undefined) { + throw new Error(`Failed to persist cleaned Zoo Gateway profile "${entry.name}"`) + } + provider.log( `[zooCodeSignOut] Cleared zooSessionToken from "${entry.name}" profile and updated in-memory handler`, ) diff --git a/src/eslint-suppressions.json b/src/eslint-suppressions.json index 24db0bf433..88bd3b47c1 100644 --- a/src/eslint-suppressions.json +++ b/src/eslint-suppressions.json @@ -1121,7 +1121,7 @@ }, "core/webview/__tests__/webviewMessageHandler.spec.ts": { "@typescript-eslint/no-explicit-any": { - "count": 35 + "count": 29 } }, "core/webview/messageEnhancer.ts": { diff --git a/src/shared/ProfileValidator.ts b/src/shared/ProfileValidator.ts index 923582bd06..ed0a7ccedf 100644 --- a/src/shared/ProfileValidator.ts +++ b/src/shared/ProfileValidator.ts @@ -80,6 +80,8 @@ export class ProfileValidator { return profile.requestyModelId case providerIdentifiers.unbound: return profile.unboundModelId + case providerIdentifiers.zooGateway: + return profile.zooGatewayModelId case providerIdentifiers.fakeAi: default: return undefined diff --git a/src/shared/__tests__/ProfileValidator.spec.ts b/src/shared/__tests__/ProfileValidator.spec.ts index 865fa5cf51..1eaad49d87 100644 --- a/src/shared/__tests__/ProfileValidator.spec.ts +++ b/src/shared/__tests__/ProfileValidator.spec.ts @@ -341,6 +341,42 @@ describe("ProfileValidator", () => { expect(ProfileValidator.isProfileAllowed(profile, allowList)).toBe(true) }) + it("should extract zooGatewayModelId for zoo-gateway provider", () => { + const allowList: OrganizationAllowList = { + allowAll: false, + providers: { + [providerIdentifiers.zooGateway]: { + allowAll: false, + models: ["anthropic/claude-sonnet-4"], + }, + }, + } + const profile: ProviderSettings = { + apiProvider: providerIdentifiers.zooGateway, + zooGatewayModelId: "anthropic/claude-sonnet-4", + } + + expect(ProfileValidator.isProfileAllowed(profile, allowList)).toBe(true) + }) + + it("should reject a zoo-gateway profile whose model is not on the allow-list", () => { + const allowList: OrganizationAllowList = { + allowAll: false, + providers: { + [providerIdentifiers.zooGateway]: { + allowAll: false, + models: ["anthropic/claude-sonnet-4"], + }, + }, + } + const profile: ProviderSettings = { + apiProvider: providerIdentifiers.zooGateway, + zooGatewayModelId: "anthropic/claude-opus-4", + } + + expect(ProfileValidator.isProfileAllowed(profile, allowList)).toBe(false) + }) + it("should handle providers with undefined models list gracefully", () => { 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..0dccecd9da --- /dev/null +++ b/webview-ui/src/components/chat/ChatModelSelector.tsx @@ -0,0 +1,193 @@ +import { useMemo, useState, useCallback } 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 { 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 portalContainer = useRooPortal("roo-portal") + + // Filter deprecated models but always keep the currently selected one visible. + const modelIds = useMemo( + () => + Object.entries(filterModels(models, provider, organizationAllowList) ?? {}) + .filter(([modelId, modelInfo]) => modelId === selectedModelId || !modelInfo.deprecated) + .map(([modelId]) => modelId) + .sort((a, b) => a.localeCompare(b)), + [models, provider, organizationAllowList, selectedModelId], + ) + + // Gate arbitrary model ids against the organization allow-list. The webview is not a security + // boundary, but this keeps the UI from offering a custom model the host would reject on save. + const isModelAllowed = useCallback( + (modelId: string): boolean => { + if (!organizationAllowList || organizationAllowList.allowAll) { + return true + } + if (!provider) { + return false + } + const providerConfig = organizationAllowList.providers[provider] + if (!providerConfig) { + return false + } + if (providerConfig.allowAll) { + return true + } + return providerConfig.models?.includes(modelId) ?? false + }, + [organizationAllowList, provider], + ) + + // Only offer a custom (non-listed) model when the allow-list permits that exact model id. + const customModelAllowed = !!searchValue && isModelAllowed(searchValue) + + // Resolve the display value (custom transform for compound config values like VSCode LM). + const displayValue = useMemo(() => { + if (displayTransform && modelIdKey) { + const storedValue = apiConfiguration?.[modelIdKey] + return storedValue ? displayTransform(storedValue) : undefined + } + // Only the saved model is reflected; fall back to undefined so the placeholder shows. + 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]) + + const onSelect = useCallback( + (modelId: string) => { + // Defense in depth: never persist a model the allow-list forbids, even if a caller + // bypasses the rendered list (e.g. custom search); also require the lookup to be ready. + if (!modelId || !modelIdKey || !apiConfiguration || !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 to the current API configuration profile; the backend saves and activates + // it and broadcasts the updated apiConfiguration back to the webview. + vscode.postMessage({ + type: "upsertApiConfiguration", + text: currentApiConfigName, + apiConfiguration: { ...apiConfiguration, [modelIdKey]: valueToStore }, + }) + }, + [modelIdKey, apiConfiguration, valueTransform, currentApiConfigName, isModelAllowed], + ) + + const onClearSearch = useCallback(() => setSearchValue(""), []) + + 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 && ( +
+
+ )} +
+ + {modelIds.length === 0 && !isLoading && ( +
{t("chat:modelListEmpty")}
+ )} + + {modelIds.length > 0 && ( +
+ {filteredModelIds.map((modelId) => ( + + ))} +
+ )} + + {searchValue && customModelAllowed && !modelIds.includes(searchValue) && ( + + )} +
+
+
+ ) +} diff --git a/webview-ui/src/components/chat/ChatTextArea.tsx b/webview-ui/src/components/chat/ChatTextArea.tsx index 2c539766fe..27150a9e92 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..4221cb3c04 --- /dev/null +++ b/webview-ui/src/components/chat/__tests__/ChatModelSelector.spec.tsx @@ -0,0 +1,500 @@ +import type { ReactNode } from "react" +import { render, screen, fireEvent } from "@/utils/test-utils" +import { providerIdentifiers } from "@roo-code/types" +import { vscode } from "@/utils/vscode" + +import { ChatModelSelector } from "../ChatModelSelector" +import { useChatModelSelector } from "../hooks/useChatModelSelector" + +// 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(), +})) + +// Mock Popover components to be testable +vi.mock("@/components/ui", () => { + type PopoverProps = { children: ReactNode; open?: boolean; onOpenChange?: (open: boolean) => void } + type TriggerProps = { + children: ReactNode + disabled?: boolean + className?: string + onClick?: () => void + } + return { + Popover: ({ children, open }: PopoverProps) => ( +
+ {children} +
+ ), + PopoverTrigger: ({ children, disabled, className, onClick }: TriggerProps) => ( + + ), + PopoverContent: ({ children }: { children: ReactNode }) =>
{children}
, + 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, + }, +} + +// Minimal ModelInfo factory for ad-hoc fixtures. +const modelInfo = (overrides: Record = {}) => ({ + maxTokens: 32000, + contextWindow: 200000, + supportsPromptCache: true, + ...overrides, +}) + +describe("ChatModelSelector", () => { + const defaultProps = { + title: "Select model", + } + + /** Helper: render the selector and open the popover so options are visible. */ + const renderOpen = (disabled = false) => { + render() + fireEvent.click(screen.getByTestId("chat-model-selector-trigger")) + } + + /** Helper: read the visible option ids in DOM order. */ + const visibleOptionIds = () => + screen + .getAllByTestId(/^chat-model-option-/) + .map((element) => element.getAttribute("data-testid")!.replace("chat-model-option-", "")) + + beforeEach(() => { + vi.clearAllMocks() + mockUseExtensionState.mockReturnValue({ + apiConfiguration: { + apiProvider: providerIdentifiers.anthropic, + apiModelId: "claude-opus-4-20250514", + }, + organizationAllowList: { allowAll: true, providers: {} }, + currentApiConfigName: "default", + }) + mockUseChatModelSelector.mockReturnValue({ + provider: providerIdentifiers.anthropic, + models: anthropicModels, + modelIdKey: "apiModelId", + defaultModelId: "claude-opus-4-20250514", + isLoading: false, + }) + }) + + 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) + + // The mocked Popover always renders content, so the models are visible. + // 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("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", + 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 (present in the always-rendered PopoverContent) + 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", + 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 = searchInput.parentElement?.querySelector(".codicon-close") as HTMLElement | null + expect(clearButton).toBeTruthy() + fireEvent.click(clearButton!) + + expect((screen.getByPlaceholderText("chat:searchModel") as HTMLInputElement).value).toBe("") + }) + + test("filters model ids case-insensitively", () => { + mockUseChatModelSelector.mockReturnValue({ + provider: providerIdentifiers.anthropic, + models: { + "Claude-Opus-4": modelInfo(), + "Claude-Sonnet-4": modelInfo(), + }, + modelIdKey: "apiModelId", + defaultModelId: "Claude-Opus-4", + isLoading: false, + }) + + renderOpen() + + // Uppercase query must match the lowercase portion of the model id. + fireEvent.change(screen.getByPlaceholderText("chat:searchModel"), { target: { value: "SONNET" } }) + + expect(screen.getByTestId("chat-model-option-Claude-Sonnet-4")).toBeInTheDocument() + expect(screen.queryByTestId("chat-model-option-Claude-Opus-4")).not.toBeInTheDocument() + }) + + test("renders the empty state when no models are available", () => { + mockUseChatModelSelector.mockReturnValue({ + provider: providerIdentifiers.anthropic, + models: {}, + modelIdKey: "apiModelId", + defaultModelId: "claude-opus-4-20250514", + isLoading: false, + }) + + renderOpen() + + expect(screen.getByText("chat:modelListEmpty")).toBeInTheDocument() + expect(screen.queryByTestId(/^chat-model-option-/)).not.toBeInTheDocument() + }) + + test("shows a loading label and hides the empty state while loading", () => { + mockUseChatModelSelector.mockReturnValue({ + provider: providerIdentifiers.anthropic, + models: null, + modelIdKey: "apiModelId", + defaultModelId: "claude-opus-4-20250514", + isLoading: true, + }) + + render() + + // The trigger collapses to an ellipsis while loading, regardless of the + // resolved display/default value. + expect(screen.getByTestId("chat-model-selector-trigger")).toHaveTextContent("…") + // The empty-state message must not appear while a load is in flight. + expect(screen.queryByText("chat:modelListEmpty")).not.toBeInTheDocument() + }) + + test("does not post a message when selecting with no apiConfiguration", () => { + // Keep modelIdKey defined so the guard under test is the `!apiConfiguration` + // early return, not the `!modelIdKey` guard. + mockUseExtensionState.mockReturnValue({ + apiConfiguration: undefined, + organizationAllowList: { allowAll: true, providers: {} }, + currentApiConfigName: "default", + }) + + renderOpen() + + fireEvent.click(screen.getByTestId("chat-model-option-claude-opus-4-20250514")) + + 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") + }) + + describe("organization allow-list", () => { + test("filters models down to the provider allow-list", () => { + mockUseExtensionState.mockReturnValue({ + apiConfiguration: { + apiProvider: providerIdentifiers.anthropic, + apiModelId: "test-model-not-real", + }, + organizationAllowList: { + allowAll: false, + providers: { + [providerIdentifiers.anthropic]: { + allowAll: false, + models: ["claude-sonnet-4-20250514"], + }, + }, + }, + currentApiConfigName: "default", + }) + + renderOpen() + + expect(screen.getByTestId("chat-model-option-claude-sonnet-4-20250514")).toBeInTheDocument() + expect(screen.queryByTestId("chat-model-option-claude-opus-4-20250514")).not.toBeInTheDocument() + }) + + test("shows the empty state when the provider is absent from the allow-list", () => { + mockUseExtensionState.mockReturnValue({ + apiConfiguration: { + apiProvider: providerIdentifiers.anthropic, + apiModelId: "test-model-not-real", + }, + organizationAllowList: { allowAll: false, providers: {} }, + currentApiConfigName: "default", + }) + + renderOpen() + + expect(screen.queryByTestId(/^chat-model-option-/)).not.toBeInTheDocument() + expect(screen.getByText("chat:modelListEmpty")).toBeInTheDocument() + }) + + test("hides the custom-model row when the allow-list forbids the typed id", () => { + mockUseExtensionState.mockReturnValue({ + apiConfiguration: { + apiProvider: providerIdentifiers.anthropic, + apiModelId: "test-model-not-real", + }, + organizationAllowList: { + allowAll: false, + providers: { + [providerIdentifiers.anthropic]: { + allowAll: false, + models: ["claude-opus-4-20250514"], + }, + }, + }, + currentApiConfigName: "default", + }) + + renderOpen() + + fireEvent.change(screen.getByPlaceholderText("chat:searchModel"), { + target: { value: "unauthorized-model" }, + }) + + expect(screen.queryByTestId("chat-model-use-custom")).not.toBeInTheDocument() + }) + + test("offers and posts the custom-model row when the allow-list permits it", () => { + mockUseExtensionState.mockReturnValue({ + apiConfiguration: { + apiProvider: providerIdentifiers.anthropic, + apiModelId: "test-model-not-real", + }, + organizationAllowList: { + allowAll: false, + providers: { + [providerIdentifiers.anthropic]: { + allowAll: false, + models: ["claude-opus-4-20250514", "approved-custom"], + }, + }, + }, + currentApiConfigName: "default", + }) + mockUseChatModelSelector.mockReturnValue({ + provider: providerIdentifiers.anthropic, + models: { "claude-opus-4-20250514": modelInfo() }, + modelIdKey: "apiModelId", + defaultModelId: "claude-opus-4-20250514", + isLoading: false, + }) + + renderOpen() + + fireEvent.change(screen.getByPlaceholderText("chat:searchModel"), { + target: { value: "approved-custom" }, + }) + + fireEvent.click(screen.getByTestId("chat-model-use-custom")) + + expect(vscode.postMessage).toHaveBeenCalledWith({ + type: "upsertApiConfiguration", + text: "default", + apiConfiguration: { + apiProvider: providerIdentifiers.anthropic, + apiModelId: "approved-custom", + }, + }) + }) + }) + + test("guards the visible list: sorts ids, drops unselected deprecated models and keeps the selected one", () => { + mockUseChatModelSelector.mockReturnValue({ + provider: providerIdentifiers.anthropic, + // Intentionally unsorted, with one deprecated model that is also the + // selected one (must survive) and one deprecated model that is not. + models: { + "zeta-model": modelInfo(), + "claude-opus-4-20250514": modelInfo({ deprecated: true }), + "legacy-model": modelInfo({ deprecated: true }), + "alpha-model": modelInfo(), + "middle-model": modelInfo(), + }, + modelIdKey: "apiModelId", + defaultModelId: "claude-opus-4-20250514", + isLoading: false, + }) + + renderOpen() + + // Rendered ids are sorted alphabetically and exclude the unselected deprecated model. + expect(visibleOptionIds()).toEqual(["alpha-model", "claude-opus-4-20250514", "middle-model", "zeta-model"]) + expect(screen.queryByTestId("chat-model-option-legacy-model")).not.toBeInTheDocument() + expect(screen.getByTestId("chat-model-option-claude-opus-4-20250514")).toBeInTheDocument() + }) + + test("disables the trigger when there is no modelIdKey", () => { + mockUseChatModelSelector.mockReturnValue({ + provider: undefined, + models: null, + modelIdKey: undefined, + defaultModelId: "", + isLoading: false, + }) + + render() + + const trigger = screen.getByTestId("chat-model-selector-trigger") + expect(trigger).toBeDisabled() + expect(trigger).toHaveClass("cursor-not-allowed") + }) + + test("exposes model options as keyboard-activatable buttons", () => { + renderOpen() + + const option = screen.getByTestId("chat-model-option-claude-sonnet-4-20250514") + // A native