mirror of
https://github.com/RooVetGit/Roo-Code.git
synced 2026-10-10 03:28:03 +00:00
Merge branch 'RooVetGit:main' into human-relay
This commit is contained in:
commit
626827ab3f
76 changed files with 6846 additions and 1782 deletions
|
|
@ -1,5 +0,0 @@
|
|||
---
|
||||
"roo-cline": patch
|
||||
---
|
||||
|
||||
Delete task confirmation enhancements
|
||||
5
.changeset/healthy-buckets-attack.md
Normal file
5
.changeset/healthy-buckets-attack.md
Normal file
|
|
@ -0,0 +1,5 @@
|
|||
---
|
||||
"roo-cline": patch
|
||||
---
|
||||
|
||||
ExtensionStateContext does not correctly merge state
|
||||
5
.changeset/weak-cameras-hope.md
Normal file
5
.changeset/weak-cameras-hope.md
Normal file
|
|
@ -0,0 +1,5 @@
|
|||
---
|
||||
"roo-cline": patch
|
||||
---
|
||||
|
||||
Default middle-out compression to on for OpenRouter
|
||||
|
|
@ -1,5 +0,0 @@
|
|||
---
|
||||
"roo-cline": patch
|
||||
---
|
||||
|
||||
Prettier thinking blocks
|
||||
42
.github/pull_request_template.md
vendored
42
.github/pull_request_template.md
vendored
|
|
@ -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. -->
|
||||
|
|
|
|||
2
.github/workflows/changeset-release.yml
vendored
2
.github/workflows/changeset-release.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
18
.github/workflows/code-qa.yml
vendored
18
.github/workflows/code-qa.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
5
.github/workflows/marketplace-publish.yml
vendored
5
.github/workflows/marketplace-publish.yml
vendored
|
|
@ -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 }}
|
||||
|
|
|
|||
|
|
@ -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/**
|
||||
|
|
|
|||
22
CHANGELOG.md
22
CHANGELOG.md
|
|
@ -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!)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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',
|
||||
|
|
@ -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
2387
e2e/package-lock.json
generated
Normal file
File diff suppressed because it is too large
Load diff
21
e2e/package.json
Normal file
21
e2e/package.json
Normal 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"
|
||||
}
|
||||
}
|
||||
|
|
@ -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 {
|
||||
|
|
@ -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"]
|
||||
}
|
||||
|
|
@ -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
948
package-lock.json
generated
File diff suppressed because it is too large
Load diff
37
package.json
37
package.json
|
|
@ -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",
|
||||
|
|
|
|||
257
src/api/__tests__/index.test.ts
Normal file
257
src/api/__tests__/index.test.ts
Normal 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,
|
||||
})
|
||||
})
|
||||
})
|
||||
|
|
@ -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 }
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
})
|
||||
})
|
||||
})
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
})
|
||||
})
|
||||
|
|
|
|||
|
|
@ -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" }])
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
64
src/api/providers/base-provider.ts
Normal file
64
src/api/providers/base-provider.ts
Normal 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)
|
||||
}
|
||||
}
|
||||
|
|
@ -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
|
||||
|
|
|
|||
3
src/api/providers/constants.ts
Normal file
3
src/api/providers/constants.ts
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
export const ANTHROPIC_DEFAULT_MAX_TOKENS = 8192
|
||||
|
||||
export const DEEP_SEEK_DEFAULT_TEMPERATURE = 0.6
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 = {
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 || ""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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) {
|
||||
|
|
|
|||
|
|
@ -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}`)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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 = {
|
||||
|
|
|
|||
338
src/api/transform/__tests__/vertex-gemini-format.test.ts
Normal file
338
src/api/transform/__tests__/vertex-gemini-format.test.ts
Normal 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",
|
||||
)
|
||||
})
|
||||
})
|
||||
83
src/api/transform/vertex-gemini-format.ts
Normal file
83
src/api/transform/vertex-gemini-format.ts
Normal 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),
|
||||
}
|
||||
}
|
||||
|
|
@ -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) {
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}),
|
||||
)
|
||||
})
|
||||
})
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
93
src/exports/cline.d.ts
vendored
93
src/exports/cline.d.ts
vendored
|
|
@ -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>
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
})
|
||||
})
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
},
|
||||
})
|
||||
})
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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++
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
1127
webview-ui/package-lock.json
generated
1127
webview-ui/package-lock.json
generated
File diff suppressed because it is too large
Load diff
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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" />
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -45,6 +45,7 @@ describe("ChatTextArea", () => {
|
|||
onHeightChange: jest.fn(),
|
||||
mode: defaultModeSlug,
|
||||
setMode: jest.fn(),
|
||||
modeShortcutText: "(⌘. for next mode)",
|
||||
}
|
||||
|
||||
beforeEach(() => {
|
||||
|
|
|
|||
|
|
@ -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} />
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
|
|
|
|||
227
webview-ui/src/components/common/MermaidBlock.tsx
Normal file
227
webview-ui/src/components/common/MermaidBlock.tsx
Normal 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, "<").replace(/>/g, ">")
|
||||
})
|
||||
.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;
|
||||
`
|
||||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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", () => {
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
/>,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
})
|
||||
})
|
||||
})
|
||||
|
|
|
|||
42
webview-ui/src/utils/useDebounceEffect.ts
Normal file
42
webview-ui/src/utils/useDebounceEffect.ts
Normal 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])
|
||||
}
|
||||
Loading…
Add table
Reference in a new issue