diff --git a/apps/cli/src/commands/cli/__tests__/list.test.ts b/apps/cli/src/commands/cli/__tests__/list.test.ts index 71bdc4266b..78db9752d8 100644 --- a/apps/cli/src/commands/cli/__tests__/list.test.ts +++ b/apps/cli/src/commands/cli/__tests__/list.test.ts @@ -1,7 +1,45 @@ +import fs from "fs" +import os from "os" +import path from "path" +import { EventEmitter } from "events" + +import { openRouterDefaultModelId, providerIdentifiers } from "@roo-code/types" + import { readWorkspaceTaskSessions } from "@/lib/task-history/index.js" -import { isRecord } from "@/lib/utils/guards.js" -import { listSessions, parseFormat } from "../list.js" +import { listModels, listSessions, parseFormat } from "../list.js" + +const extensionHostMock = vi.hoisted(() => ({ + activate: vi.fn(async () => undefined), + dispose: vi.fn(async () => undefined), + options: [] as unknown[], + responses: [] as unknown[], + sendToExtension: vi.fn(), +})) + +vi.mock("@/agent/index.js", () => ({ + ExtensionHost: class extends EventEmitter { + client = { + isInitialized: () => true, + on: vi.fn(() => () => undefined), + } + + constructor(options: unknown) { + super() + extensionHostMock.options.push(options) + } + + activate = extensionHostMock.activate + dispose = extensionHostMock.dispose + + sendToExtension(message: unknown): void { + extensionHostMock.sendToExtension(message) + for (const response of extensionHostMock.responses) { + this.emit("extensionWebviewMessage", response) + } + } + }, +})) vi.mock("@/lib/task-history/index.js", async (importOriginal) => { const actual = await importOriginal() @@ -39,30 +77,88 @@ describe("parseFormat", () => { }) }) -describe("router model extraction", () => { - // This mirrors the extraction logic in requestOpenRouterModels (list.ts:226-228) - const extractOpenRouterModels = (routerModelsRaw: unknown) => { - const routerModels = isRecord(routerModelsRaw) ? routerModelsRaw : {} - const openRouterModels = routerModels.openrouter - return isRecord(openRouterModels) ? openRouterModels : {} - } +describe("listModels", () => { + let tempDir: string + let workspacePath: string + let extensionPath: string - it("extracts openrouter models from valid routerModels", () => { - const models = { "openai/gpt-4.1": { contextWindow: 128000, supportsPromptCache: false } } - const result = extractOpenRouterModels({ openrouter: models }) - expect(result).toEqual(models) + beforeEach(() => { + vi.clearAllMocks() + extensionHostMock.options.length = 0 + extensionHostMock.responses.length = 0 + + tempDir = fs.mkdtempSync(path.join(os.tmpdir(), "roo-list-test-")) + workspacePath = path.join(tempDir, "workspace") + extensionPath = path.join(tempDir, "extension") + fs.mkdirSync(workspacePath) + fs.mkdirSync(extensionPath) + fs.writeFileSync(path.join(extensionPath, "extension.js"), "") }) - it("returns empty object when routerModels is null", () => { - expect(extractOpenRouterModels(null)).toEqual({}) + afterEach(() => { + fs.rmSync(tempDir, { recursive: true, force: true }) + vi.restoreAllMocks() }) - it("returns empty object when openrouter key is missing", () => { - expect(extractOpenRouterModels({ requesty: {} })).toEqual({}) + const captureStdout = async (fn: () => Promise): Promise => { + const stdoutSpy = vi.spyOn(process.stdout, "write").mockImplementation(() => true) + await fn() + return stdoutSpy.mock.calls.map(([chunk]) => String(chunk)).join("") + } + + it("creates a host with resolved paths and returns OpenRouter models", async () => { + const models = { "openai/gpt-4.1": { contextWindow: 128000, supportsPromptCache: false } } + extensionHostMock.responses.push( + { type: "unrelatedMessage" }, + { type: "routerModels", routerModels: { [providerIdentifiers.openrouter]: models } }, + ) + + const output = await captureStdout(() => + listModels({ + format: "json", + workspace: path.relative(process.cwd(), workspacePath), + extension: path.relative(process.cwd(), extensionPath), + apiKey: "test-api-key", + debug: true, + }), + ) + + expect(extensionHostMock.options).toEqual([ + expect.objectContaining({ + mode: "code", + provider: providerIdentifiers.openrouter, + model: openRouterDefaultModelId, + apiKey: "test-api-key", + workspacePath, + extensionPath, + nonInteractive: true, + ephemeral: true, + debug: true, + exitOnComplete: true, + exitOnError: false, + disableOutput: true, + }), + ]) + expect(extensionHostMock.activate).toHaveBeenCalledOnce() + expect(extensionHostMock.sendToExtension).toHaveBeenCalledWith({ + type: "requestRouterModels", + values: { provider: providerIdentifiers.openrouter }, + }) + expect(extensionHostMock.dispose).toHaveBeenCalledOnce() + expect(JSON.parse(output)).toEqual({ models }) }) - it("returns empty object when openrouter value is not a record", () => { - expect(extractOpenRouterModels({ openrouter: "invalid" })).toEqual({}) + it.each([ + ["a malformed routerModels value", null], + ["a malformed OpenRouter value", { [providerIdentifiers.openrouter]: "invalid" }], + ])("returns an empty model record for %s", async (_description, routerModels) => { + extensionHostMock.responses.push({ type: "routerModels", routerModels }) + + const output = await captureStdout(() => + listModels({ format: "json", workspace: workspacePath, extension: extensionPath }), + ) + + expect(JSON.parse(output)).toEqual({ models: {} }) }) }) diff --git a/apps/cli/src/commands/cli/__tests__/run.test.ts b/apps/cli/src/commands/cli/__tests__/run.test.ts index 7b7693a39c..e20d0672c3 100644 --- a/apps/cli/src/commands/cli/__tests__/run.test.ts +++ b/apps/cli/src/commands/cli/__tests__/run.test.ts @@ -2,6 +2,152 @@ import fs from "fs" import path from "path" import os from "os" +import { providerIdentifiers } from "@roo-code/types" +import { DEFAULT_FLAGS, FlagOptions } from "@/types/index.js" +import { + resolveLegacyRequireApproval, + resolveModel, + resolveProvider, + resolveReasoningEffort, + resolveWorkspacePath, + run, +} from "../run.js" + +const runCommandMocks = vi.hoisted(() => ({ + activate: vi.fn(async () => undefined), + dispose: vi.fn(async () => undefined), + loadSettings: vi.fn(), + options: [] as unknown[], + runTask: vi.fn(async () => undefined), +})) + +vi.mock("@/lib/storage/index.js", () => ({ + loadSettings: runCommandMocks.loadSettings, +})) + +vi.mock("@/agent/index.js", () => ({ + ExtensionHost: class { + client = {} + + constructor(options: unknown) { + runCommandMocks.options.push(options) + } + + activate = runCommandMocks.activate + dispose = runCommandMocks.dispose + runTask = runCommandMocks.runTask + }, +})) + +describe("resolveModel", () => { + it("uses the CLI flag before the settings model", () => { + expect(resolveModel("flag-model", "settings-model")).toBe("flag-model") + }) + + it("uses the settings model when the CLI flag is absent", () => { + expect(resolveModel(undefined, "settings-model")).toBe("settings-model") + }) + + it("uses the default model when neither the CLI flag nor settings provide one", () => { + expect(resolveModel()).toBe(DEFAULT_FLAGS.model) + }) +}) + +describe("resolveReasoningEffort", () => { + it("uses CLI, settings, and default values in priority order", () => { + expect(resolveReasoningEffort("high", "low")).toBe("high") + expect(resolveReasoningEffort(undefined, "low")).toBe("low") + expect(resolveReasoningEffort()).toBe(DEFAULT_FLAGS.reasoningEffort) + }) +}) + +describe("resolveProvider", () => { + it("uses CLI, settings, and openrouter values in priority order", () => { + expect(resolveProvider(providerIdentifiers.anthropic, providerIdentifiers.gemini)).toBe( + providerIdentifiers.anthropic, + ) + expect(resolveProvider(undefined, providerIdentifiers.gemini)).toBe(providerIdentifiers.gemini) + expect(resolveProvider()).toBe(providerIdentifiers.openrouter) + }) +}) + +describe("resolveWorkspacePath", () => { + it("resolves the provided workspace path", () => { + expect(resolveWorkspacePath("relative/workspace")).toBe(path.resolve("relative/workspace")) + }) + + it("uses the current working directory when workspace is absent", () => { + expect(resolveWorkspacePath()).toBe(process.cwd()) + }) +}) + +describe("resolveLegacyRequireApproval", () => { + it.each([ + { requireApproval: true, dangerouslySkipPermissions: true, expected: true }, + { requireApproval: false, dangerouslySkipPermissions: false, expected: false }, + { requireApproval: undefined, dangerouslySkipPermissions: false, expected: true }, + { requireApproval: undefined, dangerouslySkipPermissions: true, expected: false }, + { requireApproval: undefined, dangerouslySkipPermissions: undefined, expected: undefined }, + ])( + "resolves requireApproval=$requireApproval and dangerouslySkipPermissions=$dangerouslySkipPermissions", + ({ requireApproval, dangerouslySkipPermissions, expected }) => { + expect(resolveLegacyRequireApproval(requireApproval, dangerouslySkipPermissions)).toBe(expected) + }, + ) +}) + +describe("run command option resolution", () => { + let workspacePath: string + + beforeEach(() => { + vi.clearAllMocks() + runCommandMocks.options.length = 0 + workspacePath = fs.mkdtempSync(path.join(os.tmpdir(), "roo-run-test-")) + }) + + afterEach(() => { + fs.rmSync(workspacePath, { recursive: true, force: true }) + vi.restoreAllMocks() + }) + + it("passes resolved settings and workspace values to the extension host", async () => { + runCommandMocks.loadSettings.mockResolvedValue({ + model: "settings-model", + reasoningEffort: "high", + provider: providerIdentifiers.anthropic, + dangerouslySkipPermissions: false, + }) + const exitSpy = vi.spyOn(process, "exit").mockImplementation(() => undefined as never) + const flags: FlagOptions = { + continue: false, + workspace: path.relative(process.cwd(), workspacePath), + print: true, + stdinPromptStream: false, + signalOnlyExit: false, + debug: false, + requireApproval: false, + exitOnError: false, + apiKey: "test-api-key", + ephemeral: true, + oneshot: false, + } + + await run("test prompt", flags) + + expect(runCommandMocks.options).toEqual([ + expect.objectContaining({ + model: "settings-model", + reasoningEffort: "high", + provider: providerIdentifiers.anthropic, + workspacePath, + nonInteractive: false, + }), + ]) + expect(runCommandMocks.runTask).toHaveBeenCalledWith("test prompt", undefined) + expect(exitSpy).toHaveBeenCalledWith(0) + }) +}) + describe("run command --prompt-file option", () => { let tempDir: string let promptFilePath: string diff --git a/apps/cli/src/commands/cli/list.ts b/apps/cli/src/commands/cli/list.ts index fbd33da2cc..c5fbb4dba9 100644 --- a/apps/cli/src/commands/cli/list.ts +++ b/apps/cli/src/commands/cli/list.ts @@ -6,7 +6,7 @@ import pWaitFor from "p-wait-for" import type { TaskSessionEntry } from "@roo-code/core/cli" import type { Command, ModelRecord, WebviewMessage } from "@roo-code/types" -import { openRouterDefaultModelId } from "@roo-code/types" +import { openRouterDefaultModelId, providerIdentifiers } from "@roo-code/types" import { ExtensionHost, type ExtensionHostOptions } from "@/agent/index.js" import { readWorkspaceTaskSessions } from "@/lib/task-history/index.js" @@ -105,13 +105,13 @@ function outputSessionsText(sessions: SessionLike[]): void { async function createListHost(options: BaseListOptions, hostOptions: ListHostOptions): Promise { const workspacePath = resolveWorkspacePath(options.workspace) const extensionPath = resolveExtensionPath(options.extension) - const apiKey = options.apiKey || getApiKeyFromEnv("openrouter") + const apiKey = options.apiKey || getApiKeyFromEnv(providerIdentifiers.openrouter) const extensionHostOptions: ExtensionHostOptions = { mode: "code", reasoningEffort: undefined, user: null, - provider: "openrouter", + provider: providerIdentifiers.openrouter, model: openRouterDefaultModelId, apiKey, workspacePath, @@ -217,14 +217,14 @@ function requestModes(host: ExtensionHost): Promise { function requestOpenRouterModels(host: ExtensionHost): Promise { return requestFromExtension( host, - { type: "requestRouterModels", values: { provider: "openrouter" } }, + { type: "requestRouterModels", values: { provider: providerIdentifiers.openrouter } }, (message) => { if (message.type !== "routerModels") { return undefined } const routerModels = isRecord(message.routerModels) ? message.routerModels : {} - const openRouterModels = routerModels.openrouter + const openRouterModels = routerModels[providerIdentifiers.openrouter] return isRecord(openRouterModels) ? (openRouterModels as ModelRecord) : {} }, ) diff --git a/apps/cli/src/commands/cli/run.ts b/apps/cli/src/commands/cli/run.ts index 908df9938b..bedb520ed4 100644 --- a/apps/cli/src/commands/cli/run.ts +++ b/apps/cli/src/commands/cli/run.ts @@ -5,10 +5,13 @@ import { fileURLToPath } from "url" import { createElement } from "react" import pWaitFor from "p-wait-for" +import { providerIdentifiers } from "@roo-code/types" import { setLogger } from "@roo-code/vscode-shim" import { FlagOptions, + ReasoningEffortFlagOptions, + SupportedProvider, isSupportedProvider, supportedProviders, DEFAULT_FLAGS, @@ -49,6 +52,35 @@ function normalizeError(error: unknown): Error { return error instanceof Error ? error : new Error(String(error)) } +export function resolveModel(flagModel?: string, settingsModel?: string): string { + return flagModel || settingsModel || DEFAULT_FLAGS.model +} + +export function resolveReasoningEffort( + flagReasoningEffort?: ReasoningEffortFlagOptions, + settingsReasoningEffort?: ReasoningEffortFlagOptions, +): ReasoningEffortFlagOptions { + return flagReasoningEffort || settingsReasoningEffort || DEFAULT_FLAGS.reasoningEffort +} + +export function resolveProvider( + flagProvider?: SupportedProvider, + settingsProvider?: SupportedProvider, +): SupportedProvider { + return flagProvider ?? settingsProvider ?? providerIdentifiers.openrouter +} + +export function resolveWorkspacePath(workspace?: string): string { + return workspace ? path.resolve(workspace) : process.cwd() +} + +export function resolveLegacyRequireApproval( + requireApproval?: boolean, + dangerouslySkipPermissions?: boolean, +): boolean | undefined { + return requireApproval ?? (dangerouslySkipPermissions === undefined ? undefined : !dangerouslySkipPermissions) +} + export async function run(promptArg: string | undefined, flagOptions: FlagOptions) { setLogger({ info: () => {}, @@ -119,14 +151,14 @@ export async function run(promptArg: string | undefined, flagOptions: FlagOption // Determine effective values: CLI flags > settings file > DEFAULT_FLAGS. const effectiveMode = flagOptions.mode || settings.mode || DEFAULT_FLAGS.mode - const effectiveModel = flagOptions.model || settings.model || DEFAULT_FLAGS.model - const effectiveReasoningEffort = - flagOptions.reasoningEffort || settings.reasoningEffort || DEFAULT_FLAGS.reasoningEffort - const effectiveProvider = flagOptions.provider ?? settings.provider ?? "openrouter" - const effectiveWorkspacePath = flagOptions.workspace ? path.resolve(flagOptions.workspace) : process.cwd() - const legacyRequireApprovalFromSettings = - settings.requireApproval ?? - (settings.dangerouslySkipPermissions === undefined ? undefined : !settings.dangerouslySkipPermissions) + const effectiveModel = resolveModel(flagOptions.model, settings.model) + const effectiveReasoningEffort = resolveReasoningEffort(flagOptions.reasoningEffort, settings.reasoningEffort) + const effectiveProvider = resolveProvider(flagOptions.provider, settings.provider) + const effectiveWorkspacePath = resolveWorkspacePath(flagOptions.workspace) + const legacyRequireApprovalFromSettings = resolveLegacyRequireApproval( + settings.requireApproval, + settings.dangerouslySkipPermissions, + ) const effectiveRequireApproval = flagOptions.requireApproval || legacyRequireApprovalFromSettings || false const effectiveExitOnComplete = flagOptions.print || flagOptions.oneshot || settings.oneshot || false const rawConsecutiveMistakeLimit = diff --git a/apps/cli/src/lib/utils/__tests__/context-window.test.ts b/apps/cli/src/lib/utils/__tests__/context-window.test.ts new file mode 100644 index 0000000000..8d33ef5e2b --- /dev/null +++ b/apps/cli/src/lib/utils/__tests__/context-window.test.ts @@ -0,0 +1,40 @@ +import { providerIdentifiers, type ProviderSettings } from "@roo-code/types" + +import { DEFAULT_CONTEXT_WINDOW, getContextWindow } from "../context-window.js" + +describe("getContextWindow", () => { + it.each([ + [providerIdentifiers.openrouter, "openRouterModelId"], + [providerIdentifiers.ollama, "ollamaModelId"], + [providerIdentifiers.lmstudio, "lmStudioModelId"], + [providerIdentifiers.openai, "openAiModelId"], + [providerIdentifiers.requesty, "requestyModelId"], + [providerIdentifiers.unbound, "unboundModelId"], + [providerIdentifiers.litellm, "litellmModelId"], + [providerIdentifiers.vercelAiGateway, "vercelAiGatewayModelId"], + [providerIdentifiers.opencodeGo, "opencodeGoModelId"], + [providerIdentifiers.kenari, "kenariModelId"], + [providerIdentifiers.zooGateway, "zooGatewayModelId"], + ] as const)("uses the provider-specific model field for %s", (provider, modelField) => { + const config = { apiProvider: provider, [modelField]: "selected-model" } as ProviderSettings + const routerModels = { [provider]: { "selected-model": { contextWindow: 123_456 } } } + + expect(getContextWindow(routerModels, config)).toBe(123_456) + }) + + it("uses apiModelId for providers without a specialized model field", () => { + const config: ProviderSettings = { + apiProvider: providerIdentifiers.anthropic, + apiModelId: "selected-model", + } + const routerModels = { + [providerIdentifiers.anthropic]: { "selected-model": { contextWindow: 64_000 } }, + } + + expect(getContextWindow(routerModels, config)).toBe(64_000) + }) + + it("returns the default when the selected model is unavailable", () => { + expect(getContextWindow({}, { apiProvider: providerIdentifiers.openrouter })).toBe(DEFAULT_CONTEXT_WINDOW) + }) +}) diff --git a/apps/cli/src/lib/utils/__tests__/provider.test.ts b/apps/cli/src/lib/utils/__tests__/provider.test.ts index 70d8a2a555..db44174f45 100644 --- a/apps/cli/src/lib/utils/__tests__/provider.test.ts +++ b/apps/cli/src/lib/utils/__tests__/provider.test.ts @@ -1,4 +1,47 @@ -import { getApiKeyFromEnv } from "../provider.js" +import { providerIdentifiers } from "@roo-code/types" + +import { getApiKeyFromEnv, getEnvVarName, getProviderSettings } from "../provider.js" + +describe("provider configuration", () => { + it.each([ + [providerIdentifiers.anthropic, "ANTHROPIC_API_KEY"], + [providerIdentifiers.openaiNative, "OPENAI_API_KEY"], + [providerIdentifiers.gemini, "GOOGLE_API_KEY"], + [providerIdentifiers.openrouter, "OPENROUTER_API_KEY"], + [providerIdentifiers.vercelAiGateway, "VERCEL_AI_GATEWAY_API_KEY"], + ] as const)("maps canonical provider %s to %s", (provider, envVarName) => { + expect(getEnvVarName(provider)).toBe(envVarName) + }) + + it.each([ + [ + providerIdentifiers.anthropic, + { apiProvider: providerIdentifiers.anthropic, apiKey: "key", apiModelId: "model" }, + ], + [ + providerIdentifiers.openaiNative, + { apiProvider: providerIdentifiers.openaiNative, openAiNativeApiKey: "key", apiModelId: "model" }, + ], + [ + providerIdentifiers.gemini, + { apiProvider: providerIdentifiers.gemini, geminiApiKey: "key", apiModelId: "model" }, + ], + [ + providerIdentifiers.openrouter, + { apiProvider: providerIdentifiers.openrouter, openRouterApiKey: "key", openRouterModelId: "model" }, + ], + [ + providerIdentifiers.vercelAiGateway, + { + apiProvider: providerIdentifiers.vercelAiGateway, + vercelAiGatewayApiKey: "key", + vercelAiGatewayModelId: "model", + }, + ], + ] as const)("builds settings for canonical provider %s", (provider, expected) => { + expect(getProviderSettings(provider, "key", "model")).toEqual(expected) + }) +}) describe("getApiKeyFromEnv", () => { const originalEnv = process.env diff --git a/apps/cli/src/lib/utils/context-window.ts b/apps/cli/src/lib/utils/context-window.ts index 5cd58b55a8..1d6402c525 100644 --- a/apps/cli/src/lib/utils/context-window.ts +++ b/apps/cli/src/lib/utils/context-window.ts @@ -1,4 +1,4 @@ -import type { ProviderSettings } from "@roo-code/types" +import { providerIdentifiers, retiredProviderIdentifiers, type ProviderSettings } from "@roo-code/types" import type { RouterModels } from "@/ui/store.js" @@ -36,24 +36,61 @@ export function getContextWindow(routerModels: RouterModels | null, apiConfigura */ function getModelIdForProvider(config: ProviderSettings): string | undefined { switch (config.apiProvider) { - case "openrouter": + case providerIdentifiers.openrouter: return config.openRouterModelId - case "ollama": + case providerIdentifiers.ollama: return config.ollamaModelId - case "lmstudio": + case providerIdentifiers.lmstudio: return config.lmStudioModelId - case "openai": + case providerIdentifiers.openai: return config.openAiModelId - case "requesty": + case providerIdentifiers.requesty: return config.requestyModelId - case "unbound": + case providerIdentifiers.unbound: return config.unboundModelId - case "litellm": + case providerIdentifiers.litellm: return config.litellmModelId - case "vercel-ai-gateway": + case providerIdentifiers.vercelAiGateway: return config.vercelAiGatewayModelId - default: - // For anthropic, bedrock, vertex, gemini, xai, etc. + case providerIdentifiers.opencodeGo: + return config.opencodeGoModelId + case providerIdentifiers.kenari: + return config.kenariModelId + case providerIdentifiers.zooGateway: + return config.zooGatewayModelId + case providerIdentifiers.anthropic: + case providerIdentifiers.bedrock: + case providerIdentifiers.baseten: + case providerIdentifiers.deepseek: + case providerIdentifiers.fireworks: + case providerIdentifiers.friendli: + case providerIdentifiers.gemini: + case providerIdentifiers.geminiCli: + case providerIdentifiers.mistral: + case providerIdentifiers.moonshot: + case providerIdentifiers.kimiCode: + case providerIdentifiers.minimax: + case providerIdentifiers.mimo: + case providerIdentifiers.openaiCodex: + case providerIdentifiers.openaiNative: + case providerIdentifiers.poe: + case providerIdentifiers.qwenCode: + case providerIdentifiers.sambanova: + case providerIdentifiers.vertex: + case providerIdentifiers.xai: + case providerIdentifiers.zai: + case retiredProviderIdentifiers.cerebras: + case retiredProviderIdentifiers.chutes: + case retiredProviderIdentifiers.deepinfra: + case retiredProviderIdentifiers.doubao: + case retiredProviderIdentifiers.featherless: + case retiredProviderIdentifiers.groq: + case retiredProviderIdentifiers.huggingface: + case retiredProviderIdentifiers.ioIntelligence: + case retiredProviderIdentifiers.roo: + case providerIdentifiers.vscodeLm: + case providerIdentifiers.fakeAi: + case undefined: return config.apiModelId } } diff --git a/apps/cli/src/lib/utils/provider.ts b/apps/cli/src/lib/utils/provider.ts index 26beaf90c4..7cb7b30ffb 100644 --- a/apps/cli/src/lib/utils/provider.ts +++ b/apps/cli/src/lib/utils/provider.ts @@ -1,13 +1,13 @@ -import { RooCodeSettings } from "@roo-code/types" +import { providerIdentifiers, type RooCodeSettings } from "@roo-code/types" import type { SupportedProvider } from "@/types/index.js" const envVarMap: Record = { - anthropic: "ANTHROPIC_API_KEY", - "openai-native": "OPENAI_API_KEY", - gemini: "GOOGLE_API_KEY", - openrouter: "OPENROUTER_API_KEY", - "vercel-ai-gateway": "VERCEL_AI_GATEWAY_API_KEY", + [providerIdentifiers.anthropic]: "ANTHROPIC_API_KEY", + [providerIdentifiers.openaiNative]: "OPENAI_API_KEY", + [providerIdentifiers.gemini]: "GOOGLE_API_KEY", + [providerIdentifiers.openrouter]: "OPENROUTER_API_KEY", + [providerIdentifiers.vercelAiGateway]: "VERCEL_AI_GATEWAY_API_KEY", } export function getEnvVarName(provider: SupportedProvider): string { @@ -27,23 +27,23 @@ export function getProviderSettings( const config: RooCodeSettings = { apiProvider: provider } switch (provider) { - case "anthropic": + case providerIdentifiers.anthropic: if (apiKey) config.apiKey = apiKey if (model) config.apiModelId = model break - case "openai-native": + case providerIdentifiers.openaiNative: if (apiKey) config.openAiNativeApiKey = apiKey if (model) config.apiModelId = model break - case "gemini": + case providerIdentifiers.gemini: if (apiKey) config.geminiApiKey = apiKey if (model) config.apiModelId = model break - case "openrouter": + case providerIdentifiers.openrouter: if (apiKey) config.openRouterApiKey = apiKey if (model) config.openRouterModelId = model break - case "vercel-ai-gateway": + case providerIdentifiers.vercelAiGateway: if (apiKey) config.vercelAiGatewayApiKey = apiKey if (model) config.vercelAiGatewayModelId = model break diff --git a/apps/cli/src/types/__tests__/types.test.ts b/apps/cli/src/types/__tests__/types.test.ts index 1e54b5069e..5ed0c84016 100644 --- a/apps/cli/src/types/__tests__/types.test.ts +++ b/apps/cli/src/types/__tests__/types.test.ts @@ -1,5 +1,19 @@ +import { providerIdentifiers } from "@roo-code/types" + import { isSupportedProvider, supportedProviders } from "../types.js" +describe("supportedProviders", () => { + it("contains the canonical identifiers for the CLI provider subset", () => { + expect(supportedProviders).toEqual([ + providerIdentifiers.anthropic, + providerIdentifiers.openaiNative, + providerIdentifiers.gemini, + providerIdentifiers.openrouter, + providerIdentifiers.vercelAiGateway, + ]) + }) +}) + describe("isSupportedProvider", () => { it.each(supportedProviders)("returns true for supported provider '%s'", (provider) => { expect(isSupportedProvider(provider)).toBe(true) @@ -22,25 +36,25 @@ describe("provider resolution fallback", () => { it("defaults to openrouter when no flag or setting is provided", () => { const flagProvider = undefined const settingsProvider = undefined - const effectiveProvider = flagProvider ?? settingsProvider ?? "openrouter" + const effectiveProvider = flagProvider ?? settingsProvider ?? providerIdentifiers.openrouter - expect(effectiveProvider).toBe("openrouter") + expect(effectiveProvider).toBe(providerIdentifiers.openrouter) expect(isSupportedProvider(effectiveProvider)).toBe(true) }) it("uses flag provider over settings and default", () => { - const flagProvider = "anthropic" - const settingsProvider = "gemini" - const effectiveProvider = flagProvider ?? settingsProvider ?? "openrouter" + const flagProvider = providerIdentifiers.anthropic + const settingsProvider = providerIdentifiers.gemini + const effectiveProvider = flagProvider ?? settingsProvider ?? providerIdentifiers.openrouter - expect(effectiveProvider).toBe("anthropic") + expect(effectiveProvider).toBe(providerIdentifiers.anthropic) }) it("uses settings provider when flag is not provided", () => { const flagProvider = undefined - const settingsProvider = "gemini" - const effectiveProvider = flagProvider ?? settingsProvider ?? "openrouter" + const settingsProvider = providerIdentifiers.gemini + const effectiveProvider = flagProvider ?? settingsProvider ?? providerIdentifiers.openrouter - expect(effectiveProvider).toBe("gemini") + expect(effectiveProvider).toBe(providerIdentifiers.gemini) }) }) diff --git a/apps/cli/src/types/types.ts b/apps/cli/src/types/types.ts index 0a9f3d2259..999c7b655a 100644 --- a/apps/cli/src/types/types.ts +++ b/apps/cli/src/types/types.ts @@ -1,12 +1,12 @@ -import type { ProviderName, ReasoningEffortExtended } from "@roo-code/types" +import { providerIdentifiers, type ProviderName, type ReasoningEffortExtended } from "@roo-code/types" import type { OutputFormat } from "./json-events.js" export const supportedProviders = [ - "anthropic", - "openai-native", - "gemini", - "openrouter", - "vercel-ai-gateway", + providerIdentifiers.anthropic, + providerIdentifiers.openaiNative, + providerIdentifiers.gemini, + providerIdentifiers.openrouter, + providerIdentifiers.vercelAiGateway, ] as const satisfies ProviderName[] export type SupportedProvider = (typeof supportedProviders)[number] diff --git a/packages/types/src/__tests__/provider-identifiers.test.ts b/packages/types/src/__tests__/provider-identifiers.test.ts index b3640a8f5d..870ce77d78 100644 --- a/packages/types/src/__tests__/provider-identifiers.test.ts +++ b/packages/types/src/__tests__/provider-identifiers.test.ts @@ -11,6 +11,7 @@ import { isProviderName, isRetiredProvider, localProviders, + MODELS_BY_PROVIDER, providerIdentifiers, providerNames, providerNamesSchema, @@ -113,6 +114,12 @@ describe("provider identifiers", () => { expect(fauxProviders).toEqual([providerIdentifiers.fakeAi]) }) + it("keeps model provider ids aligned with their keys", () => { + for (const [identifier, providerModels] of Object.entries(MODELS_BY_PROVIDER)) { + expect(providerModels.id).toBe(identifier) + } + }) + it("preserves provider category type guards", () => { for (const identifier of dynamicProviders) { expect(isDynamicProvider(identifier)).toBe(true) diff --git a/packages/types/src/provider-settings.ts b/packages/types/src/provider-settings.ts index e17cd5ddbc..99b75de2e4 100644 --- a/packages/types/src/provider-settings.ts +++ b/packages/types/src/provider-settings.ts @@ -428,40 +428,40 @@ const defaultSchema = z.object({ }) export const providerSettingsSchemaDiscriminated = z.discriminatedUnion("apiProvider", [ - anthropicSchema.merge(z.object({ apiProvider: z.literal("anthropic") })), - openRouterSchema.merge(z.object({ apiProvider: z.literal("openrouter") })), - bedrockSchema.merge(z.object({ apiProvider: z.literal("bedrock") })), - vertexSchema.merge(z.object({ apiProvider: z.literal("vertex") })), - openAiSchema.merge(z.object({ apiProvider: z.literal("openai") })), - ollamaSchema.merge(z.object({ apiProvider: z.literal("ollama") })), - vsCodeLmSchema.merge(z.object({ apiProvider: z.literal("vscode-lm") })), - lmStudioSchema.merge(z.object({ apiProvider: z.literal("lmstudio") })), - geminiSchema.merge(z.object({ apiProvider: z.literal("gemini") })), - geminiCliSchema.merge(z.object({ apiProvider: z.literal("gemini-cli") })), - openAiCodexSchema.merge(z.object({ apiProvider: z.literal("openai-codex") })), - openAiNativeSchema.merge(z.object({ apiProvider: z.literal("openai-native") })), - mistralSchema.merge(z.object({ apiProvider: z.literal("mistral") })), - deepSeekSchema.merge(z.object({ apiProvider: z.literal("deepseek") })), - poeSchema.merge(z.object({ apiProvider: z.literal("poe") })), - moonshotSchema.merge(z.object({ apiProvider: z.literal("moonshot") })), - kimiCodeSchema.merge(z.object({ apiProvider: z.literal("kimi-code") })), - minimaxSchema.merge(z.object({ apiProvider: z.literal("minimax") })), - mimoSchema.merge(z.object({ apiProvider: z.literal("mimo") })), - requestySchema.merge(z.object({ apiProvider: z.literal("requesty") })), - unboundSchema.merge(z.object({ apiProvider: z.literal("unbound") })), - fakeAiSchema.merge(z.object({ apiProvider: z.literal("fake-ai") })), - xaiSchema.merge(z.object({ apiProvider: z.literal("xai") })), - basetenSchema.merge(z.object({ apiProvider: z.literal("baseten") })), - litellmSchema.merge(z.object({ apiProvider: z.literal("litellm") })), - sambaNovaSchema.merge(z.object({ apiProvider: z.literal("sambanova") })), - zaiSchema.merge(z.object({ apiProvider: z.literal("zai") })), - fireworksSchema.merge(z.object({ apiProvider: z.literal("fireworks") })), - friendliSchema.merge(z.object({ apiProvider: z.literal("friendli") })), - qwenCodeSchema.merge(z.object({ apiProvider: z.literal("qwen-code") })), - vercelAiGatewaySchema.merge(z.object({ apiProvider: z.literal("vercel-ai-gateway") })), - opencodeGoSchema.merge(z.object({ apiProvider: z.literal("opencode-go") })), - kenariSchema.merge(z.object({ apiProvider: z.literal("kenari") })), - zooGatewaySchema.merge(z.object({ apiProvider: z.literal("zoo-gateway") })), + anthropicSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.anthropic) })), + openRouterSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.openrouter) })), + bedrockSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.bedrock) })), + vertexSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.vertex) })), + openAiSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.openai) })), + ollamaSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.ollama) })), + vsCodeLmSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.vscodeLm) })), + lmStudioSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.lmstudio) })), + geminiSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.gemini) })), + geminiCliSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.geminiCli) })), + openAiCodexSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.openaiCodex) })), + openAiNativeSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.openaiNative) })), + mistralSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.mistral) })), + deepSeekSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.deepseek) })), + poeSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.poe) })), + moonshotSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.moonshot) })), + kimiCodeSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.kimiCode) })), + minimaxSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.minimax) })), + mimoSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.mimo) })), + requestySchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.requesty) })), + unboundSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.unbound) })), + fakeAiSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.fakeAi) })), + xaiSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.xai) })), + basetenSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.baseten) })), + litellmSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.litellm) })), + sambaNovaSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.sambanova) })), + zaiSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.zai) })), + fireworksSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.fireworks) })), + friendliSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.friendli) })), + qwenCodeSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.qwenCode) })), + vercelAiGatewaySchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.vercelAiGateway) })), + opencodeGoSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.opencodeGo) })), + kenariSchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.kenari) })), + zooGatewaySchema.merge(z.object({ apiProvider: z.literal(providerIdentifiers.zooGateway) })), defaultSchema, ]) @@ -553,37 +553,37 @@ export const isTypicalProvider = (key: unknown): key is TypicalProvider => isProviderName(key) && !isInternalProvider(key) && !isCustomProvider(key) && !isFauxProvider(key) export const modelIdKeysByProvider: Record = { - anthropic: "apiModelId", - openrouter: "openRouterModelId", - bedrock: "apiModelId", - vertex: "apiModelId", - "openai-codex": "apiModelId", - "openai-native": "openAiModelId", - ollama: "ollamaModelId", - lmstudio: "lmStudioModelId", - gemini: "apiModelId", - "gemini-cli": "apiModelId", - mistral: "apiModelId", - moonshot: "apiModelId", - "kimi-code": "apiModelId", - minimax: "apiModelId", - mimo: "apiModelId", - deepseek: "apiModelId", - poe: "apiModelId", - "qwen-code": "apiModelId", - requesty: "requestyModelId", - unbound: "unboundModelId", - xai: "apiModelId", - baseten: "apiModelId", - litellm: "litellmModelId", - sambanova: "apiModelId", - zai: "apiModelId", - fireworks: "apiModelId", - friendli: "apiModelId", - "vercel-ai-gateway": "vercelAiGatewayModelId", - "opencode-go": "opencodeGoModelId", - kenari: "kenariModelId", - "zoo-gateway": "zooGatewayModelId", + [providerIdentifiers.anthropic]: "apiModelId", + [providerIdentifiers.openrouter]: "openRouterModelId", + [providerIdentifiers.bedrock]: "apiModelId", + [providerIdentifiers.vertex]: "apiModelId", + [providerIdentifiers.openaiCodex]: "apiModelId", + [providerIdentifiers.openaiNative]: "openAiModelId", + [providerIdentifiers.ollama]: "ollamaModelId", + [providerIdentifiers.lmstudio]: "lmStudioModelId", + [providerIdentifiers.gemini]: "apiModelId", + [providerIdentifiers.geminiCli]: "apiModelId", + [providerIdentifiers.mistral]: "apiModelId", + [providerIdentifiers.moonshot]: "apiModelId", + [providerIdentifiers.kimiCode]: "apiModelId", + [providerIdentifiers.minimax]: "apiModelId", + [providerIdentifiers.mimo]: "apiModelId", + [providerIdentifiers.deepseek]: "apiModelId", + [providerIdentifiers.poe]: "apiModelId", + [providerIdentifiers.qwenCode]: "apiModelId", + [providerIdentifiers.requesty]: "requestyModelId", + [providerIdentifiers.unbound]: "unboundModelId", + [providerIdentifiers.xai]: "apiModelId", + [providerIdentifiers.baseten]: "apiModelId", + [providerIdentifiers.litellm]: "litellmModelId", + [providerIdentifiers.sambanova]: "apiModelId", + [providerIdentifiers.zai]: "apiModelId", + [providerIdentifiers.fireworks]: "apiModelId", + [providerIdentifiers.friendli]: "apiModelId", + [providerIdentifiers.vercelAiGateway]: "vercelAiGatewayModelId", + [providerIdentifiers.opencodeGo]: "opencodeGoModelId", + [providerIdentifiers.kenari]: "kenariModelId", + [providerIdentifiers.zooGateway]: "zooGatewayModelId", } /** @@ -653,106 +653,125 @@ export const getApiProtocol = (provider: ProviderName | undefined, modelId?: str */ export const MODELS_BY_PROVIDER: Record< - Exclude, + Exclude< + ProviderName, + typeof providerIdentifiers.fakeAi | typeof providerIdentifiers.geminiCli | typeof providerIdentifiers.openai + >, { id: ProviderName; label: string; models: string[] } > = { - anthropic: { - id: "anthropic", + [providerIdentifiers.anthropic]: { + id: providerIdentifiers.anthropic, label: "Anthropic", models: Object.keys(anthropicModels), }, - bedrock: { - id: "bedrock", + [providerIdentifiers.bedrock]: { + id: providerIdentifiers.bedrock, label: "Amazon Bedrock", models: Object.keys(bedrockModels), }, - deepseek: { - id: "deepseek", + [providerIdentifiers.deepseek]: { + id: providerIdentifiers.deepseek, label: "DeepSeek", models: Object.keys(deepSeekModels), }, - fireworks: { - id: "fireworks", + [providerIdentifiers.fireworks]: { + id: providerIdentifiers.fireworks, label: "Fireworks", models: Object.keys(fireworksModels), }, - friendli: { - id: "friendli", + [providerIdentifiers.friendli]: { + id: providerIdentifiers.friendli, label: "Friendli", models: Object.keys(friendliModels), }, - gemini: { - id: "gemini", + [providerIdentifiers.gemini]: { + id: providerIdentifiers.gemini, label: "Google Gemini", models: Object.keys(geminiModels), }, - mistral: { - id: "mistral", + [providerIdentifiers.mistral]: { + id: providerIdentifiers.mistral, label: "Mistral", models: Object.keys(mistralModels), }, - moonshot: { - id: "moonshot", + [providerIdentifiers.moonshot]: { + id: providerIdentifiers.moonshot, label: "Moonshot", models: Object.keys(moonshotModels), }, - "kimi-code": { - id: "kimi-code", + [providerIdentifiers.kimiCode]: { + id: providerIdentifiers.kimiCode, label: "Kimi Code", models: [], }, - minimax: { - id: "minimax", + [providerIdentifiers.minimax]: { + id: providerIdentifiers.minimax, label: "MiniMax", models: Object.keys(minimaxModels), }, - mimo: { - id: "mimo", + [providerIdentifiers.mimo]: { + id: providerIdentifiers.mimo, label: "Xiaomi MiMo", models: Object.keys(mimoModels), }, - "openai-codex": { - id: "openai-codex", + [providerIdentifiers.openaiCodex]: { + id: providerIdentifiers.openaiCodex, label: "OpenAI - ChatGPT Plus/Pro", models: Object.keys(openAiCodexModels), }, - "openai-native": { - id: "openai-native", + [providerIdentifiers.openaiNative]: { + id: providerIdentifiers.openaiNative, label: "OpenAI", models: Object.keys(openAiNativeModels), }, - "qwen-code": { id: "qwen-code", label: "Qwen Code", models: Object.keys(qwenCodeModels) }, - sambanova: { - id: "sambanova", + [providerIdentifiers.qwenCode]: { + id: providerIdentifiers.qwenCode, + label: "Qwen Code", + models: Object.keys(qwenCodeModels), + }, + [providerIdentifiers.sambanova]: { + id: providerIdentifiers.sambanova, label: "SambaNova", models: Object.keys(sambaNovaModels), }, - vertex: { - id: "vertex", + [providerIdentifiers.vertex]: { + id: providerIdentifiers.vertex, label: "GCP Vertex AI", models: Object.keys(vertexModels), }, - "vscode-lm": { - id: "vscode-lm", + [providerIdentifiers.vscodeLm]: { + id: providerIdentifiers.vscodeLm, label: "VS Code LM API", models: Object.keys(vscodeLlmModels), }, - xai: { id: "xai", label: "xAI (Grok)", models: Object.keys(xaiModels) }, - zai: { id: "zai", label: "Z.ai", models: Object.keys(internationalZAiModels) }, - baseten: { id: "baseten", label: "Baseten", models: Object.keys(basetenModels) }, + [providerIdentifiers.xai]: { id: providerIdentifiers.xai, label: "xAI (Grok)", models: Object.keys(xaiModels) }, + [providerIdentifiers.zai]: { + id: providerIdentifiers.zai, + label: "Z.ai", + models: Object.keys(internationalZAiModels), + }, + [providerIdentifiers.baseten]: { + id: providerIdentifiers.baseten, + label: "Baseten", + models: Object.keys(basetenModels), + }, // Dynamic providers; models pulled from remote APIs. - poe: { id: "poe", label: "Poe", models: [] }, - litellm: { id: "litellm", label: "LiteLLM", models: [] }, - openrouter: { id: "openrouter", label: "OpenRouter", models: [] }, - requesty: { id: "requesty", label: "Requesty", models: [] }, - unbound: { id: "unbound", label: "Unbound", models: [] }, - "vercel-ai-gateway": { id: "vercel-ai-gateway", label: "Vercel AI Gateway", models: [] }, - "opencode-go": { id: "opencode-go", label: "Opencode Go", models: [] }, - kenari: { id: "kenari", label: "Kenari", models: [] }, - "zoo-gateway": { id: "zoo-gateway", label: "Zoo Gateway", models: [] }, + [providerIdentifiers.poe]: { id: providerIdentifiers.poe, label: "Poe", models: [] }, + [providerIdentifiers.litellm]: { id: providerIdentifiers.litellm, label: "LiteLLM", models: [] }, + [providerIdentifiers.openrouter]: { id: providerIdentifiers.openrouter, label: "OpenRouter", models: [] }, + [providerIdentifiers.requesty]: { id: providerIdentifiers.requesty, label: "Requesty", models: [] }, + [providerIdentifiers.unbound]: { id: providerIdentifiers.unbound, label: "Unbound", models: [] }, + [providerIdentifiers.vercelAiGateway]: { + id: providerIdentifiers.vercelAiGateway, + label: "Vercel AI Gateway", + models: [], + }, + [providerIdentifiers.opencodeGo]: { id: providerIdentifiers.opencodeGo, label: "Opencode Go", models: [] }, + [providerIdentifiers.kenari]: { id: providerIdentifiers.kenari, label: "Kenari", models: [] }, + [providerIdentifiers.zooGateway]: { id: providerIdentifiers.zooGateway, label: "Zoo Gateway", models: [] }, // Local providers; models discovered from localhost endpoints. - lmstudio: { id: "lmstudio", label: "LM Studio", models: [] }, - ollama: { id: "ollama", label: "Ollama", models: [] }, + [providerIdentifiers.lmstudio]: { id: providerIdentifiers.lmstudio, label: "LM Studio", models: [] }, + [providerIdentifiers.ollama]: { id: providerIdentifiers.ollama, label: "Ollama", models: [] }, }