Merge branch 'RooVetGit:main' into human-relay

This commit is contained in:
Felix NyxJae 2025-03-03 10:19:02 +08:00 • committed by GitHub
commit 626827ab3f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
76 changed files with 6846 additions and 1782 deletions

View file

@ -1,5 +0,0 @@
---
"roo-cline": patch
---
Delete task confirmation enhancements

View file

@ -0,0 +1,5 @@
---
"roo-cline": patch
---
ExtensionStateContext does not correctly merge state

View file

@ -0,0 +1,5 @@
---
"roo-cline": patch
---
Default middle-out compression to on for OpenRouter

View file

@ -1,5 +0,0 @@
---
"roo-cline": patch
---
Prettier thinking blocks

View file

@ -1,37 +1,35 @@
<!-- **Note:** Consider creating PRs as a DRAFT. For early feedback and self-review. -->
## Context
## Description
<!-- Brief description of WHAT you’re doing and WHY. -->
## Type of change
## Implementation
<!-- Please ignore options that are not relevant -->
<!--
- [ ] Bug fix (non-breaking change which fixes an issue)
- [ ] New feature
- [ ] Breaking change (fix or feature that would cause existing functionality to not work as expected)
- [ ] This change requires a documentation update
Some description of HOW you achieved it. Perhaps give a high level description of the program flow. Did you need to refactor something? What tradeoffs did you take? Are there things in here which you’d particularly like people to pay close attention to?
## How Has This Been Tested?
-->
<!-- Please describe the tests that you ran to verify your changes -->
## Screenshots
## Checklist:
| before | after |
| ------ | ----- |
| | |
<!-- Go over all the following points, and put an `x` in all the boxes that apply -->
## How to Test
- [ ] My code follows the patterns of this project
- [ ] I have performed a self-review of my own code
- [ ] I have commented my code, particularly in hard-to-understand areas
- [ ] I have made corresponding changes to the documentation
<!--
## Additional context
A straightforward scenario of how to test your changes will help reviewers that are not familiar with the part of the code that you are changing but want to see it in action. This section can include a description or step-by-step instructions of how to get to the state of v2 that your change affects.
<!-- Add any other context or screenshots about the pull request here -->
A "How To Test" section can look something like this:
## Related Issues
- Sign in with a user with tracks
- Activate `show_awesome_cat_gifs` feature (add `?feature.show_awesome_cat_gifs=1` to your URL)
- You should see a GIF with cats dancing
<!-- List any related issues here. Use the GitHub issue linking syntax: #issue-number -->
-->
## Reviewers
## Get in Touch
<!-- @mention specific team members or individuals who should review this PR -->
<!-- We'd love to have a way to chat with you about your changes if necessary. If you're in the [Roo Code Discord](https://discord.gg/roocode), please share your handle here. -->

View file

@ -37,7 +37,7 @@ jobs:
cache: 'npm'
- name: Install Dependencies
run: npm run install:all
run: npm run install:ci
# Check if there are any new changesets to process
- name: Check for changesets

View file

@ -20,7 +20,7 @@ jobs:
node-version: '18'
cache: 'npm'
- name: Install dependencies
run: npm run install:all
run: npm run install:ci
- 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:all
run: npm run install:ci
- 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:all
run: npm run install:ci
- 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:all
run: npm run install:ci
- name: Run unit tests
working-directory: webview-ui
run: npx jest --silent
@ -108,9 +108,11 @@ jobs:
with:
node-version: '18'
cache: 'npm'
- name: Create env.integration file
run: echo "OPENROUTER_API_KEY=${{ secrets.OPENROUTER_API_KEY }}" > .env.integration
- name: Install dependencies
run: npm run install:all
run: npm run install:ci
- name: Create env.integration file
working-directory: e2e
run: echo "OPENROUTER_API_KEY=${{ secrets.OPENROUTER_API_KEY }}" > .env.integration
- name: Run integration tests
run: xvfb-run -a npm run test:integration
working-directory: e2e
run: xvfb-run -a npm run ci

View file

@ -29,10 +29,7 @@ jobs:
- name: Install Dependencies
run: |
npm install -g vsce ovsx
npm install
cd webview-ui
npm install
cd ..
npm run install:ci
- name: Package and Publish Extension
env:
VSCE_PAT: ${{ secrets.VSCE_PAT }}

View file

@ -4,6 +4,8 @@
.vscode/**
.vscode-test/**
out/**
out-integration/**
e2e/**
node_modules/**
src/**
.gitignore
@ -25,7 +27,6 @@ demo.gif
.roomodes
cline_docs/**
coverage/**
out-integration/**
# Ignore all webview-ui files except the build directory (https://github.com/microsoft/vscode-webview-ui-toolkit-samples/blob/main/frameworks/hello-world-react-cra/.vscodeignore)
webview-ui/src/**

View file

@ -1,5 +1,27 @@
# Roo Code Changelog
## [3.7.11]
- Don't honor custom max tokens for non thinking models
- Include custom modes in mode switching keyboard shortcut
- Support read-only modes that can run commands
## [3.7.10]
- Add Gemini models on Vertex AI (thanks @ashktn!)
- Keyboard shortcuts to switch modes (thanks @aheizi!)
- Add support for Mermaid diagrams (thanks Cline!)
## [3.7.9]
- Delete task confirmation enhancements
- Smarter context window management
- Prettier thinking blocks
- Fix maxTokens defaults for Claude 3.7 Sonnet models
- Terminal output parsing improvements (thanks @KJ7LNW!)
- UI fix to dropdown hover colors (thanks @SamirSaji!)
- Add support for Claude Sonnet 3.7 thinking via Vertex AI (thanks @lupuletic!)
## [3.7.8]
- Add Vertex AI prompt caching support for Claude models (thanks @aitoroses and @lupuletic!)

View file

@ -119,6 +119,12 @@ Make Roo Code work your way with:
```bash
npm run install:all
```
if that fails, try:
```bash
npm run install:ci
```
3. **Build** the extension:
```bash
npm run build

View file

@ -6,7 +6,7 @@ import { defineConfig } from '@vscode/test-cli';
export default defineConfig({
label: 'integrationTest',
files: 'out-integration/test/**/*.test.js',
files: 'out/suite/**/*.test.js',
workspaceFolder: '.',
mocha: {
ui: 'tdd',

View file

@ -11,8 +11,8 @@ The integration tests use the `@vscode/test-electron` package to run tests in a
### Directory Structure
```
src/test/
├── runTest.ts # Main test runner
e2e/src/
├── runTest.ts # Main test runner
├── suite/
│ ├── index.ts # Test suite configuration
│ ├── modes.test.ts # Mode switching tests

2387
e2e/package-lock.json generated Normal file

File diff suppressed because it is too large Load diff

21
e2e/package.json Normal file
View file

@ -0,0 +1,21 @@
{
"name": "e2e",
"version": "0.1.0",
"private": true,
"scripts": {
"build": "cd .. && npm run build",
"compile": "tsc -p tsconfig.json",
"lint": "eslint src --ext ts",
"check-types": "tsc --noEmit",
"test": "npm run compile && npx dotenvx run -f .env.integration -- node ./out/runTest.js",
"ci": "npm run build && npm run test"
},
"dependencies": {},
"devDependencies": {
"@types/mocha": "^10.0.10",
"@vscode/test-cli": "^0.0.9",
"@vscode/test-electron": "^2.4.0",
"mocha": "^11.1.0",
"typescript": "^5.4.5"
}
}

View file

@ -1,8 +1,7 @@
import * as path from "path"
import Mocha from "mocha"
import { glob } from "glob"
import { ClineAPI } from "../../exports/cline"
import { ClineProvider } from "../../core/webview/ClineProvider"
import { ClineAPI, ClineProvider } from "../../../src/exports/cline"
import * as vscode from "vscode"
declare global {

View file

@ -9,9 +9,8 @@
"strict": true,
"skipLibCheck": true,
"useUnknownInCatchVariables": false,
"rootDir": "src",
"outDir": "out-integration"
"outDir": "out"
},
"include": ["**/*.ts"],
"exclude": [".vscode-test", "benchmark", "dist", "**/node_modules/**", "out", "out-integration", "webview-ui"]
"include": ["src", "../src/exports/cline.d.ts"],
"exclude": [".vscode-test", "**/node_modules/**", "out"]
}

View file

@ -16,7 +16,9 @@
"src/activate/**",
"src/exports/**",
"src/extension.ts",
".vscode-test.mjs"
"e2e/.vscode-test.mjs",
"e2e/src/runTest.ts",
"e2e/src/suite/index.ts"
],
"workspaces": {
"webview-ui": {

948
package-lock.json generated

File diff suppressed because it is too large Load diff

View file

@ -3,7 +3,7 @@
"displayName": "Roo Code (prev. Roo Cline)",
"description": "A whole dev team of AI agents in your editor.",
"publisher": "RooVeterinaryInc",
"version": "3.7.8",
"version": "3.7.11",
"icon": "assets/icons/rocket.png",
"galleryBanner": {
"color": "#617A91",
@ -276,21 +276,24 @@
"scripts": {
"build": "npm run build:webview && npm run vsix",
"build:webview": "cd webview-ui && npm run build",
"changeset": "changeset",
"check-types": "tsc --noEmit && cd webview-ui && npm run check-types",
"compile": "tsc -p . --outDir out && node esbuild.js",
"compile:integration": "tsc -p tsconfig.integration.json",
"install:all": "npm install && cd webview-ui && npm install",
"knip": "knip --include files",
"lint": "eslint src --ext ts && npm run lint --prefix webview-ui",
"lint-local": "eslint -c .eslintrc.local.json src --ext ts && npm run lint --prefix webview-ui",
"lint-fix": "eslint src --ext ts --fix && npm run lint-fix --prefix webview-ui",
"lint-fix-local": "eslint -c .eslintrc.local.json src --ext ts --fix && npm run lint-fix --prefix webview-ui",
"install:all": "npm-run-all -p install-*",
"install:ci": "npm install npm-run-all && npm run install:all",
"install-extension": "npm install",
"install-webview-ui": "cd webview-ui && npm install",
"install-e2e": "cd e2e && 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",
"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",
"package": "npm run build:webview && npm run check-types && npm run lint && node esbuild.js --production",
"pretest": "npm run compile && npm run compile:integration",
"pretest": "npm run compile",
"dev": "cd webview-ui && npm run dev",
"test": "jest && cd webview-ui && npm run test",
"test:integration": "npm run build && npm run compile:integration && npx dotenvx run -f .env.integration -- node ./out-integration/test/runTest.js",
"prepare": "husky",
"publish:marketplace": "vsce publish && ovsx publish",
"publish": "npm run build && changeset publish && npm install --package-lock-only",
@ -300,7 +303,9 @@
"watch": "npm-run-all -p watch:*",
"watch:esbuild": "node esbuild.js --watch",
"watch:tsc": "tsc --noEmit --watch --project tsconfig.json",
"watch-tests": "tsc -p . -w --outDir out"
"watch-tests": "tsc -p . -w --outDir out",
"changeset": "changeset",
"knip": "knip --include files"
},
"dependencies": {
"@anthropic-ai/bedrock-sdk": "^0.10.2",
@ -308,6 +313,7 @@
"@anthropic-ai/vertex-sdk": "^0.7.0",
"@aws-sdk/client-bedrock-runtime": "^3.706.0",
"@google/generative-ai": "^0.18.0",
"@google-cloud/vertexai": "^1.9.3",
"@mistralai/mistralai": "^1.3.6",
"@modelcontextprotocol/sdk": "^1.0.1",
"@types/clone-deep": "^4.0.4",
@ -329,6 +335,7 @@
"get-folder-size": "^5.0.0",
"globby": "^14.0.2",
"isbinaryfile": "^5.0.2",
"js-tiktoken": "^1.0.19",
"mammoth": "^1.8.0",
"monaco-vscode-textmate-theme-converter": "^0.1.7",
"openai": "^4.78.1",
@ -358,13 +365,10 @@
"@types/diff-match-patch": "^1.0.36",
"@types/glob": "^8.1.0",
"@types/jest": "^29.5.14",
"@types/mocha": "^10.0.10",
"@types/node": "20.x",
"@types/string-similarity": "^4.0.2",
"@typescript-eslint/eslint-plugin": "^7.14.1",
"@typescript-eslint/parser": "^7.11.0",
"@vscode/test-cli": "^0.0.9",
"@vscode/test-electron": "^2.4.0",
"esbuild": "^0.24.0",
"eslint": "^8.57.0",
"glob": "^11.0.1",
@ -374,7 +378,6 @@
"knip": "^5.44.4",
"lint-staged": "^15.2.11",
"mkdirp": "^3.0.1",
"mocha": "^11.1.0",
"npm-run-all": "^4.1.5",
"prettier": "^3.4.2",
"rimraf": "^6.0.1",

View file

@ -0,0 +1,257 @@
// npx jest src/api/__tests__/index.test.ts
import { BetaThinkingConfigParam } from "@anthropic-ai/sdk/resources/beta/messages/index.mjs"
import { getModelParams } from "../index"
import { ANTHROPIC_DEFAULT_MAX_TOKENS } from "../providers/constants"
describe("getModelParams", () => {
it("should return default values when no custom values are provided", () => {
const options = {}
const model = {
id: "test-model",
contextWindow: 16000,
supportsPromptCache: true,
}
const result = getModelParams({
options,
model,
defaultMaxTokens: 1000,
defaultTemperature: 0.5,
})
expect(result).toEqual({
maxTokens: 1000,
thinking: undefined,
temperature: 0.5,
})
})
it("should use custom temperature from options when provided", () => {
const options = { modelTemperature: 0.7 }
const model = {
id: "test-model",
contextWindow: 16000,
supportsPromptCache: true,
}
const result = getModelParams({
options,
model,
defaultMaxTokens: 1000,
defaultTemperature: 0.5,
})
expect(result).toEqual({
maxTokens: 1000,
thinking: undefined,
temperature: 0.7,
})
})
it("should use model maxTokens when available", () => {
const options = {}
const model = {
id: "test-model",
maxTokens: 2000,
contextWindow: 16000,
supportsPromptCache: true,
}
const result = getModelParams({
options,
model,
defaultMaxTokens: 1000,
})
expect(result).toEqual({
maxTokens: 2000,
thinking: undefined,
temperature: 0,
})
})
it("should handle thinking models correctly", () => {
const options = {}
const model = {
id: "test-model",
thinking: true,
maxTokens: 2000,
contextWindow: 16000,
supportsPromptCache: true,
}
const result = getModelParams({
options,
model,
})
const expectedThinking: BetaThinkingConfigParam = {
type: "enabled",
budget_tokens: 1600, // 80% of 2000
}
expect(result).toEqual({
maxTokens: 2000,
thinking: expectedThinking,
temperature: 1.0, // Thinking models require temperature 1.0.
})
})
it("should honor customMaxTokens for thinking models", () => {
const options = { modelMaxTokens: 3000 }
const model = {
id: "test-model",
thinking: true,
contextWindow: 16000,
supportsPromptCache: true,
}
const result = getModelParams({
options,
model,
defaultMaxTokens: 2000,
})
const expectedThinking: BetaThinkingConfigParam = {
type: "enabled",
budget_tokens: 2400, // 80% of 3000
}
expect(result).toEqual({
maxTokens: 3000,
thinking: expectedThinking,
temperature: 1.0,
})
})
it("should honor customMaxThinkingTokens for thinking models", () => {
const options = { modelMaxThinkingTokens: 1500 }
const model = {
id: "test-model",
thinking: true,
maxTokens: 4000,
contextWindow: 16000,
supportsPromptCache: true,
}
const result = getModelParams({
options,
model,
})
const expectedThinking: BetaThinkingConfigParam = {
type: "enabled",
budget_tokens: 1500, // Using the custom value
}
expect(result).toEqual({
maxTokens: 4000,
thinking: expectedThinking,
temperature: 1.0,
})
})
it("should not honor customMaxThinkingTokens for non-thinking models", () => {
const options = { modelMaxThinkingTokens: 1500 }
const model = {
id: "test-model",
maxTokens: 4000,
contextWindow: 16000,
supportsPromptCache: true,
// Note: model.thinking is not set (so it's falsey).
}
const result = getModelParams({
options,
model,
})
expect(result).toEqual({
maxTokens: 4000,
thinking: undefined, // Should remain undefined despite customMaxThinkingTokens being set.
temperature: 0, // Using default temperature.
})
})
it("should clamp thinking budget to at least 1024 tokens", () => {
const options = { modelMaxThinkingTokens: 500 }
const model = {
id: "test-model",
thinking: true,
maxTokens: 2000,
contextWindow: 16000,
supportsPromptCache: true,
}
const result = getModelParams({
options,
model,
})
const expectedThinking: BetaThinkingConfigParam = {
type: "enabled",
budget_tokens: 1024, // Minimum is 1024
}
expect(result).toEqual({
maxTokens: 2000,
thinking: expectedThinking,
temperature: 1.0,
})
})
it("should clamp thinking budget to at most 80% of max tokens", () => {
const options = { modelMaxThinkingTokens: 5000 }
const model = {
id: "test-model",
thinking: true,
maxTokens: 4000,
contextWindow: 16000,
supportsPromptCache: true,
}
const result = getModelParams({
options,
model,
})
const expectedThinking: BetaThinkingConfigParam = {
type: "enabled",
budget_tokens: 3200, // 80% of 4000
}
expect(result).toEqual({
maxTokens: 4000,
thinking: expectedThinking,
temperature: 1.0,
})
})
it("should use ANTHROPIC_DEFAULT_MAX_TOKENS when no maxTokens is provided for thinking models", () => {
const options = {}
const model = {
id: "test-model",
thinking: true,
contextWindow: 16000,
supportsPromptCache: true,
}
const result = getModelParams({
options,
model,
})
const expectedThinking: BetaThinkingConfigParam = {
type: "enabled",
budget_tokens: Math.floor(ANTHROPIC_DEFAULT_MAX_TOKENS * 0.8),
}
expect(result).toEqual({
maxTokens: undefined,
thinking: expectedThinking,
temperature: 1.0,
})
})
})

View file

@ -1,6 +1,9 @@
import { Anthropic } from "@anthropic-ai/sdk"
import { BetaThinkingConfigParam } from "@anthropic-ai/sdk/resources/beta/messages/index.mjs"
import { ApiConfiguration, ModelInfo, ApiHandlerOptions } from "../shared/api"
import { ANTHROPIC_DEFAULT_MAX_TOKENS } from "./providers/constants"
import { GlamaHandler } from "./providers/glama"
import { ApiConfiguration, ModelInfo } from "../shared/api"
import { AnthropicHandler } from "./providers/anthropic"
import { AwsBedrockHandler } from "./providers/bedrock"
import { OpenRouterHandler } from "./providers/openrouter"
@ -25,6 +28,16 @@ export interface SingleCompletionHandler {
export interface ApiHandler {
createMessage(systemPrompt: string, messages: Anthropic.Messages.MessageParam[]): ApiStream
getModel(): { id: string; info: ModelInfo }
/**
* Counts tokens for content blocks
* All providers extend BaseProvider which provides a default tiktoken implementation,
* but they can override this to use their native token counting endpoints
*
* @param content The content to count tokens for
* @returns A promise resolving to the token count
*/
countTokens(content: Array<Anthropic.Messages.ContentBlockParam>): Promise<number>
}
export function buildApiHandler(configuration: ApiConfiguration): ApiHandler {
@ -66,3 +79,41 @@ export function buildApiHandler(configuration: ApiConfiguration): ApiHandler {
return new AnthropicHandler(options)
}
}
export function getModelParams({
options,
model,
defaultMaxTokens,
defaultTemperature = 0,
}: {
options: ApiHandlerOptions
model: ModelInfo
defaultMaxTokens?: number
defaultTemperature?: number
}) {
const {
modelMaxTokens: customMaxTokens,
modelMaxThinkingTokens: customMaxThinkingTokens,
modelTemperature: customTemperature,
} = options
let maxTokens = model.maxTokens ?? defaultMaxTokens
let thinking: BetaThinkingConfigParam | undefined = undefined
let temperature = customTemperature ?? defaultTemperature
if (model.thinking) {
// Only honor `customMaxTokens` for thinking models.
maxTokens = customMaxTokens ?? maxTokens
// Clamp the thinking budget to be at most 80% of max tokens and at
// least 1024 tokens.
const maxBudgetTokens = Math.floor((maxTokens || ANTHROPIC_DEFAULT_MAX_TOKENS) * 0.8)
const budgetTokens = Math.max(Math.min(customMaxThinkingTokens ?? maxBudgetTokens, maxBudgetTokens), 1024)
thinking = { type: "enabled", budget_tokens: budgetTokens }
// Anthropic "Thinking" models require a temperature of 1.0.
temperature = 1.0
}
return { maxTokens, thinking, temperature }
}

View file

@ -194,5 +194,33 @@ describe("AnthropicHandler", () => {
expect(model.info.supportsImages).toBe(true)
expect(model.info.supportsPromptCache).toBe(true)
})
it("honors custom maxTokens for thinking models", () => {
const handler = new AnthropicHandler({
apiKey: "test-api-key",
apiModelId: "claude-3-7-sonnet-20250219:thinking",
modelMaxTokens: 32_768,
modelMaxThinkingTokens: 16_384,
})
const result = handler.getModel()
expect(result.maxTokens).toBe(32_768)
expect(result.thinking).toEqual({ type: "enabled", budget_tokens: 16_384 })
expect(result.temperature).toBe(1.0)
})
it("does not honor custom maxTokens for non-thinking models", () => {
const handler = new AnthropicHandler({
apiKey: "test-api-key",
apiModelId: "claude-3-7-sonnet-20250219",
modelMaxTokens: 32_768,
modelMaxThinkingTokens: 16_384,
})
const result = handler.getModel()
expect(result.maxTokens).toBe(16_384)
expect(result.thinking).toBeUndefined()
expect(result.temperature).toBe(0)
})
})
})

View file

@ -1,29 +1,30 @@
// npx jest src/api/providers/__tests__/openrouter.test.ts
import axios from "axios"
import { Anthropic } from "@anthropic-ai/sdk"
import OpenAI from "openai"
import { OpenRouterHandler } from "../openrouter"
import { ApiHandlerOptions, ModelInfo } from "../../../shared/api"
import OpenAI from "openai"
import axios from "axios"
import { Anthropic } from "@anthropic-ai/sdk"
// Mock dependencies
jest.mock("openai")
jest.mock("axios")
jest.mock("delay", () => jest.fn(() => Promise.resolve()))
const mockOpenRouterModelInfo: ModelInfo = {
maxTokens: 1000,
contextWindow: 2000,
supportsPromptCache: true,
inputPrice: 0.01,
outputPrice: 0.02,
}
describe("OpenRouterHandler", () => {
const mockOptions: ApiHandlerOptions = {
openRouterApiKey: "test-key",
openRouterModelId: "test-model",
openRouterModelInfo: {
name: "Test Model",
description: "Test Description",
maxTokens: 1000,
contextWindow: 2000,
supportsPromptCache: true,
inputPrice: 0.01,
outputPrice: 0.02,
} as ModelInfo,
openRouterModelInfo: mockOpenRouterModelInfo,
}
beforeEach(() => {
@ -50,6 +51,10 @@ describe("OpenRouterHandler", () => {
expect(result).toEqual({
id: mockOptions.openRouterModelId,
info: mockOptions.openRouterModelInfo,
maxTokens: 1000,
temperature: 0,
thinking: undefined,
topP: undefined,
})
})
@ -61,6 +66,38 @@ describe("OpenRouterHandler", () => {
expect(result.info.supportsPromptCache).toBe(true)
})
test("getModel honors custom maxTokens for thinking models", () => {
const handler = new OpenRouterHandler({
openRouterApiKey: "test-key",
openRouterModelId: "test-model",
openRouterModelInfo: {
...mockOpenRouterModelInfo,
maxTokens: 64_000,
thinking: true,
},
modelMaxTokens: 32_768,
modelMaxThinkingTokens: 16_384,
})
const result = handler.getModel()
expect(result.maxTokens).toBe(32_768)
expect(result.thinking).toEqual({ type: "enabled", budget_tokens: 16_384 })
expect(result.temperature).toBe(1.0)
})
test("getModel does not honor custom maxTokens for non-thinking models", () => {
const handler = new OpenRouterHandler({
...mockOptions,
modelMaxTokens: 32_768,
modelMaxThinkingTokens: 16_384,
})
const result = handler.getModel()
expect(result.maxTokens).toBe(1000)
expect(result.thinking).toBeUndefined()
expect(result.temperature).toBe(0)
})
test("createMessage generates correct stream chunks", async () => {
const handler = new OpenRouterHandler(mockOptions)
const mockStream = {
@ -242,15 +279,7 @@ describe("OpenRouterHandler", () => {
test("completePrompt returns correct response", async () => {
const handler = new OpenRouterHandler(mockOptions)
const mockResponse = {
choices: [
{
message: {
content: "test completion",
},
},
],
}
const mockResponse = { choices: [{ message: { content: "test completion" } }] }
const mockCreate = jest.fn().mockResolvedValue(mockResponse)
;(OpenAI as jest.MockedClass<typeof OpenAI>).prototype.chat = {
@ -260,10 +289,13 @@ describe("OpenRouterHandler", () => {
const result = await handler.completePrompt("test prompt")
expect(result).toBe("test completion")
expect(mockCreate).toHaveBeenCalledWith({
model: mockOptions.openRouterModelId,
messages: [{ role: "user", content: "test prompt" }],
max_tokens: 1000,
thinking: undefined,
temperature: 0,
messages: [{ role: "user", content: "test prompt" }],
stream: false,
})
})
@ -292,8 +324,6 @@ describe("OpenRouterHandler", () => {
completions: { create: mockCreate },
} as any
await expect(handler.completePrompt("test prompt")).rejects.toThrow(
"OpenRouter completion error: Unexpected error",
)
await expect(handler.completePrompt("test prompt")).rejects.toThrow("Unexpected error")
})
})

View file

@ -6,6 +6,7 @@ import { BetaThinkingConfigParam } from "@anthropic-ai/sdk/resources/beta"
import { VertexHandler } from "../vertex"
import { ApiStreamChunk } from "../../transform/stream"
import { VertexAI } from "@google-cloud/vertexai"
// Mock Vertex SDK
jest.mock("@anthropic-ai/vertex-sdk", () => ({
@ -49,24 +50,100 @@ jest.mock("@anthropic-ai/vertex-sdk", () => ({
})),
}))
// Mock Vertex Gemini SDK
jest.mock("@google-cloud/vertexai", () => {
const mockGenerateContentStream = jest.fn().mockImplementation(() => {
return {
stream: {
async *[Symbol.asyncIterator]() {
yield {
candidates: [
{
content: {
parts: [{ text: "Test Gemini response" }],
},
},
],
}
},
},
response: {
usageMetadata: {
promptTokenCount: 5,
candidatesTokenCount: 10,
},
},
}
})
const mockGenerateContent = jest.fn().mockResolvedValue({
response: {
candidates: [
{
content: {
parts: [{ text: "Test Gemini response" }],
},
},
],
},
})
const mockGenerativeModel = jest.fn().mockImplementation(() => {
return {
generateContentStream: mockGenerateContentStream,
generateContent: mockGenerateContent,
}
})
return {
VertexAI: jest.fn().mockImplementation(() => {
return {
getGenerativeModel: mockGenerativeModel,
}
}),
GenerativeModel: mockGenerativeModel,
}
})
describe("VertexHandler", () => {
let handler: VertexHandler
beforeEach(() => {
handler = new VertexHandler({
apiModelId: "claude-3-5-sonnet-v2@20241022",
vertexProjectId: "test-project",
vertexRegion: "us-central1",
})
})
describe("constructor", () => {
it("should initialize with provided config", () => {
it("should initialize with provided config for Claude", () => {
handler = new VertexHandler({
apiModelId: "claude-3-5-sonnet-v2@20241022",
vertexProjectId: "test-project",
vertexRegion: "us-central1",
})
expect(AnthropicVertex).toHaveBeenCalledWith({
projectId: "test-project",
region: "us-central1",
})
})
it("should initialize with provided config for Gemini", () => {
handler = new VertexHandler({
apiModelId: "gemini-1.5-pro-001",
vertexProjectId: "test-project",
vertexRegion: "us-central1",
})
expect(VertexAI).toHaveBeenCalledWith({
project: "test-project",
location: "us-central1",
})
})
it("should throw error for invalid model", () => {
expect(() => {
new VertexHandler({
apiModelId: "invalid-model",
vertexProjectId: "test-project",
vertexRegion: "us-central1",
})
}).toThrow("Unknown model ID: invalid-model")
})
})
describe("createMessage", () => {
@ -83,7 +160,13 @@ describe("VertexHandler", () => {
const systemPrompt = "You are a helpful assistant"
it("should handle streaming responses correctly", async () => {
it("should handle streaming responses correctly for Claude", async () => {
handler = new VertexHandler({
apiModelId: "claude-3-5-sonnet-v2@20241022",
vertexProjectId: "test-project",
vertexRegion: "us-central1",
})
const mockStream = [
{
type: "message_start",
@ -127,7 +210,7 @@ describe("VertexHandler", () => {
}
const mockCreate = jest.fn().mockResolvedValue(asyncIterator)
;(handler["client"].messages as any).create = mockCreate
;(handler["anthropicClient"].messages as any).create = mockCreate
const stream = handler.createMessage(systemPrompt, mockMessages)
const chunks: ApiStreamChunk[] = []
@ -187,7 +270,58 @@ describe("VertexHandler", () => {
})
})
it("should handle multiple content blocks with line breaks", async () => {
it("should handle streaming responses correctly for Gemini", async () => {
const mockGemini = require("@google-cloud/vertexai")
const mockGenerateContentStream = mockGemini.VertexAI().getGenerativeModel().generateContentStream
handler = new VertexHandler({
apiModelId: "gemini-1.5-pro-001",
vertexProjectId: "test-project",
vertexRegion: "us-central1",
})
const stream = handler.createMessage(systemPrompt, mockMessages)
const chunks: ApiStreamChunk[] = []
for await (const chunk of stream) {
chunks.push(chunk)
}
expect(chunks.length).toBe(2)
expect(chunks[0]).toEqual({
type: "text",
text: "Test Gemini response",
})
expect(chunks[1]).toEqual({
type: "usage",
inputTokens: 5,
outputTokens: 10,
})
expect(mockGenerateContentStream).toHaveBeenCalledWith({
contents: [
{
role: "user",
parts: [{ text: "Hello" }],
},
{
role: "model",
parts: [{ text: "Hi there!" }],
},
],
generationConfig: {
maxOutputTokens: 16384,
temperature: 0,
},
})
})
it("should handle multiple content blocks with line breaks for Claude", async () => {
handler = new VertexHandler({
apiModelId: "claude-3-5-sonnet-v2@20241022",
vertexProjectId: "test-project",
vertexRegion: "us-central1",
})
const mockStream = [
{
type: "content_block_start",
@ -216,7 +350,7 @@ describe("VertexHandler", () => {
}
const mockCreate = jest.fn().mockResolvedValue(asyncIterator)
;(handler["client"].messages as any).create = mockCreate
;(handler["anthropicClient"].messages as any).create = mockCreate
const stream = handler.createMessage(systemPrompt, mockMessages)
const chunks: ApiStreamChunk[] = []
@ -240,10 +374,16 @@ describe("VertexHandler", () => {
})
})
it("should handle API errors", async () => {
it("should handle API errors for Claude", async () => {
handler = new VertexHandler({
apiModelId: "claude-3-5-sonnet-v2@20241022",
vertexProjectId: "test-project",
vertexRegion: "us-central1",
})
const mockError = new Error("Vertex API error")
const mockCreate = jest.fn().mockRejectedValue(mockError)
;(handler["client"].messages as any).create = mockCreate
;(handler["anthropicClient"].messages as any).create = mockCreate
const stream = handler.createMessage(systemPrompt, mockMessages)
@ -254,7 +394,13 @@ describe("VertexHandler", () => {
}).rejects.toThrow("Vertex API error")
})
it("should handle prompt caching for supported models", async () => {
it("should handle prompt caching for supported models for Claude", async () => {
handler = new VertexHandler({
apiModelId: "claude-3-5-sonnet-v2@20241022",
vertexProjectId: "test-project",
vertexRegion: "us-central1",
})
const mockStream = [
{
type: "message_start",
@ -299,7 +445,7 @@ describe("VertexHandler", () => {
}
const mockCreate = jest.fn().mockResolvedValue(asyncIterator)
;(handler["client"].messages as any).create = mockCreate
;(handler["anthropicClient"].messages as any).create = mockCreate
const stream = handler.createMessage(systemPrompt, [
{
@ -383,7 +529,13 @@ describe("VertexHandler", () => {
)
})
it("should handle cache-related usage metrics", async () => {
it("should handle cache-related usage metrics for Claude", async () => {
handler = new VertexHandler({
apiModelId: "claude-3-5-sonnet-v2@20241022",
vertexProjectId: "test-project",
vertexRegion: "us-central1",
})
const mockStream = [
{
type: "message_start",
@ -415,7 +567,7 @@ describe("VertexHandler", () => {
}
const mockCreate = jest.fn().mockResolvedValue(asyncIterator)
;(handler["client"].messages as any).create = mockCreate
;(handler["anthropicClient"].messages as any).create = mockCreate
const stream = handler.createMessage(systemPrompt, mockMessages)
const chunks: ApiStreamChunk[] = []
@ -442,7 +594,13 @@ describe("VertexHandler", () => {
const systemPrompt = "You are a helpful assistant"
it("should handle thinking content blocks and deltas", async () => {
it("should handle thinking content blocks and deltas for Claude", async () => {
handler = new VertexHandler({
apiModelId: "claude-3-5-sonnet-v2@20241022",
vertexProjectId: "test-project",
vertexRegion: "us-central1",
})
const mockStream = [
{
type: "message_start",
@ -488,7 +646,7 @@ describe("VertexHandler", () => {
}
const mockCreate = jest.fn().mockResolvedValue(asyncIterator)
;(handler["client"].messages as any).create = mockCreate
;(handler["anthropicClient"].messages as any).create = mockCreate
const stream = handler.createMessage(systemPrompt, mockMessages)
const chunks: ApiStreamChunk[] = []
@ -510,7 +668,13 @@ describe("VertexHandler", () => {
expect(textChunks[1].text).toBe("Here's my answer:")
})
it("should handle multiple thinking blocks with line breaks", async () => {
it("should handle multiple thinking blocks with line breaks for Claude", async () => {
handler = new VertexHandler({
apiModelId: "claude-3-5-sonnet-v2@20241022",
vertexProjectId: "test-project",
vertexRegion: "us-central1",
})
const mockStream = [
{
type: "content_block_start",
@ -539,7 +703,7 @@ describe("VertexHandler", () => {
}
const mockCreate = jest.fn().mockResolvedValue(asyncIterator)
;(handler["client"].messages as any).create = mockCreate
;(handler["anthropicClient"].messages as any).create = mockCreate
const stream = handler.createMessage(systemPrompt, mockMessages)
const chunks: ApiStreamChunk[] = []
@ -565,10 +729,16 @@ describe("VertexHandler", () => {
})
describe("completePrompt", () => {
it("should complete prompt successfully", async () => {
it("should complete prompt successfully for Claude", async () => {
handler = new VertexHandler({
apiModelId: "claude-3-5-sonnet-v2@20241022",
vertexProjectId: "test-project",
vertexRegion: "us-central1",
})
const result = await handler.completePrompt("Test prompt")
expect(result).toBe("Test response")
expect(handler["client"].messages.create).toHaveBeenCalledWith({
expect(handler["anthropicClient"].messages.create).toHaveBeenCalledWith({
model: "claude-3-5-sonnet-v2@20241022",
max_tokens: 8192,
temperature: 0,
@ -583,31 +753,109 @@ describe("VertexHandler", () => {
})
})
it("should handle API errors", async () => {
it("should complete prompt successfully for Gemini", async () => {
const mockGemini = require("@google-cloud/vertexai")
const mockGenerateContent = mockGemini.VertexAI().getGenerativeModel().generateContent
handler = new VertexHandler({
apiModelId: "gemini-1.5-pro-001",
vertexProjectId: "test-project",
vertexRegion: "us-central1",
})
const result = await handler.completePrompt("Test prompt")
expect(result).toBe("Test Gemini response")
expect(mockGenerateContent).toHaveBeenCalled()
expect(mockGenerateContent).toHaveBeenCalledWith({
contents: [{ role: "user", parts: [{ text: "Test prompt" }] }],
generationConfig: {
temperature: 0,
},
})
})
it("should handle API errors for Claude", async () => {
handler = new VertexHandler({
apiModelId: "claude-3-5-sonnet-v2@20241022",
vertexProjectId: "test-project",
vertexRegion: "us-central1",
})
const mockError = new Error("Vertex API error")
const mockCreate = jest.fn().mockRejectedValue(mockError)
;(handler["client"].messages as any).create = mockCreate
;(handler["anthropicClient"].messages as any).create = mockCreate
await expect(handler.completePrompt("Test prompt")).rejects.toThrow(
"Vertex completion error: Vertex API error",
)
})
it("should handle non-text content", async () => {
it("should handle API errors for Gemini", async () => {
const mockGemini = require("@google-cloud/vertexai")
const mockGenerateContent = mockGemini.VertexAI().getGenerativeModel().generateContent
mockGenerateContent.mockRejectedValue(new Error("Vertex API error"))
handler = new VertexHandler({
apiModelId: "gemini-1.5-pro-001",
vertexProjectId: "test-project",
vertexRegion: "us-central1",
})
await expect(handler.completePrompt("Test prompt")).rejects.toThrow(
"Vertex completion error: Vertex API error",
)
})
it("should handle non-text content for Claude", async () => {
handler = new VertexHandler({
apiModelId: "claude-3-5-sonnet-v2@20241022",
vertexProjectId: "test-project",
vertexRegion: "us-central1",
})
const mockCreate = jest.fn().mockResolvedValue({
content: [{ type: "image" }],
})
;(handler["client"].messages as any).create = mockCreate
;(handler["anthropicClient"].messages as any).create = mockCreate
const result = await handler.completePrompt("Test prompt")
expect(result).toBe("")
})
it("should handle empty response", async () => {
it("should handle empty response for Claude", async () => {
handler = new VertexHandler({
apiModelId: "claude-3-5-sonnet-v2@20241022",
vertexProjectId: "test-project",
vertexRegion: "us-central1",
})
const mockCreate = jest.fn().mockResolvedValue({
content: [{ type: "text", text: "" }],
})
;(handler["client"].messages as any).create = mockCreate
;(handler["anthropicClient"].messages as any).create = mockCreate
const result = await handler.completePrompt("Test prompt")
expect(result).toBe("")
})
it("should handle empty response for Gemini", async () => {
const mockGemini = require("@google-cloud/vertexai")
const mockGenerateContent = mockGemini.VertexAI().getGenerativeModel().generateContent
mockGenerateContent.mockResolvedValue({
response: {
candidates: [
{
content: {
parts: [{ text: "" }],
},
},
],
},
})
handler = new VertexHandler({
apiModelId: "gemini-1.5-pro-001",
vertexProjectId: "test-project",
vertexRegion: "us-central1",
})
const result = await handler.completePrompt("Test prompt")
expect(result).toBe("")
@ -615,7 +863,13 @@ describe("VertexHandler", () => {
})
describe("getModel", () => {
it("should return correct model info", () => {
it("should return correct model info for Claude", () => {
handler = new VertexHandler({
apiModelId: "claude-3-5-sonnet-v2@20241022",
vertexProjectId: "test-project",
vertexRegion: "us-central1",
})
const modelInfo = handler.getModel()
expect(modelInfo.id).toBe("claude-3-5-sonnet-v2@20241022")
expect(modelInfo.info).toBeDefined()
@ -623,14 +877,46 @@ describe("VertexHandler", () => {
expect(modelInfo.info.contextWindow).toBe(200_000)
})
it("should return default model if invalid model specified", () => {
const invalidHandler = new VertexHandler({
apiModelId: "invalid-model",
it("should return correct model info for Gemini", () => {
handler = new VertexHandler({
apiModelId: "gemini-2.0-flash-001",
vertexProjectId: "test-project",
vertexRegion: "us-central1",
})
const modelInfo = invalidHandler.getModel()
expect(modelInfo.id).toBe("claude-3-7-sonnet@20250219") // Default model
const modelInfo = handler.getModel()
expect(modelInfo.id).toBe("gemini-2.0-flash-001")
expect(modelInfo.info).toBeDefined()
expect(modelInfo.info.maxTokens).toBe(8192)
expect(modelInfo.info.contextWindow).toBe(1048576)
})
it("honors custom maxTokens for thinking models", () => {
const handler = new VertexHandler({
apiKey: "test-api-key",
apiModelId: "claude-3-7-sonnet@20250219:thinking",
modelMaxTokens: 32_768,
modelMaxThinkingTokens: 16_384,
})
const result = handler.getModel()
expect(result.maxTokens).toBe(32_768)
expect(result.thinking).toEqual({ type: "enabled", budget_tokens: 16_384 })
expect(result.temperature).toBe(1.0)
})
it("does not honor custom maxTokens for non-thinking models", () => {
const handler = new VertexHandler({
apiKey: "test-api-key",
apiModelId: "claude-3-7-sonnet@20250219",
modelMaxTokens: 32_768,
modelMaxThinkingTokens: 16_384,
})
const result = handler.getModel()
expect(result.maxTokens).toBe(16_384)
expect(result.thinking).toBeUndefined()
expect(result.temperature).toBe(0)
})
})
@ -724,7 +1010,7 @@ describe("VertexHandler", () => {
},
}
})
;(thinkingHandler["client"].messages as any).create = mockCreate
;(thinkingHandler["anthropicClient"].messages as any).create = mockCreate
await thinkingHandler
.createMessage("You are a helpful assistant", [{ role: "user", content: "Hello" }])

View file

@ -1,7 +1,6 @@
import { Anthropic } from "@anthropic-ai/sdk"
import { Stream as AnthropicStream } from "@anthropic-ai/sdk/streaming"
import { CacheControlEphemeral } from "@anthropic-ai/sdk/resources"
import { BetaThinkingConfigParam } from "@anthropic-ai/sdk/resources/beta"
import {
anthropicDefaultModelId,
AnthropicModelId,
@ -9,18 +8,18 @@ import {
ApiHandlerOptions,
ModelInfo,
} from "../../shared/api"
import { ApiHandler, SingleCompletionHandler } from "../index"
import { ApiStream } from "../transform/stream"
import { BaseProvider } from "./base-provider"
import { ANTHROPIC_DEFAULT_MAX_TOKENS } from "./constants"
import { SingleCompletionHandler, getModelParams } from "../index"
const ANTHROPIC_DEFAULT_TEMPERATURE = 0
export class AnthropicHandler implements ApiHandler, SingleCompletionHandler {
export class AnthropicHandler extends BaseProvider implements SingleCompletionHandler {
private options: ApiHandlerOptions
private client: Anthropic
constructor(options: ApiHandlerOptions) {
super()
this.options = options
this.client = new Anthropic({
apiKey: this.options.apiKey,
baseURL: this.options.anthropicBaseUrl || undefined,
@ -30,7 +29,7 @@ export class AnthropicHandler implements ApiHandler, SingleCompletionHandler {
async *createMessage(systemPrompt: string, messages: Anthropic.Messages.MessageParam[]): ApiStream {
let stream: AnthropicStream<Anthropic.Messages.RawMessageStreamEvent>
const cacheControl: CacheControlEphemeral = { type: "ephemeral" }
let { id: modelId, temperature, maxTokens, thinking } = this.getModel()
let { id: modelId, maxTokens, thinking, temperature } = this.getModel()
switch (modelId) {
case "claude-3-7-sonnet-20250219":
@ -53,7 +52,7 @@ export class AnthropicHandler implements ApiHandler, SingleCompletionHandler {
stream = await this.client.messages.create(
{
model: modelId,
max_tokens: maxTokens,
max_tokens: maxTokens ?? ANTHROPIC_DEFAULT_MAX_TOKENS,
temperature,
thinking,
// Setting cache breakpoint for system prompt so new tasks can reuse it.
@ -101,7 +100,7 @@ export class AnthropicHandler implements ApiHandler, SingleCompletionHandler {
default: {
stream = (await this.client.messages.create({
model: modelId,
max_tokens: maxTokens,
max_tokens: maxTokens ?? ANTHROPIC_DEFAULT_MAX_TOKENS,
temperature,
system: [{ text: systemPrompt, type: "text" }],
messages,
@ -182,55 +181,31 @@ export class AnthropicHandler implements ApiHandler, SingleCompletionHandler {
getModel() {
const modelId = this.options.apiModelId
let temperature = this.options.modelTemperature ?? ANTHROPIC_DEFAULT_TEMPERATURE
let thinking: BetaThinkingConfigParam | undefined = undefined
let id = modelId && modelId in anthropicModels ? (modelId as AnthropicModelId) : anthropicDefaultModelId
const info: ModelInfo = anthropicModels[id]
if (modelId && modelId in anthropicModels) {
let id = modelId as AnthropicModelId
const info: ModelInfo = anthropicModels[id]
// The `:thinking` variant is a virtual identifier for the
// `claude-3-7-sonnet-20250219` model with a thinking budget.
// We can handle this more elegantly in the future.
if (id === "claude-3-7-sonnet-20250219:thinking") {
id = "claude-3-7-sonnet-20250219"
}
const maxTokens = this.options.modelMaxTokens || info.maxTokens || 8192
if (info.thinking) {
// Anthropic "Thinking" models require a temperature of 1.0.
temperature = 1.0
// Clamp the thinking budget to be at most 80% of max tokens and at
// least 1024 tokens.
const maxBudgetTokens = Math.floor(maxTokens * 0.8)
const budgetTokens = Math.max(
Math.min(this.options.modelMaxThinkingTokens ?? maxBudgetTokens, maxBudgetTokens),
1024,
)
thinking = { type: "enabled", budget_tokens: budgetTokens }
}
return { id, info, temperature, maxTokens, thinking }
// The `:thinking` variant is a virtual identifier for the
// `claude-3-7-sonnet-20250219` model with a thinking budget.
// We can handle this more elegantly in the future.
if (id === "claude-3-7-sonnet-20250219:thinking") {
id = "claude-3-7-sonnet-20250219"
}
const id = anthropicDefaultModelId
const info: ModelInfo = anthropicModels[id]
const maxTokens = this.options.modelMaxTokens || info.maxTokens || 8192
return { id, info, temperature, maxTokens, thinking }
return {
id,
info,
...getModelParams({ options: this.options, model: info, defaultMaxTokens: ANTHROPIC_DEFAULT_MAX_TOKENS }),
}
}
async completePrompt(prompt: string) {
let { id: modelId, temperature, maxTokens, thinking } = this.getModel()
let { id: modelId, maxTokens, thinking, temperature } = this.getModel()
const message = await this.client.messages.create({
model: modelId,
max_tokens: maxTokens,
temperature,
max_tokens: maxTokens ?? ANTHROPIC_DEFAULT_MAX_TOKENS,
thinking,
temperature,
messages: [{ role: "user", content: prompt }],
stream: false,
})
@ -238,4 +213,35 @@ export class AnthropicHandler implements ApiHandler, SingleCompletionHandler {
const content = message.content.find(({ type }) => type === "text")
return content?.type === "text" ? content.text : ""
}
/**
* Counts tokens for the given content using Anthropic's API
*
* @param content The content blocks to count tokens for
* @returns A promise resolving to the token count
*/
override async countTokens(content: Array<Anthropic.Messages.ContentBlockParam>): Promise<number> {
try {
// Use the current model
const actualModelId = this.getModel().id
const response = await this.client.messages.countTokens({
model: actualModelId,
messages: [
{
role: "user",
content: content,
},
],
})
return response.input_tokens
} catch (error) {
// Log error but fallback to tiktoken estimation
console.warn("Anthropic token counting failed, using fallback", error)
// Use the base provider's implementation as fallback
return super.countTokens(content)
}
}
}

View file

@ -0,0 +1,64 @@
import { Anthropic } from "@anthropic-ai/sdk"
import { ApiHandler } from ".."
import { ModelInfo } from "../../shared/api"
import { ApiStream } from "../transform/stream"
import { Tiktoken } from "js-tiktoken/lite"
import o200kBase from "js-tiktoken/ranks/o200k_base"
// Reuse the fudge factor used in the original code
const TOKEN_FUDGE_FACTOR = 1.5
/**
* Base class for API providers that implements common functionality
*/
export abstract class BaseProvider implements ApiHandler {
// Cache the Tiktoken encoder instance since it's stateless
private encoder: Tiktoken | null = null
abstract createMessage(systemPrompt: string, messages: Anthropic.Messages.MessageParam[]): ApiStream
abstract getModel(): { id: string; info: ModelInfo }
/**
* Default token counting implementation using tiktoken
* Providers can override this to use their native token counting endpoints
*
* Uses a cached Tiktoken encoder instance for performance since it's stateless.
* The encoder is created lazily on first use and reused for subsequent calls.
*
* @param content The content to count tokens for
* @returns A promise resolving to the token count
*/
async countTokens(content: Array<Anthropic.Messages.ContentBlockParam>): Promise<number> {
if (!content || content.length === 0) return 0
let totalTokens = 0
// Lazily create and cache the encoder if it doesn't exist
if (!this.encoder) {
this.encoder = new Tiktoken(o200kBase)
}
// Process each content block using the cached encoder
for (const block of content) {
if (block.type === "text") {
// Use tiktoken for text token counting
const text = block.text || ""
if (text.length > 0) {
const tokens = this.encoder.encode(text)
totalTokens += tokens.length
}
} else if (block.type === "image") {
// For images, calculate based on data size
const imageSource = block.source
if (imageSource && typeof imageSource === "object" && "data" in imageSource) {
const base64Data = imageSource.data as string
totalTokens += Math.ceil(Math.sqrt(base64Data.length))
} else {
totalTokens += 300 // Conservative estimate for unknown images
}
}
}
// Add a fudge factor to account for the fact that tiktoken is not always accurate
return Math.ceil(totalTokens * TOKEN_FUDGE_FACTOR)
}
}

View file

@ -6,10 +6,11 @@ import {
} from "@aws-sdk/client-bedrock-runtime"
import { fromIni } from "@aws-sdk/credential-providers"
import { Anthropic } from "@anthropic-ai/sdk"
import { ApiHandler, SingleCompletionHandler } from "../"
import { SingleCompletionHandler } from "../"
import { ApiHandlerOptions, BedrockModelId, ModelInfo, bedrockDefaultModelId, bedrockModels } from "../../shared/api"
import { ApiStream } from "../transform/stream"
import { convertToBedrockConverseMessages } from "../transform/bedrock-converse-format"
import { BaseProvider } from "./base-provider"
const BEDROCK_DEFAULT_TEMPERATURE = 0.3
@ -46,11 +47,12 @@ export interface StreamEvent {
}
}
export class AwsBedrockHandler implements ApiHandler, SingleCompletionHandler {
private options: ApiHandlerOptions
export class AwsBedrockHandler extends BaseProvider implements SingleCompletionHandler {
protected options: ApiHandlerOptions
private client: BedrockRuntimeClient
constructor(options: ApiHandlerOptions) {
super()
this.options = options
const clientConfig: BedrockRuntimeClientConfig = {
@ -74,7 +76,7 @@ export class AwsBedrockHandler implements ApiHandler, SingleCompletionHandler {
this.client = new BedrockRuntimeClient(clientConfig)
}
async *createMessage(systemPrompt: string, messages: Anthropic.Messages.MessageParam[]): ApiStream {
override async *createMessage(systemPrompt: string, messages: Anthropic.Messages.MessageParam[]): ApiStream {
const modelConfig = this.getModel()
// Handle cross-region inference
@ -205,7 +207,7 @@ export class AwsBedrockHandler implements ApiHandler, SingleCompletionHandler {
}
}
getModel(): { id: BedrockModelId | string; info: ModelInfo } {
override getModel(): { id: BedrockModelId | string; info: ModelInfo } {
const modelId = this.options.apiModelId
if (modelId) {
// For tests, allow any model ID

View file

@ -0,0 +1,3 @@
export const ANTHROPIC_DEFAULT_MAX_TOKENS = 8192
export const DEEP_SEEK_DEFAULT_TEMPERATURE = 0.6

View file

@ -1,22 +1,24 @@
import { Anthropic } from "@anthropic-ai/sdk"
import { GoogleGenerativeAI } from "@google/generative-ai"
import { ApiHandler, SingleCompletionHandler } from "../"
import { SingleCompletionHandler } from "../"
import { ApiHandlerOptions, geminiDefaultModelId, GeminiModelId, geminiModels, ModelInfo } from "../../shared/api"
import { convertAnthropicMessageToGemini } from "../transform/gemini-format"
import { ApiStream } from "../transform/stream"
import { BaseProvider } from "./base-provider"
const GEMINI_DEFAULT_TEMPERATURE = 0
export class GeminiHandler implements ApiHandler, SingleCompletionHandler {
private options: ApiHandlerOptions
export class GeminiHandler extends BaseProvider implements SingleCompletionHandler {
protected options: ApiHandlerOptions
private client: GoogleGenerativeAI
constructor(options: ApiHandlerOptions) {
super()
this.options = options
this.client = new GoogleGenerativeAI(options.geminiApiKey ?? "not-provided")
}
async *createMessage(systemPrompt: string, messages: Anthropic.Messages.MessageParam[]): ApiStream {
override async *createMessage(systemPrompt: string, messages: Anthropic.Messages.MessageParam[]): ApiStream {
const model = this.client.getGenerativeModel({
model: this.getModel().id,
systemInstruction: systemPrompt,
@ -44,7 +46,7 @@ export class GeminiHandler implements ApiHandler, SingleCompletionHandler {
}
}
getModel(): { id: GeminiModelId; info: ModelInfo } {
override getModel(): { id: GeminiModelId; info: ModelInfo } {
const modelId = this.options.apiModelId
if (modelId && modelId in geminiModels) {
const id = modelId as GeminiModelId

View file

@ -6,22 +6,39 @@ import { ApiHandlerOptions, ModelInfo, glamaDefaultModelId, glamaDefaultModelInf
import { parseApiPrice } from "../../utils/cost"
import { convertToOpenAiMessages } from "../transform/openai-format"
import { ApiStream } from "../transform/stream"
import { ApiHandler, SingleCompletionHandler } from "../"
import { SingleCompletionHandler } from "../"
import { BaseProvider } from "./base-provider"
const GLAMA_DEFAULT_TEMPERATURE = 0
export class GlamaHandler implements ApiHandler, SingleCompletionHandler {
private options: ApiHandlerOptions
export class GlamaHandler extends BaseProvider implements SingleCompletionHandler {
protected options: ApiHandlerOptions
private client: OpenAI
constructor(options: ApiHandlerOptions) {
super()
this.options = options
const baseURL = "https://glama.ai/api/gateway/openai/v1"
const apiKey = this.options.glamaApiKey ?? "not-provided"
this.client = new OpenAI({ baseURL, apiKey })
}
async *createMessage(systemPrompt: string, messages: Anthropic.Messages.MessageParam[]): ApiStream {
private supportsTemperature(): boolean {
return !this.getModel().id.startsWith("openai/o3-mini")
}
override getModel(): { id: string; info: ModelInfo } {
const modelId = this.options.glamaModelId
const modelInfo = this.options.glamaModelInfo
if (modelId && modelInfo) {
return { id: modelId, info: modelInfo }
}
return { id: glamaDefaultModelId, info: glamaDefaultModelInfo }
}
override async *createMessage(systemPrompt: string, messages: Anthropic.Messages.MessageParam[]): ApiStream {
// Convert Anthropic messages to OpenAI format
const openAiMessages: OpenAI.Chat.ChatCompletionMessageParam[] = [
{ role: "system", content: systemPrompt },
@ -152,21 +169,6 @@ export class GlamaHandler implements ApiHandler, SingleCompletionHandler {
}
}
private supportsTemperature(): boolean {
return !this.getModel().id.startsWith("openai/o3-mini")
}
getModel(): { id: string; info: ModelInfo } {
const modelId = this.options.glamaModelId
const modelInfo = this.options.glamaModelInfo
if (modelId && modelInfo) {
return { id: modelId, info: modelInfo }
}
return { id: glamaDefaultModelId, info: glamaDefaultModelInfo }
}
async completePrompt(prompt: string): Promise<string> {
try {
const requestOptions: OpenAI.Chat.Completions.ChatCompletionCreateParamsNonStreaming = {

View file

@ -2,18 +2,20 @@ import { Anthropic } from "@anthropic-ai/sdk"
import OpenAI from "openai"
import axios from "axios"
import { ApiHandler, SingleCompletionHandler } from "../"
import { SingleCompletionHandler } from "../"
import { ApiHandlerOptions, ModelInfo, openAiModelInfoSaneDefaults } from "../../shared/api"
import { convertToOpenAiMessages } from "../transform/openai-format"
import { ApiStream } from "../transform/stream"
import { BaseProvider } from "./base-provider"
const LMSTUDIO_DEFAULT_TEMPERATURE = 0
export class LmStudioHandler implements ApiHandler, SingleCompletionHandler {
private options: ApiHandlerOptions
export class LmStudioHandler extends BaseProvider implements SingleCompletionHandler {
protected options: ApiHandlerOptions
private client: OpenAI
constructor(options: ApiHandlerOptions) {
super()
this.options = options
this.client = new OpenAI({
baseURL: (this.options.lmStudioBaseUrl || "http://localhost:1234") + "/v1",
@ -21,7 +23,7 @@ export class LmStudioHandler implements ApiHandler, SingleCompletionHandler {
})
}
async *createMessage(systemPrompt: string, messages: Anthropic.Messages.MessageParam[]): ApiStream {
override async *createMessage(systemPrompt: string, messages: Anthropic.Messages.MessageParam[]): ApiStream {
const openAiMessages: OpenAI.Chat.ChatCompletionMessageParam[] = [
{ role: "system", content: systemPrompt },
...convertToOpenAiMessages(messages),
@ -51,7 +53,7 @@ export class LmStudioHandler implements ApiHandler, SingleCompletionHandler {
}
}
getModel(): { id: string; info: ModelInfo } {
override getModel(): { id: string; info: ModelInfo } {
return {
id: this.options.lmStudioModelId || "",
info: openAiModelInfoSaneDefaults,

View file

@ -1,6 +1,6 @@
import { Anthropic } from "@anthropic-ai/sdk"
import { Mistral } from "@mistralai/mistralai"
import { ApiHandler } from "../"
import { SingleCompletionHandler } from "../"
import {
ApiHandlerOptions,
mistralDefaultModelId,
@ -13,14 +13,16 @@ import {
} from "../../shared/api"
import { convertToMistralMessages } from "../transform/mistral-format"
import { ApiStream } from "../transform/stream"
import { BaseProvider } from "./base-provider"
const MISTRAL_DEFAULT_TEMPERATURE = 0
export class MistralHandler implements ApiHandler {
private options: ApiHandlerOptions
export class MistralHandler extends BaseProvider implements SingleCompletionHandler {
protected options: ApiHandlerOptions
private client: Mistral
constructor(options: ApiHandlerOptions) {
super()
if (!options.mistralApiKey) {
throw new Error("Mistral API key is required")
}
@ -48,7 +50,7 @@ export class MistralHandler implements ApiHandler {
return "https://api.mistral.ai"
}
async *createMessage(systemPrompt: string, messages: Anthropic.Messages.MessageParam[]): ApiStream {
override async *createMessage(systemPrompt: string, messages: Anthropic.Messages.MessageParam[]): ApiStream {
const response = await this.client.chat.stream({
model: this.options.apiModelId || mistralDefaultModelId,
messages: [{ role: "system", content: systemPrompt }, ...convertToMistralMessages(messages)],
@ -81,7 +83,7 @@ export class MistralHandler implements ApiHandler {
}
}
getModel(): { id: MistralModelId; info: ModelInfo } {
override getModel(): { id: MistralModelId; info: ModelInfo } {
const modelId = this.options.apiModelId
if (modelId && modelId in mistralModels) {
const id = modelId as MistralModelId

View file

@ -2,21 +2,21 @@ import { Anthropic } from "@anthropic-ai/sdk"
import OpenAI from "openai"
import axios from "axios"
import { ApiHandler, SingleCompletionHandler } from "../"
import { SingleCompletionHandler } from "../"
import { ApiHandlerOptions, ModelInfo, openAiModelInfoSaneDefaults } from "../../shared/api"
import { convertToOpenAiMessages } from "../transform/openai-format"
import { convertToR1Format } from "../transform/r1-format"
import { ApiStream } from "../transform/stream"
import { DEEP_SEEK_DEFAULT_TEMPERATURE } from "./openai"
import { DEEP_SEEK_DEFAULT_TEMPERATURE } from "./constants"
import { XmlMatcher } from "../../utils/xml-matcher"
import { BaseProvider } from "./base-provider"
const OLLAMA_DEFAULT_TEMPERATURE = 0
export class OllamaHandler implements ApiHandler, SingleCompletionHandler {
private options: ApiHandlerOptions
export class OllamaHandler extends BaseProvider implements SingleCompletionHandler {
protected options: ApiHandlerOptions
private client: OpenAI
constructor(options: ApiHandlerOptions) {
super()
this.options = options
this.client = new OpenAI({
baseURL: (this.options.ollamaBaseUrl || "http://localhost:11434") + "/v1",
@ -24,7 +24,7 @@ export class OllamaHandler implements ApiHandler, SingleCompletionHandler {
})
}
async *createMessage(systemPrompt: string, messages: Anthropic.Messages.MessageParam[]): ApiStream {
override async *createMessage(systemPrompt: string, messages: Anthropic.Messages.MessageParam[]): ApiStream {
const modelId = this.getModel().id
const useR1Format = modelId.toLowerCase().includes("deepseek-r1")
const openAiMessages: OpenAI.Chat.ChatCompletionMessageParam[] = [
@ -35,7 +35,7 @@ export class OllamaHandler implements ApiHandler, SingleCompletionHandler {
const stream = await this.client.chat.completions.create({
model: this.getModel().id,
messages: openAiMessages,
temperature: this.options.modelTemperature ?? OLLAMA_DEFAULT_TEMPERATURE,
temperature: this.options.modelTemperature ?? 0,
stream: true,
})
const matcher = new XmlMatcher(
@ -60,7 +60,7 @@ export class OllamaHandler implements ApiHandler, SingleCompletionHandler {
}
}
getModel(): { id: string; info: ModelInfo } {
override getModel(): { id: string; info: ModelInfo } {
return {
id: this.options.ollamaModelId || "",
info: openAiModelInfoSaneDefaults,
@ -76,9 +76,7 @@ export class OllamaHandler implements ApiHandler, SingleCompletionHandler {
messages: useR1Format
? convertToR1Format([{ role: "user", content: prompt }])
: [{ role: "user", content: prompt }],
temperature:
this.options.modelTemperature ??
(useR1Format ? DEEP_SEEK_DEFAULT_TEMPERATURE : OLLAMA_DEFAULT_TEMPERATURE),
temperature: this.options.modelTemperature ?? (useR1Format ? DEEP_SEEK_DEFAULT_TEMPERATURE : 0),
stream: false,
})
return response.choices[0]?.message.content || ""

View file

@ -1,6 +1,6 @@
import { Anthropic } from "@anthropic-ai/sdk"
import OpenAI from "openai"
import { ApiHandler, SingleCompletionHandler } from "../"
import { SingleCompletionHandler } from "../"
import {
ApiHandlerOptions,
ModelInfo,
@ -10,20 +10,22 @@ import {
} from "../../shared/api"
import { convertToOpenAiMessages } from "../transform/openai-format"
import { ApiStream } from "../transform/stream"
import { BaseProvider } from "./base-provider"
const OPENAI_NATIVE_DEFAULT_TEMPERATURE = 0
export class OpenAiNativeHandler implements ApiHandler, SingleCompletionHandler {
private options: ApiHandlerOptions
export class OpenAiNativeHandler extends BaseProvider implements SingleCompletionHandler {
protected options: ApiHandlerOptions
private client: OpenAI
constructor(options: ApiHandlerOptions) {
super()
this.options = options
const apiKey = this.options.openAiNativeApiKey ?? "not-provided"
this.client = new OpenAI({ apiKey })
}
async *createMessage(systemPrompt: string, messages: Anthropic.Messages.MessageParam[]): ApiStream {
override async *createMessage(systemPrompt: string, messages: Anthropic.Messages.MessageParam[]): ApiStream {
const modelId = this.getModel().id
if (modelId.startsWith("o1")) {
@ -133,7 +135,7 @@ export class OpenAiNativeHandler implements ApiHandler, SingleCompletionHandler
}
}
getModel(): { id: OpenAiNativeModelId; info: ModelInfo } {
override getModel(): { id: OpenAiNativeModelId; info: ModelInfo } {
const modelId = this.options.apiModelId
if (modelId && modelId in openAiNativeModels) {
const id = modelId as OpenAiNativeModelId

View file

@ -8,24 +8,24 @@ import {
ModelInfo,
openAiModelInfoSaneDefaults,
} from "../../shared/api"
import { ApiHandler, SingleCompletionHandler } from "../index"
import { SingleCompletionHandler } from "../index"
import { convertToOpenAiMessages } from "../transform/openai-format"
import { convertToR1Format } from "../transform/r1-format"
import { convertToSimpleMessages } from "../transform/simple-format"
import { ApiStream, ApiStreamUsageChunk } from "../transform/stream"
import { BaseProvider } from "./base-provider"
const DEEP_SEEK_DEFAULT_TEMPERATURE = 0.6
export interface OpenAiHandlerOptions extends ApiHandlerOptions {
defaultHeaders?: Record<string, string>
}
export const DEEP_SEEK_DEFAULT_TEMPERATURE = 0.6
const OPENAI_DEFAULT_TEMPERATURE = 0
export class OpenAiHandler implements ApiHandler, SingleCompletionHandler {
export class OpenAiHandler extends BaseProvider implements SingleCompletionHandler {
protected options: OpenAiHandlerOptions
private client: OpenAI
constructor(options: OpenAiHandlerOptions) {
super()
this.options = options
const baseURL = this.options.openAiBaseUrl ?? "https://api.openai.com/v1"
@ -53,7 +53,7 @@ export class OpenAiHandler implements ApiHandler, SingleCompletionHandler {
}
}
async *createMessage(systemPrompt: string, messages: Anthropic.Messages.MessageParam[]): ApiStream {
override async *createMessage(systemPrompt: string, messages: Anthropic.Messages.MessageParam[]): ApiStream {
const modelInfo = this.getModel().info
const modelUrl = this.options.openAiBaseUrl ?? ""
const modelId = this.options.openAiModelId ?? ""
@ -78,9 +78,7 @@ export class OpenAiHandler implements ApiHandler, SingleCompletionHandler {
const requestOptions: OpenAI.Chat.Completions.ChatCompletionCreateParamsStreaming = {
model: modelId,
temperature:
this.options.modelTemperature ??
(deepseekReasoner ? DEEP_SEEK_DEFAULT_TEMPERATURE : OPENAI_DEFAULT_TEMPERATURE),
temperature: this.options.modelTemperature ?? (deepseekReasoner ? DEEP_SEEK_DEFAULT_TEMPERATURE : 0),
messages: convertedMessages,
stream: true as const,
stream_options: { include_usage: true },
@ -143,7 +141,7 @@ export class OpenAiHandler implements ApiHandler, SingleCompletionHandler {
}
}
getModel(): { id: string; info: ModelInfo } {
override getModel(): { id: string; info: ModelInfo } {
return {
id: this.options.openAiModelId ?? "",
info: this.options.openAiCustomModelInfo ?? openAiModelInfoSaneDefaults,

View file

@ -9,10 +9,10 @@ import { parseApiPrice } from "../../utils/cost"
import { convertToOpenAiMessages } from "../transform/openai-format"
import { ApiStreamChunk, ApiStreamUsageChunk } from "../transform/stream"
import { convertToR1Format } from "../transform/r1-format"
import { DEEP_SEEK_DEFAULT_TEMPERATURE } from "./openai"
import { ApiHandler, SingleCompletionHandler } from ".."
const OPENROUTER_DEFAULT_TEMPERATURE = 0
import { DEEP_SEEK_DEFAULT_TEMPERATURE } from "./constants"
import { getModelParams, SingleCompletionHandler } from ".."
import { BaseProvider } from "./base-provider"
// Add custom interface for OpenRouter params.
type OpenRouterChatCompletionParams = OpenAI.Chat.ChatCompletionCreateParams & {
@ -26,11 +26,12 @@ interface OpenRouterApiStreamUsageChunk extends ApiStreamUsageChunk {
fullResponseText: string
}
export class OpenRouterHandler implements ApiHandler, SingleCompletionHandler {
private options: ApiHandlerOptions
export class OpenRouterHandler extends BaseProvider implements SingleCompletionHandler {
protected options: ApiHandlerOptions
private client: OpenAI
constructor(options: ApiHandlerOptions) {
super()
this.options = options
const baseURL = this.options.openRouterBaseUrl || "https://openrouter.ai/api/v1"
@ -44,17 +45,22 @@ export class OpenRouterHandler implements ApiHandler, SingleCompletionHandler {
this.client = new OpenAI({ baseURL, apiKey, defaultHeaders })
}
async *createMessage(
override async *createMessage(
systemPrompt: string,
messages: Anthropic.Messages.MessageParam[],
): AsyncGenerator<ApiStreamChunk> {
// Convert Anthropic messages to OpenAI format
let { id: modelId, maxTokens, thinking, temperature, topP } = this.getModel()
// Convert Anthropic messages to OpenAI format.
let openAiMessages: OpenAI.Chat.ChatCompletionMessageParam[] = [
{ role: "system", content: systemPrompt },
...convertToOpenAiMessages(messages),
]
const { id: modelId, info: modelInfo } = this.getModel()
// DeepSeek highly recommends using user instead of system role.
if (modelId.startsWith("deepseek/deepseek-r1") || modelId === "perplexity/sonar-reasoning") {
openAiMessages = convertToR1Format([{ role: "user", content: systemPrompt }, ...messages])
}
// prompt caching: https://openrouter.ai/docs/prompt-caching
// this is specifically for claude models (some models may 'support prompt caching' automatically without this)
@ -95,42 +101,12 @@ export class OpenRouterHandler implements ApiHandler, SingleCompletionHandler {
break
}
let defaultTemperature = OPENROUTER_DEFAULT_TEMPERATURE
let topP: number | undefined = undefined
// Handle models based on deepseek-r1
if (modelId.startsWith("deepseek/deepseek-r1") || modelId === "perplexity/sonar-reasoning") {
// Recommended temperature for DeepSeek reasoning models
defaultTemperature = DEEP_SEEK_DEFAULT_TEMPERATURE
// DeepSeek highly recommends using user instead of system role
openAiMessages = convertToR1Format([{ role: "user", content: systemPrompt }, ...messages])
// Some provider support topP and 0.95 is value that Deepseek used in their benchmarks
topP = 0.95
}
const maxTokens = this.options.modelMaxTokens || modelInfo.maxTokens
let temperature = this.options.modelTemperature ?? defaultTemperature
let thinking: BetaThinkingConfigParam | undefined = undefined
if (modelInfo.thinking) {
// Clamp the thinking budget to be at most 80% of max tokens and at
// least 1024 tokens.
const maxBudgetTokens = Math.floor((maxTokens || 8192) * 0.8)
const budgetTokens = Math.max(
Math.min(this.options.modelMaxThinkingTokens ?? maxBudgetTokens, maxBudgetTokens),
1024,
)
thinking = { type: "enabled", budget_tokens: budgetTokens }
temperature = 1.0
}
// https://openrouter.ai/docs/transforms
let fullResponseText = ""
const completionParams: OpenRouterChatCompletionParams = {
model: modelId,
max_tokens: modelInfo.maxTokens,
max_tokens: maxTokens,
temperature,
thinking, // OpenRouter is temporarily supporting this.
top_p: topP,
@ -138,7 +114,7 @@ export class OpenRouterHandler implements ApiHandler, SingleCompletionHandler {
stream: true,
include_reasoning: true,
// This way, the transforms field will only be included in the parameters when openRouterUseMiddleOutTransform is true.
...(this.options.openRouterUseMiddleOutTransform && { transforms: ["middle-out"] }),
...((this.options.openRouterUseMiddleOutTransform ?? true) && { transforms: ["middle-out"] }),
}
const stream = await this.client.chat.completions.create(completionParams)
@ -218,37 +194,46 @@ export class OpenRouterHandler implements ApiHandler, SingleCompletionHandler {
}
}
getModel() {
override getModel() {
const modelId = this.options.openRouterModelId
const modelInfo = this.options.openRouterModelInfo
return modelId && modelInfo
? { id: modelId, info: modelInfo }
: { id: openRouterDefaultModelId, info: openRouterDefaultModelInfo }
let id = modelId ?? openRouterDefaultModelId
const info = modelInfo ?? openRouterDefaultModelInfo
const isDeepSeekR1 = id.startsWith("deepseek/deepseek-r1") || modelId === "perplexity/sonar-reasoning"
const defaultTemperature = isDeepSeekR1 ? DEEP_SEEK_DEFAULT_TEMPERATURE : 0
const topP = isDeepSeekR1 ? 0.95 : undefined
return {
id,
info,
...getModelParams({ options: this.options, model: info, defaultTemperature }),
topP,
}
}
async completePrompt(prompt: string): Promise<string> {
try {
const response = await this.client.chat.completions.create({
model: this.getModel().id,
messages: [{ role: "user", content: prompt }],
temperature: this.options.modelTemperature ?? OPENROUTER_DEFAULT_TEMPERATURE,
stream: false,
})
async completePrompt(prompt: string) {
let { id: modelId, maxTokens, thinking, temperature } = this.getModel()
if ("error" in response) {
const error = response.error as { message?: string; code?: number }
throw new Error(`OpenRouter API Error ${error?.code}: ${error?.message}`)
}
const completion = response as OpenAI.Chat.ChatCompletion
return completion.choices[0]?.message?.content || ""
} catch (error) {
if (error instanceof Error) {
throw new Error(`OpenRouter completion error: ${error.message}`)
}
throw error
const completionParams: OpenRouterChatCompletionParams = {
model: modelId,
max_tokens: maxTokens,
thinking,
temperature,
messages: [{ role: "user", content: prompt }],
stream: false,
}
const response = await this.client.chat.completions.create(completionParams)
if ("error" in response) {
const error = response.error as { message?: string; code?: number }
throw new Error(`OpenRouter API Error ${error?.code}: ${error?.message}`)
}
const completion = response as OpenAI.Chat.ChatCompletion
return completion.choices[0]?.message?.content || ""
}
}
@ -278,7 +263,7 @@ export async function getOpenRouterModels() {
modelInfo.supportsPromptCache = true
modelInfo.cacheWritesPrice = 3.75
modelInfo.cacheReadsPrice = 0.3
modelInfo.maxTokens = 64_000
modelInfo.maxTokens = rawModel.id === "anthropic/claude-3.7-sonnet:thinking" ? 64_000 : 16_384
break
case rawModel.id.startsWith("anthropic/claude-3.5-sonnet-20240620"):
modelInfo.supportsPromptCache = true

View file

@ -5,25 +5,27 @@ import OpenAI from "openai"
import { ApiHandlerOptions, ModelInfo, unboundDefaultModelId, unboundDefaultModelInfo } from "../../shared/api"
import { convertToOpenAiMessages } from "../transform/openai-format"
import { ApiStream, ApiStreamUsageChunk } from "../transform/stream"
import { ApiHandler, SingleCompletionHandler } from "../"
import { SingleCompletionHandler } from "../"
import { BaseProvider } from "./base-provider"
interface UnboundUsage extends OpenAI.CompletionUsage {
cache_creation_input_tokens?: number
cache_read_input_tokens?: number
}
export class UnboundHandler implements ApiHandler, SingleCompletionHandler {
private options: ApiHandlerOptions
export class UnboundHandler extends BaseProvider implements SingleCompletionHandler {
protected options: ApiHandlerOptions
private client: OpenAI
constructor(options: ApiHandlerOptions) {
super()
this.options = options
const baseURL = "https://api.getunbound.ai/v1"
const apiKey = this.options.unboundApiKey ?? "not-provided"
this.client = new OpenAI({ baseURL, apiKey })
}
async *createMessage(systemPrompt: string, messages: Anthropic.Messages.MessageParam[]): ApiStream {
override async *createMessage(systemPrompt: string, messages: Anthropic.Messages.MessageParam[]): ApiStream {
// Convert Anthropic messages to OpenAI format
const openAiMessages: OpenAI.Chat.ChatCompletionMessageParam[] = [
{ role: "system", content: systemPrompt },
@ -131,7 +133,7 @@ export class UnboundHandler implements ApiHandler, SingleCompletionHandler {
}
}
getModel(): { id: string; info: ModelInfo } {
override getModel(): { id: string; info: ModelInfo } {
const modelId = this.options.unboundModelId
const modelInfo = this.options.unboundModelInfo
if (modelId && modelInfo) {

View file

@ -1,10 +1,16 @@
import { Anthropic } from "@anthropic-ai/sdk"
import { AnthropicVertex } from "@anthropic-ai/vertex-sdk"
import { Stream as AnthropicStream } from "@anthropic-ai/sdk/streaming"
import { ApiHandler, SingleCompletionHandler } from "../"
import { BetaThinkingConfigParam } from "@anthropic-ai/sdk/resources/beta"
import { VertexAI } from "@google-cloud/vertexai"
import { ApiHandlerOptions, ModelInfo, vertexDefaultModelId, VertexModelId, vertexModels } from "../../shared/api"
import { ApiStream } from "../transform/stream"
import { convertAnthropicMessageToVertexGemini } from "../transform/vertex-gemini-format"
import { BaseProvider } from "./base-provider"
import { ANTHROPIC_DEFAULT_MAX_TOKENS } from "./constants"
import { getModelParams, SingleCompletionHandler } from "../"
// Types for Vertex SDK
@ -93,17 +99,37 @@ interface VertexMessageStreamEvent {
}
// https://docs.anthropic.com/en/api/claude-on-vertex-ai
export class VertexHandler implements ApiHandler, SingleCompletionHandler {
private options: ApiHandlerOptions
private client: AnthropicVertex
export class VertexHandler extends BaseProvider implements SingleCompletionHandler {
MODEL_CLAUDE = "claude"
MODEL_GEMINI = "gemini"
protected options: ApiHandlerOptions
private anthropicClient: AnthropicVertex
private geminiClient: VertexAI
private modelType: string
constructor(options: ApiHandlerOptions) {
super()
this.options = options
this.client = new AnthropicVertex({
if (this.options.apiModelId?.startsWith(this.MODEL_CLAUDE)) {
this.modelType = this.MODEL_CLAUDE
} else if (this.options.apiModelId?.startsWith(this.MODEL_GEMINI)) {
this.modelType = this.MODEL_GEMINI
} else {
throw new Error(`Unknown model ID: ${this.options.apiModelId}`)
}
this.anthropicClient = new AnthropicVertex({
projectId: this.options.vertexProjectId ?? "not-provided",
// https://cloud.google.com/vertex-ai/generative-ai/docs/partner-models/use-claude#regions
region: this.options.vertexRegion ?? "us-east5",
})
this.geminiClient = new VertexAI({
project: this.options.vertexProjectId ?? "not-provided",
location: this.options.vertexRegion ?? "us-east5",
})
}
private formatMessageForCache(message: Anthropic.Messages.MessageParam, shouldCache: boolean): VertexMessage {
@ -154,7 +180,43 @@ export class VertexHandler implements ApiHandler, SingleCompletionHandler {
}
}
async *createMessage(systemPrompt: string, messages: Anthropic.Messages.MessageParam[]): ApiStream {
private async *createGeminiMessage(systemPrompt: string, messages: Anthropic.Messages.MessageParam[]): ApiStream {
const model = this.geminiClient.getGenerativeModel({
model: this.getModel().id,
systemInstruction: systemPrompt,
})
const result = await model.generateContentStream({
contents: messages.map(convertAnthropicMessageToVertexGemini),
generationConfig: {
maxOutputTokens: this.getModel().info.maxTokens,
temperature: this.options.modelTemperature ?? 0,
},
})
for await (const chunk of result.stream) {
if (chunk.candidates?.[0]?.content?.parts) {
for (const part of chunk.candidates[0].content.parts) {
if (part.text) {
yield {
type: "text",
text: part.text,
}
}
}
}
}
const response = await result.response
yield {
type: "usage",
inputTokens: response.usageMetadata?.promptTokenCount ?? 0,
outputTokens: response.usageMetadata?.candidatesTokenCount ?? 0,
}
}
private async *createClaudeMessage(systemPrompt: string, messages: Anthropic.Messages.MessageParam[]): ApiStream {
const model = this.getModel()
let { id, info, temperature, maxTokens, thinking } = model
const useCache = model.info.supportsPromptCache
@ -192,7 +254,7 @@ export class VertexHandler implements ApiHandler, SingleCompletionHandler {
stream: true,
}
const stream = (await this.client.messages.create(
const stream = (await this.anthropicClient.messages.create(
params as Anthropic.Messages.MessageCreateParamsStreaming,
)) as unknown as AnthropicStream<VertexMessageStreamEvent>
@ -272,58 +334,77 @@ export class VertexHandler implements ApiHandler, SingleCompletionHandler {
}
}
getModel(): {
id: VertexModelId
info: ModelInfo
temperature: number
maxTokens: number
thinking?: BetaThinkingConfigParam
} {
const modelId = this.options.apiModelId
let temperature = this.options.modelTemperature ?? 0
let thinking: BetaThinkingConfigParam | undefined = undefined
if (modelId && modelId in vertexModels) {
const id = modelId as VertexModelId
const info: ModelInfo = vertexModels[id]
// The `:thinking` variant is a virtual identifier for thinking-enabled models
// Similar to how it's handled in the Anthropic provider
let actualId = id
if (id.endsWith(":thinking")) {
actualId = id.replace(":thinking", "") as VertexModelId
override async *createMessage(systemPrompt: string, messages: Anthropic.Messages.MessageParam[]): ApiStream {
switch (this.modelType) {
case this.MODEL_CLAUDE: {
yield* this.createClaudeMessage(systemPrompt, messages)
break
}
const maxTokens = this.options.modelMaxTokens || info.maxTokens || 8192
if (info.thinking) {
temperature = 1.0 // Thinking requires temperature 1.0
const maxBudgetTokens = Math.floor(maxTokens * 0.8)
const budgetTokens = Math.max(
Math.min(this.options.modelMaxThinkingTokens ?? maxBudgetTokens, maxBudgetTokens),
1024,
)
thinking = { type: "enabled", budget_tokens: budgetTokens }
case this.MODEL_GEMINI: {
yield* this.createGeminiMessage(systemPrompt, messages)
break
}
default: {
throw new Error(`Invalid model type: ${this.modelType}`)
}
return { id: actualId, info, temperature, maxTokens, thinking }
}
const id = vertexDefaultModelId
const info = vertexModels[id]
const maxTokens = this.options.modelMaxTokens || info.maxTokens || 8192
return { id, info, temperature, maxTokens, thinking }
}
async completePrompt(prompt: string): Promise<string> {
getModel() {
const modelId = this.options.apiModelId
let id = modelId && modelId in vertexModels ? (modelId as VertexModelId) : vertexDefaultModelId
const info: ModelInfo = vertexModels[id]
// The `:thinking` variant is a virtual identifier for thinking-enabled
// models (similar to how it's handled in the Anthropic provider.)
if (id.endsWith(":thinking")) {
id = id.replace(":thinking", "") as VertexModelId
}
return {
id,
info,
...getModelParams({ options: this.options, model: info, defaultMaxTokens: ANTHROPIC_DEFAULT_MAX_TOKENS }),
}
}
private async completePromptGemini(prompt: string) {
try {
const model = this.geminiClient.getGenerativeModel({
model: this.getModel().id,
})
const result = await model.generateContent({
contents: [{ role: "user", parts: [{ text: prompt }] }],
generationConfig: {
temperature: this.options.modelTemperature ?? 0,
},
})
let text = ""
result.response.candidates?.forEach((candidate) => {
candidate.content.parts.forEach((part) => {
text += part.text
})
})
return text
} catch (error) {
if (error instanceof Error) {
throw new Error(`Vertex completion error: ${error.message}`)
}
throw error
}
}
private async completePromptClaude(prompt: string) {
try {
let { id, info, temperature, maxTokens, thinking } = this.getModel()
const useCache = info.supportsPromptCache
const params = {
const params: Anthropic.Messages.MessageCreateParamsNonStreaming = {
model: id,
max_tokens: maxTokens,
max_tokens: maxTokens ?? ANTHROPIC_DEFAULT_MAX_TOKENS,
temperature,
thinking,
system: "", // No system prompt needed for single completions
@ -344,20 +425,34 @@ export class VertexHandler implements ApiHandler, SingleCompletionHandler {
stream: false,
}
const response = (await this.client.messages.create(
params as Anthropic.Messages.MessageCreateParamsNonStreaming,
)) as unknown as VertexMessageResponse
const response = (await this.anthropicClient.messages.create(params)) as unknown as VertexMessageResponse
const content = response.content[0]
if (content.type === "text") {
return content.text
}
return ""
} catch (error) {
if (error instanceof Error) {
throw new Error(`Vertex completion error: ${error.message}`)
}
throw error
}
}
async completePrompt(prompt: string) {
switch (this.modelType) {
case this.MODEL_CLAUDE: {
return this.completePromptClaude(prompt)
}
case this.MODEL_GEMINI: {
return this.completePromptGemini(prompt)
}
default: {
throw new Error(`Invalid model type: ${this.modelType}`)
}
}
}
}

View file

@ -1,18 +1,19 @@
import { Anthropic } from "@anthropic-ai/sdk"
import * as vscode from "vscode"
import { ApiHandler, SingleCompletionHandler } from "../"
import { SingleCompletionHandler } from "../"
import { calculateApiCost } from "../../utils/cost"
import { ApiStream } from "../transform/stream"
import { convertToVsCodeLmMessages } from "../transform/vscode-lm-format"
import { SELECTOR_SEPARATOR, stringifyVsCodeLmModelSelector } from "../../shared/vsCodeSelectorUtils"
import { ApiHandlerOptions, ModelInfo, openAiModelInfoSaneDefaults } from "../../shared/api"
import { BaseProvider } from "./base-provider"
/**
* Handles interaction with VS Code's Language Model API for chat-based operations.
* This handler implements the ApiHandler interface to provide VS Code LM specific functionality.
* This handler extends BaseProvider to provide VS Code LM specific functionality.
*
* @implements {ApiHandler}
* @extends {BaseProvider}
*
* @remarks
* The handler manages a VS Code language model chat client and provides methods to:
@ -35,13 +36,14 @@ import { ApiHandlerOptions, ModelInfo, openAiModelInfoSaneDefaults } from "../..
* }
* ```
*/
export class VsCodeLmHandler implements ApiHandler, SingleCompletionHandler {
private options: ApiHandlerOptions
export class VsCodeLmHandler extends BaseProvider implements SingleCompletionHandler {
protected options: ApiHandlerOptions
private client: vscode.LanguageModelChat | null
private disposable: vscode.Disposable | null
private currentRequestCancellation: vscode.CancellationTokenSource | null
constructor(options: ApiHandlerOptions) {
super()
this.options = options
this.client = null
this.disposable = null
@ -145,7 +147,33 @@ export class VsCodeLmHandler implements ApiHandler, SingleCompletionHandler {
}
}
private async countTokens(text: string | vscode.LanguageModelChatMessage): Promise<number> {
/**
* Implements the ApiHandler countTokens interface method
* Provides token counting for Anthropic content blocks
*
* @param content The content blocks to count tokens for
* @returns A promise resolving to the token count
*/
override async countTokens(content: Array<Anthropic.Messages.ContentBlockParam>): Promise<number> {
// Convert Anthropic content blocks to a string for VSCode LM token counting
let textContent = ""
for (const block of content) {
if (block.type === "text") {
textContent += block.text || ""
} else if (block.type === "image") {
// VSCode LM doesn't support images directly, so we'll just use a placeholder
textContent += "[IMAGE]"
}
}
return this.internalCountTokens(textContent)
}
/**
* Private implementation of token counting used internally by VsCodeLmHandler
*/
private async internalCountTokens(text: string | vscode.LanguageModelChatMessage): Promise<number> {
// Check for required dependencies
if (!this.client) {
console.warn("Roo Code <Language Model API>: No client available for token counting")
@ -216,9 +244,9 @@ export class VsCodeLmHandler implements ApiHandler, SingleCompletionHandler {
systemPrompt: string,
vsCodeLmMessages: vscode.LanguageModelChatMessage[],
): Promise<number> {
const systemTokens: number = await this.countTokens(systemPrompt)
const systemTokens: number = await this.internalCountTokens(systemPrompt)
const messageTokens: number[] = await Promise.all(vsCodeLmMessages.map((msg) => this.countTokens(msg)))
const messageTokens: number[] = await Promise.all(vsCodeLmMessages.map((msg) => this.internalCountTokens(msg)))
return systemTokens + messageTokens.reduce((sum: number, tokens: number): number => sum + tokens, 0)
}
@ -319,7 +347,7 @@ export class VsCodeLmHandler implements ApiHandler, SingleCompletionHandler {
return content
}
async *createMessage(systemPrompt: string, messages: Anthropic.Messages.MessageParam[]): ApiStream {
override async *createMessage(systemPrompt: string, messages: Anthropic.Messages.MessageParam[]): ApiStream {
// Ensure clean state before starting a new request
this.ensureCleanState()
const client: vscode.LanguageModelChat = await this.getClient()
@ -427,7 +455,7 @@ export class VsCodeLmHandler implements ApiHandler, SingleCompletionHandler {
}
// Count tokens in the accumulated text after stream completion
const totalOutputTokens: number = await this.countTokens(accumulatedText)
const totalOutputTokens: number = await this.internalCountTokens(accumulatedText)
// Report final usage after stream completion
yield {
@ -467,7 +495,7 @@ export class VsCodeLmHandler implements ApiHandler, SingleCompletionHandler {
}
// Return model information based on the current client state
getModel(): { id: string; info: ModelInfo } {
override getModel(): { id: string; info: ModelInfo } {
if (this.client) {
// Validate client properties
const requiredProps = {

View file

@ -0,0 +1,338 @@
// npx jest src/api/transform/__tests__/vertex-gemini-format.test.ts
import { Anthropic } from "@anthropic-ai/sdk"
import { convertAnthropicMessageToVertexGemini } from "../vertex-gemini-format"
describe("convertAnthropicMessageToVertexGemini", () => {
it("should convert a simple text message", () => {
const anthropicMessage: Anthropic.Messages.MessageParam = {
role: "user",
content: "Hello, world!",
}
const result = convertAnthropicMessageToVertexGemini(anthropicMessage)
expect(result).toEqual({
role: "user",
parts: [{ text: "Hello, world!" }],
})
})
it("should convert assistant role to model role", () => {
const anthropicMessage: Anthropic.Messages.MessageParam = {
role: "assistant",
content: "I'm an assistant",
}
const result = convertAnthropicMessageToVertexGemini(anthropicMessage)
expect(result).toEqual({
role: "model",
parts: [{ text: "I'm an assistant" }],
})
})
it("should convert a message with text blocks", () => {
const anthropicMessage: Anthropic.Messages.MessageParam = {
role: "user",
content: [
{ type: "text", text: "First paragraph" },
{ type: "text", text: "Second paragraph" },
],
}
const result = convertAnthropicMessageToVertexGemini(anthropicMessage)
expect(result).toEqual({
role: "user",
parts: [{ text: "First paragraph" }, { text: "Second paragraph" }],
})
})
it("should convert a message with an image", () => {
const anthropicMessage: Anthropic.Messages.MessageParam = {
role: "user",
content: [
{ type: "text", text: "Check out this image:" },
{
type: "image",
source: {
type: "base64",
media_type: "image/jpeg",
data: "base64encodeddata",
},
},
],
}
const result = convertAnthropicMessageToVertexGemini(anthropicMessage)
expect(result).toEqual({
role: "user",
parts: [
{ text: "Check out this image:" },
{
inlineData: {
data: "base64encodeddata",
mimeType: "image/jpeg",
},
},
],
})
})
it("should throw an error for unsupported image source type", () => {
const anthropicMessage: Anthropic.Messages.MessageParam = {
role: "user",
content: [
{
type: "image",
source: {
type: "url", // Not supported
url: "https://example.com/image.jpg",
} as any,
},
],
}
expect(() => convertAnthropicMessageToVertexGemini(anthropicMessage)).toThrow("Unsupported image source type")
})
it("should convert a message with tool use", () => {
const anthropicMessage: Anthropic.Messages.MessageParam = {
role: "assistant",
content: [
{ type: "text", text: "Let me calculate that for you." },
{
type: "tool_use",
id: "calc-123",
name: "calculator",
input: { operation: "add", numbers: [2, 3] },
},
],
}
const result = convertAnthropicMessageToVertexGemini(anthropicMessage)
expect(result).toEqual({
role: "model",
parts: [
{ text: "Let me calculate that for you." },
{
functionCall: {
name: "calculator",
args: { operation: "add", numbers: [2, 3] },
},
},
],
})
})
it("should convert a message with tool result as string", () => {
const anthropicMessage: Anthropic.Messages.MessageParam = {
role: "user",
content: [
{ type: "text", text: "Here's the result:" },
{
type: "tool_result",
tool_use_id: "calculator-123",
content: "The result is 5",
},
],
}
const result = convertAnthropicMessageToVertexGemini(anthropicMessage)
expect(result).toEqual({
role: "user",
parts: [
{ text: "Here's the result:" },
{
functionResponse: {
name: "calculator",
response: {
name: "calculator",
content: "The result is 5",
},
},
},
],
})
})
it("should handle empty tool result content", () => {
const anthropicMessage: Anthropic.Messages.MessageParam = {
role: "user",
content: [
{
type: "tool_result",
tool_use_id: "calculator-123",
content: null as any, // Empty content
},
],
}
const result = convertAnthropicMessageToVertexGemini(anthropicMessage)
// Should skip the empty tool result
expect(result).toEqual({
role: "user",
parts: [],
})
})
it("should convert a message with tool result as array with text only", () => {
const anthropicMessage: Anthropic.Messages.MessageParam = {
role: "user",
content: [
{
type: "tool_result",
tool_use_id: "search-123",
content: [
{ type: "text", text: "First result" },
{ type: "text", text: "Second result" },
],
},
],
}
const result = convertAnthropicMessageToVertexGemini(anthropicMessage)
expect(result).toEqual({
role: "user",
parts: [
{
functionResponse: {
name: "search",
response: {
name: "search",
content: "First result\n\nSecond result",
},
},
},
],
})
})
it("should convert a message with tool result as array with text and images", () => {
const anthropicMessage: Anthropic.Messages.MessageParam = {
role: "user",
content: [
{
type: "tool_result",
tool_use_id: "search-123",
content: [
{ type: "text", text: "Search results:" },
{
type: "image",
source: {
type: "base64",
media_type: "image/png",
data: "image1data",
},
},
{
type: "image",
source: {
type: "base64",
media_type: "image/jpeg",
data: "image2data",
},
},
],
},
],
}
const result = convertAnthropicMessageToVertexGemini(anthropicMessage)
expect(result).toEqual({
role: "user",
parts: [
{
functionResponse: {
name: "search",
response: {
name: "search",
content: "Search results:\n\n(See next part for image)",
},
},
},
{
inlineData: {
data: "image1data",
mimeType: "image/png",
},
},
{
inlineData: {
data: "image2data",
mimeType: "image/jpeg",
},
},
],
})
})
it("should convert a message with tool result containing only images", () => {
const anthropicMessage: Anthropic.Messages.MessageParam = {
role: "user",
content: [
{
type: "tool_result",
tool_use_id: "imagesearch-123",
content: [
{
type: "image",
source: {
type: "base64",
media_type: "image/png",
data: "onlyimagedata",
},
},
],
},
],
}
const result = convertAnthropicMessageToVertexGemini(anthropicMessage)
expect(result).toEqual({
role: "user",
parts: [
{
functionResponse: {
name: "imagesearch",
response: {
name: "imagesearch",
content: "\n\n(See next part for image)",
},
},
},
{
inlineData: {
data: "onlyimagedata",
mimeType: "image/png",
},
},
],
})
})
it("should throw an error for unsupported content block type", () => {
const anthropicMessage: Anthropic.Messages.MessageParam = {
role: "user",
content: [
{
type: "unknown_type", // Unsupported type
data: "some data",
} as any,
],
}
expect(() => convertAnthropicMessageToVertexGemini(anthropicMessage)).toThrow(
"Unsupported content block type: unknown_type",
)
})
})

View file

@ -0,0 +1,83 @@
import { Anthropic } from "@anthropic-ai/sdk"
import { Content, FunctionCallPart, FunctionResponsePart, InlineDataPart, Part, TextPart } from "@google-cloud/vertexai"
function convertAnthropicContentToVertexGemini(content: Anthropic.Messages.MessageParam["content"]): Part[] {
if (typeof content === "string") {
return [{ text: content } as TextPart]
}
return content.flatMap((block) => {
switch (block.type) {
case "text":
return { text: block.text } as TextPart
case "image":
if (block.source.type !== "base64") {
throw new Error("Unsupported image source type")
}
return {
inlineData: {
data: block.source.data,
mimeType: block.source.media_type,
},
} as InlineDataPart
case "tool_use":
return {
functionCall: {
name: block.name,
args: block.input,
},
} as FunctionCallPart
case "tool_result":
const name = block.tool_use_id.split("-")[0]
if (!block.content) {
return []
}
if (typeof block.content === "string") {
return {
functionResponse: {
name,
response: {
name,
content: block.content,
},
},
} as FunctionResponsePart
} else {
// The only case when tool_result could be array is when the tool failed and we're providing ie user feedback potentially with images
const textParts = block.content.filter((part) => part.type === "text")
const imageParts = block.content.filter((part) => part.type === "image")
const text = textParts.length > 0 ? textParts.map((part) => part.text).join("\n\n") : ""
const imageText = imageParts.length > 0 ? "\n\n(See next part for image)" : ""
return [
{
functionResponse: {
name,
response: {
name,
content: text + imageText,
},
},
} as FunctionResponsePart,
...imageParts.map(
(part) =>
({
inlineData: {
data: part.source.data,
mimeType: part.source.media_type,
},
}) as InlineDataPart,
),
]
}
default:
throw new Error(`Unsupported content block type: ${(block as any).type}`)
}
})
}
export function convertAnthropicMessageToVertexGemini(message: Anthropic.Messages.MessageParam): Content {
return {
role: message.role === "assistant" ? "model" : "user",
parts: convertAnthropicContentToVertexGemini(message.content),
}
}

View file

@ -22,7 +22,7 @@ import {
everyLineHasLineNumbers,
truncateOutput,
} from "../integrations/misc/extract-text"
import { TerminalManager } from "../integrations/terminal/TerminalManager"
import { TerminalManager, ExitCodeDetails } from "../integrations/terminal/TerminalManager"
import { UrlContentFetcher } from "../services/browser/UrlContentFetcher"
import { listFiles } from "../services/glob/list-files"
import { regexSearchFiles } from "../services/ripgrep"
@ -148,7 +148,8 @@ export class Cline {
throw new Error("Either historyItem or task/images must be provided")
}
this.taskId = crypto.randomUUID()
this.taskId = historyItem ? historyItem.id : crypto.randomUUID()
this.apiConfiguration = apiConfiguration
this.api = buildApiHandler(apiConfiguration)
this.terminalManager = new TerminalManager()
@ -161,10 +162,6 @@ export class Cline {
this.diffViewProvider = new DiffViewProvider(cwd)
this.enableCheckpoints = enableCheckpoints ?? false
if (historyItem) {
this.taskId = historyItem.id
}
// Initialize diffStrategy based on current state
this.updateDiffStrategy(Experiments.isEnabled(experiments ?? {}, EXPERIMENT_IDS.DIFF_STRATEGY))
@ -834,10 +831,21 @@ export class Cline {
})
let completed = false
process.once("completed", () => {
let exitDetails: ExitCodeDetails | undefined
process.once("completed", (output?: string) => {
// Use provided output if available, otherwise keep existing result.
if (output) {
lines = output.split("\n")
}
completed = true
})
process.once("shell_execution_complete", (id: number, details: ExitCodeDetails) => {
if (id === terminalInfo.id) {
exitDetails = details
}
})
process.once("no_shell_integration", async () => {
await this.say("shell_integration_warning")
})
@ -869,7 +877,18 @@ export class Cline {
}
if (completed) {
return [false, `Command executed.${result.length > 0 ? `\nOutput:\n${result}` : ""}`]
let exitStatus = "No exit code available"
if (exitDetails !== undefined) {
if (exitDetails.signal) {
exitStatus = `Process terminated by signal ${exitDetails.signal} (${exitDetails.signalName})`
if (exitDetails.coreDumpPossible) {
exitStatus += " - core dump possible"
}
} else {
exitStatus = `Exit code: ${exitDetails.exitCode}`
}
}
return [false, `Command executed. ${exitStatus}${result.length > 0 ? `\nOutput:\n${result}` : ""}`]
} else {
return [
false,
@ -971,12 +990,12 @@ export class Cline {
? this.apiConfiguration.modelMaxTokens || modelInfo.maxTokens
: modelInfo.maxTokens
const contextWindow = modelInfo.contextWindow
const trimmedMessages = truncateConversationIfNeeded({
const trimmedMessages = await truncateConversationIfNeeded({
messages: this.apiConversationHistory,
totalTokens,
maxTokens,
contextWindow,
apiHandler: this.api,
})
if (trimmedMessages !== this.apiConversationHistory) {
@ -3315,7 +3334,7 @@ export class Cline {
) {
const currentModeName = getModeBySlug(currentMode, customModes)?.name ?? currentMode
const defaultModeName = getModeBySlug(defaultModeSlug, customModes)?.name ?? defaultModeSlug
details += `\n\nNOTE: You are currently in '${currentModeName}' mode which only allows read-only operations. To write files or execute commands, the user will need to switch to '${defaultModeName}' mode. Note that only the user can switch modes.`
details += `\n\nNOTE: You are currently in '${currentModeName}' mode, which does not allow write operations. To write files, the user will need to switch to a mode that supports file writing, such as '${defaultModeName}' mode.`
}
if (includeFileDetails) {

View file

@ -3899,9 +3899,17 @@ USER'S CUSTOM INSTRUCTIONS
The following additional instructions are provided by the user, and should be followed to the best of your ability without interfering with the TOOL USE guidelines.
Mode-specific Instructions:
Depending on the user's request, you may need to do some information gathering (for example using read_file or search_files) to get more context about the task. You may also ask the user clarifying questions to get a better understanding of the task. Once you've gained more context about the user's request, you should create a detailed plan for how to accomplish the task. (You can write the plan to a markdown file if it seems appropriate.)
1. Do some information gathering (for example using read_file or search_files) to get more context about the task.
Then you might ask the user if they are pleased with this plan, or if they would like to make any changes. Think of this as a brainstorming session where you can discuss the task and plan the best way to accomplish it. Finally once it seems like you've reached a good plan, use the switch_mode tool to request that the user switch to another mode to implement the solution.
2. You should also ask the user clarifying questions to get a better understanding of the task.
3. Once you've gained more context about the user's request, you should create a detailed plan for how to accomplish the task. Include Mermaid diagrams if they help make your plan clearer.
4. Ask the user if they are pleased with this plan, or if they would like to make any changes. Think of this as a brainstorming session where you can discuss the task and plan the best way to accomplish it.
5. Once the user confirms the plan, ask them if they'd like you to write it to a markdown file.
6. Use the switch_mode tool to request that the user switch to another mode to implement the solution.
Rules:
# Rules from .clinerules-architect:
@ -4176,7 +4184,7 @@ USER'S CUSTOM INSTRUCTIONS
The following additional instructions are provided by the user, and should be followed to the best of your ability without interfering with the TOOL USE guidelines.
Mode-specific Instructions:
You can analyze code, explain concepts, and access external resources. Make sure to answer the user's questions and don't rush to switch to implementing code.
You can analyze code, explain concepts, and access external resources. Make sure to answer the user's questions and don't rush to switch to implementing code. Include Mermaid diagrams if they help make your response clearer.
Rules:
# Rules from .clinerules-ask:

View file

@ -3,7 +3,35 @@
import { Anthropic } from "@anthropic-ai/sdk"
import { ModelInfo } from "../../../shared/api"
import { truncateConversation, truncateConversationIfNeeded } from "../index"
import { ApiHandler } from "../../../api"
import { BaseProvider } from "../../../api/providers/base-provider"
import { TOKEN_BUFFER_PERCENTAGE } from "../index"
import { estimateTokenCount, truncateConversation, truncateConversationIfNeeded } from "../index"
// Create a mock ApiHandler for testing
class MockApiHandler extends BaseProvider {
createMessage(): any {
throw new Error("Method not implemented.")
}
getModel(): { id: string; info: ModelInfo } {
return {
id: "test-model",
info: {
contextWindow: 100000,
maxTokens: 50000,
supportsPromptCache: true,
supportsImages: false,
inputPrice: 0,
outputPrice: 0,
description: "Test model",
},
}
}
}
// Create a singleton instance for tests
const mockApiHandler = new MockApiHandler()
/**
* Tests for the truncateConversation function
@ -94,6 +122,296 @@ describe("truncateConversation", () => {
})
})
/**
* Tests for the estimateTokenCount function
*/
describe("estimateTokenCount", () => {
it("should return 0 for empty or undefined content", async () => {
expect(await estimateTokenCount([], mockApiHandler)).toBe(0)
// @ts-ignore - Testing with undefined
expect(await estimateTokenCount(undefined, mockApiHandler)).toBe(0)
})
it("should estimate tokens for text blocks", async () => {
const content: Array<Anthropic.Messages.ContentBlockParam> = [
{ type: "text", text: "This is a text block with 36 characters" },
]
// With tiktoken, the exact token count may differ from character-based estimation
// Instead of expecting an exact number, we verify it's a reasonable positive number
const result = await estimateTokenCount(content, mockApiHandler)
expect(result).toBeGreaterThan(0)
// We can also verify that longer text results in more tokens
const longerContent: Array<Anthropic.Messages.ContentBlockParam> = [
{
type: "text",
text: "This is a longer text block with significantly more characters to encode into tokens",
},
]
const longerResult = await estimateTokenCount(longerContent, mockApiHandler)
expect(longerResult).toBeGreaterThan(result)
})
it("should estimate tokens for image blocks based on data size", async () => {
// Small image
const smallImage: Array<Anthropic.Messages.ContentBlockParam> = [
{ type: "image", source: { type: "base64", media_type: "image/jpeg", data: "small_dummy_data" } },
]
// Larger image with more data
const largerImage: Array<Anthropic.Messages.ContentBlockParam> = [
{ type: "image", source: { type: "base64", media_type: "image/png", data: "X".repeat(1000) } },
]
// Verify the token count scales with the size of the image data
const smallImageTokens = await estimateTokenCount(smallImage, mockApiHandler)
const largerImageTokens = await estimateTokenCount(largerImage, mockApiHandler)
// Small image should have some tokens
expect(smallImageTokens).toBeGreaterThan(0)
// Larger image should have proportionally more tokens
expect(largerImageTokens).toBeGreaterThan(smallImageTokens)
// Verify the larger image calculation matches our formula including the 50% fudge factor
expect(largerImageTokens).toBe(48)
})
it("should estimate tokens for mixed content blocks", async () => {
const content: Array<Anthropic.Messages.ContentBlockParam> = [
{ type: "text", text: "A text block with 30 characters" },
{ type: "image", source: { type: "base64", media_type: "image/jpeg", data: "dummy_data" } },
{ type: "text", text: "Another text with 24 chars" },
]
// We know image tokens calculation should be consistent
const imageTokens = Math.ceil(Math.sqrt("dummy_data".length)) * 1.5
// With tiktoken, we can't predict exact text token counts,
// but we can verify the total is greater than just the image tokens
const result = await estimateTokenCount(content, mockApiHandler)
expect(result).toBeGreaterThan(imageTokens)
// Also test against a version with only the image to verify text adds tokens
const imageOnlyContent: Array<Anthropic.Messages.ContentBlockParam> = [
{ type: "image", source: { type: "base64", media_type: "image/jpeg", data: "dummy_data" } },
]
const imageOnlyResult = await estimateTokenCount(imageOnlyContent, mockApiHandler)
expect(result).toBeGreaterThan(imageOnlyResult)
})
it("should handle empty text blocks", async () => {
const content: Array<Anthropic.Messages.ContentBlockParam> = [{ type: "text", text: "" }]
expect(await estimateTokenCount(content, mockApiHandler)).toBe(0)
})
it("should handle plain string messages", async () => {
const content = "This is a plain text message"
expect(await estimateTokenCount([{ type: "text", text: content }], mockApiHandler)).toBeGreaterThan(0)
})
})
/**
* Tests for the truncateConversationIfNeeded function
*/
describe("truncateConversationIfNeeded", () => {
const createModelInfo = (contextWindow: number, maxTokens?: number): ModelInfo => ({
contextWindow,
supportsPromptCache: true,
maxTokens,
})
const messages: Anthropic.Messages.MessageParam[] = [
{ role: "user", content: "First message" },
{ role: "assistant", content: "Second message" },
{ role: "user", content: "Third message" },
{ role: "assistant", content: "Fourth message" },
{ role: "user", content: "Fifth message" },
]
it("should not truncate if tokens are below max tokens threshold", async () => {
const modelInfo = createModelInfo(100000, 30000)
const maxTokens = 100000 - 30000 // 70000
const dynamicBuffer = modelInfo.contextWindow * TOKEN_BUFFER_PERCENTAGE // 10000
const totalTokens = 70000 - dynamicBuffer - 1 // Just below threshold - buffer
// Create messages with very small content in the last one to avoid token overflow
const messagesWithSmallContent = [...messages.slice(0, -1), { ...messages[messages.length - 1], content: "" }]
const result = await truncateConversationIfNeeded({
messages: messagesWithSmallContent,
totalTokens,
contextWindow: modelInfo.contextWindow,
maxTokens: modelInfo.maxTokens,
apiHandler: mockApiHandler,
})
expect(result).toEqual(messagesWithSmallContent) // No truncation occurs
})
it("should truncate if tokens are above max tokens threshold", async () => {
const modelInfo = createModelInfo(100000, 30000)
const maxTokens = 100000 - 30000 // 70000
const totalTokens = 70001 // Above threshold
// Create messages with very small content in the last one to avoid token overflow
const messagesWithSmallContent = [...messages.slice(0, -1), { ...messages[messages.length - 1], content: "" }]
// When truncating, always uses 0.5 fraction
// With 4 messages after the first, 0.5 fraction means remove 2 messages
const expectedResult = [messagesWithSmallContent[0], messagesWithSmallContent[3], messagesWithSmallContent[4]]
const result = await truncateConversationIfNeeded({
messages: messagesWithSmallContent,
totalTokens,
contextWindow: modelInfo.contextWindow,
maxTokens: modelInfo.maxTokens,
apiHandler: mockApiHandler,
})
expect(result).toEqual(expectedResult)
})
it("should work with non-prompt caching models the same as prompt caching models", async () => {
// The implementation no longer differentiates between prompt caching and non-prompt caching models
const modelInfo1 = createModelInfo(100000, 30000)
const modelInfo2 = createModelInfo(100000, 30000)
// Create messages with very small content in the last one to avoid token overflow
const messagesWithSmallContent = [...messages.slice(0, -1), { ...messages[messages.length - 1], content: "" }]
// Test below threshold
const belowThreshold = 69999
const result1 = await truncateConversationIfNeeded({
messages: messagesWithSmallContent,
totalTokens: belowThreshold,
contextWindow: modelInfo1.contextWindow,
maxTokens: modelInfo1.maxTokens,
apiHandler: mockApiHandler,
})
const result2 = await truncateConversationIfNeeded({
messages: messagesWithSmallContent,
totalTokens: belowThreshold,
contextWindow: modelInfo2.contextWindow,
maxTokens: modelInfo2.maxTokens,
apiHandler: mockApiHandler,
})
expect(result1).toEqual(result2)
// Test above threshold
const aboveThreshold = 70001
const result3 = await truncateConversationIfNeeded({
messages: messagesWithSmallContent,
totalTokens: aboveThreshold,
contextWindow: modelInfo1.contextWindow,
maxTokens: modelInfo1.maxTokens,
apiHandler: mockApiHandler,
})
const result4 = await truncateConversationIfNeeded({
messages: messagesWithSmallContent,
totalTokens: aboveThreshold,
contextWindow: modelInfo2.contextWindow,
maxTokens: modelInfo2.maxTokens,
apiHandler: mockApiHandler,
})
expect(result3).toEqual(result4)
})
it("should consider incoming content when deciding to truncate", async () => {
const modelInfo = createModelInfo(100000, 30000)
const maxTokens = 30000
const availableTokens = modelInfo.contextWindow - maxTokens
// Test case 1: Small content that won't push us over the threshold
const smallContent = [{ type: "text" as const, text: "Small content" }]
const smallContentTokens = await estimateTokenCount(smallContent, mockApiHandler)
const messagesWithSmallContent: Anthropic.Messages.MessageParam[] = [
...messages.slice(0, -1),
{ role: messages[messages.length - 1].role, content: smallContent },
]
// Set base tokens so total is well below threshold + buffer even with small content added
const dynamicBuffer = modelInfo.contextWindow * TOKEN_BUFFER_PERCENTAGE
const baseTokensForSmall = availableTokens - smallContentTokens - dynamicBuffer - 10
const resultWithSmall = await truncateConversationIfNeeded({
messages: messagesWithSmallContent,
totalTokens: baseTokensForSmall,
contextWindow: modelInfo.contextWindow,
maxTokens,
apiHandler: mockApiHandler,
})
expect(resultWithSmall).toEqual(messagesWithSmallContent) // No truncation
// Test case 2: Large content that will push us over the threshold
const largeContent = [
{
type: "text" as const,
text: "A very large incoming message that would consume a significant number of tokens and push us over the threshold",
},
]
const largeContentTokens = await estimateTokenCount(largeContent, mockApiHandler)
const messagesWithLargeContent: Anthropic.Messages.MessageParam[] = [
...messages.slice(0, -1),
{ role: messages[messages.length - 1].role, content: largeContent },
]
// Set base tokens so we're just below threshold without content, but over with content
const baseTokensForLarge = availableTokens - Math.floor(largeContentTokens / 2)
const resultWithLarge = await truncateConversationIfNeeded({
messages: messagesWithLargeContent,
totalTokens: baseTokensForLarge,
contextWindow: modelInfo.contextWindow,
maxTokens,
apiHandler: mockApiHandler,
})
expect(resultWithLarge).not.toEqual(messagesWithLargeContent) // Should truncate
// Test case 3: Very large content that will definitely exceed threshold
const veryLargeContent = [{ type: "text" as const, text: "X".repeat(1000) }]
const veryLargeContentTokens = await estimateTokenCount(veryLargeContent, mockApiHandler)
const messagesWithVeryLargeContent: Anthropic.Messages.MessageParam[] = [
...messages.slice(0, -1),
{ role: messages[messages.length - 1].role, content: veryLargeContent },
]
// Set base tokens so we're just below threshold without content
const baseTokensForVeryLarge = availableTokens - Math.floor(veryLargeContentTokens / 2)
const resultWithVeryLarge = await truncateConversationIfNeeded({
messages: messagesWithVeryLargeContent,
totalTokens: baseTokensForVeryLarge,
contextWindow: modelInfo.contextWindow,
maxTokens,
apiHandler: mockApiHandler,
})
expect(resultWithVeryLarge).not.toEqual(messagesWithVeryLargeContent) // Should truncate
})
it("should truncate if tokens are within TOKEN_BUFFER_PERCENTAGE of the threshold", async () => {
const modelInfo = createModelInfo(100000, 30000)
const maxTokens = 100000 - 30000 // 70000
const dynamicBuffer = modelInfo.contextWindow * TOKEN_BUFFER_PERCENTAGE // 10% of 100000 = 10000
const totalTokens = 70000 - dynamicBuffer + 1 // Just within the dynamic buffer of threshold (70000)
// Create messages with very small content in the last one to avoid token overflow
const messagesWithSmallContent = [...messages.slice(0, -1), { ...messages[messages.length - 1], content: "" }]
// When truncating, always uses 0.5 fraction
// With 4 messages after the first, 0.5 fraction means remove 2 messages
const expectedResult = [messagesWithSmallContent[0], messagesWithSmallContent[3], messagesWithSmallContent[4]]
const result = await truncateConversationIfNeeded({
messages: messagesWithSmallContent,
totalTokens,
contextWindow: modelInfo.contextWindow,
maxTokens: modelInfo.maxTokens,
apiHandler: mockApiHandler,
})
expect(result).toEqual(expectedResult)
})
})
/**
* Tests for the getMaxTokens function (private but tested through truncateConversationIfNeeded)
*/
@ -114,192 +432,122 @@ describe("getMaxTokens", () => {
{ role: "user", content: "Fifth message" },
]
it("should use maxTokens as buffer when specified", () => {
it("should use maxTokens as buffer when specified", async () => {
const modelInfo = createModelInfo(100000, 50000)
// Max tokens = 100000 - 50000 = 50000
// Below max tokens - no truncation
const result1 = truncateConversationIfNeeded({
messages,
totalTokens: 49999,
// Create messages with very small content in the last one to avoid token overflow
const messagesWithSmallContent = [...messages.slice(0, -1), { ...messages[messages.length - 1], content: "" }]
// Account for the dynamic buffer which is 10% of context window (10,000 tokens)
// Below max tokens and buffer - no truncation
const result1 = await truncateConversationIfNeeded({
messages: messagesWithSmallContent,
totalTokens: 39999, // Well below threshold + dynamic buffer
contextWindow: modelInfo.contextWindow,
maxTokens: modelInfo.maxTokens,
apiHandler: mockApiHandler,
})
expect(result1).toEqual(messages)
expect(result1).toEqual(messagesWithSmallContent)
// Above max tokens - truncate
const result2 = truncateConversationIfNeeded({
messages,
totalTokens: 50001,
const result2 = await truncateConversationIfNeeded({
messages: messagesWithSmallContent,
totalTokens: 50001, // Above threshold
contextWindow: modelInfo.contextWindow,
maxTokens: modelInfo.maxTokens,
apiHandler: mockApiHandler,
})
expect(result2).not.toEqual(messages)
expect(result2).not.toEqual(messagesWithSmallContent)
expect(result2.length).toBe(3) // Truncated with 0.5 fraction
})
it("should use 20% of context window as buffer when maxTokens is undefined", () => {
it("should use 20% of context window as buffer when maxTokens is undefined", async () => {
const modelInfo = createModelInfo(100000, undefined)
// Max tokens = 100000 - (100000 * 0.2) = 80000
// Below max tokens - no truncation
const result1 = truncateConversationIfNeeded({
messages,
totalTokens: 79999,
// Create messages with very small content in the last one to avoid token overflow
const messagesWithSmallContent = [...messages.slice(0, -1), { ...messages[messages.length - 1], content: "" }]
// Account for the dynamic buffer which is 10% of context window (10,000 tokens)
// Below max tokens and buffer - no truncation
const result1 = await truncateConversationIfNeeded({
messages: messagesWithSmallContent,
totalTokens: 69999, // Well below threshold + dynamic buffer
contextWindow: modelInfo.contextWindow,
maxTokens: modelInfo.maxTokens,
apiHandler: mockApiHandler,
})
expect(result1).toEqual(messages)
expect(result1).toEqual(messagesWithSmallContent)
// Above max tokens - truncate
const result2 = truncateConversationIfNeeded({
messages,
totalTokens: 80001,
const result2 = await truncateConversationIfNeeded({
messages: messagesWithSmallContent,
totalTokens: 80001, // Above threshold
contextWindow: modelInfo.contextWindow,
maxTokens: modelInfo.maxTokens,
apiHandler: mockApiHandler,
})
expect(result2).not.toEqual(messages)
expect(result2).not.toEqual(messagesWithSmallContent)
expect(result2.length).toBe(3) // Truncated with 0.5 fraction
})
it("should handle small context windows appropriately", () => {
it("should handle small context windows appropriately", async () => {
const modelInfo = createModelInfo(50000, 10000)
// Max tokens = 50000 - 10000 = 40000
// Below max tokens - no truncation
const result1 = truncateConversationIfNeeded({
messages,
totalTokens: 39999,
// Create messages with very small content in the last one to avoid token overflow
const messagesWithSmallContent = [...messages.slice(0, -1), { ...messages[messages.length - 1], content: "" }]
// Below max tokens and buffer - no truncation
const result1 = await truncateConversationIfNeeded({
messages: messagesWithSmallContent,
totalTokens: 34999, // Well below threshold + buffer
contextWindow: modelInfo.contextWindow,
maxTokens: modelInfo.maxTokens,
apiHandler: mockApiHandler,
})
expect(result1).toEqual(messages)
expect(result1).toEqual(messagesWithSmallContent)
// Above max tokens - truncate
const result2 = truncateConversationIfNeeded({
messages,
totalTokens: 40001,
const result2 = await truncateConversationIfNeeded({
messages: messagesWithSmallContent,
totalTokens: 40001, // Above threshold
contextWindow: modelInfo.contextWindow,
maxTokens: modelInfo.maxTokens,
apiHandler: mockApiHandler,
})
expect(result2).not.toEqual(messages)
expect(result2).not.toEqual(messagesWithSmallContent)
expect(result2.length).toBe(3) // Truncated with 0.5 fraction
})
it("should handle large context windows appropriately", () => {
it("should handle large context windows appropriately", async () => {
const modelInfo = createModelInfo(200000, 30000)
// Max tokens = 200000 - 30000 = 170000
// Below max tokens - no truncation
const result1 = truncateConversationIfNeeded({
messages,
totalTokens: 169999,
// Create messages with very small content in the last one to avoid token overflow
const messagesWithSmallContent = [...messages.slice(0, -1), { ...messages[messages.length - 1], content: "" }]
// Account for the dynamic buffer which is 10% of context window (20,000 tokens for this test)
// Below max tokens and buffer - no truncation
const result1 = await truncateConversationIfNeeded({
messages: messagesWithSmallContent,
totalTokens: 149999, // Well below threshold + dynamic buffer
contextWindow: modelInfo.contextWindow,
maxTokens: modelInfo.maxTokens,
apiHandler: mockApiHandler,
})
expect(result1).toEqual(messages)
expect(result1).toEqual(messagesWithSmallContent)
// Above max tokens - truncate
const result2 = truncateConversationIfNeeded({
messages,
totalTokens: 170001,
const result2 = await truncateConversationIfNeeded({
messages: messagesWithSmallContent,
totalTokens: 170001, // Above threshold
contextWindow: modelInfo.contextWindow,
maxTokens: modelInfo.maxTokens,
apiHandler: mockApiHandler,
})
expect(result2).not.toEqual(messages)
expect(result2).not.toEqual(messagesWithSmallContent)
expect(result2.length).toBe(3) // Truncated with 0.5 fraction
})
})
/**
* Tests for the truncateConversationIfNeeded function
*/
describe("truncateConversationIfNeeded", () => {
const createModelInfo = (contextWindow: number, supportsPromptCache: boolean, maxTokens?: number): ModelInfo => ({
contextWindow,
supportsPromptCache,
maxTokens,
})
const messages: Anthropic.Messages.MessageParam[] = [
{ role: "user", content: "First message" },
{ role: "assistant", content: "Second message" },
{ role: "user", content: "Third message" },
{ role: "assistant", content: "Fourth message" },
{ role: "user", content: "Fifth message" },
]
it("should not truncate if tokens are below max tokens threshold", () => {
const modelInfo = createModelInfo(100000, true, 30000)
const maxTokens = 100000 - 30000 // 70000
const totalTokens = 69999 // Below threshold
const result = truncateConversationIfNeeded({
messages,
totalTokens,
contextWindow: modelInfo.contextWindow,
maxTokens: modelInfo.maxTokens,
})
expect(result).toEqual(messages) // No truncation occurs
})
it("should truncate if tokens are above max tokens threshold", () => {
const modelInfo = createModelInfo(100000, true, 30000)
const maxTokens = 100000 - 30000 // 70000
const totalTokens = 70001 // Above threshold
// When truncating, always uses 0.5 fraction
// With 4 messages after the first, 0.5 fraction means remove 2 messages
const expectedResult = [messages[0], messages[3], messages[4]]
const result = truncateConversationIfNeeded({
messages,
totalTokens,
contextWindow: modelInfo.contextWindow,
maxTokens: modelInfo.maxTokens,
})
expect(result).toEqual(expectedResult)
})
it("should work with non-prompt caching models the same as prompt caching models", () => {
// The implementation no longer differentiates between prompt caching and non-prompt caching models
const modelInfo1 = createModelInfo(100000, true, 30000)
const modelInfo2 = createModelInfo(100000, false, 30000)
// Test below threshold
const belowThreshold = 69999
expect(
truncateConversationIfNeeded({
messages,
totalTokens: belowThreshold,
contextWindow: modelInfo1.contextWindow,
maxTokens: modelInfo1.maxTokens,
}),
).toEqual(
truncateConversationIfNeeded({
messages,
totalTokens: belowThreshold,
contextWindow: modelInfo2.contextWindow,
maxTokens: modelInfo2.maxTokens,
}),
)
// Test above threshold
const aboveThreshold = 70001
expect(
truncateConversationIfNeeded({
messages,
totalTokens: aboveThreshold,
contextWindow: modelInfo1.contextWindow,
maxTokens: modelInfo1.maxTokens,
}),
).toEqual(
truncateConversationIfNeeded({
messages,
totalTokens: aboveThreshold,
contextWindow: modelInfo2.contextWindow,
maxTokens: modelInfo2.maxTokens,
}),
)
})
})

View file

@ -1,4 +1,25 @@
import { Anthropic } from "@anthropic-ai/sdk"
import { ApiHandler } from "../../api"
/**
* Default percentage of the context window to use as a buffer when deciding when to truncate
*/
export const TOKEN_BUFFER_PERCENTAGE = 0.1
/**
* Counts tokens for user content using the provider's token counting implementation.
*
* @param {Array<Anthropic.Messages.ContentBlockParam>} content - The content to count tokens for
* @param {ApiHandler} apiHandler - The API handler to use for token counting
* @returns {Promise<number>} A promise resolving to the token count
*/
export async function estimateTokenCount(
content: Array<Anthropic.Messages.ContentBlockParam>,
apiHandler: ApiHandler,
): Promise<number> {
if (!content || content.length === 0) return 0
return apiHandler.countTokens(content)
}
/**
* Truncates a conversation by removing a fraction of the messages.
@ -25,12 +46,13 @@ export function truncateConversation(
/**
* Conditionally truncates the conversation messages if the total token count
* exceeds the model's limit.
* exceeds the model's limit, considering the size of incoming content.
*
* @param {Anthropic.Messages.MessageParam[]} messages - The conversation messages.
* @param {number} totalTokens - The total number of tokens in the conversation.
* @param {number} totalTokens - The total number of tokens in the conversation (excluding the last user message).
* @param {number} contextWindow - The context window size.
* @param {number} maxTokens - The maximum number of tokens allowed.
* @param {ApiHandler} apiHandler - The API handler to use for token counting.
* @returns {Anthropic.Messages.MessageParam[]} The original or truncated conversation messages.
*/
@ -39,14 +61,40 @@ type TruncateOptions = {
totalTokens: number
contextWindow: number
maxTokens?: number
apiHandler: ApiHandler
}
export function truncateConversationIfNeeded({
/**
* Conditionally truncates the conversation messages if the total token count
* exceeds the model's limit, considering the size of incoming content.
*
* @param {TruncateOptions} options - The options for truncation
* @returns {Promise<Anthropic.Messages.MessageParam[]>} The original or truncated conversation messages.
*/
export async function truncateConversationIfNeeded({
messages,
totalTokens,
contextWindow,
maxTokens,
}: TruncateOptions): Anthropic.Messages.MessageParam[] {
const allowedTokens = contextWindow - (maxTokens || contextWindow * 0.2)
return totalTokens < allowedTokens ? messages : truncateConversation(messages, 0.5)
apiHandler,
}: TruncateOptions): Promise<Anthropic.Messages.MessageParam[]> {
// Calculate the maximum tokens reserved for response
const reservedTokens = maxTokens || contextWindow * 0.2
// Estimate tokens for the last message (which is always a user message)
const lastMessage = messages[messages.length - 1]
const lastMessageContent = lastMessage.content
const lastMessageTokens = Array.isArray(lastMessageContent)
? await estimateTokenCount(lastMessageContent, apiHandler)
: await estimateTokenCount([{ type: "text", text: lastMessageContent as string }], apiHandler)
// Calculate total effective tokens (totalTokens never includes the last message)
const effectiveTokens = totalTokens + lastMessageTokens
// Calculate available tokens for conversation history
// Truncate if we're within TOKEN_BUFFER_PERCENTAGE of the context window
const allowedTokens = contextWindow * (1 - TOKEN_BUFFER_PERCENTAGE) - reservedTokens
// Determine if truncation is needed and apply if necessary
return effectiveTokens > allowedTokens ? truncateConversation(messages, 0.5) : messages
}

View file

@ -519,7 +519,7 @@ export class ClineProvider implements vscode.WebviewViewProvider {
<body>
<noscript>You need to enable JavaScript to run this app.</noscript>
<div id="root"></div>
<script nonce="${nonce}" src="${scriptUri}"></script>
<script nonce="${nonce}" type="module" src="${scriptUri}"></script>
</body>
</html>
`
@ -2459,6 +2459,7 @@ export class ClineProvider implements vscode.WebviewViewProvider {
autoApprovalEnabled: autoApprovalEnabled ?? false,
customModes,
maxOpenTabsContext: maxOpenTabsContext ?? 20,
openRouterUseMiddleOutTransform: openRouterUseMiddleOutTransform ?? true,
}
}

View file

@ -40,3 +40,96 @@ export interface ClineAPI {
*/
sidebarProvider: ClineSidebarProvider
}
export interface ClineProvider {
readonly context: vscode.ExtensionContext
readonly viewLaunched: boolean
readonly messages: ClineMessage[]
/**
* Resolves the webview view for the provider
* @param webviewView The webview view or panel to resolve
*/
resolveWebviewView(webviewView: vscode.WebviewView | vscode.WebviewPanel): Promise<void>
/**
* Initializes Cline with a task
*/
initClineWithTask(task?: string, images?: string[]): Promise<void>
/**
* Initializes Cline with a history item
*/
initClineWithHistoryItem(historyItem: HistoryItem): Promise<void>
/**
* Posts a message to the webview
*/
postMessageToWebview(message: ExtensionMessage): Promise<void>
/**
* Handles mode switching
*/
handleModeSwitch(newMode: Mode): Promise<void>
/**
* Updates custom instructions
*/
updateCustomInstructions(instructions?: string): Promise<void>
/**
* Cancels the current task
*/
cancelTask(): Promise<void>
/**
* Clears the current task
*/
clearTask(): Promise<void>
/**
* Gets the current state
*/
getState(): Promise<any>
/**
* Updates a value in the global state
* @param key The key to update
* @param value The value to set
*/
updateGlobalState(key: GlobalStateKey, value: any): Promise<void>
/**
* Gets a value from the global state
* @param key The key to get
*/
getGlobalState(key: GlobalStateKey): Promise<any>
/**
* Stores a secret value in secure storage
* @param key The key to store the secret under
* @param value The secret value to store, or undefined to remove the secret
*/
storeSecret(key: SecretKey, value?: string): Promise<void>
/**
* Retrieves a secret value from secure storage
* @param key The key of the secret to retrieve
*/
getSecret(key: SecretKey): Promise<string | undefined>
/**
* Resets the state
*/
resetState(): Promise<void>
/**
* Logs a message
*/
log(message: string): void
/**
* Disposes of the provider
*/
dispose(): Promise<void>
}

View file

@ -70,6 +70,15 @@ Interestingly, some environments like Cursor enable these APIs even without the
This approach allows us to leverage advanced features when available while ensuring broad compatibility.
*/
declare module "vscode" {
// https://github.com/microsoft/vscode/blob/f0417069c62e20f3667506f4b7e53ca0004b4e3e/src/vscode-dts/vscode.d.ts#L7442
// interface Terminal {
// shellIntegration?: {
// cwd?: vscode.Uri
// executeCommand?: (command: string) => {
// read: () => AsyncIterable<string>
// }
// }
// }
// https://github.com/microsoft/vscode/blob/f0417069c62e20f3667506f4b7e53ca0004b4e3e/src/vscode-dts/vscode.d.ts#L10794
interface Window {
onDidStartTerminalShellExecution?: (
@ -77,17 +86,19 @@ declare module "vscode" {
thisArgs?: any,
disposables?: vscode.Disposable[],
) => vscode.Disposable
onDidEndTerminalShellExecution?: (
listener: (e: { terminal: vscode.Terminal; exitCode?: number; shellType?: string }) => any,
thisArgs?: any,
disposables?: vscode.Disposable[],
) => vscode.Disposable
}
}
// Extend the Terminal type to include our custom properties
type ExtendedTerminal = vscode.Terminal & {
shellIntegration?: {
cwd?: vscode.Uri
executeCommand?: (command: string) => {
read: () => AsyncIterable<string>
}
}
export interface ExitCodeDetails {
exitCode: number | undefined
signal?: number | undefined
signalName?: string
coreDumpPossible?: boolean
}
export class TerminalManager {
@ -95,18 +106,156 @@ export class TerminalManager {
private processes: Map<number, TerminalProcess> = new Map()
private disposables: vscode.Disposable[] = []
private interpretExitCode(exitCode: number | undefined): ExitCodeDetails {
if (exitCode === undefined) {
return { exitCode }
}
if (exitCode <= 128) {
return { exitCode }
}
const signal = exitCode - 128
const signals: Record<number, string> = {
// Standard signals
1: "SIGHUP",
2: "SIGINT",
3: "SIGQUIT",
4: "SIGILL",
5: "SIGTRAP",
6: "SIGABRT",
7: "SIGBUS",
8: "SIGFPE",
9: "SIGKILL",
10: "SIGUSR1",
11: "SIGSEGV",
12: "SIGUSR2",
13: "SIGPIPE",
14: "SIGALRM",
15: "SIGTERM",
16: "SIGSTKFLT",
17: "SIGCHLD",
18: "SIGCONT",
19: "SIGSTOP",
20: "SIGTSTP",
21: "SIGTTIN",
22: "SIGTTOU",
23: "SIGURG",
24: "SIGXCPU",
25: "SIGXFSZ",
26: "SIGVTALRM",
27: "SIGPROF",
28: "SIGWINCH",
29: "SIGIO",
30: "SIGPWR",
31: "SIGSYS",
// Real-time signals base
34: "SIGRTMIN",
// SIGRTMIN+n signals
35: "SIGRTMIN+1",
36: "SIGRTMIN+2",
37: "SIGRTMIN+3",
38: "SIGRTMIN+4",
39: "SIGRTMIN+5",
40: "SIGRTMIN+6",
41: "SIGRTMIN+7",
42: "SIGRTMIN+8",
43: "SIGRTMIN+9",
44: "SIGRTMIN+10",
45: "SIGRTMIN+11",
46: "SIGRTMIN+12",
47: "SIGRTMIN+13",
48: "SIGRTMIN+14",
49: "SIGRTMIN+15",
// SIGRTMAX-n signals
50: "SIGRTMAX-14",
51: "SIGRTMAX-13",
52: "SIGRTMAX-12",
53: "SIGRTMAX-11",
54: "SIGRTMAX-10",
55: "SIGRTMAX-9",
56: "SIGRTMAX-8",
57: "SIGRTMAX-7",
58: "SIGRTMAX-6",
59: "SIGRTMAX-5",
60: "SIGRTMAX-4",
61: "SIGRTMAX-3",
62: "SIGRTMAX-2",
63: "SIGRTMAX-1",
64: "SIGRTMAX",
}
// These signals may produce core dumps:
// SIGQUIT, SIGILL, SIGABRT, SIGBUS, SIGFPE, SIGSEGV
const coreDumpPossible = new Set([3, 4, 6, 7, 8, 11])
return {
exitCode,
signal,
signalName: signals[signal] || `Unknown Signal (${signal})`,
coreDumpPossible: coreDumpPossible.has(signal),
}
}
constructor() {
let disposable: vscode.Disposable | undefined
let startDisposable: vscode.Disposable | undefined
let endDisposable: vscode.Disposable | undefined
try {
disposable = (vscode.window as vscode.Window).onDidStartTerminalShellExecution?.(async (e) => {
// Creating a read stream here results in a more consistent output. This is most obvious when running the `date` command.
e?.execution?.read()
// onDidStartTerminalShellExecution
startDisposable = (vscode.window as vscode.Window).onDidStartTerminalShellExecution?.(async (e) => {
// Get a handle to the stream as early as possible:
const stream = e?.execution.read()
const terminalInfo = TerminalRegistry.getTerminalInfoByTerminal(e.terminal)
if (stream && terminalInfo) {
const process = this.processes.get(terminalInfo.id)
if (process) {
terminalInfo.stream = stream
terminalInfo.running = true
terminalInfo.streamClosed = false
process.emit("stream_available", terminalInfo.id, stream)
}
} else {
console.error("[TerminalManager] Stream failed, not registered for terminal")
}
console.info("[TerminalManager] Shell execution started:", {
hasExecution: !!e?.execution,
command: e?.execution?.commandLine?.value,
terminalId: terminalInfo?.id,
})
})
// onDidEndTerminalShellExecution
endDisposable = (vscode.window as vscode.Window).onDidEndTerminalShellExecution?.(async (e) => {
const exitDetails = this.interpretExitCode(e?.exitCode)
console.info("[TerminalManager] Shell execution ended:", {
...exitDetails,
})
// Signal completion to any waiting processes
for (const id of this.terminalIds) {
const info = TerminalRegistry.getTerminal(id)
if (info && info.terminal === e.terminal) {
info.running = false
const process = this.processes.get(id)
if (process) {
process.emit("shell_execution_complete", id, exitDetails)
}
break
}
}
})
} catch (error) {
// console.error("Error setting up onDidEndTerminalShellExecution", error)
console.error("[TerminalManager] Error setting up shell execution handlers:", error)
}
if (disposable) {
this.disposables.push(disposable)
if (startDisposable) {
this.disposables.push(startDisposable)
}
if (endDisposable) {
this.disposables.push(endDisposable)
}
}
@ -140,19 +289,16 @@ export class TerminalManager {
})
// if shell integration is already active, run the command immediately
const terminal = terminalInfo.terminal as ExtendedTerminal
if (terminal.shellIntegration) {
if (terminalInfo.terminal.shellIntegration) {
process.waitForShellIntegration = false
process.run(terminal, command)
process.run(terminalInfo.terminal, command)
} else {
// docs recommend waiting 3s for shell integration to activate
pWaitFor(() => (terminalInfo.terminal as ExtendedTerminal).shellIntegration !== undefined, {
timeout: 4000,
}).finally(() => {
pWaitFor(() => terminalInfo.terminal.shellIntegration !== undefined, { timeout: 4000 }).finally(() => {
const existingProcess = this.processes.get(terminalInfo.id)
if (existingProcess && existingProcess.waitForShellIntegration) {
existingProcess.waitForShellIntegration = false
existingProcess.run(terminal, command)
existingProcess.run(terminalInfo.terminal, command)
}
})
}
@ -168,8 +314,7 @@ export class TerminalManager {
if (t.busy) {
return false
}
const terminal = t.terminal as ExtendedTerminal
const terminalCwd = terminal.shellIntegration?.cwd // one of cline's commands could have changed the cwd of the terminal
const terminalCwd = t.terminal.shellIntegration?.cwd // one of cline's commands could have changed the cwd of the terminal
if (!terminalCwd) {
return false
}

View file

@ -1,13 +1,24 @@
import { EventEmitter } from "events"
import stripAnsi from "strip-ansi"
import * as vscode from "vscode"
import { inspect } from "util"
import { ExitCodeDetails } from "./TerminalManager"
import { TerminalInfo, TerminalRegistry } from "./TerminalRegistry"
export interface TerminalProcessEvents {
line: [line: string]
continue: []
completed: []
completed: [output?: string]
error: [error: Error]
no_shell_integration: []
/**
* Emitted when a shell execution completes
* @param id The terminal ID
* @param exitDetails Contains exit code and signal information if process was terminated by signal
*/
shell_execution_complete: [id: number, exitDetails: ExitCodeDetails]
stream_available: [id: number, stream: AsyncIterable<string>]
}
// how long to wait after a process outputs anything before we consider it "cool" again
@ -17,104 +28,99 @@ const PROCESS_HOT_TIMEOUT_COMPILING = 15_000
export class TerminalProcess extends EventEmitter<TerminalProcessEvents> {
waitForShellIntegration: boolean = true
private isListening: boolean = true
private buffer: string = ""
private terminalInfo: TerminalInfo | undefined
private lastEmitTime_ms: number = 0
private fullOutput: string = ""
private lastRetrievedIndex: number = 0
isHot: boolean = false
private hotTimer: NodeJS.Timeout | null = null
// constructor() {
// super()
async run(terminal: vscode.Terminal, command: string) {
if (terminal.shellIntegration && terminal.shellIntegration.executeCommand) {
const execution = terminal.shellIntegration.executeCommand(command)
const stream = execution.read()
// todo: need to handle errors
let isFirstChunk = true
let didOutputNonCommand = false
let didEmitEmptyLine = false
// Get terminal info to access stream
const terminalInfo = TerminalRegistry.getTerminalInfoByTerminal(terminal)
if (!terminalInfo) {
console.error("[TerminalProcess] Terminal not found in registry")
this.emit("no_shell_integration")
this.emit("completed")
this.emit("continue")
return
}
// When executeCommand() is called, onDidStartTerminalShellExecution will fire in TerminalManager
// which creates a new stream via execution.read() and emits 'stream_available'
const streamAvailable = new Promise<AsyncIterable<string>>((resolve) => {
this.once("stream_available", (id: number, stream: AsyncIterable<string>) => {
if (id === terminalInfo.id) {
resolve(stream)
}
})
})
// Create promise that resolves when shell execution completes for this terminal
const shellExecutionComplete = new Promise<ExitCodeDetails>((resolve) => {
this.once("shell_execution_complete", (id: number, exitDetails: ExitCodeDetails) => {
if (id === terminalInfo.id) {
resolve(exitDetails)
}
})
})
// getUnretrievedOutput needs to know if streamClosed, so store this for later
this.terminalInfo = terminalInfo
// Execute command
terminal.shellIntegration.executeCommand(command)
this.isHot = true
// Wait for stream to be available
const stream = await streamAvailable
let preOutput = ""
let commandOutputStarted = false
/*
* Extract clean output from raw accumulated output. FYI:
* ]633 is a custom sequence number used by VSCode shell integration:
* - OSC 633 ; A ST - Mark prompt start
* - OSC 633 ; B ST - Mark prompt end
* - OSC 633 ; C ST - Mark pre-execution (start of command output)
* - OSC 633 ; D [; <exitcode>] ST - Mark execution finished with optional exit code
* - OSC 633 ; E ; <commandline> [; <nonce>] ST - Explicitly set command line with optional nonce
*/
// Process stream data
for await (let data of stream) {
// 1. Process chunk and remove artifacts
if (isFirstChunk) {
/*
The first chunk we get from this stream needs to be processed to be more human readable, ie remove vscode's custom escape sequences and identifiers, removing duplicate first char bug, etc.
*/
// bug where sometimes the command output makes its way into vscode shell integration metadata
/*
]633 is a custom sequence number used by VSCode shell integration:
- OSC 633 ; A ST - Mark prompt start
- OSC 633 ; B ST - Mark prompt end
- OSC 633 ; C ST - Mark pre-execution (start of command output)
- OSC 633 ; D [; <exitcode>] ST - Mark execution finished with optional exit code
- OSC 633 ; E ; <commandline> [; <nonce>] ST - Explicitly set command line with optional nonce
*/
// if you print this data you might see something like "eecho hello worldo hello world;5ba85d14-e92a-40c4-b2fd-71525581eeb0]633;C" but this is actually just a bunch of escape sequences, ignore up to the first ;C
/* ddateb15026-6a64-40db-b21f-2a621a9830f0]633;CTue Sep 17 06:37:04 EDT 2024 % ]633;D;0]633;P;Cwd=/Users/saoud/Repositories/test */
// Gets output between ]633;C (command start) and ]633;D (command end)
const outputBetweenSequences = this.removeLastLineArtifacts(
data.match(/\]633;C([\s\S]*?)\]633;D/)?.[1] || "",
).trim()
// Once we've retrieved any potential output between sequences, we can remove everything up to end of the last sequence
// https://code.visualstudio.com/docs/terminal/shell-integration#_vs-code-custom-sequences-osc-633-st
const vscodeSequenceRegex = /\x1b\]633;.[^\x07]*\x07/g
const lastMatch = [...data.matchAll(vscodeSequenceRegex)].pop()
if (lastMatch && lastMatch.index !== undefined) {
data = data.slice(lastMatch.index + lastMatch[0].length)
// Check for command output start marker
if (!commandOutputStarted) {
preOutput += data
const match = this.matchAfterVsceStartMarkers(data)
if (match !== undefined) {
commandOutputStarted = true
data = match
this.fullOutput = "" // Reset fullOutput when command actually starts
} else {
continue
}
// Place output back after removing vscode sequences
if (outputBetweenSequences) {
data = outputBetweenSequences + "\n" + data
}
// remove ansi
data = stripAnsi(data)
// Split data by newlines
let lines = data ? data.split("\n") : []
// Remove non-human readable characters from the first line
if (lines.length > 0) {
lines[0] = lines[0].replace(/[^\x20-\x7E]/g, "")
}
// Check if first two characters are the same, if so remove the first character
if (lines.length > 0 && lines[0].length >= 2 && lines[0][0] === lines[0][1]) {
lines[0] = lines[0].slice(1)
}
// Remove everything up to the first alphanumeric character for first two lines
if (lines.length > 0) {
lines[0] = lines[0].replace(/^[^a-zA-Z0-9]*/, "")
}
if (lines.length > 1) {
lines[1] = lines[1].replace(/^[^a-zA-Z0-9]*/, "")
}
// Join lines back
data = lines.join("\n")
isFirstChunk = false
} else {
data = stripAnsi(data)
}
// first few chunks could be the command being echoed back, so we must ignore
// note this means that 'echo' commands wont work
if (!didOutputNonCommand) {
const lines = data.split("\n")
for (let i = 0; i < lines.length; i++) {
if (command.includes(lines[i].trim())) {
lines.splice(i, 1)
i-- // Adjust index after removal
} else {
didOutputNonCommand = true
break
}
}
data = lines.join("\n")
// Command output started, accumulate data without filtering.
// notice to future programmers: do not add escape sequence
// filtering here: fullOutput cannot change in length (see getUnretrievedOutput),
// and chunks may not be complete so you cannot rely on detecting or removing escape sequences mid-stream.
this.fullOutput += data
// For non-immediately returning commands we want to show loading spinner
// right away but this wouldnt happen until it emits a line break, so
// as soon as we get any output we emit to let webview know to show spinner
const now = Date.now()
if (this.isListening && (now - this.lastEmitTime_ms > 100 || this.lastEmitTime_ms === 0)) {
this.emitRemainingBufferIfListening()
this.lastEmitTime_ms = now
}
// FIXME: right now it seems that data chunks returned to us from the shell integration stream contains random commas, which from what I can tell is not the expected behavior. There has to be a better solution here than just removing all commas.
data = data.replace(/,/g, "")
// 2. Set isHot depending on the command
// Set to hot to stall API requests until terminal is cool again
// 2. Set isHot depending on the command.
// This stalls API requests until terminal is cool again.
this.isHot = true
if (this.hotTimer) {
clearTimeout(this.hotTimer)
@ -144,21 +150,37 @@ export class TerminalProcess extends EventEmitter<TerminalProcessEvents> {
},
isCompiling ? PROCESS_HOT_TIMEOUT_COMPILING : PROCESS_HOT_TIMEOUT_NORMAL,
)
// For non-immediately returning commands we want to show loading spinner right away but this wouldnt happen until it emits a line break, so as soon as we get any output we emit "" to let webview know to show spinner
if (!didEmitEmptyLine && !this.fullOutput && data) {
this.emit("line", "") // empty line to indicate start of command output stream
didEmitEmptyLine = true
}
this.fullOutput += data
if (this.isListening) {
this.emitIfEol(data)
this.lastRetrievedIndex = this.fullOutput.length - this.buffer.length
}
}
this.emitRemainingBufferIfListening()
// Set streamClosed immediately after stream ends
if (this.terminalInfo) {
this.terminalInfo.streamClosed = true
}
// Wait for shell execution to complete and handle exit details
const exitDetails = await shellExecutionComplete
this.isHot = false
if (commandOutputStarted) {
// Emit any remaining output before completing
this.emitRemainingBufferIfListening()
} else {
console.error(
"[Terminal Process] VSCE output start escape sequence (]633;C or ]133;C) not received! VSCE Bug? preOutput: " +
inspect(preOutput, { colors: false, breakLength: Infinity }),
)
}
// console.debug("[Terminal Process] raw output: " + inspect(output, { colors: false, breakLength: Infinity }))
// fullOutput begins after C marker so we only need to trim off D marker
// (if D exists, see VSCode bug# 237208):
const match = this.matchBeforeVsceEndMarkers(this.fullOutput)
if (match !== undefined) {
this.fullOutput = match
}
// console.debug(`[Terminal Process] processed output via ${matchSource}: ` + inspect(output, { colors: false, breakLength: Infinity }))
// for now we don't want this delaying requests since we don't send diagnostics automatically anymore (previous: "even though the command is finished, we still want to consider it 'hot' in case so that api request stalls to let diagnostics catch up")
if (this.hotTimer) {
@ -166,7 +188,7 @@ export class TerminalProcess extends EventEmitter<TerminalProcessEvents> {
}
this.isHot = false
this.emit("completed")
this.emit("completed", this.removeEscapeSequences(this.fullOutput))
this.emit("continue")
} else {
terminal.sendText(command, true)
@ -182,29 +204,12 @@ export class TerminalProcess extends EventEmitter<TerminalProcessEvents> {
}
}
// Inspired by https://github.com/sindresorhus/execa/blob/main/lib/transform/split.js
private emitIfEol(chunk: string) {
this.buffer += chunk
let lineEndIndex: number
while ((lineEndIndex = this.buffer.indexOf("\n")) !== -1) {
let line = this.buffer.slice(0, lineEndIndex).trimEnd() // removes trailing \r
// Remove \r if present (for Windows-style line endings)
// if (line.endsWith("\r")) {
// line = line.slice(0, -1)
// }
this.emit("line", line)
this.buffer = this.buffer.slice(lineEndIndex + 1)
}
}
private emitRemainingBufferIfListening() {
if (this.buffer && this.isListening) {
const remainingBuffer = this.removeLastLineArtifacts(this.buffer)
if (remainingBuffer) {
if (this.isListening) {
const remainingBuffer = this.getUnretrievedOutput()
if (remainingBuffer !== "") {
this.emit("line", remainingBuffer)
}
this.buffer = ""
this.lastRetrievedIndex = this.fullOutput.length
}
}
@ -215,22 +220,180 @@ export class TerminalProcess extends EventEmitter<TerminalProcessEvents> {
this.emit("continue")
}
// Returns complete lines with their carriage returns.
// The final line may lack a carriage return if the program didn't send one.
getUnretrievedOutput(): string {
const unretrieved = this.fullOutput.slice(this.lastRetrievedIndex)
this.lastRetrievedIndex = this.fullOutput.length
return this.removeLastLineArtifacts(unretrieved)
// Get raw unretrieved output
let outputToProcess = this.fullOutput.slice(this.lastRetrievedIndex)
// Check for VSCE command end markers
const index633 = outputToProcess.indexOf("\x1b]633;D")
const index133 = outputToProcess.indexOf("\x1b]133;D")
let endIndex = -1
if (index633 !== -1 && index133 !== -1) {
endIndex = Math.min(index633, index133)
} else if (index633 !== -1) {
endIndex = index633
} else if (index133 !== -1) {
endIndex = index133
}
// If no end markers were found yet (possibly due to VSCode bug#237208):
// For active streams: return only complete lines (up to last \n).
// For closed streams: return all remaining content.
if (endIndex === -1) {
if (!this.terminalInfo?.streamClosed) {
// Stream still running - only process complete lines
endIndex = outputToProcess.lastIndexOf("\n")
if (endIndex === -1) {
// No complete lines
return ""
}
// Include carriage return
endIndex++
} else {
// Stream closed - process all remaining output
endIndex = outputToProcess.length
}
}
// Update index and slice output
this.lastRetrievedIndex += endIndex
outputToProcess = outputToProcess.slice(0, endIndex)
// Clean and return output
return this.removeEscapeSequences(outputToProcess)
}
// some processing to remove artifacts like '%' at the end of the buffer (it seems that since vsode uses % at the beginning of newlines in terminal, it makes its way into the stream)
// This modification will remove '%', '$', '#', or '>' followed by optional whitespace
removeLastLineArtifacts(output: string) {
const lines = output.trimEnd().split("\n")
if (lines.length > 0) {
const lastLine = lines[lines.length - 1]
// Remove prompt characters and trailing whitespace from the last line
lines[lines.length - 1] = lastLine.replace(/[%$#>]\s*$/, "")
private stringIndexMatch(
data: string,
prefix?: string,
suffix?: string,
bell: string = "\x07",
): string | undefined {
let startIndex: number
let endIndex: number
let prefixLength: number
if (prefix === undefined) {
startIndex = 0
prefixLength = 0
} else {
startIndex = data.indexOf(prefix)
if (startIndex === -1) {
return undefined
}
if (bell.length > 0) {
// Find the bell character after the prefix
const bellIndex = data.indexOf(bell, startIndex + prefix.length)
if (bellIndex === -1) {
return undefined
}
const distanceToBell = bellIndex - startIndex
prefixLength = distanceToBell + bell.length
} else {
prefixLength = prefix.length
}
}
return lines.join("\n").trimEnd()
const contentStart = startIndex + prefixLength
if (suffix === undefined) {
// When suffix is undefined, match to end
endIndex = data.length
} else {
endIndex = data.indexOf(suffix, contentStart)
if (endIndex === -1) {
return undefined
}
}
return data.slice(contentStart, endIndex)
}
// Removes ANSI escape sequences and VSCode-specific terminal control codes from output.
// While stripAnsi handles most ANSI codes, VSCode's shell integration adds custom
// escape sequences (OSC 633) that need special handling. These sequences control
// terminal features like marking command start/end and setting prompts.
//
// This method could be extended to handle other escape sequences, but any additions
// should be carefully considered to ensure they only remove control codes and don't
// alter the actual content or behavior of the output stream.
private removeEscapeSequences(str: string): string {
return stripAnsi(str.replace(/\x1b\]633;[^\x07]+\x07/gs, "").replace(/\x1b\]133;[^\x07]+\x07/gs, ""))
}
/**
* Helper function to match VSCode shell integration start markers (C).
* Looks for content after ]633;C or ]133;C markers.
* If both exist, takes the content after the last marker found.
*/
private matchAfterVsceStartMarkers(data: string): string | undefined {
return this.matchVsceMarkers(data, "\x1b]633;C", "\x1b]133;C", undefined, undefined)
}
/**
* Helper function to match VSCode shell integration end markers (D).
* Looks for content before ]633;D or ]133;D markers.
* If both exist, takes the content before the first marker found.
*/
private matchBeforeVsceEndMarkers(data: string): string | undefined {
return this.matchVsceMarkers(data, undefined, undefined, "\x1b]633;D", "\x1b]133;D")
}
/**
* Handles VSCode shell integration markers for command output:
*
* For C (Command Start):
* - Looks for content after ]633;C or ]133;C markers
* - These markers indicate the start of command output
* - If both exist, takes the content after the last marker found
* - This ensures we get the actual command output after any shell integration prefixes
*
* For D (Command End):
* - Looks for content before ]633;D or ]133;D markers
* - These markers indicate command completion
* - If both exist, takes the content before the first marker found
* - This ensures we don't include shell integration suffixes in the output
*
* In both cases, checks 633 first since it's more commonly used in VSCode shell integration
*
* @param data The string to search for markers in
* @param prefix633 The 633 marker to match after (for C markers)
* @param prefix133 The 133 marker to match after (for C markers)
* @param suffix633 The 633 marker to match before (for D markers)
* @param suffix133 The 133 marker to match before (for D markers)
* @returns The content between/after markers, or undefined if no markers found
*
* Note: Always makes exactly 2 calls to stringIndexMatch regardless of match results.
* Using string indexOf matching is ~500x faster than regular expressions, so even
* matching twice is still very efficient comparatively.
*/
private matchVsceMarkers(
data: string,
prefix633: string | undefined,
prefix133: string | undefined,
suffix633: string | undefined,
suffix133: string | undefined,
): string | undefined {
// Support both VSCode shell integration markers (633 and 133)
// Check 633 first since it's more commonly used in VSCode shell integration
let match133: string | undefined
const match633 = this.stringIndexMatch(data, prefix633, suffix633)
// Must check explicitly for undefined because stringIndexMatch can return empty strings
// that are valid matches (e.g., when a marker exists but has no content between markers)
if (match633 !== undefined) {
match133 = this.stringIndexMatch(match633, prefix133, suffix133)
} else {
match133 = this.stringIndexMatch(data, prefix133, suffix133)
}
return match133 !== undefined ? match133 : match633
}
}

View file

@ -5,6 +5,9 @@ export interface TerminalInfo {
busy: boolean
lastCommand: string
id: number
stream?: AsyncIterable<string>
running: boolean
streamClosed: boolean
}
// Although vscode.window.terminals provides a list of all open terminals, there's no way to know whether they're busy or not (exitStatus does not provide useful information for most commands). In order to prevent creating too many terminals, we need to keep track of terminals through the life of the extension, as well as session specific terminals for the life of a task (to get latest unretrieved output).
@ -20,34 +23,61 @@ export class TerminalRegistry {
iconPath: new vscode.ThemeIcon("rocket"),
env: {
PAGER: "cat",
// VSCode bug#237208: Command output can be lost due to a race between completion
// sequences and consumers. Add 50ms delay via PROMPT_COMMAND to ensure the
// \x1b]633;D escape sequence arrives after command output is processed.
PROMPT_COMMAND: "sleep 0.050",
// VTE must be disabled because it prevents the prompt command above from executing
// See https://wiki.gnome.org/Apps/Terminal/VTE
VTE_VERSION: "0",
},
})
const newInfo: TerminalInfo = {
terminal,
busy: false,
lastCommand: "",
id: this.nextTerminalId++,
running: false,
streamClosed: false,
}
this.terminals.push(newInfo)
return newInfo
}
static getTerminal(id: number): TerminalInfo | undefined {
const terminalInfo = this.terminals.find((t) => t.id === id)
if (terminalInfo && this.isTerminalClosed(terminalInfo.terminal)) {
this.removeTerminal(id)
return undefined
}
return terminalInfo
}
static updateTerminal(id: number, updates: Partial<TerminalInfo>) {
const terminal = this.getTerminal(id)
if (terminal) {
Object.assign(terminal, updates)
}
}
static getTerminalInfoByTerminal(terminal: vscode.Terminal): TerminalInfo | undefined {
const terminalInfo = this.terminals.find((t) => t.terminal === terminal)
if (terminalInfo && this.isTerminalClosed(terminalInfo.terminal)) {
this.removeTerminal(terminalInfo.id)
return undefined
}
return terminalInfo
}
static removeTerminal(id: number) {
this.terminals = this.terminals.filter((t) => t.id !== id)
}

View file

@ -1,9 +1,24 @@
import { TerminalProcess, mergePromise } from "../TerminalProcess"
import * as vscode from "vscode"
import { EventEmitter } from "events"
// npx jest src/integrations/terminal/__tests__/TerminalProcess.test.ts
// Mock vscode
jest.mock("vscode")
import * as vscode from "vscode"
import { TerminalProcess, mergePromise } from "../TerminalProcess"
import { TerminalInfo, TerminalRegistry } from "../TerminalRegistry"
// Mock vscode.window.createTerminal
const mockCreateTerminal = jest.fn()
jest.mock("vscode", () => ({
window: {
createTerminal: (...args: any[]) => {
mockCreateTerminal(...args)
return {
exitStatus: undefined,
}
},
},
ThemeIcon: jest.fn(),
}))
describe("TerminalProcess", () => {
let terminalProcess: TerminalProcess
@ -14,6 +29,7 @@ describe("TerminalProcess", () => {
}
}
>
let mockTerminalInfo: TerminalInfo
let mockExecution: any
let mockStream: AsyncIterableIterator<string>
@ -25,7 +41,7 @@ describe("TerminalProcess", () => {
shellIntegration: {
executeCommand: jest.fn(),
},
name: "Mock Terminal",
name: "Roo Code",
processId: Promise.resolve(123),
creationOptions: {},
exitStatus: undefined,
@ -42,27 +58,39 @@ describe("TerminalProcess", () => {
}
>
mockTerminalInfo = {
terminal: mockTerminal,
busy: false,
lastCommand: "",
id: 1,
running: false,
streamClosed: false,
}
TerminalRegistry["terminals"].push(mockTerminalInfo)
// Reset event listeners
terminalProcess.removeAllListeners()
})
describe("run", () => {
it("handles shell integration commands correctly", async () => {
const lines: string[] = []
terminalProcess.on("line", (line) => {
// Skip empty lines used for loading spinner
if (line !== "") {
lines.push(line)
let lines: string[] = []
terminalProcess.on("completed", (output) => {
if (output) {
lines = output.split("\n")
}
})
// Mock stream data with shell integration sequences
// Mock stream data with shell integration sequences.
mockStream = (async function* () {
// The first chunk contains the command start sequence
yield "\x1b]633;C\x07" // The first chunk contains the command start sequence with bell character.
yield "Initial output\n"
yield "More output\n"
// The last chunk contains the command end sequence
yield "Final output"
yield "\x1b]633;D\x07" // The last chunk contains the command end sequence with bell character.
terminalProcess.emit("shell_execution_complete", mockTerminalInfo.id, { exitCode: 0 })
})()
mockExecution = {
@ -71,12 +99,9 @@ describe("TerminalProcess", () => {
mockTerminal.shellIntegration.executeCommand.mockReturnValue(mockExecution)
const completedPromise = new Promise<void>((resolve) => {
terminalProcess.once("completed", resolve)
})
await terminalProcess.run(mockTerminal, "test command")
await completedPromise
const runPromise = terminalProcess.run(mockTerminal, "test command")
terminalProcess.emit("stream_available", mockTerminalInfo.id, mockStream)
await runPromise
expect(lines).toEqual(["Initial output", "More output", "Final output"])
expect(terminalProcess.isHot).toBe(false)
@ -99,95 +124,41 @@ describe("TerminalProcess", () => {
})
it("sets hot state for compiling commands", async () => {
const lines: string[] = []
terminalProcess.on("line", (line) => {
if (line !== "") {
lines.push(line)
let lines: string[] = []
terminalProcess.on("completed", (output) => {
if (output) {
lines = output.split("\n")
}
})
// Create a promise that resolves when the first chunk is processed
const firstChunkProcessed = new Promise<void>((resolve) => {
terminalProcess.on("line", () => resolve())
const completePromise = new Promise<void>((resolve) => {
terminalProcess.on("shell_execution_complete", () => resolve())
})
mockStream = (async function* () {
yield "\x1b]633;C\x07" // The first chunk contains the command start sequence with bell character.
yield "compiling...\n"
// Wait to ensure hot state check happens after first chunk
await new Promise((resolve) => setTimeout(resolve, 10))
yield "still compiling...\n"
yield "done"
yield "\x1b]633;D\x07" // The last chunk contains the command end sequence with bell character.
terminalProcess.emit("shell_execution_complete", mockTerminalInfo.id, { exitCode: 0 })
})()
mockExecution = {
mockTerminal.shellIntegration.executeCommand.mockReturnValue({
read: jest.fn().mockReturnValue(mockStream),
}
mockTerminal.shellIntegration.executeCommand.mockReturnValue(mockExecution)
// Start the command execution
const runPromise = terminalProcess.run(mockTerminal, "npm run build")
// Wait for the first chunk to be processed
await firstChunkProcessed
// Hot state should be true while compiling
expect(terminalProcess.isHot).toBe(true)
// Complete the execution
const completedPromise = new Promise<void>((resolve) => {
terminalProcess.once("completed", resolve)
})
const runPromise = terminalProcess.run(mockTerminal, "npm run build")
terminalProcess.emit("stream_available", mockTerminalInfo.id, mockStream)
expect(terminalProcess.isHot).toBe(true)
await runPromise
await completedPromise
expect(lines).toEqual(["compiling...", "still compiling...", "done"])
})
})
describe("buffer processing", () => {
it("correctly processes and emits lines", () => {
const lines: string[] = []
terminalProcess.on("line", (line) => lines.push(line))
// Simulate incoming chunks
terminalProcess["emitIfEol"]("first line\n")
terminalProcess["emitIfEol"]("second")
terminalProcess["emitIfEol"](" line\n")
terminalProcess["emitIfEol"]("third line")
expect(lines).toEqual(["first line", "second line"])
// Process remaining buffer
terminalProcess["emitRemainingBufferIfListening"]()
expect(lines).toEqual(["first line", "second line", "third line"])
})
it("handles Windows-style line endings", () => {
const lines: string[] = []
terminalProcess.on("line", (line) => lines.push(line))
terminalProcess["emitIfEol"]("line1\r\nline2\r\n")
expect(lines).toEqual(["line1", "line2"])
})
})
describe("removeLastLineArtifacts", () => {
it("removes terminal artifacts from output", () => {
const cases = [
["output%", "output"],
["output$ ", "output"],
["output#", "output"],
["output> ", "output"],
["multi\nline%", "multi\nline"],
["no artifacts", "no artifacts"],
]
for (const [input, expected] of cases) {
expect(terminalProcess["removeLastLineArtifacts"](input)).toBe(expected)
}
await completePromise
expect(terminalProcess.isHot).toBe(false)
})
})
@ -205,13 +176,13 @@ describe("TerminalProcess", () => {
describe("getUnretrievedOutput", () => {
it("returns and clears unretrieved output", () => {
terminalProcess["fullOutput"] = "previous\nnew output"
terminalProcess["lastRetrievedIndex"] = 9 // After "previous\n"
terminalProcess["fullOutput"] = `\x1b]633;C\x07previous\nnew output\x1b]633;D\x07`
terminalProcess["lastRetrievedIndex"] = 17 // After "previous\n"
const unretrieved = terminalProcess.getUnretrievedOutput()
expect(unretrieved).toBe("new output")
expect(terminalProcess["lastRetrievedIndex"]).toBe(terminalProcess["fullOutput"].length)
expect(terminalProcess["lastRetrievedIndex"]).toBe(terminalProcess["fullOutput"].length - "previous".length)
})
})

View file

@ -1,4 +1,5 @@
import * as vscode from "vscode"
// npx jest src/integrations/terminal/__tests__/TerminalRegistry.test.ts
import { TerminalRegistry } from "../TerminalRegistry"
// Mock vscode.window.createTerminal
@ -30,6 +31,8 @@ describe("TerminalRegistry", () => {
iconPath: expect.any(Object),
env: {
PAGER: "cat",
PROMPT_COMMAND: "sleep 0.050",
VTE_VERSION: "0",
},
})
})

View file

@ -112,7 +112,7 @@ export const anthropicModels = {
thinking: true,
},
"claude-3-7-sonnet-20250219": {
maxTokens: 64_000,
maxTokens: 16_384,
contextWindow: 200_000,
supportsImages: true,
supportsComputerUse: true,
@ -437,8 +437,48 @@ export const openRouterDefaultModelInfo: ModelInfo = {
export type VertexModelId = keyof typeof vertexModels
export const vertexDefaultModelId: VertexModelId = "claude-3-7-sonnet@20250219"
export const vertexModels = {
"gemini-2.0-flash-001": {
maxTokens: 8192,
contextWindow: 1_048_576,
supportsImages: true,
supportsPromptCache: false,
inputPrice: 0.15,
outputPrice: 0.6,
},
"gemini-2.0-flash-lite-001": {
maxTokens: 8192,
contextWindow: 1_048_576,
supportsImages: true,
supportsPromptCache: false,
inputPrice: 0.075,
outputPrice: 0.3,
},
"gemini-2.0-flash-thinking-exp-01-21": {
maxTokens: 8192,
contextWindow: 32_768,
supportsImages: true,
supportsPromptCache: false,
inputPrice: 0,
outputPrice: 0,
},
"gemini-1.5-flash-002": {
maxTokens: 8192,
contextWindow: 1_048_576,
supportsImages: true,
supportsPromptCache: false,
inputPrice: 0.075,
outputPrice: 0.3,
},
"gemini-1.5-pro-002": {
maxTokens: 8192,
contextWindow: 2_097_152,
supportsImages: true,
supportsPromptCache: false,
inputPrice: 1.25,
outputPrice: 5,
},
"claude-3-7-sonnet@20250219:thinking": {
maxTokens: 64000,
maxTokens: 64_000,
contextWindow: 200_000,
supportsImages: true,
supportsComputerUse: true,
@ -450,7 +490,7 @@ export const vertexModels = {
thinking: true,
},
"claude-3-7-sonnet@20250219": {
maxTokens: 8192,
maxTokens: 16_384,
contextWindow: 200_000,
supportsImages: true,
supportsComputerUse: true,

View file

@ -44,7 +44,7 @@ export function combineCommandSequences(messages: ClineMessage[]): ClineMessage[
// handle cases where we receive empty command_output (ie when extension is relinquishing control over exit command button)
const output = messages[j].text || ""
if (output.length > 0) {
combinedText += "\n" + output
combinedText += output
}
}
j++

View file

@ -36,7 +36,11 @@ export type CustomModePrompts = {
// Helper to extract group name regardless of format
export function getGroupName(group: GroupEntry): ToolGroup {
return Array.isArray(group) ? group[0] : group
if (typeof group === "string") {
return group
}
return group[0]
}
// Helper to get group options if they exist
@ -88,7 +92,7 @@ export const modes: readonly ModeConfig[] = [
"You are Roo, an experienced technical leader who is inquisitive and an excellent planner. Your goal is to gather information and get context to create a detailed plan for accomplishing the user's task, which the user will review and approve before they switch into another mode to implement the solution.",
groups: ["read", ["edit", { fileRegex: "\\.md$", description: "Markdown files only" }], "browser", "mcp"],
customInstructions:
"Depending on the user's request, you may need to do some information gathering (for example using read_file or search_files) to get more context about the task. You may also ask the user clarifying questions to get a better understanding of the task. Once you've gained more context about the user's request, you should create a detailed plan for how to accomplish the task. (You can write the plan to a markdown file if it seems appropriate.)\n\nThen you might ask the user if they are pleased with this plan, or if they would like to make any changes. Think of this as a brainstorming session where you can discuss the task and plan the best way to accomplish it. Finally once it seems like you've reached a good plan, use the switch_mode tool to request that the user switch to another mode to implement the solution.",
"1. Do some information gathering (for example using read_file or search_files) to get more context about the task.\n\n2. You should also ask the user clarifying questions to get a better understanding of the task.\n\n3. Once you've gained more context about the user's request, you should create a detailed plan for how to accomplish the task. Include Mermaid diagrams if they help make your plan clearer.\n\n4. Ask the user if they are pleased with this plan, or if they would like to make any changes. Think of this as a brainstorming session where you can discuss the task and plan the best way to accomplish it.\n\n5. Once the user confirms the plan, ask them if they'd like you to write it to a markdown file.\n\n6. Use the switch_mode tool to request that the user switch to another mode to implement the solution.",
},
{
slug: "ask",
@ -97,7 +101,7 @@ export const modes: readonly ModeConfig[] = [
"You are Roo, a knowledgeable technical assistant focused on answering questions and providing information about software development, technology, and related topics.",
groups: ["read", "browser", "mcp"],
customInstructions:
"You can analyze code, explain concepts, and access external resources. Make sure to answer the user's questions and don't rush to switch to implementing code.",
"You can analyze code, explain concepts, and access external resources. Make sure to answer the user's questions and don't rush to switch to implementing code. Include Mermaid diagrams if they help make your response clearer.",
},
{
slug: "debug",

File diff suppressed because it is too large Load diff

View file

@ -35,6 +35,7 @@
"fast-deep-equal": "^3.1.3",
"fzf": "^0.5.2",
"lucide-react": "^0.475.0",
"mermaid": "^11.4.1",
"react": "^18.3.1",
"react-dom": "^18.3.1",
"react-markdown": "^9.0.3",

View file

@ -31,6 +31,7 @@ interface ChatTextAreaProps {
onHeightChange?: (height: number) => void
mode: Mode
setMode: (value: Mode) => void
modeShortcutText: string
}
const ChatTextArea = forwardRef<HTMLTextAreaElement, ChatTextAreaProps>(
@ -48,6 +49,7 @@ const ChatTextArea = forwardRef<HTMLTextAreaElement, ChatTextAreaProps>(
onHeightChange,
mode,
setMode,
modeShortcutText,
},
ref,
) => {
@ -816,6 +818,11 @@ const ChatTextArea = forwardRef<HTMLTextAreaElement, ChatTextAreaProps>(
minWidth: "70px",
flex: "0 0 auto",
}}>
<option
disabled
style={{ ...optionStyle, fontStyle: "italic", opacity: 0.6, padding: "2px 8px" }}>
{modeShortcutText}
</option>
{getAllModes(customModes).map((mode) => (
<option key={mode.slug} value={mode.slug} style={{ ...optionStyle }}>
{mode.name}

View file

@ -28,6 +28,7 @@ import TaskHeader from "./TaskHeader"
import AutoApproveMenu from "./AutoApproveMenu"
import { AudioType } from "../../../../src/shared/WebviewMessage"
import { validateCommand } from "../../utils/command-validation"
import { getAllModes } from "../../../../src/shared/modes"
interface ChatViewProps {
isHidden: boolean
@ -38,6 +39,9 @@ interface ChatViewProps {
export const MAX_IMAGES_PER_MESSAGE = 20 // Anthropic limits to 20 images
const isMac = navigator.platform.toUpperCase().indexOf("MAC") >= 0
const modeShortcutText = `${isMac ? "⌘" : "Ctrl"} + . for next mode`
const ChatView = ({ isHidden, showAnnouncement, hideAnnouncement, showHistoryView }: ChatViewProps) => {
const {
version,
@ -56,6 +60,7 @@ const ChatView = ({ isHidden, showAnnouncement, hideAnnouncement, showHistoryVie
setMode,
autoApprovalEnabled,
alwaysAllowModeSwitch,
customModes,
} = useExtensionState()
//const task = messages.length > 0 ? (messages[0].say === "task" ? messages[0] : undefined) : undefined) : undefined
@ -880,7 +885,7 @@ const ChatView = ({ isHidden, showAnnouncement, hideAnnouncement, showHistoryVie
const placeholderText = useMemo(() => {
const baseText = task ? "Type a message..." : "Type your task here..."
const contextText = "(@ to add context, / to switch modes"
const imageText = shouldDisableImages ? "hold shift to drag in files" : ", hold shift to drag in files/images"
const imageText = shouldDisableImages ? ", hold shift to drag in files" : ", hold shift to drag in files/images"
return baseText + `\n${contextText}${imageText})`
}, [task, shouldDisableImages])
@ -963,6 +968,34 @@ const ChatView = ({ isHidden, showAnnouncement, hideAnnouncement, showHistoryVie
isWriteToolAction,
])
// Function to handle mode switching
const switchToNextMode = useCallback(() => {
const allModes = getAllModes(customModes)
const currentModeIndex = allModes.findIndex((m) => m.slug === mode)
const nextModeIndex = (currentModeIndex + 1) % allModes.length
setMode(allModes[nextModeIndex].slug)
}, [mode, setMode, customModes])
// Add keyboard event handler
const handleKeyDown = useCallback(
(event: KeyboardEvent) => {
// Check for Command + . (period)
if ((event.metaKey || event.ctrlKey) && event.key === ".") {
event.preventDefault() // Prevent default browser behavior
switchToNextMode()
}
},
[switchToNextMode],
)
// Add event listener
useEffect(() => {
window.addEventListener("keydown", handleKeyDown)
return () => {
window.removeEventListener("keydown", handleKeyDown)
}
}, [handleKeyDown])
return (
<div
style={{
@ -1171,6 +1204,7 @@ const ChatView = ({ isHidden, showAnnouncement, hideAnnouncement, showHistoryVie
}}
mode={mode}
setMode={setMode}
modeShortcutText={modeShortcutText}
/>
<div id="chat-view-portal" />

View file

@ -187,10 +187,12 @@ const ContextMenu: React.FC<ContextMenuProps> = ({
display: "flex",
alignItems: "center",
justifyContent: "space-between",
backgroundColor:
index === selectedIndex && isOptionSelectable(option)
? "var(--vscode-list-activeSelectionBackground)"
: "",
...(index === selectedIndex && isOptionSelectable(option)
? {
backgroundColor: "var(--vscode-list-activeSelectionBackground)",
color: "var(--vscode-list-activeSelectionForeground)",
}
: {}),
}}
onMouseEnter={() => isOptionSelectable(option) && setSelectedIndex(index)}>
<div

View file

@ -45,6 +45,7 @@ describe("ChatTextArea", () => {
onHeightChange: jest.fn(),
mode: defaultModeSlug,
setMode: jest.fn(),
modeShortcutText: "(⌘. for next mode)",
}
beforeEach(() => {

View file

@ -1,10 +1,11 @@
import { memo, useEffect } from "react"
import React, { memo, useEffect } from "react"
import { useRemark } from "react-remark"
import rehypeHighlight, { Options } from "rehype-highlight"
import styled from "styled-components"
import { visit } from "unist-util-visit"
import { useExtensionState } from "../../context/ExtensionStateContext"
import { CODE_BLOCK_BG_COLOR } from "./CodeBlock"
import MermaidBlock from "./MermaidBlock"
interface MarkdownBlockProps {
markdown?: string
@ -182,7 +183,27 @@ const MarkdownBlock = memo(({ markdown }: MarkdownBlockProps) => {
],
rehypeReactOptions: {
components: {
pre: ({ node, ...preProps }: any) => <StyledPre {...preProps} theme={theme} />,
pre: ({ node, children, ...preProps }: any) => {
if (Array.isArray(children) && children.length === 1 && React.isValidElement(children[0])) {
const child = children[0] as React.ReactElement<{ className?: string }>
if (child.props?.className?.includes("language-mermaid")) {
return child
}
}
return (
<StyledPre {...preProps} theme={theme}>
{children}
</StyledPre>
)
},
code: (props: any) => {
const className = props.className || ""
if (className.includes("language-mermaid")) {
const codeText = String(props.children || "")
return <MermaidBlock code={codeText} />
}
return <code {...props} />
},
},
},
})

View file

@ -0,0 +1,227 @@
import { useEffect, useRef, useState } from "react"
import mermaid from "mermaid"
import { useDebounceEffect } from "../../utils/useDebounceEffect"
import styled from "styled-components"
import { vscode } from "../../utils/vscode"
const MERMAID_THEME = {
background: "#1e1e1e", // VS Code dark theme background
textColor: "#ffffff", // Main text color
mainBkg: "#2d2d2d", // Background for nodes
nodeBorder: "#888888", // Border color for nodes
lineColor: "#cccccc", // Lines connecting nodes
primaryColor: "#3c3c3c", // Primary color for highlights
primaryTextColor: "#ffffff", // Text in primary colored elements
primaryBorderColor: "#888888",
secondaryColor: "#2d2d2d", // Secondary color for alternate elements
tertiaryColor: "#454545", // Third color for special elements
// Class diagram specific
classText: "#ffffff",
// State diagram specific
labelColor: "#ffffff",
// Sequence diagram specific
actorLineColor: "#cccccc",
actorBkg: "#2d2d2d",
actorBorder: "#888888",
actorTextColor: "#ffffff",
// Flow diagram specific
fillType0: "#2d2d2d",
fillType1: "#3c3c3c",
fillType2: "#454545",
}
mermaid.initialize({
startOnLoad: false,
securityLevel: "loose",
theme: "dark",
themeVariables: {
...MERMAID_THEME,
fontSize: "16px",
fontFamily: "var(--vscode-font-family, 'Segoe UI', Tahoma, Geneva, Verdana, sans-serif)",
// Additional styling
noteTextColor: "#ffffff",
noteBkgColor: "#454545",
noteBorderColor: "#888888",
// Improve contrast for special elements
critBorderColor: "#ff9580",
critBkgColor: "#803d36",
// Task diagram specific
taskTextColor: "#ffffff",
taskTextOutsideColor: "#ffffff",
taskTextLightColor: "#ffffff",
// Numbers/sections
sectionBkgColor: "#2d2d2d",
sectionBkgColor2: "#3c3c3c",
// Alt sections in sequence diagrams
altBackground: "#2d2d2d",
// Links
linkColor: "#6cb6ff",
// Borders and lines
compositeBackground: "#2d2d2d",
compositeBorder: "#888888",
titleColor: "#ffffff",
},
})
interface MermaidBlockProps {
code: string
}
export default function MermaidBlock({ code }: MermaidBlockProps) {
const containerRef = useRef<HTMLDivElement>(null)
const [isLoading, setIsLoading] = useState(false)
// 1) Whenever `code` changes, mark that we need to re-render a new chart
useEffect(() => {
setIsLoading(true)
}, [code])
// 2) Debounce the actual parse/render
useDebounceEffect(
() => {
if (containerRef.current) {
containerRef.current.innerHTML = ""
}
mermaid
.parse(code, { suppressErrors: true })
.then((isValid) => {
if (!isValid) {
throw new Error("Invalid or incomplete Mermaid code")
}
const id = `mermaid-${Math.random().toString(36).substring(2)}`
return mermaid.render(id, code)
})
.then(({ svg }) => {
if (containerRef.current) {
containerRef.current.innerHTML = svg
}
})
.catch((err) => {
console.warn("Mermaid parse/render failed:", err)
containerRef.current!.innerHTML = code.replace(/</g, "&lt;").replace(/>/g, "&gt;")
})
.finally(() => {
setIsLoading(false)
})
},
500, // Delay 500ms
[code], // Dependencies for scheduling
)
/**
* Called when user clicks the rendered diagram.
* Converts the <svg> to a PNG and sends it to the extension.
*/
const handleClick = async () => {
if (!containerRef.current) return
const svgEl = containerRef.current.querySelector("svg")
if (!svgEl) return
try {
const pngDataUrl = await svgToPng(svgEl)
vscode.postMessage({
type: "openImage",
text: pngDataUrl,
})
} catch (err) {
console.error("Error converting SVG to PNG:", err)
}
}
return (
<MermaidBlockContainer>
{isLoading && <LoadingMessage>Generating mermaid diagram...</LoadingMessage>}
{/* The container for the final <svg> or raw code. */}
<SvgContainer onClick={handleClick} ref={containerRef} $isLoading={isLoading} />
</MermaidBlockContainer>
)
}
async function svgToPng(svgEl: SVGElement): Promise<string> {
console.log("svgToPng function called")
// Clone the SVG to avoid modifying the original
const svgClone = svgEl.cloneNode(true) as SVGElement
// Get the original viewBox
const viewBox = svgClone.getAttribute("viewBox")?.split(" ").map(Number) || []
const originalWidth = viewBox[2] || svgClone.clientWidth
const originalHeight = viewBox[3] || svgClone.clientHeight
// Calculate the scale factor to fit editor width while maintaining aspect ratio
// Unless we can find a way to get the actual editor window dimensions through the VS Code API (which might be possible but would require changes to the extension side),
// the fixed width seems like a reliable approach.
const editorWidth = 3_600
const scale = editorWidth / originalWidth
const scaledHeight = originalHeight * scale
// Update SVG dimensions
svgClone.setAttribute("width", `${editorWidth}`)
svgClone.setAttribute("height", `${scaledHeight}`)
const serializer = new XMLSerializer()
const svgString = serializer.serializeToString(svgClone)
const svgDataUrl = "data:image/svg+xml;base64," + btoa(decodeURIComponent(encodeURIComponent(svgString)))
return new Promise((resolve, reject) => {
const img = new Image()
img.onload = () => {
const canvas = document.createElement("canvas")
canvas.width = editorWidth
canvas.height = scaledHeight
const ctx = canvas.getContext("2d")
if (!ctx) return reject("Canvas context not available")
// Fill background with Mermaid's dark theme background color
ctx.fillStyle = MERMAID_THEME.background
ctx.fillRect(0, 0, canvas.width, canvas.height)
ctx.imageSmoothingEnabled = true
ctx.imageSmoothingQuality = "high"
ctx.drawImage(img, 0, 0, editorWidth, scaledHeight)
resolve(canvas.toDataURL("image/png", 1.0))
}
img.onerror = reject
img.src = svgDataUrl
})
}
const MermaidBlockContainer = styled.div`
position: relative;
margin: 8px 0;
`
const LoadingMessage = styled.div`
padding: 8px 0;
color: var(--vscode-descriptionForeground);
font-style: italic;
font-size: 0.9em;
`
interface SvgContainerProps {
$isLoading: boolean
}
const SvgContainer = styled.div<SvgContainerProps>`
opacity: ${(props) => (props.$isLoading ? 0.3 : 1)};
min-height: 20px;
transition: opacity 0.2s ease;
cursor: pointer;
display: flex;
justify-content: center;
`

View file

@ -7,7 +7,6 @@ import * as vscodemodels from "vscode"
import {
ApiConfiguration,
ModelInfo,
ApiProvider,
anthropicDefaultModelId,
anthropicModels,
azureOpenAiDefaultApiVersion,
@ -500,7 +499,7 @@ const ApiOptions = ({
/>
)}
<Checkbox
checked={apiConfiguration?.openRouterUseMiddleOutTransform || false}
checked={apiConfiguration?.openRouterUseMiddleOutTransform ?? true}
onChange={handleInputChange("openRouterUseMiddleOutTransform", noTransform)}>
Compress prompts and message chains to the context size (
<a href="https://openrouter.ai/docs/transforms">OpenRouter Transforms</a>)
@ -1412,7 +1411,6 @@ const ApiOptions = ({
apiConfiguration={apiConfiguration}
setApiConfigurationField={setApiConfigurationField}
modelInfo={selectedModelInfo}
provider={selectedProvider as ApiProvider}
/>
<ModelInfoView
selectedModelId={selectedModelId}

View file

@ -1,5 +1,5 @@
import { useEffect, useMemo } from "react"
import { ApiProvider } from "../../../../src/shared/api"
import { Slider } from "@/components/ui"
import { ApiConfiguration, ModelInfo } from "../../../../src/shared/api"
@ -8,16 +8,10 @@ interface ThinkingBudgetProps {
apiConfiguration: ApiConfiguration
setApiConfigurationField: <K extends keyof ApiConfiguration>(field: K, value: ApiConfiguration[K]) => void
modelInfo?: ModelInfo
provider?: ApiProvider
}
export const ThinkingBudget = ({
apiConfiguration,
setApiConfigurationField,
modelInfo,
provider,
}: ThinkingBudgetProps) => {
const tokens = apiConfiguration?.modelMaxTokens || modelInfo?.maxTokens || 64_000
export const ThinkingBudget = ({ apiConfiguration, setApiConfigurationField, modelInfo }: ThinkingBudgetProps) => {
const tokens = apiConfiguration?.modelMaxTokens || 16_384
const tokensMin = 8192
const tokensMax = modelInfo?.maxTokens || 64_000

View file

@ -92,7 +92,6 @@ describe("ApiOptions", () => {
})
expect(screen.getByTestId("thinking-budget")).toBeInTheDocument()
expect(screen.getByTestId("thinking-budget")).toHaveAttribute("data-provider", "anthropic")
})
it("should show ThinkingBudget for Vertex models that support thinking", () => {
@ -104,7 +103,6 @@ describe("ApiOptions", () => {
})
expect(screen.getByTestId("thinking-budget")).toBeInTheDocument()
expect(screen.getByTestId("thinking-budget")).toHaveAttribute("data-provider", "vertex")
})
it("should not show ThinkingBudget for models that don't support thinking", () => {

View file

@ -1,7 +1,6 @@
import React from "react"
import { render, screen, fireEvent } from "@testing-library/react"
import { ThinkingBudget } from "../ThinkingBudget"
import { ApiProvider, ModelInfo } from "../../../../../src/shared/api"
import { ModelInfo } from "../../../../../src/shared/api"
// Mock Slider component
jest.mock("@/components/ui", () => ({
@ -25,11 +24,11 @@ describe("ThinkingBudget", () => {
supportsPromptCache: true,
supportsImages: true,
}
const defaultProps = {
apiConfiguration: {},
setApiConfigurationField: jest.fn(),
modelInfo: mockModelInfo,
provider: "anthropic" as ApiProvider,
}
beforeEach(() => {
@ -60,7 +59,7 @@ describe("ThinkingBudget", () => {
expect(screen.getAllByTestId("slider")).toHaveLength(2)
})
it("should use modelMaxThinkingTokens field for Anthropic provider", () => {
it("should update modelMaxThinkingTokens", () => {
const setApiConfigurationField = jest.fn()
render(
@ -68,25 +67,6 @@ describe("ThinkingBudget", () => {
{...defaultProps}
apiConfiguration={{ modelMaxThinkingTokens: 4096 }}
setApiConfigurationField={setApiConfigurationField}
provider="anthropic"
/>,
)
const sliders = screen.getAllByTestId("slider")
fireEvent.change(sliders[1], { target: { value: "5000" } })
expect(setApiConfigurationField).toHaveBeenCalledWith("modelMaxThinkingTokens", 5000)
})
it("should use modelMaxThinkingTokens field for Vertex provider", () => {
const setApiConfigurationField = jest.fn()
render(
<ThinkingBudget
{...defaultProps}
apiConfiguration={{ modelMaxThinkingTokens: 4096 }}
setApiConfigurationField={setApiConfigurationField}
provider="vertex"
/>,
)

View file

@ -69,6 +69,32 @@ export interface ExtensionStateContextType extends ExtensionState {
export const ExtensionStateContext = createContext<ExtensionStateContextType | undefined>(undefined)
export const mergeExtensionState = (prevState: ExtensionState, newState: ExtensionState) => {
const {
apiConfiguration: prevApiConfiguration,
customModePrompts: prevCustomModePrompts,
customSupportPrompts: prevCustomSupportPrompts,
experiments: prevExperiments,
...prevRest
} = prevState
const {
apiConfiguration: newApiConfiguration,
customModePrompts: newCustomModePrompts,
customSupportPrompts: newCustomSupportPrompts,
experiments: newExperiments,
...newRest
} = newState
const apiConfiguration = { ...prevApiConfiguration, ...newApiConfiguration }
const customModePrompts = { ...prevCustomModePrompts, ...newCustomModePrompts }
const customSupportPrompts = { ...prevCustomSupportPrompts, ...newCustomSupportPrompts }
const experiments = { ...prevExperiments, ...newExperiments }
const rest = { ...prevRest, ...newRest }
return { ...rest, apiConfiguration, customModePrompts, customSupportPrompts, experiments }
}
export const ExtensionStateContextProvider: React.FC<{ children: React.ReactNode }> = ({ children }) => {
const [state, setState] = useState<ExtensionState>({
version: "",
@ -123,13 +149,8 @@ export const ExtensionStateContextProvider: React.FC<{ children: React.ReactNode
switch (message.type) {
case "state": {
const newState = message.state!
setState((prevState) => ({
...prevState,
...newState,
}))
const config = newState.apiConfiguration
const hasKey = checkExistKey(config)
setShowWelcome(!hasKey)
setState((prevState) => mergeExtensionState(prevState, newState))
setShowWelcome(!checkExistKey(newState.apiConfiguration))
setDidHydrateState(true)
break
}

View file

@ -1,6 +1,11 @@
import React from "react"
// npx jest webview-ui/src/context/__tests__/ExtensionStateContext.test.tsx
import { render, screen, act } from "@testing-library/react"
import { ExtensionStateContextProvider, useExtensionState } from "../ExtensionStateContext"
import { ExtensionState } from "../../../../src/shared/ExtensionMessage"
import { ExtensionStateContextProvider, useExtensionState, mergeExtensionState } from "../ExtensionStateContext"
import { ExperimentId } from "../../../../src/shared/experiments"
import { ApiConfiguration } from "../../../../src/shared/api"
// Test component that consumes the context
const TestComponent = () => {
@ -63,3 +68,43 @@ describe("ExtensionStateContext", () => {
consoleSpy.mockRestore()
})
})
describe("mergeExtensionState", () => {
it("should correctly merge extension states", () => {
const baseState: ExtensionState = {
version: "",
mcpEnabled: false,
enableMcpServerCreation: false,
clineMessages: [],
taskHistory: [],
shouldShowAnnouncement: false,
enableCheckpoints: true,
preferredLanguage: "English",
writeDelayMs: 1000,
requestDelaySeconds: 5,
rateLimitSeconds: 0,
mode: "default",
experiments: {} as Record<ExperimentId, boolean>,
customModes: [],
maxOpenTabsContext: 20,
apiConfiguration: { providerId: "openrouter" } as ApiConfiguration,
}
const prevState: ExtensionState = {
...baseState,
apiConfiguration: { modelMaxTokens: 1234, modelMaxThinkingTokens: 123 },
}
const newState: ExtensionState = {
...baseState,
apiConfiguration: { modelMaxThinkingTokens: 456, modelTemperature: 0.3 },
}
const result = mergeExtensionState(prevState, newState)
expect(result.apiConfiguration).toEqual({
modelMaxTokens: 1234,
modelMaxThinkingTokens: 456,
modelTemperature: 0.3,
})
})
})

View file

@ -0,0 +1,42 @@
import { useEffect, useRef } from "react"
type VoidFn = () => void
/**
* Runs `effectRef.current()` after `delay` ms whenever any of the `deps` change,
* but cancels/re-schedules if they change again before the delay.
*/
export function useDebounceEffect(effect: VoidFn, delay: number, deps: any[]) {
const callbackRef = useRef<VoidFn>(effect)
const timeoutRef = useRef<NodeJS.Timeout | null>(null)
// Keep callbackRef current
useEffect(() => {
callbackRef.current = effect
}, [effect])
useEffect(() => {
// Clear any queued call
if (timeoutRef.current) {
clearTimeout(timeoutRef.current)
}
// Schedule a new call
timeoutRef.current = setTimeout(() => {
// always call the *latest* version of effect
callbackRef.current()
}, delay)
// Cleanup on unmount or next effect
return () => {
if (timeoutRef.current) {
clearTimeout(timeoutRef.current)
}
}
// We want to re‐schedule if any item in `deps` changed,
// or if `delay` changed.
// eslint-disable-next-line react-hooks/exhaustive-deps
}, [delay, ...deps])
}