From 22f51304a26bcaca031f52270f0b2ceb68467fb2 Mon Sep 17 00:00:00 2001 From: Afshawn Lotfi Date: Sun, 9 Mar 2025 20:11:52 +0000 Subject: [PATCH 01/20] Add support for remote browser connection and settings --- src/core/webview/ClineProvider.ts | 7 +++ src/services/browser/BrowserSession.ts | 47 ++++++++++++++++++- src/shared/ExtensionMessage.ts | 1 + src/shared/WebviewMessage.ts | 1 + src/shared/globalState.ts | 1 + .../components/settings/BrowserSettings.tsx | 33 ++++++++++++- .../src/components/settings/SettingsView.tsx | 3 ++ 7 files changed, 90 insertions(+), 3 deletions(-) diff --git a/src/core/webview/ClineProvider.ts b/src/core/webview/ClineProvider.ts index 657e4a9ab6..3e8ed51668 100644 --- a/src/core/webview/ClineProvider.ts +++ b/src/core/webview/ClineProvider.ts @@ -1262,6 +1262,10 @@ 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 "fuzzyMatchThreshold": await this.updateGlobalState("fuzzyMatchThreshold", message.value) await this.postStateToWebview() @@ -2188,6 +2192,7 @@ export class ClineProvider implements vscode.WebviewViewProvider { soundVolume, browserViewportSize, screenshotQuality, + remoteBrowserHost, preferredLanguage, writeDelayMs, terminalOutputLimit, @@ -2246,6 +2251,7 @@ export class ClineProvider implements vscode.WebviewViewProvider { soundVolume: soundVolume ?? 0.5, browserViewportSize: browserViewportSize ?? "900x600", screenshotQuality: screenshotQuality ?? 75, + remoteBrowserHost, preferredLanguage: preferredLanguage ?? "English", writeDelayMs: writeDelayMs ?? 1000, terminalOutputLimit: terminalOutputLimit ?? TERMINAL_OUTPUT_LIMIT, @@ -2399,6 +2405,7 @@ export class ClineProvider implements vscode.WebviewViewProvider { soundVolume: stateValues.soundVolume, browserViewportSize: stateValues.browserViewportSize ?? "900x600", screenshotQuality: stateValues.screenshotQuality ?? 75, + remoteBrowserHost: stateValues.remoteBrowserHost, fuzzyMatchThreshold: stateValues.fuzzyMatchThreshold ?? 1.0, writeDelayMs: stateValues.writeDelayMs ?? 1000, terminalOutputLimit: stateValues.terminalOutputLimit ?? TERMINAL_OUTPUT_LIMIT, diff --git a/src/services/browser/BrowserSession.ts b/src/services/browser/BrowserSession.ts index bed0332244..6fe25ddf5b 100644 --- a/src/services/browser/BrowserSession.ts +++ b/src/services/browser/BrowserSession.ts @@ -1,11 +1,12 @@ 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" @@ -52,6 +53,41 @@ export class BrowserSession { await this.closeBrowser() // this may happen when the model launches a browser again after having used it already before } + const remoteBrowserHost = this.context.globalState.get("remoteBrowserHost") as string | undefined + + if (remoteBrowserHost) { + 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) + const browserWSEndpoint = response.data.webSocketDebuggerUrl + + if (!browserWSEndpoint) { + throw new Error("Could not find webSocketDebuggerUrl in the response") + } + + console.log(`Found WebSocket endpoint: ${browserWSEndpoint}`) + + this.browser = await connect({ + browserWSEndpoint, + defaultViewport: (() => { + const size = + (this.context.globalState.get("browserViewportSize") as string | undefined) || "900x600" + const [width, height] = size.split("x").map(Number) + return { width, height } + })(), + }) + this.page = await this.browser?.newPage() + return + } catch (error) { + console.error(`Failed to connect to remote browser: ${error}`) + // Fall back to local browser if remote connection fails + } + } + const stats = await this.ensureChromiumExists() this.browser = await stats.puppeteer.launch({ args: [ @@ -72,7 +108,14 @@ export class BrowserSession { async closeBrowser(): Promise { if (this.browser || this.page) { console.log("closing browser...") - await this.browser?.close().catch(() => {}) + + const remoteBrowserHost = this.context.globalState.get("remoteBrowserHost") as string | undefined + if (remoteBrowserHost && 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/shared/ExtensionMessage.ts b/src/shared/ExtensionMessage.ts index b7e3d850cf..5b638e6534 100644 --- a/src/shared/ExtensionMessage.ts +++ b/src/shared/ExtensionMessage.ts @@ -123,6 +123,7 @@ export interface ExtensionState { checkpointStorage: CheckpointStorage browserViewportSize?: string screenshotQuality?: number + remoteBrowserHost?: string fuzzyMatchThreshold?: number preferredLanguage: string writeDelayMs: number diff --git a/src/shared/WebviewMessage.ts b/src/shared/WebviewMessage.ts index 216c7588d7..df5c7d29f9 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" diff --git a/src/shared/globalState.ts b/src/shared/globalState.ts index 35e53bbe9c..05f54bfb8b 100644 --- a/src/shared/globalState.ts +++ b/src/shared/globalState.ts @@ -66,6 +66,7 @@ export const GLOBAL_STATE_KEYS = [ "checkpointStorage", "browserViewportSize", "screenshotQuality", + "remoteBrowserHost", "fuzzyMatchThreshold", "preferredLanguage", // Language setting for Cline's communication "writeDelayMs", diff --git a/webview-ui/src/components/settings/BrowserSettings.tsx b/webview-ui/src/components/settings/BrowserSettings.tsx index 1ff0a36b9c..ab4a88113a 100644 --- a/webview-ui/src/components/settings/BrowserSettings.tsx +++ b/webview-ui/src/components/settings/BrowserSettings.tsx @@ -12,13 +12,17 @@ type BrowserSettingsProps = HTMLAttributes & { browserToolEnabled?: boolean browserViewportSize?: string screenshotQuality?: number - setCachedStateField: SetCachedStateField<"browserToolEnabled" | "browserViewportSize" | "screenshotQuality"> + remoteBrowserHost?: string + setCachedStateField: SetCachedStateField< + "browserToolEnabled" | "browserViewportSize" | "screenshotQuality" | "remoteBrowserHost" + > } export const BrowserSettings = ({ browserToolEnabled, browserViewportSize, screenshotQuality, + remoteBrowserHost, setCachedStateField, ...props }: BrowserSettingsProps) => { @@ -96,6 +100,33 @@ export const BrowserSettings = ({ screenshots but increase token usage.

+
+ + + setCachedStateField("remoteBrowserHost", e.target.value || undefined) + } + /> +

+ Connect to a remote Chrome browser by providing the DevTools Protocol host address. + Roo will automatically fetch the WebSocket endpoint from this address. If provided, + Roo will use this browser instead of launching a local one. Leave empty to use the + built-in browser. +

+
)} diff --git a/webview-ui/src/components/settings/SettingsView.tsx b/webview-ui/src/components/settings/SettingsView.tsx index 7cafbd6663..0b8cce9b87 100644 --- a/webview-ui/src/components/settings/SettingsView.tsx +++ b/webview-ui/src/components/settings/SettingsView.tsx @@ -77,6 +77,7 @@ const SettingsView = forwardRef(({ onDone }, mcpEnabled, rateLimitSeconds, requestDelaySeconds, + remoteBrowserHost, screenshotQuality, soundEnabled, soundVolume, @@ -172,6 +173,7 @@ 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: "fuzzyMatchThreshold", value: fuzzyMatchThreshold ?? 1.0 }) vscode.postMessage({ type: "writeDelayMs", value: writeDelayMs }) vscode.postMessage({ type: "screenshotQuality", value: screenshotQuality ?? 75 }) @@ -378,6 +380,7 @@ const SettingsView = forwardRef(({ onDone }, browserToolEnabled={browserToolEnabled} browserViewportSize={browserViewportSize} screenshotQuality={screenshotQuality} + remoteBrowserHost={remoteBrowserHost} setCachedStateField={setCachedStateField} /> From 66e3b9610c3ee8658e2d0c2d6cfa3cb8fd5bd40b Mon Sep 17 00:00:00 2001 From: Afshawn Lotfi Date: Mon, 10 Mar 2025 03:41:38 +0000 Subject: [PATCH 02/20] Add remote browser connection support and related state management --- src/core/webview/ClineProvider.ts | 100 +++++++ .../webview/__tests__/ClineProvider.test.ts | 223 +++++++++++++-- src/services/browser/BrowserSession.ts | 123 +++++++-- src/services/browser/browserDiscovery.ts | 253 ++++++++++++++++++ src/shared/ExtensionMessage.ts | 5 + src/shared/WebviewMessage.ts | 4 + src/shared/globalState.ts | 1 + .../components/settings/BrowserSettings.tsx | 166 ++++++++++-- .../src/components/settings/SettingsView.tsx | 3 + .../src/context/ExtensionStateContext.tsx | 3 + 10 files changed, 810 insertions(+), 71 deletions(-) create mode 100644 src/services/browser/browserDiscovery.ts diff --git a/src/core/webview/ClineProvider.ts b/src/core/webview/ClineProvider.ts index 3e8ed51668..9267682eb5 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" @@ -1266,6 +1268,101 @@ export class ClineProvider implements vscode.WebviewViewProvider { 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() @@ -2193,6 +2290,7 @@ export class ClineProvider implements vscode.WebviewViewProvider { browserViewportSize, screenshotQuality, remoteBrowserHost, + remoteBrowserEnabled, preferredLanguage, writeDelayMs, terminalOutputLimit, @@ -2252,6 +2350,7 @@ export class ClineProvider implements vscode.WebviewViewProvider { browserViewportSize: browserViewportSize ?? "900x600", screenshotQuality: screenshotQuality ?? 75, remoteBrowserHost, + remoteBrowserEnabled: remoteBrowserEnabled ?? false, preferredLanguage: preferredLanguage ?? "English", writeDelayMs: writeDelayMs ?? 1000, terminalOutputLimit: terminalOutputLimit ?? TERMINAL_OUTPUT_LIMIT, @@ -2406,6 +2505,7 @@ export class ClineProvider implements vscode.WebviewViewProvider { 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 f9fc5d3ece..9a6c2c28f3 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") { @@ -1591,6 +1595,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/services/browser/BrowserSession.ts b/src/services/browser/BrowserSession.ts index 6fe25ddf5b..5c5f59ffeb 100644 --- a/src/services/browser/BrowserSession.ts +++ b/src/services/browser/BrowserSession.ts @@ -9,6 +9,7 @@ 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 } @@ -20,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) { @@ -53,9 +63,59 @@ export class BrowserSession { await this.closeBrowser() // this may happen when the model launches a browser again after having used it already before } - const remoteBrowserHost = this.context.globalState.get("remoteBrowserHost") as string | undefined + // 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 } + } - if (remoteBrowserHost) { + // 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 @@ -63,7 +123,7 @@ export class BrowserSession { console.log(`Fetching WebSocket endpoint from ${versionUrl}`) const response = await axios.get(versionUrl) - const browserWSEndpoint = response.data.webSocketDebuggerUrl + browserWSEndpoint = response.data.webSocketDebuggerUrl if (!browserWSEndpoint) { throw new Error("Could not find webSocketDebuggerUrl in the response") @@ -71,34 +131,63 @@ export class BrowserSession { console.log(`Found WebSocket endpoint: ${browserWSEndpoint}`) + // Cache the successful endpoint + this.cachedWebSocketEndpoint = browserWSEndpoint + this.lastConnectionAttempt = Date.now() + this.browser = await connect({ browserWSEndpoint, - 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(), }) this.page = await this.browser?.newPage() return } catch (error) { console.error(`Failed to connect to remote browser: ${error}`) - // Fall back to local browser if remote connection fails + // 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) @@ -109,8 +198,8 @@ export class BrowserSession { if (this.browser || this.page) { console.log("closing browser...") - const remoteBrowserHost = this.context.globalState.get("remoteBrowserHost") as string | undefined - if (remoteBrowserHost && this.browser) { + 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(() => {}) diff --git a/src/services/browser/browserDiscovery.ts b/src/services/browser/browserDiscovery.ts new file mode 100644 index 0000000000..a29bab5b78 --- /dev/null +++ b/src/services/browser/browserDiscovery.ts @@ -0,0 +1,253 @@ +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 + } +} + +/** + * Get Docker gateway IP + */ +export async function getDockerGatewayIP(): Promise { + try { + // Try to get the default gateway from the route table + if (process.platform === "linux") { + try { + const { stdout } = await vscode.window.withProgress( + { + location: vscode.ProgressLocation.Notification, + title: "Checking Docker gateway IP", + cancellable: false, + }, + async () => { + const result = await new Promise<{ stdout: string; stderr: string }>((resolve) => { + const cp = require("child_process") + cp.exec( + "ip route | grep default | awk '{print $3}'", + (err: any, stdout: string, stderr: string) => { + resolve({ stdout, stderr }) + }, + ) + }) + return result + }, + ) + 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 + 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 5b638e6534..5e95c0f68e 100644 --- a/src/shared/ExtensionMessage.ts +++ b/src/shared/ExtensionMessage.ts @@ -51,6 +51,8 @@ export interface ExtensionMessage { | "humanRelayResponse" | "humanRelayCancel" | "browserToolEnabled" + | "browserConnectionResult" + | "remoteBrowserEnabled" text?: string action?: | "chatButtonClicked" @@ -83,6 +85,8 @@ export interface ExtensionMessage { mode?: Mode customMode?: ModeConfig slug?: string + success?: boolean + values?: Record } export interface ApiConfigMeta { @@ -124,6 +128,7 @@ export interface ExtensionState { browserViewportSize?: string screenshotQuality?: number remoteBrowserHost?: string + remoteBrowserEnabled?: boolean fuzzyMatchThreshold?: number preferredLanguage: string writeDelayMs: number diff --git a/src/shared/WebviewMessage.ts b/src/shared/WebviewMessage.ts index df5c7d29f9..79bf5c7405 100644 --- a/src/shared/WebviewMessage.ts +++ b/src/shared/WebviewMessage.ts @@ -102,6 +102,10 @@ export interface WebviewMessage { | "browserToolEnabled" | "telemetrySetting" | "showRooIgnoredFiles" + | "testBrowserConnection" + | "discoverBrowser" + | "browserConnectionResult" + | "remoteBrowserEnabled" text?: string disabled?: boolean askResponse?: ClineAskResponse diff --git a/src/shared/globalState.ts b/src/shared/globalState.ts index 05f54bfb8b..5f4b216f6d 100644 --- a/src/shared/globalState.ts +++ b/src/shared/globalState.ts @@ -100,6 +100,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/settings/BrowserSettings.tsx b/webview-ui/src/components/settings/BrowserSettings.tsx index ab4a88113a..5c20a632c6 100644 --- a/webview-ui/src/components/settings/BrowserSettings.tsx +++ b/webview-ui/src/components/settings/BrowserSettings.tsx @@ -1,5 +1,5 @@ -import { HTMLAttributes } from "react" -import { VSCodeCheckbox } from "@vscode/webview-ui-toolkit/react" +import React, { HTMLAttributes, useState, useEffect } from "react" +import { VSCodeButton, VSCodeCheckbox } from "@vscode/webview-ui-toolkit/react" import { Dropdown, type DropdownOption } from "vscrui" import { SquareMousePointer } from "lucide-react" @@ -7,14 +7,20 @@ import { SetCachedStateField } from "./types" import { sliderLabelStyle } from "./styles" import { SectionHeader } from "./SectionHeader" import { Section } from "./Section" +import { vscode } from "../../utils/vscode" type BrowserSettingsProps = HTMLAttributes & { browserToolEnabled?: boolean browserViewportSize?: string screenshotQuality?: number remoteBrowserHost?: string + remoteBrowserEnabled?: boolean setCachedStateField: SetCachedStateField< - "browserToolEnabled" | "browserViewportSize" | "screenshotQuality" | "remoteBrowserHost" + | "browserToolEnabled" + | "browserViewportSize" + | "screenshotQuality" + | "remoteBrowserHost" + | "remoteBrowserEnabled" > } @@ -23,9 +29,74 @@ export const BrowserSettings = ({ 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 (
@@ -101,31 +172,70 @@ export const BrowserSettings = ({

- - - setCachedStateField("remoteBrowserHost", e.target.value || undefined) - } - /> -

- Connect to a remote Chrome browser by providing the DevTools Protocol host address. - Roo will automatically fetch the WebSocket endpoint from this address. If provided, - Roo will use this browser instead of launching a local one. Leave empty to use the - built-in browser. -

+
+ { + // 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, + ) + } + /> + + {testingConnection || discovering ? "Testing..." : "Test Connection"} + +
+ {testResult && ( +
+ {testResult.message} +
+ )} +

+ Enter the DevTools Protocol host address or leave empty to auto-discover + Chrome instances on your network. 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 0b8cce9b87..866f8bc280 100644 --- a/webview-ui/src/components/settings/SettingsView.tsx +++ b/webview-ui/src/components/settings/SettingsView.tsx @@ -85,6 +85,7 @@ const SettingsView = forwardRef(({ onDone }, terminalOutputLimit, writeDelayMs, showRooIgnoredFiles, + remoteBrowserEnabled, } = cachedState // Make sure apiConfiguration is initialized and managed by SettingsView. @@ -174,6 +175,7 @@ const SettingsView = forwardRef(({ onDone }, 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 }) @@ -381,6 +383,7 @@ const SettingsView = forwardRef(({ onDone }, browserViewportSize={browserViewportSize} screenshotQuality={screenshotQuality} remoteBrowserHost={remoteBrowserHost} + remoteBrowserEnabled={remoteBrowserEnabled} setCachedStateField={setCachedStateField} /> diff --git a/webview-ui/src/context/ExtensionStateContext.tsx b/webview-ui/src/context/ExtensionStateContext.tsx index 8d16f2e0f0..6426474753 100644 --- a/webview-ui/src/context/ExtensionStateContext.tsx +++ b/webview-ui/src/context/ExtensionStateContext.tsx @@ -73,6 +73,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 } @@ -281,6 +283,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} From 71795d5487302ca0b9894c58db946b273cee0c03 Mon Sep 17 00:00:00 2001 From: Afshawn Lotfi Date: Mon, 10 Mar 2025 06:51:35 +0000 Subject: [PATCH 03/20] Enhance BrowserSettings component with VSCodeTextField for remote browser URL input --- .../components/settings/BrowserSettings.tsx | 32 +++++++------------ 1 file changed, 12 insertions(+), 20 deletions(-) diff --git a/webview-ui/src/components/settings/BrowserSettings.tsx b/webview-ui/src/components/settings/BrowserSettings.tsx index 5c20a632c6..5c385a3d8d 100644 --- a/webview-ui/src/components/settings/BrowserSettings.tsx +++ b/webview-ui/src/components/settings/BrowserSettings.tsx @@ -1,5 +1,5 @@ import React, { HTMLAttributes, useState, useEffect } from "react" -import { VSCodeButton, VSCodeCheckbox } from "@vscode/webview-ui-toolkit/react" +import { VSCodeButton, VSCodeCheckbox, VSCodeTextField } from "@vscode/webview-ui-toolkit/react" import { Dropdown, type DropdownOption } from "vscrui" import { SquareMousePointer } from "lucide-react" @@ -192,28 +192,19 @@ export const BrowserSettings = ({ {remoteBrowserEnabled && ( <> -
- + + onChange={(e: any) => setCachedStateField( "remoteBrowserHost", e.target.value || undefined, ) } + placeholder="Custom URL (e.g., http://localhost:9222)" + style={{ flexGrow: 1 }} /> {testingConnection || discovering ? "Testing..." : "Test Connection"} @@ -221,7 +212,7 @@ export const BrowserSettings = ({
{testResult && (
)} -

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

+ 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.

)} From f4b3c44371af324185e5f7ad733c5649e2fc1122 Mon Sep 17 00:00:00 2001 From: shohei-ihaya Date: Tue, 11 Mar 2025 02:09:18 +0900 Subject: [PATCH 04/20] add gemini-2.0-pro-exp-02-05 model to vertex --- src/shared/api.ts | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/src/shared/api.ts b/src/shared/api.ts index 98d595cd03..f527e6a392 100644 --- a/src/shared/api.ts +++ b/src/shared/api.ts @@ -496,6 +496,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, From 5a3c20764a03371d29b3e35bbe08a927bede666c Mon Sep 17 00:00:00 2001 From: cte Date: Mon, 10 Mar 2025 11:12:11 -0700 Subject: [PATCH 05/20] Follow the established pattern for command registration --- README.md | 53 ++++++++++++++++--------------- src/activate/humanRelay.ts | 26 +++++++++++++++ src/activate/registerCommands.ts | 32 +++++++++---------- src/extension.ts | 54 +++----------------------------- 4 files changed, 75 insertions(+), 90 deletions(-) create mode 100644 src/activate/humanRelay.ts 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/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/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/extension.ts b/src/extension.ts index d88bf2a251..c85d61053f 100644 --- a/src/extension.ts +++ b/src/extension.ts @@ -11,15 +11,17 @@ 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 { createClineAPI } from "./exports" +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 } from "./activate" + /** * Built using https://github.com/microsoft/vscode-webview-ui-toolkit * @@ -31,18 +33,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 +62,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 From f8a592aac9cce5b73b1d528aa9f298854d4518f2 Mon Sep 17 00:00:00 2001 From: cte Date: Mon, 10 Mar 2025 11:16:38 -0700 Subject: [PATCH 06/20] Rename ClineAPI to RooCodeAPI and improve types --- .eslintrc.json | 2 +- e2e/VSCODE_INTEGRATION_TESTS.md | 4 +- e2e/src/suite/index.ts | 6 +- e2e/tsconfig.json | 2 +- .../index.ts => activate/createRooCodeAPI.ts} | 24 +++---- src/activate/index.ts | 1 + src/exports/README.md | 72 +++++++++---------- src/exports/{cline.d.ts => roo-code.d.ts} | 68 ++++++++++++++---- src/extension.ts | 9 +-- src/shared/ExtensionMessage.ts | 64 ++--------------- 10 files changed, 119 insertions(+), 133 deletions(-) rename src/{exports/index.ts => activate/createRooCodeAPI.ts} (79%) rename src/exports/{cline.d.ts => roo-code.d.ts} (73%) 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/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/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/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/exports/README.md b/src/exports/README.md index 03b8983b7e..876ff3bec8 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, Cline! 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 73% rename from src/exports/cline.d.ts rename to src/exports/roo-code.d.ts index e529947b6b..2004b31e8b 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,60 @@ 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 } export interface ClineProvider { @@ -82,11 +135,6 @@ export interface ClineProvider { */ cancelTask(): Promise - /** - * Clears the current task - */ - clearTask(): Promise - /** * Gets the current state */ @@ -112,12 +160,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..f4c372e67c 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 * @@ -143,7 +144,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/shared/ExtensionMessage.ts b/src/shared/ExtensionMessage.ts index 73f8127c53..0786509662 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" @@ -145,58 +146,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 +170,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 +207,6 @@ export interface ClineApiReqInfo { streamingFailedMessage?: string } -// Human relay related message types export interface ShowHumanRelayDialogMessage { type: "showHumanRelayDialog" requestId: string From 0a03c88e92efb0260a44145aa4982de9d7bc15cd Mon Sep 17 00:00:00 2001 From: Chris Estreich Date: Mon, 10 Mar 2025 11:24:21 -0700 Subject: [PATCH 07/20] Update src/exports/README.md Co-authored-by: ellipsis-dev[bot] <65095814+ellipsis-dev[bot]@users.noreply.github.com> --- src/exports/README.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/exports/README.md b/src/exports/README.md index 876ff3bec8..0554580836 100644 --- a/src/exports/README.md +++ b/src/exports/README.md @@ -27,7 +27,7 @@ const instructions = await api.getCustomInstructions() console.log("Current custom instructions:", instructions) // Start a new task with an initial message. -await api.startNewTask("Hello, Cline! Let's make a new project...") +await api.startNewTask("Hello, Roo Code API! Let's make a new project...") // Start a new task with an initial message and images. await api.startNewTask("Use this design language", ["data:image/webp;base64,..."]) From 7681a988426d2a554a102331c198b412f799ca70 Mon Sep 17 00:00:00 2001 From: cte Date: Mon, 10 Mar 2025 11:36:11 -0700 Subject: [PATCH 08/20] Simplify `npm install` by automatically installing npm-run-all --- .github/workflows/changeset-release.yml | 2 +- .github/workflows/code-qa.yml | 10 ++--- .github/workflows/marketplace-publish.yml | 10 ++--- README.md | 53 ++++++++++++----------- package.json | 11 +++-- 5 files changed, 45 insertions(+), 41 deletions(-) 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/package.json b/package.json index b32a81071d..3d11264cc7 100644 --- a/package.json +++ b/package.json @@ -230,23 +230,28 @@ "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", + "install-benchmark": "cd benchmark && npm install", "lint": "npm-run-all -p lint:*", "lint:extension": "eslint src --ext ts", "lint:webview-ui": "cd webview-ui && npm run lint", "lint:e2e": "cd e2e && npm run lint", + "lint:benchmark": "cd benchmark && npm run lint", "check-types": "npm-run-all -p check-types:*", "check-types:extension": "tsc --noEmit", "check-types:webview-ui": "cd webview-ui && npm run check-types", "check-types:e2e": "cd e2e && npm run check-types", + "check-types:benchmark": "cd benchmark && npm run check-types", "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", From 2187369ba4251e7f8e293540c394c1d898a4c63c Mon Sep 17 00:00:00 2001 From: cte Date: Mon, 10 Mar 2025 11:39:23 -0700 Subject: [PATCH 09/20] Remove benchmark --- package.json | 3 --- 1 file changed, 3 deletions(-) diff --git a/package.json b/package.json index 3d11264cc7..d6ccae86ed 100644 --- a/package.json +++ b/package.json @@ -235,17 +235,14 @@ "install-extension": "npm install", "install-webview-ui": "cd webview-ui && npm install", "install-e2e": "cd e2e && npm install", - "install-benchmark": "cd benchmark && npm install", "lint": "npm-run-all -p lint:*", "lint:extension": "eslint src --ext ts", "lint:webview-ui": "cd webview-ui && npm run lint", "lint:e2e": "cd e2e && npm run lint", - "lint:benchmark": "cd benchmark && npm run lint", "check-types": "npm-run-all -p check-types:*", "check-types:extension": "tsc --noEmit", "check-types:webview-ui": "cd webview-ui && npm run check-types", "check-types:e2e": "cd e2e && npm run check-types", - "check-types:benchmark": "cd benchmark && npm run check-types", "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", From 102a996875d25ea968b8287674a10fba6f1b094c Mon Sep 17 00:00:00 2001 From: hannesrudolph Date: Mon, 10 Mar 2025 15:18:19 -0600 Subject: [PATCH 10/20] fix: update MCP servers directory path for platform compatibility --- src/core/webview/ClineProvider.ts | 17 +++++++++++++++-- 1 file changed, 15 insertions(+), 2 deletions(-) diff --git a/src/core/webview/ClineProvider.ts b/src/core/webview/ClineProvider.ts index 72cf56d35b..7f7ffa0683 100644 --- a/src/core/webview/ClineProvider.ts +++ b/src/core/webview/ClineProvider.ts @@ -1971,11 +1971,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%\Cline\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("~", ".roo-code", "mcp") } return mcpServersDir } From 171037a93878dece89a86665d7540a67dfe283d6 Mon Sep 17 00:00:00 2001 From: Smartsheet-JB-Brown Date: Mon, 10 Mar 2025 11:10:49 -0700 Subject: [PATCH 11/20] Add enhanced error handling and logging for AWS Bedrock custom ARNs --- .../__tests__/bedrock-custom-arn.test.ts | 75 +++ src/api/providers/__tests__/bedrock.test.ts | 29 ++ src/api/providers/bedrock.ts | 491 ++++++++++++++++-- src/shared/api.ts | 2 + src/shared/globalState.ts | 1 + test-custom-arn.js | 196 +++++++ .../src/components/settings/ApiOptions.tsx | 89 +++- webview-ui/src/utils/validate.ts | 38 ++ 8 files changed, 887 insertions(+), 34 deletions(-) create mode 100644 src/api/providers/__tests__/bedrock-custom-arn.test.ts create mode 100644 test-custom-arn.js 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..f778621e9c 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(5000) + 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/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/shared/api.ts b/src/shared/api.ts index 98d595cd03..986a412354 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", diff --git a/src/shared/globalState.ts b/src/shared/globalState.ts index 540b7e72be..e3522e1c0b 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", diff --git a/test-custom-arn.js b/test-custom-arn.js new file mode 100644 index 0000000000..dd22ed69b0 --- /dev/null +++ b/test-custom-arn.js @@ -0,0 +1,196 @@ +// Test script to verify AWS Bedrock functionality with custom ARNs +// This file should be deleted after testing + +// IMPORTANT: Before running this script, make sure you have: +// 1. Configured an AWS profile in your AWS credentials file (~/.aws/credentials) +// 2. For prompt routing, created a prompt router in AWS Bedrock (https://docs.aws.amazon.com/bedrock/latest/userguide/prompt-routing.html) +// 3. For prompt routing, have the prompt router ARN in the format: arn:aws:bedrock:region:account-id:default-prompt-router/router-name + +const { BedrockRuntimeClient, ConverseCommand } = require("@aws-sdk/client-bedrock-runtime") +const { fromIni } = require("@aws-sdk/credential-providers") + +// The model ID or ARN provided by the user (not stored in source code) +const modelIdOrArn = process.env.CUSTOM_ARN +// The AWS profile to use for authentication +const awsProfile = process.env.AWS_PROFILE || "default" + +if (!modelIdOrArn) { + console.error("Please provide a model ID or ARN via the CUSTOM_ARN environment variable") + process.exit(1) +} + +console.log(`Using AWS profile: ${awsProfile}`) + +// Check if the input is an ARN or a model ID +const arnRegex = + /^arn:aws:bedrock:([^:]+):(\d+):(foundation-model|provisioned-model|default-prompt-router|prompt-router)\/(.+)$/ +const match = modelIdOrArn.match(arnRegex) +const isArn = !!match + +// If it's not an ARN, assume it's a model ID +if (!isArn) { + console.log(`Using model ID: ${modelIdOrArn}`) +} + +// Use us-west-2 region by default +const defaultRegion = "us-west-2" +// Always use the default region, ignoring the region in the ARN +const region = defaultRegion + +if (isArn) { + console.log(`Using region: ${region} with AWS profile "${awsProfile}" (overriding ARN region: ${match[1]})`) +} else { + console.log(`Using region: ${region} with AWS profile "${awsProfile}"`) +} + +// Create a client with the specified AWS profile +let client +try { + client = new BedrockRuntimeClient({ + region: region, + credentials: fromIni({ + profile: awsProfile, + }), + }) + console.log("Successfully created Bedrock client") +} catch (error) { + console.error("Error creating Bedrock client:", error) + process.exit(1) +} + +// Use the input as the model ID +if (isArn) { + console.log(`Using custom ARN as model ID: ${modelIdOrArn}`) +} else { + console.log(`Using standard model ID: ${modelIdOrArn}`) +} + +const payload = { + modelId: modelIdOrArn, + messages: [ + { + role: "user", + content: [ + { + text: isArn + ? "Hello, can you verify that this prompt router ARN is working correctly? This is a test of AWS Bedrock Intelligent Prompt Routing." + : `Hello, can you verify that this model ID is working correctly with the specified AWS profile?`, + }, + ], + }, + ], + inferenceConfig: { + // For Claude models, use appropriate token limits based on model type + // Claude 3.7 Sonnet: 8192, Claude 3.5 Sonnet: 8192, Claude 3 Opus: 4096, Claude 3 Haiku: 4096 + maxTokens: 4096, // Conservative default that works for all Claude models + temperature: 0.3, + topP: 0.1, + }, +} + +console.log( + isArn + ? "Sending request to Bedrock API using prompt router ARN..." + : "Sending request to Bedrock API using standard model ID...", +) + +async function testCustomArn() { + try { + const command = new ConverseCommand(payload) + const response = await client.send(command) + + // Handle the response format where output is an object + if (response.output && typeof response.output === "object") { + if (response.output.message && response.output.message.content) { + console.log("Success! Received response:") + console.log(JSON.stringify(response)) + console.log(response.output.message.content) + return + } + } + // Handle the response format where output is a Uint8Array + else if (response.output && response.output instanceof Uint8Array) { + try { + const outputStr = new TextDecoder().decode(response.output) + const output = JSON.parse(outputStr) + if (output.content) { + console.log("Success! Received response:") + console.log(output.content) + return + } + } catch (parseError) { + console.error("Failed to parse Bedrock response:", parseError) + } + } + console.error("No valid response content received") + } catch (error) { + console.error(isArn ? "Error occurred with custom ARN:" : "Error occurred with model ID:", error) + + if (error.message) { + const errorMessage = error.message.toLowerCase() + + // Access denied errors + if ( + errorMessage.includes("access") && + (errorMessage.includes("model") || errorMessage.includes("denied")) + ) { + if (isArn) { + console.error("\nThis appears to be a permissions issue with the prompt router ARN. Please verify:") + console.error("1. The ARN is correct and points to a valid prompt router") + console.error( + `2. Your AWS credentials (${awsProfile} profile) have permission to access this prompt router`, + ) + console.error("3. The region in the ARN matches the region where the prompt router is deployed") + console.error("4. The prompt router is properly configured and active") + } else { + console.error("\nThis appears to be a permissions issue with the model. Please verify:") + console.error( + `1. Your AWS credentials (${awsProfile} profile) have permission to access this model`, + ) + console.error("2. The model exists in the specified region") + console.error("3. The model is available for use with your account") + } + } + // Model not found errors + else if (errorMessage.includes("not found") || errorMessage.includes("does not exist")) { + if (isArn) { + console.error("\nThis appears to be an invalid prompt router ARN. Please check:") + console.error( + "1. The ARN format is correct (arn:aws:bedrock:region:account-id:default-prompt-router/router-name)", + ) + console.error("2. The prompt router exists in the specified region") + console.error("3. The account ID in the ARN is correct") + } else { + console.error("\nThis appears to be an invalid model ID. Please check:") + console.error("1. The model ID is correct") + console.error("2. The model exists in the specified region") + } + } + // Validation errors + else if (errorMessage.includes("validation")) { + if (isArn) { + console.error("\nThis appears to be a validation error with the prompt router ARN. Please check:") + console.error("1. The ARN format is correct") + console.error("2. The prompt router is properly configured") + console.error("3. The request payload is valid for prompt routing") + } else { + console.error("\nThis appears to be a validation error with the model ID. Please check:") + console.error("1. The model ID format is correct") + console.error("2. The request payload is valid for this model") + } + } + // Throttling errors + else if ( + errorMessage.includes("throttl") || + errorMessage.includes("rate") || + errorMessage.includes("limit") + ) { + console.error("\nThis appears to be a throttling or rate limit issue. Please try:") + console.error("1. Reducing the frequency of requests") + console.error("2. Contact AWS support to request a quota increase if needed") + } + } + } +} + +testCustomArn() 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 + })()} + ======= + + )} Date: Mon, 10 Mar 2025 21:40:08 +0000 Subject: [PATCH 12/20] Refactor Docker gateway IP retrieval to use a dedicated shell command execution function --- src/services/browser/browserDiscovery.ts | 37 ++++++++++-------------- 1 file changed, 15 insertions(+), 22 deletions(-) diff --git a/src/services/browser/browserDiscovery.ts b/src/services/browser/browserDiscovery.ts index a29bab5b78..187f90e299 100644 --- a/src/services/browser/browserDiscovery.ts +++ b/src/services/browser/browserDiscovery.ts @@ -55,32 +55,25 @@ export async function tryConnect(ipAddress: string): Promise<{ endpoint: string; } /** - * Get Docker gateway IP + * 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 { - // Try to get the default gateway from the route table if (process.platform === "linux") { try { - const { stdout } = await vscode.window.withProgress( - { - location: vscode.ProgressLocation.Notification, - title: "Checking Docker gateway IP", - cancellable: false, - }, - async () => { - const result = await new Promise<{ stdout: string; stderr: string }>((resolve) => { - const cp = require("child_process") - cp.exec( - "ip route | grep default | awk '{print $3}'", - (err: any, stdout: string, stderr: string) => { - resolve({ stdout, stderr }) - }, - ) - }) - return result - }, - ) + 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) @@ -159,7 +152,7 @@ export async function discoverChromeInstances(): Promise { ipAddresses.push("localhost") ipAddresses.push("127.0.0.1") - // Try to get Docker gateway IP + // Try to get Docker gateway IP (headless mode) const gatewayIP = await getDockerGatewayIP() if (gatewayIP) { console.log("Found Docker gateway IP:", gatewayIP) From abcb7c18b4c5e181e4042e0ee627fc1f6621d3b6 Mon Sep 17 00:00:00 2001 From: Hannes Rudolph Date: Mon, 10 Mar 2025 16:23:43 -0600 Subject: [PATCH 13/20] Update src/core/webview/ClineProvider.ts Co-authored-by: ellipsis-dev[bot] <65095814+ellipsis-dev[bot]@users.noreply.github.com> --- src/core/webview/ClineProvider.ts | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/core/webview/ClineProvider.ts b/src/core/webview/ClineProvider.ts index 7f7ffa0683..6433b948c7 100644 --- a/src/core/webview/ClineProvider.ts +++ b/src/core/webview/ClineProvider.ts @@ -1974,7 +1974,7 @@ export class ClineProvider implements vscode.WebviewViewProvider { // Get platform-specific application data directory let mcpServersDir: string if (process.platform === "win32") { - // Windows: %APPDATA%\Cline\MCP + // Windows: %APPDATA%\Roo-Code\MCP mcpServersDir = path.join(os.homedir(), "AppData", "Roaming", "Roo-Code", "MCP") } else if (process.platform === "darwin") { // macOS: ~/Documents/Cline/MCP From a9773b70c8b70d772f0f75ed2e726e9beb9c94fd Mon Sep 17 00:00:00 2001 From: cannuri <91494156+cannuri@users.noreply.github.com> Date: Tue, 11 Mar 2025 01:03:18 +0100 Subject: [PATCH 14/20] fix browser_action system prompt --- src/core/webview/ClineProvider.ts | 7 +- .../webview/__tests__/ClineProvider.test.ts | 274 +++++++++++++++--- 2 files changed, 238 insertions(+), 43 deletions(-) diff --git a/src/core/webview/ClineProvider.ts b/src/core/webview/ClineProvider.ts index 72cf56d35b..8c321514d9 100644 --- a/src/core/webview/ClineProvider.ts +++ b/src/core/webview/ClineProvider.ts @@ -1826,6 +1826,7 @@ export class ClineProvider implements vscode.WebviewViewProvider { fuzzyMatchThreshold, experiments, enableMcpServerCreation, + browserToolEnabled, } = await this.getState() // Create diffStrategy based on current model and settings @@ -1841,10 +1842,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", diff --git a/src/core/webview/__tests__/ClineProvider.test.ts b/src/core/webview/__tests__/ClineProvider.test.ts index 2e9fcdf336..5941c1bf83 100644 --- a/src/core/webview/__tests__/ClineProvider.test.ts +++ b/src/core/webview/__tests__/ClineProvider.test.ts @@ -1160,6 +1160,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 +1187,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 +1198,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 +1218,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 +1245,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 +1256,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 +1307,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", () => { From 7462906b7a2b4417fb5fdedfc5117fc8b6ebccef Mon Sep 17 00:00:00 2001 From: Hannes Rudolph Date: Mon, 10 Mar 2025 18:49:25 -0600 Subject: [PATCH 15/20] Update ClineProvider.ts Co-authored-by: ellipsis-dev[bot] <65095814+ellipsis-dev[bot]@users.noreply.github.com> --- src/core/webview/ClineProvider.ts | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/core/webview/ClineProvider.ts b/src/core/webview/ClineProvider.ts index 6433b948c7..80dc5e0c50 100644 --- a/src/core/webview/ClineProvider.ts +++ b/src/core/webview/ClineProvider.ts @@ -1988,7 +1988,7 @@ export class ClineProvider implements vscode.WebviewViewProvider { await fs.mkdir(mcpServersDir, { recursive: true }) } catch (error) { // Fallback to a relative path if directory creation fails - return path.join("~", ".roo-code", "mcp") + return path.join(os.homedir(), ".roo-code", "mcp") } return mcpServersDir } From 77186e8ee0523b05e065421bf70e39c94caf8aef Mon Sep 17 00:00:00 2001 From: Matt Rubens Date: Mon, 10 Mar 2025 22:15:52 -0400 Subject: [PATCH 16/20] Cleanup --- src/api/providers/__tests__/bedrock.test.ts | 2 +- test-custom-arn.js | 196 -------------------- 2 files changed, 1 insertion(+), 197 deletions(-) delete mode 100644 test-custom-arn.js diff --git a/src/api/providers/__tests__/bedrock.test.ts b/src/api/providers/__tests__/bedrock.test.ts index f778621e9c..45d5270237 100644 --- a/src/api/providers/__tests__/bedrock.test.ts +++ b/src/api/providers/__tests__/bedrock.test.ts @@ -326,7 +326,7 @@ describe("AwsBedrockHandler", () => { }) const modelInfo = customArnHandler.getModel() expect(modelInfo.id).toBe("arn:aws:bedrock:us-east-1::foundation-model/custom-model") - expect(modelInfo.info.maxTokens).toBe(5000) + expect(modelInfo.info.maxTokens).toBe(4096) expect(modelInfo.info.contextWindow).toBe(128_000) expect(modelInfo.info.supportsPromptCache).toBe(false) }) diff --git a/test-custom-arn.js b/test-custom-arn.js deleted file mode 100644 index dd22ed69b0..0000000000 --- a/test-custom-arn.js +++ /dev/null @@ -1,196 +0,0 @@ -// Test script to verify AWS Bedrock functionality with custom ARNs -// This file should be deleted after testing - -// IMPORTANT: Before running this script, make sure you have: -// 1. Configured an AWS profile in your AWS credentials file (~/.aws/credentials) -// 2. For prompt routing, created a prompt router in AWS Bedrock (https://docs.aws.amazon.com/bedrock/latest/userguide/prompt-routing.html) -// 3. For prompt routing, have the prompt router ARN in the format: arn:aws:bedrock:region:account-id:default-prompt-router/router-name - -const { BedrockRuntimeClient, ConverseCommand } = require("@aws-sdk/client-bedrock-runtime") -const { fromIni } = require("@aws-sdk/credential-providers") - -// The model ID or ARN provided by the user (not stored in source code) -const modelIdOrArn = process.env.CUSTOM_ARN -// The AWS profile to use for authentication -const awsProfile = process.env.AWS_PROFILE || "default" - -if (!modelIdOrArn) { - console.error("Please provide a model ID or ARN via the CUSTOM_ARN environment variable") - process.exit(1) -} - -console.log(`Using AWS profile: ${awsProfile}`) - -// Check if the input is an ARN or a model ID -const arnRegex = - /^arn:aws:bedrock:([^:]+):(\d+):(foundation-model|provisioned-model|default-prompt-router|prompt-router)\/(.+)$/ -const match = modelIdOrArn.match(arnRegex) -const isArn = !!match - -// If it's not an ARN, assume it's a model ID -if (!isArn) { - console.log(`Using model ID: ${modelIdOrArn}`) -} - -// Use us-west-2 region by default -const defaultRegion = "us-west-2" -// Always use the default region, ignoring the region in the ARN -const region = defaultRegion - -if (isArn) { - console.log(`Using region: ${region} with AWS profile "${awsProfile}" (overriding ARN region: ${match[1]})`) -} else { - console.log(`Using region: ${region} with AWS profile "${awsProfile}"`) -} - -// Create a client with the specified AWS profile -let client -try { - client = new BedrockRuntimeClient({ - region: region, - credentials: fromIni({ - profile: awsProfile, - }), - }) - console.log("Successfully created Bedrock client") -} catch (error) { - console.error("Error creating Bedrock client:", error) - process.exit(1) -} - -// Use the input as the model ID -if (isArn) { - console.log(`Using custom ARN as model ID: ${modelIdOrArn}`) -} else { - console.log(`Using standard model ID: ${modelIdOrArn}`) -} - -const payload = { - modelId: modelIdOrArn, - messages: [ - { - role: "user", - content: [ - { - text: isArn - ? "Hello, can you verify that this prompt router ARN is working correctly? This is a test of AWS Bedrock Intelligent Prompt Routing." - : `Hello, can you verify that this model ID is working correctly with the specified AWS profile?`, - }, - ], - }, - ], - inferenceConfig: { - // For Claude models, use appropriate token limits based on model type - // Claude 3.7 Sonnet: 8192, Claude 3.5 Sonnet: 8192, Claude 3 Opus: 4096, Claude 3 Haiku: 4096 - maxTokens: 4096, // Conservative default that works for all Claude models - temperature: 0.3, - topP: 0.1, - }, -} - -console.log( - isArn - ? "Sending request to Bedrock API using prompt router ARN..." - : "Sending request to Bedrock API using standard model ID...", -) - -async function testCustomArn() { - try { - const command = new ConverseCommand(payload) - const response = await client.send(command) - - // Handle the response format where output is an object - if (response.output && typeof response.output === "object") { - if (response.output.message && response.output.message.content) { - console.log("Success! Received response:") - console.log(JSON.stringify(response)) - console.log(response.output.message.content) - return - } - } - // Handle the response format where output is a Uint8Array - else if (response.output && response.output instanceof Uint8Array) { - try { - const outputStr = new TextDecoder().decode(response.output) - const output = JSON.parse(outputStr) - if (output.content) { - console.log("Success! Received response:") - console.log(output.content) - return - } - } catch (parseError) { - console.error("Failed to parse Bedrock response:", parseError) - } - } - console.error("No valid response content received") - } catch (error) { - console.error(isArn ? "Error occurred with custom ARN:" : "Error occurred with model ID:", error) - - if (error.message) { - const errorMessage = error.message.toLowerCase() - - // Access denied errors - if ( - errorMessage.includes("access") && - (errorMessage.includes("model") || errorMessage.includes("denied")) - ) { - if (isArn) { - console.error("\nThis appears to be a permissions issue with the prompt router ARN. Please verify:") - console.error("1. The ARN is correct and points to a valid prompt router") - console.error( - `2. Your AWS credentials (${awsProfile} profile) have permission to access this prompt router`, - ) - console.error("3. The region in the ARN matches the region where the prompt router is deployed") - console.error("4. The prompt router is properly configured and active") - } else { - console.error("\nThis appears to be a permissions issue with the model. Please verify:") - console.error( - `1. Your AWS credentials (${awsProfile} profile) have permission to access this model`, - ) - console.error("2. The model exists in the specified region") - console.error("3. The model is available for use with your account") - } - } - // Model not found errors - else if (errorMessage.includes("not found") || errorMessage.includes("does not exist")) { - if (isArn) { - console.error("\nThis appears to be an invalid prompt router ARN. Please check:") - console.error( - "1. The ARN format is correct (arn:aws:bedrock:region:account-id:default-prompt-router/router-name)", - ) - console.error("2. The prompt router exists in the specified region") - console.error("3. The account ID in the ARN is correct") - } else { - console.error("\nThis appears to be an invalid model ID. Please check:") - console.error("1. The model ID is correct") - console.error("2. The model exists in the specified region") - } - } - // Validation errors - else if (errorMessage.includes("validation")) { - if (isArn) { - console.error("\nThis appears to be a validation error with the prompt router ARN. Please check:") - console.error("1. The ARN format is correct") - console.error("2. The prompt router is properly configured") - console.error("3. The request payload is valid for prompt routing") - } else { - console.error("\nThis appears to be a validation error with the model ID. Please check:") - console.error("1. The model ID format is correct") - console.error("2. The request payload is valid for this model") - } - } - // Throttling errors - else if ( - errorMessage.includes("throttl") || - errorMessage.includes("rate") || - errorMessage.includes("limit") - ) { - console.error("\nThis appears to be a throttling or rate limit issue. Please try:") - console.error("1. Reducing the frequency of requests") - console.error("2. Contact AWS support to request a quota increase if needed") - } - } - } -} - -testCustomArn() From f306461276fbad2da4a5d693f970b2dc0b478798 Mon Sep 17 00:00:00 2001 From: Matt Rubens Date: Mon, 10 Mar 2025 22:53:40 -0400 Subject: [PATCH 17/20] Fix usage tracking for SiliconFlow etc --- .changeset/tidy-queens-pay.md | 5 + .../__tests__/openai-usage-tracking.test.ts | 235 ++++++++++++++++++ src/api/providers/openai.ts | 8 +- 3 files changed, 247 insertions(+), 1 deletion(-) create mode 100644 .changeset/tidy-queens-pay.md create mode 100644 src/api/providers/__tests__/openai-usage-tracking.test.ts 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/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/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 = { From 8c21f0ece3f31ad70f609c2693a97d8ef30d44f6 Mon Sep 17 00:00:00 2001 From: cannuri <91494156+cannuri@users.noreply.github.com> Date: Tue, 11 Mar 2025 03:59:46 +0100 Subject: [PATCH 18/20] refactor alert dialog styles, use vscode theme --- .../src/components/settings/SettingsView.tsx | 27 ++++++++---- webview-ui/src/components/ui/alert-dialog.tsx | 41 +++++++++++++------ 2 files changed, 47 insertions(+), 21 deletions(-) diff --git a/webview-ui/src/components/settings/SettingsView.tsx b/webview-ui/src/components/settings/SettingsView.tsx index df08a03971..7caf280dba 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" @@ -419,15 +428,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 { From 621fc0e867c125330f2fea6e87c18e6575db192c Mon Sep 17 00:00:00 2001 From: axb Date: Tue, 11 Mar 2025 10:38:29 +0800 Subject: [PATCH 19/20] Revert "Merge pull request #1518 from RooVetGit/revert_tool_progress_for_now" This reverts commit dba9116d26e7005c64e2562b717fce6215d744f7, reversing changes made to 85dd1a197728f32a85e358cff9ee2f628f239e8e. --- src/core/Cline.ts | 26 ++++++++++++++++++- .../diff/strategies/multi-search-replace.ts | 25 ++++++++++++++++++ src/core/diff/types.ts | 5 ++++ src/exports/roo-code.d.ts | 1 + src/shared/ExtensionMessage.ts | 5 ++++ webview-ui/src/components/chat/ChatRow.tsx | 1 + .../src/components/common/CodeAccordian.tsx | 11 ++++++++ 7 files changed, 73 insertions(+), 1 deletion(-) diff --git a/src/core/Cline.ts b/src/core/Cline.ts index ba171ce3fa..d26e75ce85 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() @@ -1703,8 +1711,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,6 +1815,14 @@ export class Cline { diff: diffContent, } satisfies ClineSayTool) + let toolProgressStatus + if (this.diffStrategy && this.diffStrategy.getProgressStatus) { + toolProgressStatus = this.diffStrategy.getProgressStatus(block, diffResult) + } + await this.ask("tool", completeMessage, block.partial, toolProgressStatus).catch( + () => {}, + ) + const didApprove = await askApproval("tool", completeMessage) if (!didApprove) { await this.diffViewProvider.revertChanges() // This likely handles closing the diff view 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/exports/roo-code.d.ts b/src/exports/roo-code.d.ts index 2004b31e8b..e310fef6a7 100644 --- a/src/exports/roo-code.d.ts +++ b/src/exports/roo-code.d.ts @@ -92,6 +92,7 @@ export interface ClineMessage { reasoning?: string conversationHistoryIndex?: number checkpoint?: Record + progressStatus?: ToolProgressStatus } export interface ClineProvider { diff --git a/src/shared/ExtensionMessage.ts b/src/shared/ExtensionMessage.ts index 33c485adff..9298b0cb1b 100644 --- a/src/shared/ExtensionMessage.ts +++ b/src/shared/ExtensionMessage.ts @@ -231,3 +231,8 @@ export interface HumanRelayCancelMessage { } export type ClineApiReqCancelReason = "streaming_failed" | "user_cancelled" + +export type ToolProgressStatus = { + icon?: string + text?: string +} 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} + + + )}
)} From 8917ab75913c429f6082cc7089c2e26b6cb01536 Mon Sep 17 00:00:00 2001 From: axb Date: Mon, 10 Mar 2025 12:12:25 +0800 Subject: [PATCH 20/20] fix duplicate ask --- src/core/Cline.ts | 13 +++++++------ 1 file changed, 7 insertions(+), 6 deletions(-) diff --git a/src/core/Cline.ts b/src/core/Cline.ts index d26e75ce85..5bc280dc42 100644 --- a/src/core/Cline.ts +++ b/src/core/Cline.ts @@ -1402,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) { @@ -1819,11 +1823,8 @@ export class Cline { if (this.diffStrategy && this.diffStrategy.getProgressStatus) { toolProgressStatus = this.diffStrategy.getProgressStatus(block, diffResult) } - await this.ask("tool", completeMessage, block.partial, toolProgressStatus).catch( - () => {}, - ) - const didApprove = await askApproval("tool", completeMessage) + const didApprove = await askApproval("tool", completeMessage, toolProgressStatus) if (!didApprove) { await this.diffViewProvider.revertChanges() // This likely handles closing the diff view break