diff --git a/.changeset/tidy-queens-pay.md b/.changeset/tidy-queens-pay.md new file mode 100644 index 0000000000..750a58c789 --- /dev/null +++ b/.changeset/tidy-queens-pay.md @@ -0,0 +1,5 @@ +--- +"roo-cline": patch +--- + +Fix usage tracking for SiliconFlow etc diff --git a/.eslintrc.json b/.eslintrc.json index bae7854a6e..e967b58a03 100644 --- a/.eslintrc.json +++ b/.eslintrc.json @@ -19,5 +19,5 @@ "no-throw-literal": "warn", "semi": "off" }, - "ignorePatterns": ["out", "dist", "**/*.d.ts"] + "ignorePatterns": ["out", "dist", "**/*.d.ts", "!roo-code.d.ts"] } diff --git a/.github/workflows/changeset-release.yml b/.github/workflows/changeset-release.yml index 462516365b..a2bcd3f039 100644 --- a/.github/workflows/changeset-release.yml +++ b/.github/workflows/changeset-release.yml @@ -37,7 +37,7 @@ jobs: cache: 'npm' - name: Install Dependencies - run: npm run install:ci + run: npm run install:all # Check if there are any new changesets to process - name: Check for changesets diff --git a/.github/workflows/code-qa.yml b/.github/workflows/code-qa.yml index b7292dd9ee..2e269cbd79 100644 --- a/.github/workflows/code-qa.yml +++ b/.github/workflows/code-qa.yml @@ -20,7 +20,7 @@ jobs: node-version: '18' cache: 'npm' - name: Install dependencies - run: npm run install:ci + run: npm run install:all - name: Compile run: npm run compile - name: Check types @@ -39,7 +39,7 @@ jobs: node-version: '18' cache: 'npm' - name: Install dependencies - run: npm run install:ci + run: npm run install:all - name: Run knip checks run: npm run knip @@ -54,7 +54,7 @@ jobs: node-version: '18' cache: 'npm' - name: Install dependencies - run: npm run install:ci + run: npm run install:all - name: Run unit tests run: npx jest --silent @@ -69,7 +69,7 @@ jobs: node-version: '18' cache: 'npm' - name: Install dependencies - run: npm run install:ci + run: npm run install:all - name: Run unit tests working-directory: webview-ui run: npx jest --silent @@ -109,7 +109,7 @@ jobs: node-version: '18' cache: 'npm' - name: Install dependencies - run: npm run install:ci + run: npm run install:all - name: Create env.integration file working-directory: e2e run: echo "OPENROUTER_API_KEY=${{ secrets.OPENROUTER_API_KEY }}" > .env.integration diff --git a/.github/workflows/marketplace-publish.yml b/.github/workflows/marketplace-publish.yml index 4ecd2af7a2..6c9fafe9c9 100644 --- a/.github/workflows/marketplace-publish.yml +++ b/.github/workflows/marketplace-publish.yml @@ -11,7 +11,7 @@ jobs: publish-extension: runs-on: ubuntu-latest permissions: - contents: write # Required for pushing tags + contents: write # Required for pushing tags. if: > ( github.event_name == 'pull_request' && github.event.pull_request.base.ref == 'main' && @@ -33,13 +33,9 @@ jobs: - name: Install Dependencies run: | npm install -g vsce ovsx - npm run install:ci - + npm run install:all - name: Create .env file - run: | - echo "# PostHog API Keys for telemetry" > .env - echo "POSTHOG_API_KEY=${{ secrets.POSTHOG_API_KEY }}" >> .env - + run: echo "POSTHOG_API_KEY=${{ secrets.POSTHOG_API_KEY }}" >> .env - name: Package Extension run: | current_package_version=$(node -p "require('./package.json').version") diff --git a/README.md b/README.md index 59ca6d699f..f907ea480e 100644 --- a/README.md +++ b/README.md @@ -115,37 +115,40 @@ Make Roo Code work your way with: ## Local Setup & Development 1. **Clone** the repo: - ```bash - git clone https://github.com/RooVetGit/Roo-Code.git - ``` + +```sh +git clone https://github.com/RooVetGit/Roo-Code.git +``` + 2. **Install dependencies**: - ```bash - npm run install:all - ``` -if that fails, try: - ```bash - npm run install:ci - ``` +```sh +npm run install:all +``` -3. **Build** the extension: - ```bash - npm run build - ``` - - A `.vsix` file will appear in the `bin/` directory. -4. **Install** the `.vsix` manually if desired: - ```bash - code --install-extension bin/roo-code-4.0.0.vsix - ``` -5. **Start the webview (Vite/React app with HMR)**: - ```bash - npm run dev - ``` -6. **Debug**: - - Press `F5` (or **Run** → **Start Debugging**) in VSCode to open a new session with Roo Code loaded. +3. **Start the webview (Vite/React app with HMR)**: + +```sh +npm run dev +``` + +4. **Debug**: + Press `F5` (or **Run** → **Start Debugging**) in VSCode to open a new session with Roo Code loaded. Changes to the webview will appear immediately. Changes to the core extension will require a restart of the extension host. +Alternatively you can build a .vsix and install it directly in VSCode: + +```sh +npm run build +``` + +A `.vsix` file will appear in the `bin/` directory which can be installed with: + +```sh +code --install-extension bin/roo-cline-.vsix +``` + We use [changesets](https://github.com/changesets/changesets) for versioning and publishing. Check our `CHANGELOG.md` for release notes. --- diff --git a/e2e/VSCODE_INTEGRATION_TESTS.md b/e2e/VSCODE_INTEGRATION_TESTS.md index 25f54492de..a36b240381 100644 --- a/e2e/VSCODE_INTEGRATION_TESTS.md +++ b/e2e/VSCODE_INTEGRATION_TESTS.md @@ -58,9 +58,9 @@ The following global objects are available in tests: ```typescript declare global { - var api: ClineAPI + var api: RooCodeAPI var provider: ClineProvider - var extension: vscode.Extension + var extension: vscode.Extension var panel: vscode.WebviewPanel } ``` diff --git a/e2e/src/suite/index.ts b/e2e/src/suite/index.ts index a9540d9600..19e00aa40c 100644 --- a/e2e/src/suite/index.ts +++ b/e2e/src/suite/index.ts @@ -1,13 +1,13 @@ import * as path from "path" import Mocha from "mocha" import { glob } from "glob" -import { ClineAPI, ClineProvider } from "../../../src/exports/cline" +import { RooCodeAPI, ClineProvider } from "../../../src/exports/roo-code" import * as vscode from "vscode" declare global { - var api: ClineAPI + var api: RooCodeAPI var provider: ClineProvider - var extension: vscode.Extension | undefined + var extension: vscode.Extension | undefined var panel: vscode.WebviewPanel | undefined } diff --git a/e2e/tsconfig.json b/e2e/tsconfig.json index 792acb14a0..4439b32b39 100644 --- a/e2e/tsconfig.json +++ b/e2e/tsconfig.json @@ -11,6 +11,6 @@ "useUnknownInCatchVariables": false, "outDir": "out" }, - "include": ["src", "../src/exports/cline.d.ts"], + "include": ["src", "../src/exports/roo-code.d.ts"], "exclude": [".vscode-test", "**/node_modules/**", "out"] } diff --git a/package.json b/package.json index 23a6e7a5c1..e8340eb8c2 100644 --- a/package.json +++ b/package.json @@ -230,8 +230,8 @@ "build": "npm run build:webview && npm run vsix", "build:webview": "cd webview-ui && npm run build", "compile": "tsc -p . --outDir out && node esbuild.js", - "install:all": "npm-run-all -p install-*", - "install:ci": "npm install npm-run-all && npm run install:all", + "install:all": "npm install npm-run-all && npm run install:_all", + "install:_all": "npm-run-all -p install-*", "install-extension": "npm install", "install-webview-ui": "cd webview-ui && npm install", "install-e2e": "cd e2e && npm install", @@ -246,7 +246,9 @@ "package": "npm run build:webview && npm run check-types && npm run lint && node esbuild.js --production", "pretest": "npm run compile", "dev": "cd webview-ui && npm run dev", - "test": "jest && cd webview-ui && npm run test", + "test": "npm-run-all -p test:*", + "test:extension": "jest", + "test:webview": "cd webview-ui && npm run test", "prepare": "husky", "publish:marketplace": "vsce publish && ovsx publish", "publish": "npm run build && changeset publish && npm install --package-lock-only", diff --git a/src/exports/index.ts b/src/activate/createRooCodeAPI.ts similarity index 79% rename from src/exports/index.ts rename to src/activate/createRooCodeAPI.ts index e4b17da484..18c696a8fd 100644 --- a/src/exports/index.ts +++ b/src/activate/createRooCodeAPI.ts @@ -1,9 +1,11 @@ import * as vscode from "vscode" -import { ClineProvider } from "../core/webview/ClineProvider" -import { ClineAPI } from "./cline" -export function createClineAPI(outputChannel: vscode.OutputChannel, sidebarProvider: ClineProvider): ClineAPI { - const api: ClineAPI = { +import { ClineProvider } from "../core/webview/ClineProvider" + +import { RooCodeAPI } from "../exports/roo-code" + +export function createRooCodeAPI(outputChannel: vscode.OutputChannel, sidebarProvider: ClineProvider): RooCodeAPI { + return { setCustomInstructions: async (value: string) => { await sidebarProvider.updateCustomInstructions(value) outputChannel.appendLine("Custom instructions set") @@ -24,6 +26,7 @@ export function createClineAPI(outputChannel: vscode.OutputChannel, sidebarProvi text: task, images: images, }) + outputChannel.appendLine( `Task started with message: ${task ? `"${task}"` : "undefined"} and ${images?.length || 0} image(s)`, ) @@ -33,6 +36,7 @@ export function createClineAPI(outputChannel: vscode.OutputChannel, sidebarProvi outputChannel.appendLine( `Sending message: ${message ? `"${message}"` : "undefined"} with ${images?.length || 0} image(s)`, ) + await sidebarProvider.postMessageToWebview({ type: "invoke", invoke: "sendMessage", @@ -43,22 +47,14 @@ export function createClineAPI(outputChannel: vscode.OutputChannel, sidebarProvi pressPrimaryButton: async () => { outputChannel.appendLine("Pressing primary button") - await sidebarProvider.postMessageToWebview({ - type: "invoke", - invoke: "primaryButtonClick", - }) + await sidebarProvider.postMessageToWebview({ type: "invoke", invoke: "primaryButtonClick" }) }, pressSecondaryButton: async () => { outputChannel.appendLine("Pressing secondary button") - await sidebarProvider.postMessageToWebview({ - type: "invoke", - invoke: "secondaryButtonClick", - }) + await sidebarProvider.postMessageToWebview({ type: "invoke", invoke: "secondaryButtonClick" }) }, sidebarProvider: sidebarProvider, } - - return api } diff --git a/src/activate/humanRelay.ts b/src/activate/humanRelay.ts new file mode 100644 index 0000000000..ed87026aa7 --- /dev/null +++ b/src/activate/humanRelay.ts @@ -0,0 +1,26 @@ +// Callback mapping of human relay response. +const humanRelayCallbacks = new Map void>() + +/** + * Register a callback function for human relay response. + * @param requestId + * @param callback + */ +export const registerHumanRelayCallback = (requestId: string, callback: (response: string | undefined) => void) => + humanRelayCallbacks.set(requestId, callback) + +export const unregisterHumanRelayCallback = (requestId: string) => humanRelayCallbacks.delete(requestId) + +export const handleHumanRelayResponse = (response: { requestId: string; text?: string; cancelled?: boolean }) => { + const callback = humanRelayCallbacks.get(response.requestId) + + if (callback) { + if (response.cancelled) { + callback(undefined) + } else { + callback(response.text) + } + + humanRelayCallbacks.delete(response.requestId) + } +} diff --git a/src/activate/index.ts b/src/activate/index.ts index 76eebd185e..8b3d91cdcb 100644 --- a/src/activate/index.ts +++ b/src/activate/index.ts @@ -1,3 +1,4 @@ export { handleUri } from "./handleUri" export { registerCommands } from "./registerCommands" export { registerCodeActions } from "./registerCodeActions" +export { createRooCodeAPI } from "./createRooCodeAPI" diff --git a/src/activate/registerCommands.ts b/src/activate/registerCommands.ts index e1c1ec93c6..1ca7ee96f8 100644 --- a/src/activate/registerCommands.ts +++ b/src/activate/registerCommands.ts @@ -3,6 +3,8 @@ import delay from "delay" import { ClineProvider } from "../core/webview/ClineProvider" +import { registerHumanRelayCallback, unregisterHumanRelayCallback, handleHumanRelayResponse } from "./humanRelay" + // Store panel references in both modes let sidebarPanel: vscode.WebviewView | undefined = undefined let tabPanel: vscode.WebviewPanel | undefined = undefined @@ -43,22 +45,6 @@ export const registerCommands = (options: RegisterCommandOptions) => { for (const [command, callback] of Object.entries(getCommandsMap(options))) { context.subscriptions.push(vscode.commands.registerCommand(command, callback)) } - - // Human Relay Dialog Command - context.subscriptions.push( - vscode.commands.registerCommand( - "roo-cline.showHumanRelayDialog", - (params: { requestId: string; promptText: string }) => { - if (getPanel()) { - getPanel()?.webview.postMessage({ - type: "showHumanRelayDialog", - requestId: params.requestId, - promptText: params.promptText, - }) - } - }, - ), - ) } const getCommandsMap = ({ context, outputChannel, provider }: RegisterCommandOptions) => { @@ -85,6 +71,20 @@ const getCommandsMap = ({ context, outputChannel, provider }: RegisterCommandOpt "roo-cline.helpButtonClicked": () => { vscode.env.openExternal(vscode.Uri.parse("https://docs.roocode.com")) }, + "roo-cline.showHumanRelayDialog": (params: { requestId: string; promptText: string }) => { + const panel = getPanel() + + if (panel) { + panel?.webview.postMessage({ + type: "showHumanRelayDialog", + requestId: params.requestId, + promptText: params.promptText, + }) + } + }, + "roo-cline.registerHumanRelayCallback": registerHumanRelayCallback, + "roo-cline.unregisterHumanRelayCallback": unregisterHumanRelayCallback, + "roo-cline.handleHumanRelayResponse": handleHumanRelayResponse, } } diff --git a/src/api/providers/__tests__/bedrock-custom-arn.test.ts b/src/api/providers/__tests__/bedrock-custom-arn.test.ts new file mode 100644 index 0000000000..f7dc2870fa --- /dev/null +++ b/src/api/providers/__tests__/bedrock-custom-arn.test.ts @@ -0,0 +1,75 @@ +import { AwsBedrockHandler } from "../bedrock" +import { ApiHandlerOptions } from "../../../shared/api" + +// Mock the AWS SDK +jest.mock("@aws-sdk/client-bedrock-runtime", () => { + const mockSend = jest.fn().mockImplementation(() => { + return Promise.resolve({ + output: new TextEncoder().encode(JSON.stringify({ content: "Test response" })), + }) + }) + + return { + BedrockRuntimeClient: jest.fn().mockImplementation(() => ({ + send: mockSend, + config: { + region: "us-east-1", + }, + })), + ConverseCommand: jest.fn(), + ConverseStreamCommand: jest.fn(), + } +}) + +describe("AwsBedrockHandler with custom ARN", () => { + const mockOptions: ApiHandlerOptions = { + apiModelId: "custom-arn", + awsCustomArn: "arn:aws:bedrock:us-east-1:123456789012:foundation-model/anthropic.claude-3-sonnet-20240229-v1:0", + awsRegion: "us-east-1", + } + + it("should use the custom ARN as the model ID", async () => { + const handler = new AwsBedrockHandler(mockOptions) + const model = handler.getModel() + + expect(model.id).toBe(mockOptions.awsCustomArn) + expect(model.info).toHaveProperty("maxTokens") + expect(model.info).toHaveProperty("contextWindow") + expect(model.info).toHaveProperty("supportsPromptCache") + }) + + it("should extract region from ARN and use it for client configuration", () => { + // Test with matching region + const handler1 = new AwsBedrockHandler(mockOptions) + expect((handler1 as any).client.config.region).toBe("us-east-1") + + // Test with mismatched region + const mismatchOptions = { + ...mockOptions, + awsRegion: "us-west-2", + } + const handler2 = new AwsBedrockHandler(mismatchOptions) + // Should use the ARN region, not the provided region + expect((handler2 as any).client.config.region).toBe("us-east-1") + }) + + it("should validate ARN format", async () => { + // Invalid ARN format + const invalidOptions = { + ...mockOptions, + awsCustomArn: "invalid-arn-format", + } + + const handler = new AwsBedrockHandler(invalidOptions) + + // completePrompt should throw an error for invalid ARN + await expect(handler.completePrompt("test")).rejects.toThrow("Invalid ARN format") + }) + + it("should complete a prompt successfully with valid ARN", async () => { + const handler = new AwsBedrockHandler(mockOptions) + const response = await handler.completePrompt("test prompt") + + expect(response).toBe("Test response") + }) +}) diff --git a/src/api/providers/__tests__/bedrock.test.ts b/src/api/providers/__tests__/bedrock.test.ts index f1b2c5527f..45d5270237 100644 --- a/src/api/providers/__tests__/bedrock.test.ts +++ b/src/api/providers/__tests__/bedrock.test.ts @@ -315,5 +315,34 @@ describe("AwsBedrockHandler", () => { expect(modelInfo.info.maxTokens).toBe(5000) expect(modelInfo.info.contextWindow).toBe(128_000) }) + + it("should use custom ARN when provided", () => { + const customArnHandler = new AwsBedrockHandler({ + apiModelId: "anthropic.claude-3-5-sonnet-20241022-v2:0", + awsAccessKey: "test-access-key", + awsSecretKey: "test-secret-key", + awsRegion: "us-east-1", + awsCustomArn: "arn:aws:bedrock:us-east-1::foundation-model/custom-model", + }) + const modelInfo = customArnHandler.getModel() + expect(modelInfo.id).toBe("arn:aws:bedrock:us-east-1::foundation-model/custom-model") + expect(modelInfo.info.maxTokens).toBe(4096) + expect(modelInfo.info.contextWindow).toBe(128_000) + expect(modelInfo.info.supportsPromptCache).toBe(false) + }) + + it("should use default model when custom-arn is selected but no ARN is provided", () => { + const customArnHandler = new AwsBedrockHandler({ + apiModelId: "custom-arn", + awsAccessKey: "test-access-key", + awsSecretKey: "test-secret-key", + awsRegion: "us-east-1", + // No awsCustomArn provided + }) + const modelInfo = customArnHandler.getModel() + // Should fall back to default model + expect(modelInfo.id).not.toBe("custom-arn") + expect(modelInfo.info).toBeDefined() + }) }) }) diff --git a/src/api/providers/__tests__/openai-usage-tracking.test.ts b/src/api/providers/__tests__/openai-usage-tracking.test.ts new file mode 100644 index 0000000000..6df9a0bca5 --- /dev/null +++ b/src/api/providers/__tests__/openai-usage-tracking.test.ts @@ -0,0 +1,235 @@ +import { OpenAiHandler } from "../openai" +import { ApiHandlerOptions } from "../../../shared/api" +import { Anthropic } from "@anthropic-ai/sdk" + +// Mock OpenAI client with multiple chunks that contain usage data +const mockCreate = jest.fn() +jest.mock("openai", () => { + return { + __esModule: true, + default: jest.fn().mockImplementation(() => ({ + chat: { + completions: { + create: mockCreate.mockImplementation(async (options) => { + if (!options.stream) { + return { + id: "test-completion", + choices: [ + { + message: { role: "assistant", content: "Test response", refusal: null }, + finish_reason: "stop", + index: 0, + }, + ], + usage: { + prompt_tokens: 10, + completion_tokens: 5, + total_tokens: 15, + }, + } + } + + // Return a stream with multiple chunks that include usage metrics + return { + [Symbol.asyncIterator]: async function* () { + // First chunk with partial usage + yield { + choices: [ + { + delta: { content: "Test " }, + index: 0, + }, + ], + usage: { + prompt_tokens: 10, + completion_tokens: 2, + total_tokens: 12, + }, + } + + // Second chunk with updated usage + yield { + choices: [ + { + delta: { content: "response" }, + index: 0, + }, + ], + usage: { + prompt_tokens: 10, + completion_tokens: 4, + total_tokens: 14, + }, + } + + // Final chunk with complete usage + yield { + choices: [ + { + delta: {}, + index: 0, + }, + ], + usage: { + prompt_tokens: 10, + completion_tokens: 5, + total_tokens: 15, + }, + } + }, + } + }), + }, + }, + })), + } +}) + +describe("OpenAiHandler with usage tracking fix", () => { + let handler: OpenAiHandler + let mockOptions: ApiHandlerOptions + + beforeEach(() => { + mockOptions = { + openAiApiKey: "test-api-key", + openAiModelId: "gpt-4", + openAiBaseUrl: "https://api.openai.com/v1", + } + handler = new OpenAiHandler(mockOptions) + mockCreate.mockClear() + }) + + describe("usage metrics with streaming", () => { + const systemPrompt = "You are a helpful assistant." + const messages: Anthropic.Messages.MessageParam[] = [ + { + role: "user", + content: [ + { + type: "text" as const, + text: "Hello!", + }, + ], + }, + ] + + it("should only yield usage metrics once at the end of the stream", async () => { + const stream = handler.createMessage(systemPrompt, messages) + const chunks: any[] = [] + for await (const chunk of stream) { + chunks.push(chunk) + } + + // Check we have text chunks + const textChunks = chunks.filter((chunk) => chunk.type === "text") + expect(textChunks).toHaveLength(2) + expect(textChunks[0].text).toBe("Test ") + expect(textChunks[1].text).toBe("response") + + // Check we only have one usage chunk and it's the last one + const usageChunks = chunks.filter((chunk) => chunk.type === "usage") + expect(usageChunks).toHaveLength(1) + expect(usageChunks[0]).toEqual({ + type: "usage", + inputTokens: 10, + outputTokens: 5, + }) + + // Check the usage chunk is the last one reported from the API + const lastChunk = chunks[chunks.length - 1] + expect(lastChunk.type).toBe("usage") + expect(lastChunk.inputTokens).toBe(10) + expect(lastChunk.outputTokens).toBe(5) + }) + + it("should handle case where usage is only in the final chunk", async () => { + // Override the mock for this specific test + mockCreate.mockImplementationOnce(async (options) => { + if (!options.stream) { + return { + id: "test-completion", + choices: [{ message: { role: "assistant", content: "Test response" } }], + usage: { prompt_tokens: 10, completion_tokens: 5, total_tokens: 15 }, + } + } + + return { + [Symbol.asyncIterator]: async function* () { + // First chunk with no usage + yield { + choices: [{ delta: { content: "Test " }, index: 0 }], + usage: null, + } + + // Second chunk with no usage + yield { + choices: [{ delta: { content: "response" }, index: 0 }], + usage: null, + } + + // Final chunk with usage data + yield { + choices: [{ delta: {}, index: 0 }], + usage: { + prompt_tokens: 10, + completion_tokens: 5, + total_tokens: 15, + }, + } + }, + } + }) + + const stream = handler.createMessage(systemPrompt, messages) + const chunks: any[] = [] + for await (const chunk of stream) { + chunks.push(chunk) + } + + // Check usage metrics + const usageChunks = chunks.filter((chunk) => chunk.type === "usage") + expect(usageChunks).toHaveLength(1) + expect(usageChunks[0]).toEqual({ + type: "usage", + inputTokens: 10, + outputTokens: 5, + }) + }) + + it("should handle case where no usage is provided", async () => { + // Override the mock for this specific test + mockCreate.mockImplementationOnce(async (options) => { + if (!options.stream) { + return { + id: "test-completion", + choices: [{ message: { role: "assistant", content: "Test response" } }], + usage: null, + } + } + + return { + [Symbol.asyncIterator]: async function* () { + yield { + choices: [{ delta: { content: "Test response" }, index: 0 }], + usage: null, + } + yield { + choices: [{ delta: {}, index: 0 }], + usage: null, + } + }, + } + }) + + const stream = handler.createMessage(systemPrompt, messages) + const chunks: any[] = [] + for await (const chunk of stream) { + chunks.push(chunk) + } + + // Check we don't have any usage chunks + const usageChunks = chunks.filter((chunk) => chunk.type === "usage") + expect(usageChunks).toHaveLength(0) + }) + }) +}) diff --git a/src/api/providers/bedrock.ts b/src/api/providers/bedrock.ts index 2deb019dc3..76d9364960 100644 --- a/src/api/providers/bedrock.ts +++ b/src/api/providers/bedrock.ts @@ -11,6 +11,47 @@ import { ApiHandlerOptions, BedrockModelId, ModelInfo, bedrockDefaultModelId, be import { ApiStream } from "../transform/stream" import { convertToBedrockConverseMessages } from "../transform/bedrock-converse-format" import { BaseProvider } from "./base-provider" +import { logger } from "../../utils/logging" + +/** + * Validates an AWS Bedrock ARN format and optionally checks if the region in the ARN matches the provided region + * @param arn The ARN string to validate + * @param region Optional region to check against the ARN's region + * @returns An object with validation results: { isValid, arnRegion, errorMessage } + */ +function validateBedrockArn(arn: string, region?: string) { + // Validate ARN format + const arnRegex = /^arn:aws:bedrock:([^:]+):(\d+):(foundation-model|provisioned-model|default-prompt-router)\/(.+)$/ + const match = arn.match(arnRegex) + + if (!match) { + return { + isValid: false, + arnRegion: undefined, + errorMessage: + "Invalid ARN format. ARN should follow the pattern: arn:aws:bedrock:region:account-id:resource-type/resource-name", + } + } + + // Extract region from ARN + const arnRegion = match[1] + + // Check if region in ARN matches provided region (if specified) + if (region && arnRegion !== region) { + return { + isValid: true, + arnRegion, + errorMessage: `Warning: The region in your ARN (${arnRegion}) does not match your selected region (${region}). This may cause access issues. The provider will use the region from the ARN.`, + } + } + + // ARN is valid and region matches (or no region was provided to check against) + return { + isValid: true, + arnRegion, + errorMessage: undefined, + } +} const BEDROCK_DEFAULT_TEMPERATURE = 0.3 @@ -55,8 +96,31 @@ export class AwsBedrockHandler extends BaseProvider implements SingleCompletionH super() this.options = options + // Extract region from custom ARN if provided + let region = this.options.awsRegion || "us-east-1" + + // If using custom ARN, extract region from the ARN + if (this.options.awsCustomArn) { + const validation = validateBedrockArn(this.options.awsCustomArn, region) + + if (validation.isValid && validation.arnRegion) { + // If there's a region mismatch warning, log it and use the ARN region + if (validation.errorMessage) { + logger.info( + `Region mismatch: Selected region is ${region}, but ARN region is ${validation.arnRegion}. Using ARN region.`, + { + ctx: "bedrock", + selectedRegion: region, + arnRegion: validation.arnRegion, + }, + ) + region = validation.arnRegion + } + } + } + const clientConfig: BedrockRuntimeClientConfig = { - region: this.options.awsRegion || "us-east-1", + region: region, } if (this.options.awsUseProfile && this.options.awsProfile) { @@ -81,7 +145,41 @@ export class AwsBedrockHandler extends BaseProvider implements SingleCompletionH // Handle cross-region inference let modelId: string - if (this.options.awsUseCrossRegionInference) { + + // For custom ARNs, use the ARN directly without modification + if (this.options.awsCustomArn) { + modelId = modelConfig.id + + // Validate ARN format and check region match + const clientRegion = this.client.config.region as string + const validation = validateBedrockArn(modelId, clientRegion) + + if (!validation.isValid) { + logger.error("Invalid ARN format", { + ctx: "bedrock", + modelId, + errorMessage: validation.errorMessage, + }) + yield { + type: "text", + text: `Error: ${validation.errorMessage}`, + } + yield { type: "usage", inputTokens: 0, outputTokens: 0 } + throw new Error("Invalid ARN format") + } + + // Extract region from ARN + const arnRegion = validation.arnRegion! + + // Log warning if there's a region mismatch + if (validation.errorMessage) { + logger.warn(validation.errorMessage, { + ctx: "bedrock", + arnRegion, + clientRegion, + }) + } + } else if (this.options.awsUseCrossRegionInference) { let regionPrefix = (this.options.awsRegion || "").slice(0, 3) switch (regionPrefix) { case "us-": @@ -107,7 +205,7 @@ export class AwsBedrockHandler extends BaseProvider implements SingleCompletionH messages: formattedMessages, system: [{ text: systemPrompt }], inferenceConfig: { - maxTokens: modelConfig.info.maxTokens || 5000, + maxTokens: modelConfig.info.maxTokens || 4096, temperature: this.options.modelTemperature ?? BEDROCK_DEFAULT_TEMPERATURE, topP: 0.1, ...(this.options.awsUsePromptCache @@ -121,6 +219,16 @@ export class AwsBedrockHandler extends BaseProvider implements SingleCompletionH } try { + // Log the payload for debugging custom ARN issues + if (this.options.awsCustomArn) { + logger.debug("Using custom ARN for Bedrock request", { + ctx: "bedrock", + customArn: this.options.awsCustomArn, + clientRegion: this.client.config.region, + payload: JSON.stringify(payload, null, 2), + }) + } + const command = new ConverseStreamCommand(payload) const response = await this.client.send(command) @@ -134,7 +242,11 @@ export class AwsBedrockHandler extends BaseProvider implements SingleCompletionH try { streamEvent = typeof chunk === "string" ? JSON.parse(chunk) : (chunk as unknown as StreamEvent) } catch (e) { - console.error("Failed to parse stream event:", e) + logger.error("Failed to parse stream event", { + ctx: "bedrock", + error: e instanceof Error ? e : String(e), + chunk: typeof chunk === "string" ? chunk : "binary data", + }) continue } @@ -177,39 +289,257 @@ export class AwsBedrockHandler extends BaseProvider implements SingleCompletionH } } } catch (error: unknown) { - console.error("Bedrock Runtime API Error:", error) - // Only access stack if error is an Error object + logger.error("Bedrock Runtime API Error", { + ctx: "bedrock", + error: error instanceof Error ? error : String(error), + }) + + // Enhanced error handling for custom ARN issues + if (this.options.awsCustomArn) { + logger.error("Error occurred with custom ARN", { + ctx: "bedrock", + customArn: this.options.awsCustomArn, + }) + + // Check for common ARN-related errors + if (error instanceof Error) { + const errorMessage = error.message.toLowerCase() + + // Access denied errors + if ( + errorMessage.includes("access") && + (errorMessage.includes("model") || errorMessage.includes("denied")) + ) { + logger.error("Permissions issue with custom ARN", { + ctx: "bedrock", + customArn: this.options.awsCustomArn, + errorType: "access_denied", + clientRegion: this.client.config.region, + }) + yield { + type: "text", + text: `Error: You don't have access to the model with the specified ARN. Please verify: + +1. The ARN is correct and points to a valid model +2. Your AWS credentials have permission to access this model (check IAM policies) +3. The region in the ARN (${this.client.config.region}) matches the region where the model is deployed +4. If using a provisioned model, ensure it's active and not in a failed state +5. If using a custom model, ensure your account has been granted access to it`, + } + } + // Model not found errors + else if (errorMessage.includes("not found") || errorMessage.includes("does not exist")) { + logger.error("Invalid ARN or non-existent model", { + ctx: "bedrock", + customArn: this.options.awsCustomArn, + errorType: "not_found", + }) + yield { + type: "text", + text: `Error: The specified ARN does not exist or is invalid. Please check: + +1. The ARN format is correct (arn:aws:bedrock:region:account-id:resource-type/resource-name) +2. The model exists in the specified region +3. The account ID in the ARN is correct +4. The resource type is one of: foundation-model, provisioned-model, or default-prompt-router`, + } + } + // Throttling errors + else if ( + errorMessage.includes("throttl") || + errorMessage.includes("rate") || + errorMessage.includes("limit") + ) { + logger.error("Throttling or rate limit issue with Bedrock", { + ctx: "bedrock", + customArn: this.options.awsCustomArn, + errorType: "throttling", + }) + yield { + type: "text", + text: `Error: Request was throttled or rate limited. Please try: + +1. Reducing the frequency of requests +2. If using a provisioned model, check its throughput settings +3. Contact AWS support to request a quota increase if needed`, + } + } + // Other errors + else { + logger.error("Unspecified error with custom ARN", { + ctx: "bedrock", + customArn: this.options.awsCustomArn, + errorStack: error.stack, + errorMessage: error.message, + }) + yield { + type: "text", + text: `Error with custom ARN: ${error.message} + +Please check: +1. Your AWS credentials are valid and have the necessary permissions +2. The ARN format is correct +3. The region in the ARN matches the region where you're making the request`, + } + } + } else { + yield { + type: "text", + text: `Unknown error occurred with custom ARN. Please check your AWS credentials and ARN format.`, + } + } + } else { + // Standard error handling for non-ARN cases + if (error instanceof Error) { + logger.error("Standard Bedrock error", { + ctx: "bedrock", + errorStack: error.stack, + errorMessage: error.message, + }) + yield { + type: "text", + text: `Error: ${error.message}`, + } + } else { + logger.error("Unknown Bedrock error", { + ctx: "bedrock", + error: String(error), + }) + yield { + type: "text", + text: "An unknown error occurred", + } + } + } + + // Always yield usage info + yield { + type: "usage", + inputTokens: 0, + outputTokens: 0, + } + + // Re-throw the error if (error instanceof Error) { - console.error("Error stack:", error.stack) - yield { - type: "text", - text: `Error: ${error.message}`, - } - yield { - type: "usage", - inputTokens: 0, - outputTokens: 0, - } throw error } else { - const unknownError = new Error("An unknown error occurred") - yield { - type: "text", - text: unknownError.message, - } - yield { - type: "usage", - inputTokens: 0, - outputTokens: 0, - } - throw unknownError + throw new Error("An unknown error occurred") } } } override getModel(): { id: BedrockModelId | string; info: ModelInfo } { + // If custom ARN is provided, use it + if (this.options.awsCustomArn) { + // Custom ARNs should not be modified with region prefixes + // as they already contain the full resource path + + // Check if the ARN contains information about the model type + // This helps set appropriate token limits for models behind prompt routers + const arnLower = this.options.awsCustomArn.toLowerCase() + + // Determine model info based on ARN content + let modelInfo: ModelInfo + + if (arnLower.includes("claude-3-7-sonnet") || arnLower.includes("claude-3.7-sonnet")) { + // Claude 3.7 Sonnet has 8192 tokens in Bedrock + modelInfo = { + maxTokens: 8192, + contextWindow: 200_000, + supportsPromptCache: false, + supportsImages: true, + supportsComputerUse: true, + } + } else if (arnLower.includes("claude-3-5-sonnet") || arnLower.includes("claude-3.5-sonnet")) { + // Claude 3.5 Sonnet has 8192 tokens in Bedrock + modelInfo = { + maxTokens: 8192, + contextWindow: 200_000, + supportsPromptCache: false, + supportsImages: true, + supportsComputerUse: true, + } + } else if (arnLower.includes("claude-3-opus") || arnLower.includes("claude-3.0-opus")) { + // Claude 3 Opus has 4096 tokens in Bedrock + modelInfo = { + maxTokens: 4096, + contextWindow: 200_000, + supportsPromptCache: false, + supportsImages: true, + } + } else if (arnLower.includes("claude-3-haiku") || arnLower.includes("claude-3.0-haiku")) { + // Claude 3 Haiku has 4096 tokens in Bedrock + modelInfo = { + maxTokens: 4096, + contextWindow: 200_000, + supportsPromptCache: false, + supportsImages: true, + } + } else if (arnLower.includes("claude-3-5-haiku") || arnLower.includes("claude-3.5-haiku")) { + // Claude 3.5 Haiku has 8192 tokens in Bedrock + modelInfo = { + maxTokens: 8192, + contextWindow: 200_000, + supportsPromptCache: false, + supportsImages: false, + } + } else if (arnLower.includes("claude")) { + // Generic Claude model with conservative token limit + modelInfo = { + maxTokens: 4096, + contextWindow: 128_000, + supportsPromptCache: false, + supportsImages: true, + } + } else if (arnLower.includes("llama3") || arnLower.includes("llama-3")) { + // Llama 3 models typically have 8192 tokens in Bedrock + modelInfo = { + maxTokens: 8192, + contextWindow: 128_000, + supportsPromptCache: false, + supportsImages: arnLower.includes("90b") || arnLower.includes("11b"), + } + } else if (arnLower.includes("nova-pro")) { + // Amazon Nova Pro + modelInfo = { + maxTokens: 5000, + contextWindow: 300_000, + supportsPromptCache: false, + supportsImages: true, + } + } else { + // Default for unknown models or prompt routers + modelInfo = { + maxTokens: 4096, + contextWindow: 128_000, + supportsPromptCache: false, + supportsImages: true, + } + } + + // If modelMaxTokens is explicitly set in options, override the default + if (this.options.modelMaxTokens && this.options.modelMaxTokens > 0) { + modelInfo.maxTokens = this.options.modelMaxTokens + } + + return { + id: this.options.awsCustomArn, + info: modelInfo, + } + } + const modelId = this.options.apiModelId if (modelId) { + // Special case for custom ARN option + if (modelId === "custom-arn") { + // This should not happen as we should have awsCustomArn set + // but just in case, return a default model + return { + id: bedrockDefaultModelId, + info: bedrockModels[bedrockDefaultModelId], + } + } + // For tests, allow any model ID if (process.env.NODE_ENV === "test") { return { @@ -239,7 +569,43 @@ export class AwsBedrockHandler extends BaseProvider implements SingleCompletionH // Handle cross-region inference let modelId: string - if (this.options.awsUseCrossRegionInference) { + + // For custom ARNs, use the ARN directly without modification + if (this.options.awsCustomArn) { + modelId = modelConfig.id + logger.debug("Using custom ARN in completePrompt", { + ctx: "bedrock", + customArn: this.options.awsCustomArn, + }) + + // Validate ARN format and check region match + const clientRegion = this.client.config.region as string + const validation = validateBedrockArn(modelId, clientRegion) + + if (!validation.isValid) { + logger.error("Invalid ARN format in completePrompt", { + ctx: "bedrock", + modelId, + errorMessage: validation.errorMessage, + }) + throw new Error( + validation.errorMessage || + "Invalid ARN format. ARN should follow the pattern: arn:aws:bedrock:region:account-id:resource-type/resource-name", + ) + } + + // Extract region from ARN + const arnRegion = validation.arnRegion! + + // Log warning if there's a region mismatch + if (validation.errorMessage) { + logger.warn(validation.errorMessage, { + ctx: "bedrock", + arnRegion, + clientRegion, + }) + } + } else if (this.options.awsUseCrossRegionInference) { let regionPrefix = (this.options.awsRegion || "").slice(0, 3) switch (regionPrefix) { case "us-": @@ -265,12 +631,21 @@ export class AwsBedrockHandler extends BaseProvider implements SingleCompletionH }, ]), inferenceConfig: { - maxTokens: modelConfig.info.maxTokens || 5000, + maxTokens: modelConfig.info.maxTokens || 4096, temperature: this.options.modelTemperature ?? BEDROCK_DEFAULT_TEMPERATURE, topP: 0.1, }, } + // Log the payload for debugging custom ARN issues + if (this.options.awsCustomArn) { + logger.debug("Bedrock completePrompt request details", { + ctx: "bedrock", + clientRegion: this.client.config.region, + payload: JSON.stringify(payload, null, 2), + }) + } + const command = new ConverseCommand(payload) const response = await this.client.send(command) @@ -282,11 +657,67 @@ export class AwsBedrockHandler extends BaseProvider implements SingleCompletionH return output.content } } catch (parseError) { - console.error("Failed to parse Bedrock response:", parseError) + logger.error("Failed to parse Bedrock response", { + ctx: "bedrock", + error: parseError instanceof Error ? parseError : String(parseError), + }) } } return "" } catch (error) { + // Enhanced error handling for custom ARN issues + if (this.options.awsCustomArn) { + logger.error("Error occurred with custom ARN in completePrompt", { + ctx: "bedrock", + customArn: this.options.awsCustomArn, + error: error instanceof Error ? error : String(error), + }) + + if (error instanceof Error) { + const errorMessage = error.message.toLowerCase() + + // Access denied errors + if ( + errorMessage.includes("access") && + (errorMessage.includes("model") || errorMessage.includes("denied")) + ) { + throw new Error( + `Bedrock custom ARN error: You don't have access to the model with the specified ARN. Please verify: +1. The ARN is correct and points to a valid model +2. Your AWS credentials have permission to access this model (check IAM policies) +3. The region in the ARN matches the region where the model is deployed +4. If using a provisioned model, ensure it's active and not in a failed state`, + ) + } + // Model not found errors + else if (errorMessage.includes("not found") || errorMessage.includes("does not exist")) { + throw new Error( + `Bedrock custom ARN error: The specified ARN does not exist or is invalid. Please check: +1. The ARN format is correct (arn:aws:bedrock:region:account-id:resource-type/resource-name) +2. The model exists in the specified region +3. The account ID in the ARN is correct +4. The resource type is one of: foundation-model, provisioned-model, or default-prompt-router`, + ) + } + // Throttling errors + else if ( + errorMessage.includes("throttl") || + errorMessage.includes("rate") || + errorMessage.includes("limit") + ) { + throw new Error( + `Bedrock custom ARN error: Request was throttled or rate limited. Please try: +1. Reducing the frequency of requests +2. If using a provisioned model, check its throughput settings +3. Contact AWS support to request a quota increase if needed`, + ) + } else { + throw new Error(`Bedrock custom ARN error: ${error.message}`) + } + } + } + + // Standard error handling if (error instanceof Error) { throw new Error(`Bedrock completion error: ${error.message}`) } diff --git a/src/api/providers/openai.ts b/src/api/providers/openai.ts index 2af3f2da05..a6e3eb1488 100644 --- a/src/api/providers/openai.ts +++ b/src/api/providers/openai.ts @@ -99,6 +99,8 @@ export class OpenAiHandler extends BaseProvider implements SingleCompletionHandl const stream = await this.client.chat.completions.create(requestOptions) + let lastUsage + for await (const chunk of stream) { const delta = chunk.choices[0]?.delta ?? {} @@ -116,9 +118,13 @@ export class OpenAiHandler extends BaseProvider implements SingleCompletionHandl } } if (chunk.usage) { - yield this.processUsageMetrics(chunk.usage, modelInfo) + lastUsage = chunk.usage } } + + if (lastUsage) { + yield this.processUsageMetrics(lastUsage, modelInfo) + } } else { // o1 for instance doesnt support streaming, non-1 temp, or system prompt const systemMessage: OpenAI.Chat.ChatCompletionUserMessageParam = { diff --git a/src/core/Cline.ts b/src/core/Cline.ts index ba171ce3fa..5bc280dc42 100644 --- a/src/core/Cline.ts +++ b/src/core/Cline.ts @@ -48,6 +48,7 @@ import { ClineSay, ClineSayBrowserAction, ClineSayTool, + ToolProgressStatus, } from "../shared/ExtensionMessage" import { getApiMetrics } from "../shared/getApiMetrics" import { HistoryItem } from "../shared/HistoryItem" @@ -408,6 +409,7 @@ export class Cline { type: ClineAsk, text?: string, partial?: boolean, + progressStatus?: ToolProgressStatus, ): Promise<{ response: ClineAskResponse; text?: string; images?: string[] }> { // If this Cline instance was aborted by the provider, then the only thing keeping us alive is a promise still running in the background, in which case we don't want to send its result to the webview as it is attached to a new instance of Cline now. So we can safely ignore the result of any active promises, and this class will be deallocated. (Although we set Cline = undefined in provider, that simply removes the reference to this instance, but the instance is still alive until this promise resolves or rejects.) if (this.abort) { @@ -423,6 +425,7 @@ export class Cline { // existing partial message, so update it lastMessage.text = text lastMessage.partial = partial + lastMessage.progressStatus = progressStatus // todo be more efficient about saving and posting only new data or one whole message at a time so ignore partial for saves, and only post parts of partial message instead of whole array in new listener // await this.saveClineMessages() // await this.providerRef.deref()?.postStateToWebview() @@ -460,6 +463,8 @@ export class Cline { // lastMessage.ts = askTs lastMessage.text = text lastMessage.partial = false + lastMessage.progressStatus = progressStatus + await this.saveClineMessages() // await this.providerRef.deref()?.postStateToWebview() await this.providerRef @@ -511,6 +516,7 @@ export class Cline { images?: string[], partial?: boolean, checkpoint?: Record, + progressStatus?: ToolProgressStatus, ): Promise { if (this.abort) { throw new Error(`Task: ${this.taskNumber} Roo Code instance aborted (#2)`) @@ -526,6 +532,7 @@ export class Cline { lastMessage.text = text lastMessage.images = images lastMessage.partial = partial + lastMessage.progressStatus = progressStatus await this.providerRef .deref() ?.postMessageToWebview({ type: "partialMessage", partialMessage: lastMessage }) @@ -545,6 +552,7 @@ export class Cline { lastMessage.text = text lastMessage.images = images lastMessage.partial = false + lastMessage.progressStatus = progressStatus // instead of streaming partialMessage events, we do a save and post like normal to persist to disk await this.saveClineMessages() @@ -1394,8 +1402,12 @@ export class Cline { isCheckpointPossible = true } - const askApproval = async (type: ClineAsk, partialMessage?: string) => { - const { response, text, images } = await this.ask(type, partialMessage, false) + const askApproval = async ( + type: ClineAsk, + partialMessage?: string, + progressStatus?: ToolProgressStatus, + ) => { + const { response, text, images } = await this.ask(type, partialMessage, false, progressStatus) if (response !== "yesButtonClicked") { // Handle both messageResponse and noButtonClicked with text if (text) { @@ -1703,8 +1715,16 @@ export class Cline { try { if (block.partial) { // update gui message + let toolProgressStatus + if (this.diffStrategy && this.diffStrategy.getProgressStatus) { + toolProgressStatus = this.diffStrategy.getProgressStatus(block) + } + const partialMessage = JSON.stringify(sharedMessageProps) - await this.ask("tool", partialMessage, block.partial).catch(() => {}) + + await this.ask("tool", partialMessage, block.partial, toolProgressStatus).catch( + () => {}, + ) break } else { if (!relPath) { @@ -1799,7 +1819,12 @@ export class Cline { diff: diffContent, } satisfies ClineSayTool) - const didApprove = await askApproval("tool", completeMessage) + let toolProgressStatus + if (this.diffStrategy && this.diffStrategy.getProgressStatus) { + toolProgressStatus = this.diffStrategy.getProgressStatus(block, diffResult) + } + + const didApprove = await askApproval("tool", completeMessage, toolProgressStatus) if (!didApprove) { await this.diffViewProvider.revertChanges() // This likely handles closing the diff view break diff --git a/src/core/diff/strategies/multi-search-replace.ts b/src/core/diff/strategies/multi-search-replace.ts index 99c22a31df..bcf2f65430 100644 --- a/src/core/diff/strategies/multi-search-replace.ts +++ b/src/core/diff/strategies/multi-search-replace.ts @@ -1,6 +1,8 @@ import { DiffStrategy, DiffResult } from "../types" import { addLineNumbers, everyLineHasLineNumbers, stripLineNumbers } from "../../../integrations/misc/extract-text" import { distance } from "fastest-levenshtein" +import { ToolProgressStatus } from "../../../shared/ExtensionMessage" +import { ToolUse } from "../../assistant-message" const BUFFER_LINES = 40 // Number of extra context lines to show before and after matches @@ -362,4 +364,27 @@ Only use a single line of '=======' between search and replacement content, beca failParts: diffResults, } } + + getProgressStatus(toolUse: ToolUse, result?: DiffResult): ToolProgressStatus { + const diffContent = toolUse.params.diff + if (diffContent) { + const icon = "diff-multiple" + const searchBlockCount = (diffContent.match(/SEARCH/g) || []).length + if (toolUse.partial) { + if (diffContent.length < 1000 || (diffContent.length / 50) % 10 === 0) { + return { icon, text: `${searchBlockCount}` } + } + } else if (result) { + if (result.failParts?.length) { + return { + icon, + text: `${searchBlockCount - result.failParts.length}/${searchBlockCount}`, + } + } else { + return { icon, text: `${searchBlockCount}` } + } + } + } + return {} + } } diff --git a/src/core/diff/types.ts b/src/core/diff/types.ts index be6d8cd311..e12a47762d 100644 --- a/src/core/diff/types.ts +++ b/src/core/diff/types.ts @@ -2,6 +2,9 @@ * Interface for implementing different diff strategies */ +import { ToolProgressStatus } from "../../shared/ExtensionMessage" +import { ToolUse } from "../assistant-message" + export type DiffResult = | { success: true; content: string; failParts?: DiffResult[] } | ({ @@ -34,4 +37,6 @@ export interface DiffStrategy { * @returns A DiffResult object containing either the successful result or error details */ applyDiff(originalContent: string, diffContent: string, startLine?: number, endLine?: number): Promise + + getProgressStatus?(toolUse: ToolUse, result?: any): ToolProgressStatus } diff --git a/src/core/webview/ClineProvider.ts b/src/core/webview/ClineProvider.ts index 72cf56d35b..f98f19e4ff 100644 --- a/src/core/webview/ClineProvider.ts +++ b/src/core/webview/ClineProvider.ts @@ -30,6 +30,8 @@ import WorkspaceTracker from "../../integrations/workspace/WorkspaceTracker" import { McpHub } from "../../services/mcp/McpHub" import { McpServerManager } from "../../services/mcp/McpServerManager" import { ShadowCheckpointService } from "../../services/checkpoints/ShadowCheckpointService" +import { BrowserSession } from "../../services/browser/BrowserSession" +import { discoverChromeInstances } from "../../services/browser/browserDiscovery" import { fileExistsAtPath } from "../../utils/fs" import { playSound, setSoundEnabled, setSoundVolume } from "../../utils/sound" import { singleCompletionHandler } from "../../utils/single-completion-handler" @@ -1262,6 +1264,105 @@ export class ClineProvider implements vscode.WebviewViewProvider { await this.updateGlobalState("browserViewportSize", browserViewportSize) await this.postStateToWebview() break + case "remoteBrowserHost": + await this.updateGlobalState("remoteBrowserHost", message.text) + await this.postStateToWebview() + break + case "remoteBrowserEnabled": + // Store the preference in global state + // remoteBrowserEnabled now means "enable remote browser connection" + await this.updateGlobalState("remoteBrowserEnabled", message.bool ?? false) + // If disabling remote browser connection, clear the remoteBrowserHost + if (!message.bool) { + await this.updateGlobalState("remoteBrowserHost", undefined) + } + await this.postStateToWebview() + break + case "testBrowserConnection": + try { + const browserSession = new BrowserSession(this.context) + // If no text is provided, try auto-discovery + if (!message.text) { + try { + const discoveredHost = await discoverChromeInstances() + if (discoveredHost) { + // Test the connection to the discovered host + const result = await browserSession.testConnection(discoveredHost) + // Send the result back to the webview + await this.postMessageToWebview({ + type: "browserConnectionResult", + success: result.success, + text: `Auto-discovered and tested connection to Chrome at ${discoveredHost}: ${result.message}`, + values: { endpoint: result.endpoint }, + }) + } else { + await this.postMessageToWebview({ + type: "browserConnectionResult", + success: false, + text: "No Chrome instances found on the network. Make sure Chrome is running with remote debugging enabled (--remote-debugging-port=9222).", + }) + } + } catch (error) { + await this.postMessageToWebview({ + type: "browserConnectionResult", + success: false, + text: `Error during auto-discovery: ${error instanceof Error ? error.message : String(error)}`, + }) + } + } else { + // Test the provided URL + const result = await browserSession.testConnection(message.text) + + // Send the result back to the webview + await this.postMessageToWebview({ + type: "browserConnectionResult", + success: result.success, + text: result.message, + values: { endpoint: result.endpoint }, + }) + } + } catch (error) { + await this.postMessageToWebview({ + type: "browserConnectionResult", + success: false, + text: `Error testing connection: ${error instanceof Error ? error.message : String(error)}`, + }) + } + break + case "discoverBrowser": + try { + const discoveredHost = await discoverChromeInstances() + + if (discoveredHost) { + // Don't update the remoteBrowserHost state when auto-discovering + // This way we don't override the user's preference + + // Test the connection to get the endpoint + const browserSession = new BrowserSession(this.context) + const result = await browserSession.testConnection(discoveredHost) + + // Send the result back to the webview + await this.postMessageToWebview({ + type: "browserConnectionResult", + success: true, + text: `Successfully discovered and connected to Chrome at ${discoveredHost}`, + values: { endpoint: result.endpoint }, + }) + } else { + await this.postMessageToWebview({ + type: "browserConnectionResult", + success: false, + text: "No Chrome instances found on the network. Make sure Chrome is running with remote debugging enabled (--remote-debugging-port=9222).", + }) + } + } catch (error) { + await this.postMessageToWebview({ + type: "browserConnectionResult", + success: false, + text: `Error discovering browser: ${error instanceof Error ? error.message : String(error)}`, + }) + } + break case "fuzzyMatchThreshold": await this.updateGlobalState("fuzzyMatchThreshold", message.value) await this.postStateToWebview() @@ -1826,6 +1927,7 @@ export class ClineProvider implements vscode.WebviewViewProvider { fuzzyMatchThreshold, experiments, enableMcpServerCreation, + browserToolEnabled, } = await this.getState() // Create diffStrategy based on current model and settings @@ -1841,10 +1943,14 @@ export class ClineProvider implements vscode.WebviewViewProvider { const rooIgnoreInstructions = this.getCurrentCline()?.rooIgnoreController?.getInstructions() + // Determine if browser tools can be used based on model support and user settings + const modelSupportsComputerUse = this.getCurrentCline()?.api.getModel().info.supportsComputerUse ?? false + const canUseBrowserTool = modelSupportsComputerUse && (browserToolEnabled ?? true) + const systemPrompt = await SYSTEM_PROMPT( this.context, cwd, - apiConfiguration.openRouterModelInfo?.supportsComputerUse ?? false, + canUseBrowserTool, mcpEnabled ? this.mcpHub : undefined, diffStrategy, browserViewportSize ?? "900x600", @@ -1971,11 +2077,24 @@ export class ClineProvider implements vscode.WebviewViewProvider { // MCP async ensureMcpServersDirectoryExists(): Promise { - const mcpServersDir = path.join(os.homedir(), "Documents", "Cline", "MCP") + // Get platform-specific application data directory + let mcpServersDir: string + if (process.platform === "win32") { + // Windows: %APPDATA%\Roo-Code\MCP + mcpServersDir = path.join(os.homedir(), "AppData", "Roaming", "Roo-Code", "MCP") + } else if (process.platform === "darwin") { + // macOS: ~/Documents/Cline/MCP + mcpServersDir = path.join(os.homedir(), "Documents", "Cline", "MCP") + } else { + // Linux: ~/.local/share/Cline/MCP + mcpServersDir = path.join(os.homedir(), ".local", "share", "Roo-Code", "MCP") + } + try { await fs.mkdir(mcpServersDir, { recursive: true }) } catch (error) { - return "~/Documents/Cline/MCP" // in case creating a directory in documents fails for whatever reason (e.g. permissions) - this is fine since this path is only ever used in the system prompt + // Fallback to a relative path if directory creation fails + return path.join(os.homedir(), ".roo-code", "mcp") } return mcpServersDir } @@ -2192,6 +2311,8 @@ export class ClineProvider implements vscode.WebviewViewProvider { soundVolume, browserViewportSize, screenshotQuality, + remoteBrowserHost, + remoteBrowserEnabled, preferredLanguage, writeDelayMs, terminalOutputLimit, @@ -2250,6 +2371,8 @@ export class ClineProvider implements vscode.WebviewViewProvider { soundVolume: soundVolume ?? 0.5, browserViewportSize: browserViewportSize ?? "900x600", screenshotQuality: screenshotQuality ?? 75, + remoteBrowserHost, + remoteBrowserEnabled: remoteBrowserEnabled ?? false, preferredLanguage: preferredLanguage ?? "English", writeDelayMs: writeDelayMs ?? 1000, terminalOutputLimit: terminalOutputLimit ?? TERMINAL_OUTPUT_LIMIT, @@ -2403,6 +2526,8 @@ export class ClineProvider implements vscode.WebviewViewProvider { soundVolume: stateValues.soundVolume, browserViewportSize: stateValues.browserViewportSize ?? "900x600", screenshotQuality: stateValues.screenshotQuality ?? 75, + remoteBrowserHost: stateValues.remoteBrowserHost, + remoteBrowserEnabled: stateValues.remoteBrowserEnabled ?? false, fuzzyMatchThreshold: stateValues.fuzzyMatchThreshold ?? 1.0, writeDelayMs: stateValues.writeDelayMs ?? 1000, terminalOutputLimit: stateValues.terminalOutputLimit ?? TERMINAL_OUTPUT_LIMIT, diff --git a/src/core/webview/__tests__/ClineProvider.test.ts b/src/core/webview/__tests__/ClineProvider.test.ts index 2e9fcdf336..be77b13c7a 100644 --- a/src/core/webview/__tests__/ClineProvider.test.ts +++ b/src/core/webview/__tests__/ClineProvider.test.ts @@ -55,6 +55,34 @@ jest.mock("../../contextProxy", () => { // Mock dependencies jest.mock("vscode") jest.mock("delay") + +// Mock BrowserSession +jest.mock("../../../services/browser/BrowserSession", () => ({ + BrowserSession: jest.fn().mockImplementation(() => ({ + testConnection: jest.fn().mockImplementation(async (url) => { + if (url === "http://localhost:9222") { + return { + success: true, + message: "Successfully connected to Chrome", + endpoint: "ws://localhost:9222/devtools/browser/123", + } + } else { + return { + success: false, + message: "Failed to connect to Chrome", + endpoint: undefined, + } + } + }), + })), +})) + +// Mock browserDiscovery +jest.mock("../../../services/browser/browserDiscovery", () => ({ + discoverChromeInstances: jest.fn().mockImplementation(async () => { + return "http://localhost:9222" + }), +})) jest.mock( "@modelcontextprotocol/sdk/types.js", () => ({ @@ -94,31 +122,7 @@ jest.mock("delay", () => { return delayFn }) -// Mock MCP-related modules -jest.mock( - "@modelcontextprotocol/sdk/types.js", - () => ({ - CallToolResultSchema: {}, - ListResourcesResultSchema: {}, - ListResourceTemplatesResultSchema: {}, - ListToolsResultSchema: {}, - ReadResourceResultSchema: {}, - ErrorCode: { - InvalidRequest: "InvalidRequest", - MethodNotFound: "MethodNotFound", - InternalError: "InternalError", - }, - McpError: class McpError extends Error { - code: string - constructor(code: string, message: string) { - super(message) - this.code = code - this.name = "McpError" - } - }, - }), - { virtual: true }, -) +// MCP-related modules are mocked once above (lines 87-109) jest.mock( "@modelcontextprotocol/sdk/client/index.js", @@ -598,7 +602,7 @@ describe("ClineProvider", () => { expect(mockPostMessage).toHaveBeenCalled() }) - test("requestDelaySeconds defaults to 5 seconds", async () => { + test("requestDelaySeconds defaults to 10 seconds", async () => { // Mock globalState.get to return undefined for requestDelaySeconds ;(mockContext.globalState.get as jest.Mock).mockImplementation((key: string) => { if (key === "requestDelaySeconds") { @@ -1160,6 +1164,17 @@ describe("ClineProvider", () => { }) test("passes diffStrategy and diffEnabled to SYSTEM_PROMPT when previewing", async () => { + // Setup Cline instance with mocked api.getModel() + const { Cline } = require("../../Cline") + const mockCline = new Cline() + mockCline.api = { + getModel: jest.fn().mockReturnValue({ + id: "claude-3-sonnet", + info: { supportsComputerUse: true }, + }), + } + await provider.addClineToStack(mockCline) + // Mock getState to return experimentalDiffStrategy, diffEnabled and fuzzyMatchThreshold jest.spyOn(provider, "getState").mockResolvedValue({ apiConfiguration: { @@ -1176,6 +1191,7 @@ describe("ClineProvider", () => { diffEnabled: true, fuzzyMatchThreshold: 0.8, experiments: experimentDefault, + browserToolEnabled: true, } as any) // Mock SYSTEM_PROMPT to verify diffStrategy and diffEnabled are passed @@ -1186,27 +1202,19 @@ describe("ClineProvider", () => { const handler = getMessageHandler() await handler({ type: "getSystemPrompt", mode: "code" }) - // Verify SYSTEM_PROMPT was called with correct arguments - expect(systemPromptSpy).toHaveBeenCalledWith( - expect.anything(), // context - expect.any(String), // cwd - true, // supportsComputerUse - undefined, // mcpHub (disabled) - expect.objectContaining({ - // diffStrategy - getToolDescription: expect.any(Function), - }), - "900x600", // browserViewportSize - "code", // mode - {}, // customModePrompts - { customModes: [] }, // customModes - undefined, // effectiveInstructions - undefined, // preferredLanguage - true, // diffEnabled - experimentDefault, - true, - undefined, // rooIgnoreInstructions - ) + // Verify SYSTEM_PROMPT was called + expect(systemPromptSpy).toHaveBeenCalled() + + // Get the actual arguments passed to SYSTEM_PROMPT + const callArgs = systemPromptSpy.mock.calls[0] + + // Verify key parameters + expect(callArgs[2]).toBe(true) // supportsComputerUse + expect(callArgs[3]).toBeUndefined() // mcpHub (disabled) + expect(callArgs[4]).toHaveProperty("getToolDescription") // diffStrategy + expect(callArgs[5]).toBe("900x600") // browserViewportSize + expect(callArgs[6]).toBe("code") // mode + expect(callArgs[11]).toBe(true) // diffEnabled // Run the test again to verify it's consistent await handler({ type: "getSystemPrompt", mode: "code" }) @@ -1214,6 +1222,17 @@ describe("ClineProvider", () => { }) test("passes diffEnabled: false to SYSTEM_PROMPT when diff is disabled", async () => { + // Setup Cline instance with mocked api.getModel() + const { Cline } = require("../../Cline") + const mockCline = new Cline() + mockCline.api = { + getModel: jest.fn().mockReturnValue({ + id: "claude-3-sonnet", + info: { supportsComputerUse: true }, + }), + } + await provider.addClineToStack(mockCline) + // Mock getState to return diffEnabled: false jest.spyOn(provider, "getState").mockResolvedValue({ apiConfiguration: { @@ -1230,6 +1249,7 @@ describe("ClineProvider", () => { fuzzyMatchThreshold: 0.8, experiments: experimentDefault, enableMcpServerCreation: true, + browserToolEnabled: true, } as any) // Mock SYSTEM_PROMPT to verify diffEnabled is passed as false @@ -1240,27 +1260,19 @@ describe("ClineProvider", () => { const handler = getMessageHandler() await handler({ type: "getSystemPrompt", mode: "code" }) - // Verify SYSTEM_PROMPT was called with diffEnabled: false - expect(systemPromptSpy).toHaveBeenCalledWith( - expect.anything(), // context - expect.any(String), // cwd - true, // supportsComputerUse - undefined, // mcpHub (disabled) - expect.objectContaining({ - // diffStrategy - getToolDescription: expect.any(Function), - }), - "900x600", // browserViewportSize - "code", // mode - {}, // customModePrompts - { customModes: [] }, // customModes - undefined, // effectiveInstructions - undefined, // preferredLanguage - false, // diffEnabled - experimentDefault, - true, - undefined, // rooIgnoreInstructions - ) + // Verify SYSTEM_PROMPT was called + expect(systemPromptSpy).toHaveBeenCalled() + + // Get the actual arguments passed to SYSTEM_PROMPT + const callArgs = systemPromptSpy.mock.calls[0] + + // Verify key parameters + expect(callArgs[2]).toBe(true) // supportsComputerUse + expect(callArgs[3]).toBeUndefined() // mcpHub (disabled) + expect(callArgs[4]).toHaveProperty("getToolDescription") // diffStrategy + expect(callArgs[5]).toBe("900x600") // browserViewportSize + expect(callArgs[6]).toBe("code") // mode + expect(callArgs[11]).toBe(false) // diffEnabled should be false }) test("uses correct mode-specific instructions when mode is specified", async () => { @@ -1299,6 +1311,188 @@ describe("ClineProvider", () => { expect.any(String), ) }) + + // Tests for browser tool support + test("correctly extracts modelSupportsComputerUse from Cline instance", async () => { + // Setup Cline instance with mocked api.getModel() + const { Cline } = require("../../Cline") + const mockCline = new Cline() + mockCline.api = { + getModel: jest.fn().mockReturnValue({ + id: "claude-3-sonnet", + info: { supportsComputerUse: true }, + }), + } + await provider.addClineToStack(mockCline) + + // Mock SYSTEM_PROMPT to verify supportsComputerUse is passed correctly + const systemPromptModule = require("../../prompts/system") + const systemPromptSpy = jest.spyOn(systemPromptModule, "SYSTEM_PROMPT") + + // Mock getState to return browserToolEnabled: true + jest.spyOn(provider, "getState").mockResolvedValue({ + apiConfiguration: { + apiProvider: "openrouter", + }, + browserToolEnabled: true, + mode: "code", + experiments: experimentDefault, + } as any) + + // Trigger getSystemPrompt + const handler = getMessageHandler() + await handler({ type: "getSystemPrompt", mode: "code" }) + + // Verify SYSTEM_PROMPT was called + expect(systemPromptSpy).toHaveBeenCalled() + + // Get the actual arguments passed to SYSTEM_PROMPT + const callArgs = systemPromptSpy.mock.calls[0] + + // Verify the supportsComputerUse parameter (3rd parameter, index 2) + expect(callArgs[2]).toBe(true) + }) + + test("correctly handles when model doesn't support computer use", async () => { + // Setup Cline instance with mocked api.getModel() that doesn't support computer use + const { Cline } = require("../../Cline") + const mockCline = new Cline() + mockCline.api = { + getModel: jest.fn().mockReturnValue({ + id: "non-computer-use-model", + info: { supportsComputerUse: false }, + }), + } + await provider.addClineToStack(mockCline) + + // Mock SYSTEM_PROMPT to verify supportsComputerUse is passed correctly + const systemPromptModule = require("../../prompts/system") + const systemPromptSpy = jest.spyOn(systemPromptModule, "SYSTEM_PROMPT") + + // Mock getState to return browserToolEnabled: true + jest.spyOn(provider, "getState").mockResolvedValue({ + apiConfiguration: { + apiProvider: "openrouter", + }, + browserToolEnabled: true, + mode: "code", + experiments: experimentDefault, + } as any) + + // Trigger getSystemPrompt + const handler = getMessageHandler() + await handler({ type: "getSystemPrompt", mode: "code" }) + + // Verify SYSTEM_PROMPT was called + expect(systemPromptSpy).toHaveBeenCalled() + + // Get the actual arguments passed to SYSTEM_PROMPT + const callArgs = systemPromptSpy.mock.calls[0] + + // Verify the supportsComputerUse parameter (3rd parameter, index 2) + // Even though browserToolEnabled is true, the model doesn't support it + expect(callArgs[2]).toBe(false) + }) + + test("correctly handles when browserToolEnabled is false", async () => { + // Setup Cline instance with mocked api.getModel() that supports computer use + const { Cline } = require("../../Cline") + const mockCline = new Cline() + mockCline.api = { + getModel: jest.fn().mockReturnValue({ + id: "claude-3-sonnet", + info: { supportsComputerUse: true }, + }), + } + await provider.addClineToStack(mockCline) + + // Mock SYSTEM_PROMPT to verify supportsComputerUse is passed correctly + const systemPromptModule = require("../../prompts/system") + const systemPromptSpy = jest.spyOn(systemPromptModule, "SYSTEM_PROMPT") + + // Mock getState to return browserToolEnabled: false + jest.spyOn(provider, "getState").mockResolvedValue({ + apiConfiguration: { + apiProvider: "openrouter", + }, + browserToolEnabled: false, + mode: "code", + experiments: experimentDefault, + } as any) + + // Trigger getSystemPrompt + const handler = getMessageHandler() + await handler({ type: "getSystemPrompt", mode: "code" }) + + // Verify SYSTEM_PROMPT was called + expect(systemPromptSpy).toHaveBeenCalled() + + // Get the actual arguments passed to SYSTEM_PROMPT + const callArgs = systemPromptSpy.mock.calls[0] + + // Verify the supportsComputerUse parameter (3rd parameter, index 2) + // Even though model supports it, browserToolEnabled is false + expect(callArgs[2]).toBe(false) + }) + + test("correctly calculates canUseBrowserTool as combination of model support and setting", async () => { + // Setup Cline instance with mocked api.getModel() + const { Cline } = require("../../Cline") + const mockCline = new Cline() + mockCline.api = { + getModel: jest.fn().mockReturnValue({ + id: "claude-3-sonnet", + info: { supportsComputerUse: true }, + }), + } + await provider.addClineToStack(mockCline) + + // Mock SYSTEM_PROMPT + const systemPromptModule = require("../../prompts/system") + const systemPromptSpy = jest.spyOn(systemPromptModule, "SYSTEM_PROMPT") + + // Test all combinations of model support and browserToolEnabled + const testCases = [ + { modelSupports: true, settingEnabled: true, expected: true }, + { modelSupports: true, settingEnabled: false, expected: false }, + { modelSupports: false, settingEnabled: true, expected: false }, + { modelSupports: false, settingEnabled: false, expected: false }, + ] + + for (const testCase of testCases) { + // Reset mocks + systemPromptSpy.mockClear() + + // Update mock Cline instance + mockCline.api.getModel = jest.fn().mockReturnValue({ + id: "test-model", + info: { supportsComputerUse: testCase.modelSupports }, + }) + + // Mock getState + jest.spyOn(provider, "getState").mockResolvedValue({ + apiConfiguration: { + apiProvider: "openrouter", + }, + browserToolEnabled: testCase.settingEnabled, + mode: "code", + experiments: experimentDefault, + } as any) + + // Trigger getSystemPrompt + const handler = getMessageHandler() + await handler({ type: "getSystemPrompt", mode: "code" }) + + // Verify SYSTEM_PROMPT was called + expect(systemPromptSpy).toHaveBeenCalled() + + // Get the actual arguments passed to SYSTEM_PROMPT + const callArgs = systemPromptSpy.mock.calls[0] + + // Verify the supportsComputerUse parameter (3rd parameter, index 2) + expect(callArgs[2]).toBe(testCase.expected) + } + }) }) describe("handleModeSwitch", () => { @@ -1591,6 +1785,173 @@ describe("ClineProvider", () => { ]) }) }) + + describe("browser connection features", () => { + beforeEach(async () => { + // Reset mocks + jest.clearAllMocks() + await provider.resolveWebviewView(mockWebviewView) + }) + + // Mock BrowserSession and discoverChromeInstances + jest.mock("../../../services/browser/BrowserSession", () => ({ + BrowserSession: jest.fn().mockImplementation(() => ({ + testConnection: jest.fn().mockImplementation(async (url) => { + if (url === "http://localhost:9222") { + return { + success: true, + message: "Successfully connected to Chrome", + endpoint: "ws://localhost:9222/devtools/browser/123", + } + } else { + return { + success: false, + message: "Failed to connect to Chrome", + endpoint: undefined, + } + } + }), + })), + })) + + jest.mock("../../../services/browser/browserDiscovery", () => ({ + discoverChromeInstances: jest.fn().mockImplementation(async () => { + return "http://localhost:9222" + }), + })) + + test("handles testBrowserConnection with provided URL", async () => { + // Get the message handler + const messageHandler = (mockWebviewView.webview.onDidReceiveMessage as jest.Mock).mock.calls[0][0] + + // Test with valid URL + await messageHandler({ + type: "testBrowserConnection", + text: "http://localhost:9222", + }) + + // Verify postMessage was called with success result + expect(mockPostMessage).toHaveBeenCalledWith( + expect.objectContaining({ + type: "browserConnectionResult", + success: true, + text: expect.stringContaining("Successfully connected to Chrome"), + }), + ) + + // Reset mock + mockPostMessage.mockClear() + + // Test with invalid URL + await messageHandler({ + type: "testBrowserConnection", + text: "http://inlocalhost:9222", + }) + + // Verify postMessage was called with failure result + expect(mockPostMessage).toHaveBeenCalledWith( + expect.objectContaining({ + type: "browserConnectionResult", + success: false, + text: expect.stringContaining("Failed to connect to Chrome"), + }), + ) + }) + + test("handles testBrowserConnection with auto-discovery", async () => { + // Get the message handler + const messageHandler = (mockWebviewView.webview.onDidReceiveMessage as jest.Mock).mock.calls[0][0] + + // Test auto-discovery (no URL provided) + await messageHandler({ + type: "testBrowserConnection", + }) + + // Verify discoverChromeInstances was called + const { discoverChromeInstances } = require("../../../services/browser/browserDiscovery") + expect(discoverChromeInstances).toHaveBeenCalled() + + // Verify postMessage was called with success result + expect(mockPostMessage).toHaveBeenCalledWith( + expect.objectContaining({ + type: "browserConnectionResult", + success: true, + text: expect.stringContaining("Auto-discovered and tested connection to Chrome"), + }), + ) + }) + + test("handles discoverBrowser message", async () => { + // Get the message handler + const messageHandler = (mockWebviewView.webview.onDidReceiveMessage as jest.Mock).mock.calls[0][0] + + // Test browser discovery + await messageHandler({ + type: "discoverBrowser", + }) + + // Verify discoverChromeInstances was called + const { discoverChromeInstances } = require("../../../services/browser/browserDiscovery") + expect(discoverChromeInstances).toHaveBeenCalled() + + // Verify postMessage was called with success result + expect(mockPostMessage).toHaveBeenCalledWith( + expect.objectContaining({ + type: "browserConnectionResult", + success: true, + text: expect.stringContaining("Successfully discovered and connected to Chrome"), + }), + ) + }) + + test("handles errors during browser discovery", async () => { + // Mock discoverChromeInstances to throw an error + const { discoverChromeInstances } = require("../../../services/browser/browserDiscovery") + discoverChromeInstances.mockImplementationOnce(() => { + throw new Error("Discovery error") + }) + + // Get the message handler + const messageHandler = (mockWebviewView.webview.onDidReceiveMessage as jest.Mock).mock.calls[0][0] + + // Test browser discovery with error + await messageHandler({ + type: "discoverBrowser", + }) + + // Verify postMessage was called with error result + expect(mockPostMessage).toHaveBeenCalledWith( + expect.objectContaining({ + type: "browserConnectionResult", + success: false, + text: expect.stringContaining("Error discovering browser"), + }), + ) + }) + + test("handles case when no browsers are discovered", async () => { + // Mock discoverChromeInstances to return null (no browsers found) + const { discoverChromeInstances } = require("../../../services/browser/browserDiscovery") + discoverChromeInstances.mockImplementationOnce(() => null) + + // Get the message handler + const messageHandler = (mockWebviewView.webview.onDidReceiveMessage as jest.Mock).mock.calls[0][0] + + // Test browser discovery with no browsers found + await messageHandler({ + type: "discoverBrowser", + }) + + // Verify postMessage was called with failure result + expect(mockPostMessage).toHaveBeenCalledWith( + expect.objectContaining({ + type: "browserConnectionResult", + success: false, + text: expect.stringContaining("No Chrome instances found"), + }), + ) + }) + }) }) describe("ContextProxy integration", () => { diff --git a/src/exports/README.md b/src/exports/README.md index 03b8983b7e..0554580836 100644 --- a/src/exports/README.md +++ b/src/exports/README.md @@ -1,55 +1,51 @@ -# Cline API +# Roo Code API -The Cline extension exposes an API that can be used by other extensions. To use this API in your extension: +The Roo Code extension exposes an API that can be used by other extensions. To use this API in your extension: -1. Copy `src/extension-api/cline.d.ts` to your extension's source directory. -2. Include `cline.d.ts` in your extension's compilation. +1. Copy `src/extension-api/roo-code.d.ts` to your extension's source directory. +2. Include `roo-code.d.ts` in your extension's compilation. 3. Get access to the API with the following code: - ```ts - const clineExtension = vscode.extensions.getExtension("rooveterinaryinc.roo-cline") +```typescript +const extension = vscode.extensions.getExtension("rooveterinaryinc.roo-cline") - if (!clineExtension?.isActive) { - throw new Error("Cline extension is not activated") - } +if (!extension?.isActive) { + throw new Error("Extension is not activated") +} - const cline = clineExtension.exports +const api = extension.exports - if (cline) { - // Now you can use the API +if (!api) { + throw new Error("API is not available") +} - // Set custom instructions - await cline.setCustomInstructions("Talk like a pirate") +// Set custom instructions. +await api.setCustomInstructions("Talk like a pirate") - // Get custom instructions - const instructions = await cline.getCustomInstructions() - console.log("Current custom instructions:", instructions) +// Get custom instructions. +const instructions = await api.getCustomInstructions() +console.log("Current custom instructions:", instructions) - // Start a new task with an initial message - await cline.startNewTask("Hello, Cline! Let's make a new project...") +// Start a new task with an initial message. +await api.startNewTask("Hello, Roo Code API! Let's make a new project...") - // Start a new task with an initial message and images - await cline.startNewTask("Use this design language", ["data:image/webp;base64,..."]) +// Start a new task with an initial message and images. +await api.startNewTask("Use this design language", ["data:image/webp;base64,..."]) - // Send a message to the current task - await cline.sendMessage("Can you fix the @problems?") +// Send a message to the current task. +await api.sendMessage("Can you fix the @problems?") - // Simulate pressing the primary button in the chat interface (e.g. 'Save' or 'Proceed While Running') - await cline.pressPrimaryButton() +// Simulate pressing the primary button in the chat interface (e.g. 'Save' or 'Proceed While Running'). +await api.pressPrimaryButton() - // Simulate pressing the secondary button in the chat interface (e.g. 'Reject') - await cline.pressSecondaryButton() - } else { - console.error("Cline API is not available") - } - ``` +// Simulate pressing the secondary button in the chat interface (e.g. 'Reject'). +await api.pressSecondaryButton() +``` - **Note:** To ensure that the `rooveterinaryinc.roo-cline` extension is activated before your extension, add it to the `extensionDependencies` in your `package.json`: +**NOTE:** To ensure that the `rooveterinaryinc.roo-cline` extension is activated before your extension, add it to the `extensionDependencies` in your `package.json`: - ```json - "extensionDependencies": [ - "rooveterinaryinc.roo-cline" - ] - ``` +```json +"extensionDependencies": ["rooveterinaryinc.roo-cline"] +``` -For detailed information on the available methods and their usage, refer to the `cline.d.ts` file. +For detailed information on the available methods and their usage, refer to the `roo-code.d.ts` file. diff --git a/src/exports/cline.d.ts b/src/exports/roo-code.d.ts similarity index 72% rename from src/exports/cline.d.ts rename to src/exports/roo-code.d.ts index e529947b6b..e310fef6a7 100644 --- a/src/exports/cline.d.ts +++ b/src/exports/roo-code.d.ts @@ -1,4 +1,4 @@ -export interface ClineAPI { +export interface RooCodeAPI { /** * Sets the custom instructions in the global storage. * @param value The custom instructions to be saved. @@ -38,7 +38,61 @@ export interface ClineAPI { /** * The sidebar provider instance. */ - sidebarProvider: ClineSidebarProvider + sidebarProvider: ClineProvider +} + +export type ClineAsk = + | "followup" + | "command" + | "command_output" + | "completion_result" + | "tool" + | "api_req_failed" + | "resume_task" + | "resume_completed_task" + | "mistake_limit_reached" + | "browser_action_launch" + | "use_mcp_server" + | "finishTask" + +export type ClineSay = + | "task" + | "error" + | "api_req_started" + | "api_req_finished" + | "api_req_retried" + | "api_req_retry_delayed" + | "api_req_deleted" + | "text" + | "reasoning" + | "completion_result" + | "user_feedback" + | "user_feedback_diff" + | "command_output" + | "tool" + | "shell_integration_warning" + | "browser_action" + | "browser_action_result" + | "command" + | "mcp_server_request_started" + | "mcp_server_response" + | "new_task_started" + | "new_task" + | "checkpoint_saved" + | "rooignore_error" + +export interface ClineMessage { + ts: number + type: "ask" | "say" + ask?: ClineAsk + say?: ClineSay + text?: string + images?: string[] + partial?: boolean + reasoning?: string + conversationHistoryIndex?: number + checkpoint?: Record + progressStatus?: ToolProgressStatus } export interface ClineProvider { @@ -82,11 +136,6 @@ export interface ClineProvider { */ cancelTask(): Promise - /** - * Clears the current task - */ - clearTask(): Promise - /** * Gets the current state */ @@ -112,12 +161,6 @@ export interface ClineProvider { */ storeSecret(key: SecretKey, value?: string): Promise - /** - * Retrieves a secret value from secure storage - * @param key The key of the secret to retrieve - */ - getSecret(key: SecretKey): Promise - /** * Resets the state */ diff --git a/src/extension.ts b/src/extension.ts index d88bf2a251..e3f9977999 100644 --- a/src/extension.ts +++ b/src/extension.ts @@ -11,15 +11,16 @@ try { console.warn("Failed to load environment variables:", e) } -import { ClineProvider } from "./core/webview/ClineProvider" -import { createClineAPI } from "./exports" import "./utils/path" // Necessary to have access to String.prototype.toPosix. + +import { ClineProvider } from "./core/webview/ClineProvider" import { CodeActionProvider } from "./core/CodeActionProvider" import { DIFF_VIEW_URI_SCHEME } from "./integrations/editor/DiffViewProvider" -import { handleUri, registerCommands, registerCodeActions } from "./activate" import { McpServerManager } from "./services/mcp/McpServerManager" import { telemetryService } from "./services/telemetry/TelemetryService" +import { handleUri, registerCommands, registerCodeActions, createRooCodeAPI } from "./activate" + /** * Built using https://github.com/microsoft/vscode-webview-ui-toolkit * @@ -31,18 +32,6 @@ import { telemetryService } from "./services/telemetry/TelemetryService" let outputChannel: vscode.OutputChannel let extensionContext: vscode.ExtensionContext -// Callback mapping of human relay response -const humanRelayCallbacks = new Map void>() - -/** - * Register a callback function for human relay response - * @param requestId - * @param callback - */ -export function registerHumanRelayCallback(requestId: string, callback: (response: string | undefined) => void): void { - humanRelayCallbacks.set(requestId, callback) -} - // This method is called when your extension is activated. // Your extension is activated the very first time the command is executed. export function activate(context: vscode.ExtensionContext) { @@ -72,40 +61,6 @@ export function activate(context: vscode.ExtensionContext) { registerCommands({ context, outputChannel, provider: sidebarProvider }) - // Register human relay callback registration command - context.subscriptions.push( - vscode.commands.registerCommand( - "roo-cline.registerHumanRelayCallback", - (requestId: string, callback: (response: string | undefined) => void) => { - registerHumanRelayCallback(requestId, callback) - }, - ), - ) - - // Register human relay response processing command - context.subscriptions.push( - vscode.commands.registerCommand( - "roo-cline.handleHumanRelayResponse", - (response: { requestId: string; text?: string; cancelled?: boolean }) => { - const callback = humanRelayCallbacks.get(response.requestId) - if (callback) { - if (response.cancelled) { - callback(undefined) - } else { - callback(response.text) - } - humanRelayCallbacks.delete(response.requestId) - } - }, - ), - ) - - context.subscriptions.push( - vscode.commands.registerCommand("roo-cline.unregisterHumanRelayCallback", (requestId: string) => { - humanRelayCallbacks.delete(requestId) - }), - ) - /** * We use the text document content provider API to show the left side for diff * view by creating a virtual document for the original content. This makes it @@ -143,7 +98,7 @@ export function activate(context: vscode.ExtensionContext) { registerCodeActions(context) - return createClineAPI(outputChannel, sidebarProvider) + return createRooCodeAPI(outputChannel, sidebarProvider) } // This method is called when your extension is deactivated. diff --git a/src/services/browser/BrowserSession.ts b/src/services/browser/BrowserSession.ts index bed0332244..5c5f59ffeb 100644 --- a/src/services/browser/BrowserSession.ts +++ b/src/services/browser/BrowserSession.ts @@ -1,13 +1,15 @@ import * as vscode from "vscode" import * as fs from "fs/promises" import * as path from "path" -import { Browser, Page, ScreenshotOptions, TimeoutError, launch } from "puppeteer-core" +import { Browser, Page, ScreenshotOptions, TimeoutError, launch, connect } from "puppeteer-core" // @ts-ignore import PCR from "puppeteer-chromium-resolver" import pWaitFor from "p-wait-for" import delay from "delay" +import axios from "axios" import { fileExistsAtPath } from "../../utils/fs" import { BrowserActionResult } from "../../shared/ExtensionMessage" +import { discoverChromeInstances, testBrowserConnection } from "./browserDiscovery" interface PCRStats { puppeteer: { launch: typeof launch } @@ -19,11 +21,20 @@ export class BrowserSession { private browser?: Browser private page?: Page private currentMousePosition?: string + private cachedWebSocketEndpoint?: string + private lastConnectionAttempt: number = 0 constructor(context: vscode.ExtensionContext) { this.context = context } + /** + * Test connection to a remote browser + */ + async testConnection(host: string): Promise<{ success: boolean; message: string; endpoint?: string }> { + return testBrowserConnection(host) + } + private async ensureChromiumExists(): Promise { const globalStoragePath = this.context?.globalStorageUri?.fsPath if (!globalStoragePath) { @@ -52,17 +63,131 @@ export class BrowserSession { await this.closeBrowser() // this may happen when the model launches a browser again after having used it already before } + // Function to get viewport size + const getViewport = () => { + const size = (this.context.globalState.get("browserViewportSize") as string | undefined) || "900x600" + const [width, height] = size.split("x").map(Number) + return { width, height } + } + + // Check if remote browser connection is enabled + const remoteBrowserEnabled = this.context.globalState.get("remoteBrowserEnabled") as boolean | undefined + + // If remote browser connection is not enabled, use local browser + if (!remoteBrowserEnabled) { + console.log("Remote browser connection is disabled, using local browser") + const stats = await this.ensureChromiumExists() + this.browser = await stats.puppeteer.launch({ + args: [ + "--user-agent=Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/128.0.0.0 Safari/537.36", + ], + executablePath: stats.executablePath, + defaultViewport: getViewport(), + // headless: false, + }) + this.page = await this.browser?.newPage() + return + } + // Remote browser connection is enabled + let remoteBrowserHost = this.context.globalState.get("remoteBrowserHost") as string | undefined + let browserWSEndpoint: string | undefined = this.cachedWebSocketEndpoint + let reconnectionAttempted = false + + // Try to connect with cached endpoint first if it exists and is recent (less than 1 hour old) + if (browserWSEndpoint && Date.now() - this.lastConnectionAttempt < 3600000) { + try { + console.log(`Attempting to connect using cached WebSocket endpoint: ${browserWSEndpoint}`) + this.browser = await connect({ + browserWSEndpoint, + defaultViewport: getViewport(), + }) + this.page = await this.browser?.newPage() + return + } catch (error) { + console.log(`Failed to connect using cached endpoint: ${error}`) + // Clear the cached endpoint since it's no longer valid + this.cachedWebSocketEndpoint = undefined + // User wants to give up after one reconnection attempt + if (remoteBrowserHost) { + reconnectionAttempted = true + } + } + } + + // If user provided a remote browser host, try to connect to it + if (remoteBrowserHost && !reconnectionAttempted) { + console.log(`Attempting to connect to remote browser at ${remoteBrowserHost}`) + try { + // Fetch the WebSocket endpoint from the Chrome DevTools Protocol + const versionUrl = `${remoteBrowserHost.replace(/\/$/, "")}/json/version` + console.log(`Fetching WebSocket endpoint from ${versionUrl}`) + + const response = await axios.get(versionUrl) + browserWSEndpoint = response.data.webSocketDebuggerUrl + + if (!browserWSEndpoint) { + throw new Error("Could not find webSocketDebuggerUrl in the response") + } + + console.log(`Found WebSocket endpoint: ${browserWSEndpoint}`) + + // Cache the successful endpoint + this.cachedWebSocketEndpoint = browserWSEndpoint + this.lastConnectionAttempt = Date.now() + + this.browser = await connect({ + browserWSEndpoint, + defaultViewport: getViewport(), + }) + this.page = await this.browser?.newPage() + return + } catch (error) { + console.error(`Failed to connect to remote browser: ${error}`) + // Fall back to auto-discovery if remote connection fails + } + } + + // Always try auto-discovery if no custom URL is specified or if connection failed + try { + console.log("Attempting auto-discovery...") + const discoveredHost = await discoverChromeInstances() + + if (discoveredHost) { + console.log(`Auto-discovered Chrome at ${discoveredHost}`) + + // Don't save the discovered host to global state to avoid overriding user preference + // We'll just use it for this session + + // Try to connect to the discovered host + const testResult = await testBrowserConnection(discoveredHost) + + if (testResult.success && testResult.endpoint) { + // Cache the successful endpoint + this.cachedWebSocketEndpoint = testResult.endpoint + this.lastConnectionAttempt = Date.now() + + this.browser = await connect({ + browserWSEndpoint: testResult.endpoint, + defaultViewport: getViewport(), + }) + this.page = await this.browser?.newPage() + return + } + } + } catch (error) { + console.error(`Auto-discovery failed: ${error}`) + // Fall back to local browser if auto-discovery fails + } + + // If all remote connection attempts fail, fall back to local browser + console.log("Falling back to local browser") const stats = await this.ensureChromiumExists() this.browser = await stats.puppeteer.launch({ args: [ "--user-agent=Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/128.0.0.0 Safari/537.36", ], executablePath: stats.executablePath, - defaultViewport: (() => { - const size = (this.context.globalState.get("browserViewportSize") as string | undefined) || "900x600" - const [width, height] = size.split("x").map(Number) - return { width, height } - })(), + defaultViewport: getViewport(), // headless: false, }) // (latest version of puppeteer does not add headless to user agent) @@ -72,7 +197,14 @@ export class BrowserSession { async closeBrowser(): Promise { if (this.browser || this.page) { console.log("closing browser...") - await this.browser?.close().catch(() => {}) + + const remoteBrowserEnabled = this.context.globalState.get("remoteBrowserEnabled") as string | undefined + if (remoteBrowserEnabled && this.browser) { + await this.browser.disconnect().catch(() => {}) + } else { + await this.browser?.close().catch(() => {}) + } + this.browser = undefined this.page = undefined this.currentMousePosition = undefined diff --git a/src/services/browser/browserDiscovery.ts b/src/services/browser/browserDiscovery.ts new file mode 100644 index 0000000000..187f90e299 --- /dev/null +++ b/src/services/browser/browserDiscovery.ts @@ -0,0 +1,246 @@ +import * as vscode from "vscode" +import * as os from "os" +import * as net from "net" +import axios from "axios" + +/** + * Check if a port is open on a given host + */ +export async function isPortOpen(host: string, port: number, timeout = 1000): Promise { + return new Promise((resolve) => { + const socket = new net.Socket() + let status = false + + // Set timeout + socket.setTimeout(timeout) + + // Handle successful connection + socket.on("connect", () => { + status = true + socket.destroy() + }) + + // Handle any errors + socket.on("error", () => { + socket.destroy() + }) + + // Handle timeout + socket.on("timeout", () => { + socket.destroy() + }) + + // Handle close + socket.on("close", () => { + resolve(status) + }) + + // Attempt to connect + socket.connect(port, host) + }) +} + +/** + * Try to connect to Chrome at a specific IP address + */ +export async function tryConnect(ipAddress: string): Promise<{ endpoint: string; ip: string } | null> { + try { + console.log(`Trying to connect to Chrome at: http://${ipAddress}:9222/json/version`) + const response = await axios.get(`http://${ipAddress}:9222/json/version`, { timeout: 1000 }) + const data = response.data + return { endpoint: data.webSocketDebuggerUrl, ip: ipAddress } + } catch (error) { + return null + } +} + +/** + * Execute a shell command and return stdout and stderr + */ +export async function executeShellCommand(command: string): Promise<{ stdout: string; stderr: string }> { + return new Promise<{ stdout: string; stderr: string }>((resolve) => { + const cp = require("child_process") + cp.exec(command, (err: any, stdout: string, stderr: string) => { + resolve({ stdout, stderr }) + }) + }) +} + +/** + * Get Docker gateway IP without UI feedback + */ +export async function getDockerGatewayIP(): Promise { + try { + if (process.platform === "linux") { + try { + const { stdout } = await executeShellCommand("ip route | grep default | awk '{print $3}'") + return stdout.trim() + } catch (error) { + console.log("Could not determine Docker gateway IP:", error) + } + } + return null + } catch (error) { + console.log("Could not determine Docker gateway IP:", error) + return null + } +} + +/** + * Get Docker host IP + */ +export async function getDockerHostIP(): Promise { + try { + // Try to resolve host.docker.internal (works on Docker Desktop) + return new Promise((resolve) => { + const dns = require("dns") + dns.lookup("host.docker.internal", (err: any, address: string) => { + if (err) { + resolve(null) + } else { + resolve(address) + } + }) + }) + } catch (error) { + console.log("Could not determine Docker host IP:", error) + return null + } +} + +/** + * Scan a network range for Chrome debugging port + */ +export async function scanNetworkForChrome(baseIP: string): Promise { + if (!baseIP || !baseIP.match(/^\d+\.\d+\.\d+\./)) { + return null + } + + // Extract the network prefix (e.g., "192.168.65.") + const networkPrefix = baseIP.split(".").slice(0, 3).join(".") + "." + + // Common Docker host IPs to try first + const priorityIPs = [ + networkPrefix + "1", // Common gateway + networkPrefix + "2", // Common host + networkPrefix + "254", // Common host in some Docker setups + ] + + console.log(`Scanning priority IPs in network ${networkPrefix}*`) + + // Check priority IPs first + for (const ip of priorityIPs) { + const isOpen = await isPortOpen(ip, 9222) + if (isOpen) { + console.log(`Found Chrome debugging port open on ${ip}`) + return ip + } + } + + return null +} + +/** + * Discover Chrome instances on the network + */ +export async function discoverChromeInstances(): Promise { + // Get all network interfaces + const networkInterfaces = os.networkInterfaces() + const ipAddresses = [] + + // Always try localhost first + ipAddresses.push("localhost") + ipAddresses.push("127.0.0.1") + + // Try to get Docker gateway IP (headless mode) + const gatewayIP = await getDockerGatewayIP() + if (gatewayIP) { + console.log("Found Docker gateway IP:", gatewayIP) + ipAddresses.push(gatewayIP) + } + + // Try to get Docker host IP + const hostIP = await getDockerHostIP() + if (hostIP) { + console.log("Found Docker host IP:", hostIP) + ipAddresses.push(hostIP) + } + + // Add all local IP addresses from network interfaces + const localIPs: string[] = [] + Object.values(networkInterfaces).forEach((interfaces) => { + if (!interfaces) return + interfaces.forEach((iface) => { + // Only consider IPv4 addresses + if (iface.family === "IPv4" || iface.family === (4 as any)) { + localIPs.push(iface.address) + } + }) + }) + + // Add local IPs to the list + ipAddresses.push(...localIPs) + + // Scan network for Chrome debugging port + for (const ip of localIPs) { + const chromeIP = await scanNetworkForChrome(ip) + if (chromeIP && !ipAddresses.includes(chromeIP)) { + console.log("Found potential Chrome host via network scan:", chromeIP) + ipAddresses.push(chromeIP) + } + } + + // Remove duplicates + const uniqueIPs = [...new Set(ipAddresses)] + console.log("IP Addresses to try:", uniqueIPs) + + // Try connecting to each IP address + for (const ip of uniqueIPs) { + const connection = await tryConnect(ip) + if (connection) { + console.log(`Successfully connected to Chrome at: ${connection.ip}`) + // Store the successful IP for future use + console.log(`✅ Found Chrome at ${connection.ip} - You can hardcode this IP if needed`) + + // Return the host URL and endpoint + return `http://${connection.ip}:9222` + } + } + + return null +} + +/** + * Test connection to a remote browser + */ +export async function testBrowserConnection( + host: string, +): Promise<{ success: boolean; message: string; endpoint?: string }> { + try { + // Fetch the WebSocket endpoint from the Chrome DevTools Protocol + const versionUrl = `${host.replace(/\/$/, "")}/json/version` + console.log(`Testing connection to ${versionUrl}`) + + const response = await axios.get(versionUrl, { timeout: 3000 }) + const browserWSEndpoint = response.data.webSocketDebuggerUrl + + if (!browserWSEndpoint) { + return { + success: false, + message: "Could not find webSocketDebuggerUrl in the response", + } + } + + return { + success: true, + message: "Successfully connected to Chrome browser", + endpoint: browserWSEndpoint, + } + } catch (error) { + console.error(`Failed to connect to remote browser: ${error}`) + return { + success: false, + message: `Failed to connect: ${error instanceof Error ? error.message : String(error)}`, + } + } +} diff --git a/src/shared/ExtensionMessage.ts b/src/shared/ExtensionMessage.ts index 73f8127c53..9298b0cb1b 100644 --- a/src/shared/ExtensionMessage.ts +++ b/src/shared/ExtensionMessage.ts @@ -1,5 +1,3 @@ -// type that represents json data that is sent from extension to webview, called ExtensionMessage and has 'type' enum which can be 'plusButtonClicked' or 'settingsButtonClicked' or 'hello' - import { ApiConfiguration, ApiProvider, ModelInfo } from "./api" import { HistoryItem } from "./HistoryItem" import { McpServer } from "./mcp" @@ -9,6 +7,7 @@ import { CustomSupportPrompts } from "./support-prompt" import { ExperimentId } from "./experiments" import { CheckpointStorage } from "./checkpoints" import { TelemetrySetting } from "./TelemetrySetting" +import { ClineMessage, ClineAsk, ClineSay } from "../exports/roo-code" export interface LanguageModelChatSelector { vendor?: string @@ -17,7 +16,9 @@ export interface LanguageModelChatSelector { id?: string } -// webview will hold state +// Represents JSON data that is sent from extension to webview, called +// ExtensionMessage and has 'type' enum which can be 'plusButtonClicked' or +// 'settingsButtonClicked' or 'hello'. Webview will hold state. export interface ExtensionMessage { type: | "action" @@ -51,6 +52,8 @@ export interface ExtensionMessage { | "humanRelayResponse" | "humanRelayCancel" | "browserToolEnabled" + | "browserConnectionResult" + | "remoteBrowserEnabled" text?: string action?: | "chatButtonClicked" @@ -83,6 +86,8 @@ export interface ExtensionMessage { mode?: Mode customMode?: ModeConfig slug?: string + success?: boolean + values?: Record } export interface ApiConfigMeta { @@ -123,6 +128,8 @@ export interface ExtensionState { checkpointStorage: CheckpointStorage browserViewportSize?: string screenshotQuality?: number + remoteBrowserHost?: string + remoteBrowserEnabled?: boolean fuzzyMatchThreshold?: number preferredLanguage: string writeDelayMs: number @@ -145,58 +152,7 @@ export interface ExtensionState { showRooIgnoredFiles: boolean // Whether to show .rooignore'd files in listings } -export interface ClineMessage { - ts: number - type: "ask" | "say" - ask?: ClineAsk - say?: ClineSay - text?: string - images?: string[] - partial?: boolean - reasoning?: string - conversationHistoryIndex?: number - checkpoint?: Record -} - -export type ClineAsk = - | "followup" - | "command" - | "command_output" - | "completion_result" - | "tool" - | "api_req_failed" - | "resume_task" - | "resume_completed_task" - | "mistake_limit_reached" - | "browser_action_launch" - | "use_mcp_server" - | "finishTask" - -export type ClineSay = - | "task" - | "error" - | "api_req_started" - | "api_req_finished" - | "api_req_retried" - | "api_req_retry_delayed" - | "api_req_deleted" - | "text" - | "reasoning" - | "completion_result" - | "user_feedback" - | "user_feedback_diff" - | "command_output" - | "tool" - | "shell_integration_warning" - | "browser_action" - | "browser_action_result" - | "command" - | "mcp_server_request_started" - | "mcp_server_response" - | "new_task_started" - | "new_task" - | "checkpoint_saved" - | "rooignore_error" +export type { ClineMessage, ClineAsk, ClineSay } export interface ClineSayTool { tool: @@ -220,8 +176,9 @@ export interface ClineSayTool { reason?: string } -// must keep in sync with system prompt +// Must keep in sync with system prompt. export const browserActions = ["launch", "click", "type", "scroll_down", "scroll_up", "close"] as const + export type BrowserAction = (typeof browserActions)[number] export interface ClineSayBrowserAction { @@ -256,7 +213,6 @@ export interface ClineApiReqInfo { streamingFailedMessage?: string } -// Human relay related message types export interface ShowHumanRelayDialogMessage { type: "showHumanRelayDialog" requestId: string @@ -275,3 +231,8 @@ export interface HumanRelayCancelMessage { } export type ClineApiReqCancelReason = "streaming_failed" | "user_cancelled" + +export type ToolProgressStatus = { + icon?: string + text?: string +} diff --git a/src/shared/WebviewMessage.ts b/src/shared/WebviewMessage.ts index 724d7e5983..dc2ddd6c77 100644 --- a/src/shared/WebviewMessage.ts +++ b/src/shared/WebviewMessage.ts @@ -57,6 +57,7 @@ export interface WebviewMessage { | "checkpointStorage" | "browserViewportSize" | "screenshotQuality" + | "remoteBrowserHost" | "openMcpSettings" | "restartMcpServer" | "toggleToolAlwaysAllow" @@ -102,6 +103,10 @@ export interface WebviewMessage { | "browserToolEnabled" | "telemetrySetting" | "showRooIgnoredFiles" + | "testBrowserConnection" + | "discoverBrowser" + | "browserConnectionResult" + | "remoteBrowserEnabled" text?: string disabled?: boolean askResponse?: ClineAskResponse diff --git a/src/shared/api.ts b/src/shared/api.ts index 98d595cd03..329531fd66 100644 --- a/src/shared/api.ts +++ b/src/shared/api.ts @@ -39,6 +39,7 @@ export interface ApiHandlerOptions { awspromptCacheId?: string awsProfile?: string awsUseProfile?: boolean + awsCustomArn?: string vertexKeyFile?: string vertexJsonCredentials?: string vertexProjectId?: string @@ -99,6 +100,7 @@ export const API_CONFIG_KEYS: GlobalStateKey[] = [ // "awspromptCacheId", // NOT exist on GlobalStateKey "awsProfile", "awsUseProfile", + "awsCustomArn", "vertexKeyFile", "vertexJsonCredentials", "vertexProjectId", @@ -496,6 +498,14 @@ export const vertexModels = { inputPrice: 0.15, outputPrice: 0.6, }, + "gemini-2.0-pro-exp-02-05": { + maxTokens: 8192, + contextWindow: 2_097_152, + supportsImages: true, + supportsPromptCache: false, + inputPrice: 0, + outputPrice: 0, + }, "gemini-2.0-flash-lite-001": { maxTokens: 8192, contextWindow: 1_048_576, diff --git a/src/shared/globalState.ts b/src/shared/globalState.ts index 540b7e72be..40973b3350 100644 --- a/src/shared/globalState.ts +++ b/src/shared/globalState.ts @@ -28,6 +28,7 @@ export const GLOBAL_STATE_KEYS = [ "awsUseCrossRegionInference", "awsProfile", "awsUseProfile", + "awsCustomArn", "vertexKeyFile", "vertexJsonCredentials", "vertexProjectId", @@ -66,6 +67,7 @@ export const GLOBAL_STATE_KEYS = [ "checkpointStorage", "browserViewportSize", "screenshotQuality", + "remoteBrowserHost", "fuzzyMatchThreshold", "preferredLanguage", // Language setting for Cline's communication "writeDelayMs", @@ -100,6 +102,7 @@ export const GLOBAL_STATE_KEYS = [ "lmStudioDraftModelId", "telemetrySetting", "showRooIgnoredFiles", + "remoteBrowserEnabled", ] as const // Derive the type from the array - creates a union of string literals diff --git a/webview-ui/src/components/chat/ChatRow.tsx b/webview-ui/src/components/chat/ChatRow.tsx index 259c03fa21..b19d67dc05 100644 --- a/webview-ui/src/components/chat/ChatRow.tsx +++ b/webview-ui/src/components/chat/ChatRow.tsx @@ -258,6 +258,7 @@ export const ChatRowContent = ({ Roo wants to edit this file: void isLoading?: boolean + progressStatus?: ToolProgressStatus } /* @@ -32,6 +34,7 @@ const CodeAccordian = ({ isExpanded, onToggleExpand, isLoading, + progressStatus, }: CodeAccordianProps) => { const inferredLanguage = useMemo( () => code && (language ?? (path ? getLanguageFromPath(path) : undefined)), @@ -95,6 +98,14 @@ const CodeAccordian = ({ )}
+ {progressStatus && progressStatus.text && ( + <> + {progressStatus.icon && } + + {progressStatus.text} + + + )} )} diff --git a/webview-ui/src/components/settings/ApiOptions.tsx b/webview-ui/src/components/settings/ApiOptions.tsx index f7982080c8..656a6831cb 100644 --- a/webview-ui/src/components/settings/ApiOptions.tsx +++ b/webview-ui/src/components/settings/ApiOptions.tsx @@ -41,7 +41,7 @@ import { VSCodeButtonLink } from "../common/VSCodeButtonLink" import { ModelInfoView } from "./ModelInfoView" import { ModelPicker } from "./ModelPicker" import { TemperatureControl } from "./TemperatureControl" -import { validateApiConfiguration, validateModelId } from "@/utils/validate" +import { validateApiConfiguration, validateModelId, validateBedrockArn } from "@/utils/validate" import { ApiErrorMessage } from "./ApiErrorMessage" import { ThinkingBudget } from "./ThinkingBudget" @@ -1267,14 +1267,82 @@ const ApiOptions = ({ { - setApiConfigurationField("apiModelId", typeof value == "string" ? value : value?.value) + const modelValue = typeof value == "string" ? value : value?.value + setApiConfigurationField("apiModelId", modelValue) + + // Clear custom ARN if not using custom ARN option + if (modelValue !== "custom-arn" && selectedProvider === "bedrock") { + setApiConfigurationField("awsCustomArn", "") + } }} - options={selectedProviderModelOptions} + options={[ + ...selectedProviderModelOptions, + ...(selectedProvider === "bedrock" + ? [{ value: "custom-arn", label: "Use custom ARN..." }] + : []), + ]} className="w-full" /> + + {selectedProvider === "bedrock" && selectedModelId === "custom-arn" && ( + <> + { + const value = (e.target as HTMLInputElement).value + setApiConfigurationField("awsCustomArn", value) + }} + placeholder="Enter ARN (e.g. arn:aws:bedrock:us-east-1:123456789012:foundation-model/my-model)" + className="w-full"> + Custom ARN + +
+ Enter a valid AWS Bedrock ARN for the model you want to use. Format examples: +
    +
  • + arn:aws:bedrock:us-east-1:123456789012:foundation-model/anthropic.claude-3-sonnet-20240229-v1:0 +
  • +
  • + arn:aws:bedrock:us-west-2:123456789012:provisioned-model/my-provisioned-model +
  • +
  • + arn:aws:bedrock:us-east-1:123456789012:default-prompt-router/anthropic.claude:1 +
  • +
+ Make sure the region in the ARN matches your selected AWS Region above. +
+ {apiConfiguration?.awsCustomArn && + (() => { + const validation = validateBedrockArn( + apiConfiguration.awsCustomArn, + apiConfiguration.awsRegion, + ) + + if (!validation.isValid) { + return ( +
+ {validation.errorMessage || + "Invalid ARN format. Please check the examples above."} +
+ ) + } + + if (validation.errorMessage) { + return ( +
+ {validation.errorMessage} +
+ ) + } + + return null + })()} + ======= + + )} & { browserToolEnabled?: boolean browserViewportSize?: string screenshotQuality?: number - setCachedStateField: SetCachedStateField<"browserToolEnabled" | "browserViewportSize" | "screenshotQuality"> + remoteBrowserHost?: string + remoteBrowserEnabled?: boolean + setCachedStateField: SetCachedStateField< + | "browserToolEnabled" + | "browserViewportSize" + | "screenshotQuality" + | "remoteBrowserHost" + | "remoteBrowserEnabled" + > } export const BrowserSettings = ({ browserToolEnabled, browserViewportSize, screenshotQuality, + remoteBrowserHost, + remoteBrowserEnabled, setCachedStateField, ...props }: BrowserSettingsProps) => { + const [testingConnection, setTestingConnection] = useState(false) + const [testResult, setTestResult] = useState<{ success: boolean; message: string } | null>(null) + const [discovering, setDiscovering] = useState(false) + // We don't need a local state for useRemoteBrowser since we're using the enableRemoteBrowser prop directly + // This ensures the checkbox always reflects the current global state + + // Set up message listener for browser connection results + useEffect(() => { + const handleMessage = (event: MessageEvent) => { + const message = event.data + + if (message.type === "browserConnectionResult") { + setTestResult({ + success: message.success, + message: message.text, + }) + setTestingConnection(false) + setDiscovering(false) + } + } + + window.addEventListener("message", handleMessage) + + return () => { + window.removeEventListener("message", handleMessage) + } + }, []) + + const testConnection = async () => { + setTestingConnection(true) + setTestResult(null) + + try { + // Send a message to the extension to test the connection + vscode.postMessage({ + type: "testBrowserConnection", + text: remoteBrowserHost, + }) + } catch (error) { + setTestResult({ + success: false, + message: `Error: ${error instanceof Error ? error.message : String(error)}`, + }) + setTestingConnection(false) + } + } + + const discoverBrowser = async () => { + setDiscovering(true) + setTestResult(null) + + try { + // Send a message to the extension to discover Chrome instances + vscode.postMessage({ + type: "discoverBrowser", + }) + } catch (error) { + setTestResult({ + success: false, + message: `Error: ${error instanceof Error ? error.message : String(error)}`, + }) + setDiscovering(false) + } + } return (
@@ -96,6 +171,64 @@ export const BrowserSettings = ({ screenshots but increase token usage.

+
+
+ { + // Update the global state - remoteBrowserEnabled now means "enable remote browser connection" + setCachedStateField("remoteBrowserEnabled", e.target.checked) + if (!e.target.checked) { + // If disabling remote browser, clear the custom URL + setCachedStateField("remoteBrowserHost", undefined) + } + }}> + Use remote browser connection + +

+ Connect to a Chrome browser running with remote debugging enabled + (--remote-debugging-port=9222). +

+
+ {remoteBrowserEnabled && ( + <> +
+ + setCachedStateField( + "remoteBrowserHost", + e.target.value || undefined, + ) + } + placeholder="Custom URL (e.g., http://localhost:9222)" + style={{ flexGrow: 1 }} + /> + + {testingConnection || discovering ? "Testing..." : "Test Connection"} + +
+ {testResult && ( +
+ {testResult.message} +
+ )} +

+ Enter the DevTools Protocol host address or + leave empty to auto-discover Chrome local instances. + The Test Connection button will try the custom URL if provided, or + auto-discover if the field is empty. +

+ + )} +
)} diff --git a/webview-ui/src/components/settings/SettingsView.tsx b/webview-ui/src/components/settings/SettingsView.tsx index df08a03971..b2e3cb26c4 100644 --- a/webview-ui/src/components/settings/SettingsView.tsx +++ b/webview-ui/src/components/settings/SettingsView.tsx @@ -1,6 +1,15 @@ import { forwardRef, memo, useCallback, useEffect, useImperativeHandle, useMemo, useRef, useState } from "react" import { Button as VSCodeButton } from "vscrui" -import { CheckCheck, SquareMousePointer, Webhook, GitBranch, Bell, Cog, FlaskConical } from "lucide-react" +import { + CheckCheck, + SquareMousePointer, + Webhook, + GitBranch, + Bell, + Cog, + FlaskConical, + AlertTriangle, +} from "lucide-react" import { ApiConfiguration } from "../../../../src/shared/api" import { ExperimentId } from "../../../../src/shared/experiments" @@ -78,6 +87,7 @@ const SettingsView = forwardRef(({ onDone }, mcpEnabled, rateLimitSeconds, requestDelaySeconds, + remoteBrowserHost, screenshotQuality, soundEnabled, soundVolume, @@ -85,6 +95,7 @@ const SettingsView = forwardRef(({ onDone }, terminalOutputLimit, writeDelayMs, showRooIgnoredFiles, + remoteBrowserEnabled, } = cachedState // Make sure apiConfiguration is initialized and managed by SettingsView. @@ -173,6 +184,8 @@ const SettingsView = forwardRef(({ onDone }, vscode.postMessage({ type: "enableCheckpoints", bool: enableCheckpoints }) vscode.postMessage({ type: "checkpointStorage", text: checkpointStorage }) vscode.postMessage({ type: "browserViewportSize", text: browserViewportSize }) + vscode.postMessage({ type: "remoteBrowserHost", text: remoteBrowserHost }) + vscode.postMessage({ type: "remoteBrowserEnabled", bool: remoteBrowserEnabled }) vscode.postMessage({ type: "fuzzyMatchThreshold", value: fuzzyMatchThreshold ?? 1.0 }) vscode.postMessage({ type: "writeDelayMs", value: writeDelayMs }) vscode.postMessage({ type: "screenshotQuality", value: screenshotQuality ?? 75 }) @@ -367,6 +380,8 @@ const SettingsView = forwardRef(({ onDone }, browserToolEnabled={browserToolEnabled} browserViewportSize={browserViewportSize} screenshotQuality={screenshotQuality} + remoteBrowserHost={remoteBrowserHost} + remoteBrowserEnabled={remoteBrowserEnabled} setCachedStateField={setCachedStateField} /> @@ -419,15 +434,17 @@ const SettingsView = forwardRef(({ onDone }, - Unsaved changes - - - Do you want to discard changes and continue? - + + + Unsaved Changes + + Do you want to discard changes and continue? - onConfirmDialogResult(true)}>Yes - onConfirmDialogResult(false)}>No + onConfirmDialogResult(false)}>Cancel + onConfirmDialogResult(true)}> + Discard changes + diff --git a/webview-ui/src/components/ui/alert-dialog.tsx b/webview-ui/src/components/ui/alert-dialog.tsx index 82a25bf8f7..6782b399c1 100644 --- a/webview-ui/src/components/ui/alert-dialog.tsx +++ b/webview-ui/src/components/ui/alert-dialog.tsx @@ -36,7 +36,7 @@ function AlertDialogContent({ className, ...props }: React.ComponentProps) { - return ( -
- ) + return
} function AlertDialogFooter({ className, ...props }: React.ComponentProps<"div">) { return (
) @@ -69,7 +63,10 @@ function AlertDialogTitle({ className, ...props }: React.ComponentProps ) @@ -82,18 +79,36 @@ function AlertDialogDescription({ return ( ) } function AlertDialogAction({ className, ...props }: React.ComponentProps) { - return + return ( + + ) } function AlertDialogCancel({ className, ...props }: React.ComponentProps) { - return + return ( + + ) } export { diff --git a/webview-ui/src/context/ExtensionStateContext.tsx b/webview-ui/src/context/ExtensionStateContext.tsx index b2172df642..d212900c03 100644 --- a/webview-ui/src/context/ExtensionStateContext.tsx +++ b/webview-ui/src/context/ExtensionStateContext.tsx @@ -75,6 +75,8 @@ export interface ExtensionStateContextType extends ExtensionState { setCustomModes: (value: ModeConfig[]) => void setMaxOpenTabsContext: (value: number) => void setTelemetrySetting: (value: TelemetrySetting) => void + remoteBrowserEnabled?: boolean + setRemoteBrowserEnabled: (value: boolean) => void machineId?: string } @@ -286,6 +288,7 @@ export const ExtensionStateContextProvider: React.FC<{ children: React.ReactNode setBrowserToolEnabled: (value) => setState((prevState) => ({ ...prevState, browserToolEnabled: value })), setTelemetrySetting: (value) => setState((prevState) => ({ ...prevState, telemetrySetting: value })), setShowRooIgnoredFiles: (value) => setState((prevState) => ({ ...prevState, showRooIgnoredFiles: value })), + setRemoteBrowserEnabled: (value) => setState((prevState) => ({ ...prevState, remoteBrowserEnabled: value })), } return {children} diff --git a/webview-ui/src/utils/validate.ts b/webview-ui/src/utils/validate.ts index 82af23ab49..f642c0a788 100644 --- a/webview-ui/src/utils/validate.ts +++ b/webview-ui/src/utils/validate.ts @@ -80,6 +80,44 @@ export function validateApiConfiguration(apiConfiguration?: ApiConfiguration): s return undefined } +/** + * Validates an AWS Bedrock ARN format and optionally checks if the region in the ARN matches the provided region + * @param arn The ARN string to validate + * @param region Optional region to check against the ARN's region + * @returns An object with validation results: { isValid, arnRegion, errorMessage } + */ +export function validateBedrockArn(arn: string, region?: string) { + // Validate ARN format + const arnRegex = /^arn:aws:bedrock:([^:]+):(\d+):(foundation-model|provisioned-model|default-prompt-router)\/(.+)$/ + const match = arn.match(arnRegex) + + if (!match) { + return { + isValid: false, + arnRegion: undefined, + errorMessage: "Invalid ARN format. Please check the format requirements.", + } + } + + // Extract region from ARN + const arnRegion = match[1] + + // Check if region in ARN matches provided region (if specified) + if (region && arnRegion !== region) { + return { + isValid: true, + arnRegion, + errorMessage: `Warning: The region in your ARN (${arnRegion}) does not match your selected region (${region}). This may cause access issues. The provider will use the region from the ARN.`, + } + } + + // ARN is valid and region matches (or no region was provided to check against) + return { + isValid: true, + arnRegion, + errorMessage: undefined, + } +} export function validateModelId( apiConfiguration?: ApiConfiguration,