Compare commits

..
36 Commits
Author SHA1 Message Date
leokun 75b7ea9cc8 chore(release): bump desktop to 0.1.6 2026-09-02 14:03:35 +08:00
leokun 0fd5e9d6f2 fix(desktop): update content security policy to allow HTTPS images
- Modified the content security policy in `tauri.conf.json` to include `https:` in the `img-src` directive, enhancing security by allowing images from secure sources.
2026-09-02 13:56:28 +08:00
leookun ddaa61c827 fix(desktop): update Windows executable in place 2026-09-02 13:02:11 +08:00
leookun 919c1d8032 fix: update ADS_ENDPOINT to production URL
- Changed the `ADS_ENDPOINT` from a local server URL to the production URL for ads.
- Commented out the local server URL for clarity and future reference.
2026-09-02 12:49:57 +08:00
leokun 669f129dcd Merge branch 'main' of github.com:leookun/cursor-byok 2026-09-02 11:03:41 +08:00
leookun f22c7b6680 feat: integrate app version into control service and update ads endpoint
- Added `app_version` field to `ControlService` and updated its initialization to include the app version.
- Modified the `ADS_ENDPOINT` to point to a local server for development purposes.
- Refactored conversation command and output handling to utilize a new `RunFinish` enum for better state management.
- Enhanced the conversation runtime to handle queued user messages after a turn has ended, ensuring smooth transitions between turns.
- Added tests to validate the new behavior of queued messages and transport handling.
2026-09-02 10:23:50 +08:00
leookun 5cdf642dd1 feat: enhance bidi request handling and observability tracing
- Updated the `append` function to include a flag for replacing closing requests, improving request handling.
- Refactored the `run_sse_handler` and `bidi_handler` functions to utilize a new tracing mechanism, enhancing observability.
- Introduced a new `trace_outcome` function to standardize tracing outcomes for requests.
- Removed the `CursorTraceRecorder` in favor of a new `CursorTraceService` for better performance and non-blocking behavior.
- Added tests to validate the new tracing functionality and ensure correct behavior during request processing.
2026-09-02 01:04:54 +08:00
leokun 2dad593263 feat: add resource limits management and network client integration
- Introduced a new module for managing process resource limits, specifically for raising the open file limit on Unix systems.
- Added a `NetworkClients` struct to handle reusable outbound HTTP clients, improving network request management.
- Updated various components, including `ControlService` and `CursorProxy`, to utilize the new network client structure for better client handling.
- Enhanced the API router to accept network clients, ensuring consistent client usage across different services.
- Added tests to validate the integration of network clients and resource limits functionality.
2026-09-01 21:03:02 +08:00
leokun 76417e005b feat: add usage snapshot event and enhance compaction logic
- Introduced `UsageSnapshot` event to track token usage during conversation runs.
- Updated `RunEngine` to emit usage snapshots, providing better visibility into token consumption.
- Refactored compaction logic to utilize a new `compaction_estimate` function for improved token budget management.
- Added tests to validate timeout constants for blob synchronization and ensure correct behavior of usage tracking during compaction.
2026-09-01 20:04:35 +08:00
leokun e768980dad feat: enhance reasoning replay functionality and integrate call recording
- Added a new test to validate the projection of reasoning response items to valid input items in the Codex API.
- Introduced `CallRecorder` to track network requests and responses during plugin interactions.
- Updated the `PluginRegistry` and `PluginWorker` to support call recording, ensuring that reasoning items are correctly processed and recorded.
- Refactored the `responses_input` function to handle reasoning items more effectively, improving the overall response handling logic.
2026-09-01 17:43:36 +08:00
leokun 2c63bd845a feat: track interaction events during automatic compaction
- Added tracking for interaction events in the `Output` struct, including `summary_started` and `token_delta`.
- Updated the `run` function to push relevant interaction events to the `interaction_events` vector.
- Enhanced the automatic compaction test to verify the immediate reset of cursor usage and the correct logging of interaction events.
2026-09-01 16:54:45 +08:00
leokun 6e74637c69 Merge branch 'main' of github.com:leookun/cursor-byok 2026-09-01 16:08:34 +08:00
leookun d004139526 feat: implement context usage anchor for improved token estimation
- Introduced `ContextUsageAnchor` struct to track context input tokens and message count for conversations.
- Updated token estimation functions to utilize the context usage anchor, enhancing accuracy in estimating tokens for projected messages.
- Refactored compaction logic to incorporate context usage anchor, allowing for more efficient management of token budgets during model runs.
- Added tests to validate the behavior of the context usage anchor across different scenarios, including model switching and message additions.
2026-09-01 16:07:51 +08:00
leokun 8c6c415a84 feat: enhance token usage merging and total token calculation
- Updated the `merge_usage` function to include `total_tokens` in the usage merging process.
- Implemented logic to calculate `total_tokens` based on the sum of `context_input_tokens` and `output_tokens`.
- Added a new test to verify that streamed usage correctly includes cached input in the total token count.
2026-09-01 11:13:00 +08:00
leokun d83e14af9a refactor: remove retry_count from ProviderConfig and enhance error handling in tool execution
- Removed the `retry_count` field from `ProviderConfig` as it is no longer needed.
- Introduced `argument_error` field in `ToolCall` to capture errors related to tool arguments.
- Updated various components to handle argument errors more gracefully, including in the `ToolDispatcher` and `ConversationOutput`.
- Enhanced tests to validate the new error handling and ensure proper functionality of tool calls.
2026-09-01 10:14:53 +08:00
leokun 29fde7d7c7 feat: enhance context token estimation and compaction logic
- Added `estimate_context_tokens` function to calculate provider-visible context size based on prompt specifications and projected messages.
- Updated `CheckpointBuilder` to record estimated context tokens during message processing.
- Refactored compaction logic to utilize the new token estimation, ensuring proper context management during model runs.
- Introduced tests to validate context estimation and compaction behavior under various scenarios.
2026-09-01 10:10:09 +08:00
leokun ee2592c469 Merge branch 'main' of github.com:leookun/cursor-byok 2026-08-31 16:15:58 +08:00
leokun 49c1fb6378 feat: add Task tool functionality
- Introduced a new `Task` presentation type in the `ToolCallStream` to handle task-related projections.
- Implemented the `TaskProjection` struct with fields for description, prompt, subagent type, model, resume, and environment.
- Added a `project` method to `TaskProjection` to process task-related events and generate interaction updates.
- Created a `task_partial` function to format task updates for the agent server message.
- Included unit tests to verify the correct behavior of task description projections.
2026-08-31 16:15:42 +08:00
leokun 788868f8b9 Merge pull request #385 from kevin9327/fix/runtime-user-message-injection-leak
fix(conversation): clear pending runtime user-message injections
2026-08-31 16:14:27 +08:00
leokun 4c3fe230ce Merge remote-tracking branch 'origin/main' into pr-385-merge
# Conflicts:
#	server/tests/interrupt.rs
2026-08-31 16:12:11 +08:00
leokun ac14245d19 Merge pull request #383 from kevin9327/fix/editnotebook-empty-old-string
fix(tools): reject empty old_string in EditNotebook
2026-08-31 16:06:42 +08:00
leokun 84addec26a Merge pull request #384 from kevin9327/fix/bash-shell-alias
fix(tools): complete the bash shell alias in the tool codec
2026-08-31 16:06:18 +08:00
leokun 5de547041c Merge branch 'main' of github.com:leookun/cursor-byok 2026-08-31 15:38:11 +08:00
leokun 8bd0d70add fix: plugin effort compress 2026-08-31 15:38:02 +08:00
leokun 3ac4402a86 Update cursor.md 2026-08-31 15:18:43 +08:00
leokun e535a98945 Update cursor.md 2026-08-31 15:10:37 +08:00
leokun 45e694fd63 Merge pull request #386 from kevin9327/fix/empty-tool-arguments
fix: handle tool calls with empty arguments
2026-08-31 13:57:45 +08:00
leookun 9120b90be7 chore: remove deprecated server_backup files
- Deleted unused build script, Cargo.toml, and migration files to clean up the project structure.
- Removed prompt files related to cursor tools and agent modes to streamline the codebase.
- This cleanup helps improve maintainability and reduces clutter in the repository.
2026-08-30 23:54:19 +08:00
leookun b807608bf3 chore(release): bump desktop to v0.1.5 2026-08-30 23:38:50 +08:00
leookun e7a1cca4c6 feat: add group name functionality to models
- Introduced a new `group_name` field in the model configuration to allow for custom provider-group display names.
- Updated the `CursorModelCards`, `CursorModelEditor`, and `CursorSettingsPage` components to support group settings.
- Enhanced the UI to include group settings options, allowing users to modify group names and associated configurations.
- Added localization strings for new group settings features in both English and Chinese.
- Implemented a database migration to add the `group_name` column to the model configurations.
2026-08-30 23:28:05 +08:00
leookun 76baa3b0e7 Merge branch 'feat/plugin' 2026-08-30 22:19:55 +08:00
leokun fc79adbb43 Merge pull request #382 from leookun/feat/plugin
feat: add plugin mode
2026-08-30 21:29:54 +08:00
kevin9327andClaude Opus 4.8 e673a034df fix: handle tool calls with empty arguments
A tool call that carries no arguments streams no argument text, so
`arguments_text` is empty and `from_str("")` fails with `EOF while parsing
a value`, aborting the whole run. The model cycle already guards this, but
two other consumers did not:

- `ConversationOutput` re-parses the streamed text on `ToolCallEnd`; and
- `create_tool_round` stored the empty text verbatim in the
  `arguments_json` column, so re-loading the round (`commit_tool_result`
  and the round loader) then failed on `from_str("")`.

Treat empty argument text as an empty object in the output projection, and
persist `{}` for it so the `arguments_json` column always holds valid JSON.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
2026-08-30 19:38:39 +09:00
kevin9327andClaude Opus 4.8 6ac666da4f fix(conversation): clear pending runtime user-message injections
Runtime user messages and context injections both queue into
`pending_injections`, but with different keys: injections use the raw
injection id (committed under `inject-context:{id}`) while user messages
use the full `user-message:{id}` event id. The commit-correlation handler
only stripped the `inject-context:` prefix, so a user message's entry was
never removed.

Consequences:
- the client never received `ContextInjectionDelivered` /
  `UserMessageAppended` for the message; and
- `pending_injections` stayed non-empty, so every later `ExecuteToolRound`
  was detached without dispatching its tools and `tool_round::execute`
  blocked forever -- a hung turn whenever the model made a tool call after
  the interruption.

Derive the lookup key by stripping the injection prefix when present and
otherwise using the event id verbatim, so both kinds are cleared and their
delivered/appended events fire.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
2026-08-30 19:23:57 +09:00
kevin9327andClaude Opus 4.8 9e2d22418b fix(tools): complete the bash shell alias in the tool codec
The tool dispatcher already treats `bash`/`Bash` as an alias of `Shell`
(routing, `is_shell_tool`, and `block_until_ms` normalization), but the
codec only matched `shell`:

- `tool_placeholder` returned `unsupported tool: bash`, which aborts the
  turn while streaming the tool call, before it ever runs;
- `request` returned `tool bash is not executed through ExecServerMessage`
  (after already reserving an exec slot); and
- `stream_closed` built its shell-specific error result only for `Shell`.

Anthropic models frequently emit `Bash` even when the tool is advertised
as `Shell`, so the alias must hold across the codec. Match `bash` wherever
the codec special-cases `shell`.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
2026-08-30 19:15:29 +09:00
kevin9327andClaude Opus 4.8 97ee138de8 fix(tools): reject empty old_string in EditNotebook
StrReplace rejects an empty `old_string`, but the EditNotebook cell-edit
path did not. Because `str::match_indices("")` matches at every byte
boundary, editing a non-empty cell with an empty `old_string` failed with
a misleading "old_string is not unique in the notebook cell; found N
occurrences" error, and editing an empty cell silently prepended
`new_string`.

Add the same guard StrReplace already uses so both edit tools reject an
empty `old_string` consistently.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
2026-08-30 19:07:28 +09:00
327 changed files with 8022 additions and 49285 deletions
@@ -0,0 +1,71 @@
import { readFile, writeFile } from "node:fs/promises";
import { basename, resolve } from "node:path";
import { pathToFileURL } from "node:url";
function readOptions(args) {
const options = new Map();
for (let index = 0; index < args.length; index += 2) {
const name = args[index];
const value = args[index + 1];
if (!name?.startsWith("--") || value === undefined) {
throw new Error(`invalid argument near ${name ?? "end of command"}`);
}
options.set(name.slice(2), value);
}
return options;
}
function required(options, name) {
const value = options.get(name)?.trim();
if (!value) throw new Error(`--${name} is required`);
return value;
}
export function generatePortableUpdate({ version, repository, assetName, signature }) {
const normalizedVersion = version.replace(/^v/, "");
if (!/^\d+\.\d+\.\d+(?:-[0-9A-Za-z.-]+)?$/.test(normalizedVersion)) {
throw new Error(`invalid semantic version: ${normalizedVersion}`);
}
if (!/^[^/\s]+\/[^/\s]+$/.test(repository)) {
throw new Error(`invalid GitHub repository: ${repository}`);
}
if (!assetName || basename(assetName) !== assetName) {
throw new Error("asset name must be a file name");
}
if (!signature.trim()) throw new Error("signature is required");
return {
version: normalizedVersion,
notes: `Cursor BYOK v${normalizedVersion}`,
pub_date: new Date().toISOString(),
platforms: {
"windows-x86_64": {
signature: signature.trim(),
url: `https://github.com/${repository}/releases/download/v${normalizedVersion}/${encodeURIComponent(assetName)}`,
},
},
};
}
async function main() {
const options = readOptions(process.argv.slice(2));
const version = required(options, "version");
const repository = required(options, "repository");
const asset = required(options, "asset");
const signaturePath = resolve(required(options, "signature"));
const output = resolve(required(options, "output"));
const manifest = generatePortableUpdate({
version,
repository,
assetName: basename(asset),
signature: await readFile(signaturePath, "utf8"),
});
await writeFile(output, `${JSON.stringify(manifest, null, 2)}\n`);
}
if (process.argv[1] && import.meta.url === pathToFileURL(resolve(process.argv[1])).href) {
main().catch((error) => {
console.error(error instanceof Error ? error.message : String(error));
process.exitCode = 1;
});
}
@@ -0,0 +1,41 @@
import assert from "node:assert/strict";
import test from "node:test";
import { generatePortableUpdate } from "./generate-portable-update.mjs";
test("generates a signed Windows portable updater manifest", () => {
const manifest = generatePortableUpdate({
version: "v1.2.3-beta.1",
repository: "owner/repository",
assetName: "cursor-byok-1.2.3-beta.1-windows-amd64.zip",
signature: "signed-payload\n",
});
assert.equal(manifest.version, "1.2.3-beta.1");
assert.deepEqual(Object.keys(manifest.platforms), ["windows-x86_64"]);
assert.equal(manifest.platforms["windows-x86_64"].signature, "signed-payload");
assert.equal(
manifest.platforms["windows-x86_64"].url,
"https://github.com/owner/repository/releases/download/v1.2.3-beta.1/cursor-byok-1.2.3-beta.1-windows-amd64.zip",
);
});
test("rejects invalid inputs", () => {
assert.throws(() => generatePortableUpdate({
version: "latest",
repository: "owner/repository",
assetName: "update.zip",
signature: "signature",
}), /semantic version/);
assert.throws(() => generatePortableUpdate({
version: "1.2.3",
repository: "owner/repository",
assetName: "../update.zip",
signature: "signature",
}), /file name/);
assert.throws(() => generatePortableUpdate({
version: "1.2.3",
repository: "owner/repository",
assetName: "update.zip",
signature: " ",
}), /signature/);
});
+29 -2
View File
@@ -158,14 +158,27 @@ jobs:
mkdir -p legacy-update
tar -czf "legacy-update/cursor-byok-${VERSION}-linux-amd64.tar.gz" -C target/release cursor-byok-desktop
- name: Package legacy Windows updater asset
- name: Package and sign legacy Windows updater asset
if: matrix.platform == 'windows-x86_64'
shell: pwsh
env:
VERSION: ${{ needs.prepare.outputs.version }}
TAURI_SIGNING_PRIVATE_KEY: ${{ secrets.TAURI_SIGNING_PRIVATE_KEY }}
TAURI_SIGNING_PRIVATE_KEY_PASSWORD: ${{ secrets.TAURI_SIGNING_PRIVATE_KEY_PASSWORD }}
run: |
New-Item -ItemType Directory -Force legacy-update | Out-Null
Compress-Archive -LiteralPath target/release/cursor-byok-desktop.exe -DestinationPath "legacy-update/cursor-byok-$env:VERSION-windows-amd64.zip"
$asset = "legacy-update/cursor-byok-$env:VERSION-windows-amd64.zip"
Compress-Archive -LiteralPath target/release/cursor-byok-desktop.exe -DestinationPath $asset
Push-Location apps/desktop
npm exec tauri signer sign -- "../../$asset"
Pop-Location
$entries = @(tar -tf $asset)
if ($entries.Count -ne 1 -or [System.IO.Path]::GetFileName($entries[0]) -ne 'cursor-byok-desktop.exe') {
throw "Windows updater archive must contain only cursor-byok-desktop.exe"
}
if (!(Test-Path "$asset.sig")) {
throw "Windows updater archive signature was not generated"
}
- name: Package legacy macOS updater asset
if: contains(matrix.platform, 'macos')
@@ -213,6 +226,20 @@ jobs:
--output legacy-update/update.json \
--notes "Cursor BYOK v${VERSION}"
- name: Generate signed Windows portable update manifest
env:
VERSION: ${{ needs.prepare.outputs.version }}
run: |
asset="cursor-byok-${VERSION}-windows-amd64.zip"
test -f "legacy-update/${asset}"
test -f "legacy-update/${asset}.sig"
node .github/scripts/generate-portable-update.mjs \
--version "${VERSION}" \
--repository "${GITHUB_REPOSITORY}" \
--asset "${asset}" \
--signature "legacy-update/${asset}.sig" \
--output legacy-update/portable-latest.json
- name: Normalize Tauri updater download URLs
env:
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
Generated
+5 -1
View File
@@ -1172,10 +1172,11 @@ checksum = "52560adf09603e58c9a7ee1fe1dcb95a16927b17c127f0ac02d6e768a0e25bc1"
[[package]]
name = "cursor-byok-desktop"
version = "0.1.5-beta.1"
version = "0.1.6"
dependencies = [
"axum",
"cursor-server",
"libc",
"rfd",
"serde",
"serde_json",
@@ -1187,12 +1188,15 @@ dependencies = [
"tauri-plugin-process",
"tauri-plugin-single-instance",
"tauri-plugin-updater",
"tempfile",
"tokio",
"tokio-util",
"tracing",
"tracing-appender",
"tracing-subscriber",
"url",
"windows-sys 0.61.2",
"zip",
]
[[package]]
+2 -2
View File
@@ -1,12 +1,12 @@
{
"name": "cursor-byok-desktop",
"version": "0.1.5-beta.1",
"version": "0.1.6",
"lockfileVersion": 3,
"requires": true,
"packages": {
"": {
"name": "cursor-byok-desktop",
"version": "0.1.5-beta.1",
"version": "0.1.6",
"license": "MIT",
"dependencies": {
"@floating-ui/dom": "^1.8.0",
+2 -2
View File
@@ -1,6 +1,6 @@
{
"name": "cursor-byok-desktop",
"version": "0.1.5-beta.1",
"version": "0.1.6",
"description": "Cursor BYOK desktop management application",
"type": "module",
"scripts": {
@@ -8,7 +8,7 @@
"dev": "vite",
"typecheck": "tsc --noEmit",
"typecheck:node": "tsc --noEmit -p tsconfig.node.json",
"i18n:scan": "STATIC_I18N_SCAN=true vite build",
"i18n:scan": "cross-env STATIC_I18N_SCAN=true vite build",
"build": "vite build",
"build:demo": "npm run typecheck && npm run typecheck:node && vite build --config vite.demo.config.ts",
"check": "npm run typecheck && npm run typecheck:node && npm run build",
+7 -1
View File
@@ -1,6 +1,6 @@
[package]
name = "cursor-byok-desktop"
version = "0.1.5-beta.1"
version = "0.1.6"
edition = "2021"
publish = false
@@ -14,6 +14,7 @@ tauri-build = { version = "2", features = [] }
[dependencies]
axum = "0.8"
cursor-server = { path = "../../../server" }
libc = "0.2"
rfd = "0.15"
serde = { version = "1", features = ["derive"] }
serde_json = "1"
@@ -24,9 +25,14 @@ tauri-plugin-opener = "2"
tauri-plugin-autostart = "2"
tauri-plugin-process = "2"
tauri-plugin-updater = "2"
tempfile = "3"
tokio = { version = "1", features = ["time"] }
tokio-util = "0.7"
tracing = "0.1"
tracing-appender = "0.2"
tracing-subscriber = { version = "0.3", features = ["env-filter"] }
url = "2"
zip = { version = "4", default-features = false, features = ["deflate"] }
[target.'cfg(windows)'.dependencies]
windows-sys = { version = "0.61", features = ["Win32_Foundation", "Win32_System_Threading"] }
+5 -1
View File
@@ -1,5 +1,9 @@
fn main() {
let manifest = tauri_build::AppManifest::new().commands(&["open_terminal_with_command"]);
let manifest = tauri_build::AppManifest::new().commands(&[
"open_terminal_with_command",
"check_portable_update",
"install_portable_update",
]);
tauri_build::try_build(tauri_build::Attributes::new().app_manifest(manifest))
.expect("failed to build Tauri application")
}
@@ -18,6 +18,8 @@
"core:window:allow-close",
"core:app:allow-set-dock-visibility",
"allow-open-terminal-with-command",
"allow-check-portable-update",
"allow-install-portable-update",
"clipboard-manager:allow-write-text",
"autostart:default",
"process:allow-restart",
@@ -0,0 +1,11 @@
# Automatically generated - DO NOT EDIT!
[[permission]]
identifier = "allow-check-portable-update"
description = "Enables the check_portable_update command without any pre-configured scope."
commands.allow = ["check_portable_update"]
[[permission]]
identifier = "deny-check-portable-update"
description = "Denies the check_portable_update command without any pre-configured scope."
commands.deny = ["check_portable_update"]
@@ -0,0 +1,11 @@
# Automatically generated - DO NOT EDIT!
[[permission]]
identifier = "allow-install-portable-update"
description = "Enables the install_portable_update command without any pre-configured scope."
commands.allow = ["install_portable_update"]
[[permission]]
identifier = "deny-install-portable-update"
description = "Denies the install_portable_update command without any pre-configured scope."
commands.deny = ["install_portable_update"]
+23 -1
View File
@@ -149,6 +149,23 @@ pub fn run() -> ExitCode {
return ExitCode::FAILURE;
}
};
#[cfg(unix)]
{
let open_file_limit = match crate::resource_limits::raise_open_file_limit() {
Ok(limit) => limit,
Err(error) => {
diagnostics.report_fatal(&error);
return ExitCode::FAILURE;
}
};
tracing::info!(
requested = crate::resource_limits::REQUESTED_OPEN_FILE_LIMIT,
previous = open_file_limit.previous,
effective = open_file_limit.effective,
hard = open_file_limit.hard,
"open file limit configured"
);
}
tracing::info!(
version = env!("CARGO_PKG_VERSION"),
os = std::env::consts::OS,
@@ -160,7 +177,11 @@ pub fn run() -> ExitCode {
let started_by_autostart = std::env::args_os().any(|arg| arg == AUTOSTART_ARG);
let app = tauri::Builder::default()
.invoke_handler(tauri::generate_handler![open_terminal_with_command])
.invoke_handler(tauri::generate_handler![
open_terminal_with_command,
crate::update::check_portable_update,
crate::update::install_portable_update,
])
.plugin(tauri_plugin_single_instance::init(|app, args, _| {
if !args.iter().any(|arg| arg == AUTOSTART_ARG) {
tray::show_main_window(app);
@@ -228,6 +249,7 @@ pub fn run() -> ExitCode {
window.set_focus()?;
}
tray::create(app)?;
crate::update::signal_ready_if_requested()?;
Ok(())
})
.build(tauri::generate_context!());
+8 -1
View File
@@ -1,7 +1,14 @@
mod desktop;
#[cfg(not(dev))]
mod frontend;
mod resource_limits;
mod startup;
mod tray;
mod update;
pub use desktop::run;
pub fn run() -> std::process::ExitCode {
if let Some(exit_code) = update::run_replacement_if_requested() {
return exit_code;
}
desktop::run()
}
@@ -0,0 +1,56 @@
//! Configures process resource limits before the desktop runtime starts.
#[cfg(unix)]
use std::io;
#[cfg(unix)]
pub(crate) const REQUESTED_OPEN_FILE_LIMIT: u64 = 65_536;
#[cfg(unix)]
pub(crate) struct OpenFileLimit {
pub(crate) previous: u64,
pub(crate) effective: u64,
pub(crate) hard: u64,
}
#[cfg(unix)]
pub(crate) fn raise_open_file_limit() -> io::Result<OpenFileLimit> {
let mut limits = libc::rlimit {
rlim_cur: 0,
rlim_max: 0,
};
// SAFETY: `limits` points to writable memory for one `rlimit` value.
if unsafe { libc::getrlimit(libc::RLIMIT_NOFILE, &mut limits) } != 0 {
return Err(io::Error::last_os_error());
}
let previous = limits.rlim_cur;
let target = limits
.rlim_max
.min(REQUESTED_OPEN_FILE_LIMIT as libc::rlim_t);
if previous < target {
let requested = libc::rlimit {
rlim_cur: target,
rlim_max: limits.rlim_max,
};
// SAFETY: `requested` is a valid `rlimit` value and does not raise the hard limit.
if unsafe { libc::setrlimit(libc::RLIMIT_NOFILE, &requested) } != 0 {
return Err(io::Error::last_os_error());
}
}
let mut effective = libc::rlimit {
rlim_cur: 0,
rlim_max: 0,
};
// SAFETY: `effective` points to writable memory for one `rlimit` value.
if unsafe { libc::getrlimit(libc::RLIMIT_NOFILE, &mut effective) } != 0 {
return Err(io::Error::last_os_error());
}
Ok(OpenFileLimit {
previous: previous as u64,
effective: effective.rlim_cur as u64,
hard: effective.rlim_max as u64,
})
}
+269
View File
@@ -0,0 +1,269 @@
use std::{
fs::{self, OpenOptions},
io::{Cursor, Read, Write},
path::{Path, PathBuf},
process::{Command, ExitCode},
};
use serde::Serialize;
use tauri::AppHandle;
#[cfg(target_os = "windows")]
use tauri_plugin_updater::UpdaterExt;
#[cfg(target_os = "windows")]
mod replacement;
const PORTABLE_UPDATE_ENDPOINT: &str =
"https://github.com/leookun/cursor-byok/releases/latest/download/portable-latest.json";
const WINDOWS_PAYLOAD_NAME: &str = "cursor-byok-desktop.exe";
#[derive(Serialize)]
#[serde(rename_all = "camelCase")]
pub(crate) struct PortableUpdateInfo {
version: String,
}
pub fn run_replacement_if_requested() -> Option<ExitCode> {
#[cfg(target_os = "windows")]
{
match replacement::request_from_args() {
Ok(Some(request)) => return Some(replacement::run(request)),
Ok(None) => {}
Err(error) => {
eprintln!("invalid portable update replacement request: {error}");
return Some(ExitCode::FAILURE);
}
}
}
None
}
pub(crate) fn signal_ready_if_requested() -> std::io::Result<()> {
#[cfg(target_os = "windows")]
if let Some(path) = replacement::ready_marker_from_args() {
if let Some(parent) = path.parent() {
fs::create_dir_all(parent)?;
}
fs::write(path, b"ready")?;
}
Ok(())
}
#[tauri::command]
pub(crate) async fn check_portable_update(
app: AppHandle,
) -> Result<Option<PortableUpdateInfo>, String> {
#[cfg(target_os = "windows")]
{
let update = portable_update(&app).await?;
return Ok(update.map(|update| PortableUpdateInfo {
version: update.version,
}));
}
#[cfg(not(target_os = "windows"))]
{
let _ = app;
Err("portable updates are only supported on Windows".into())
}
}
#[tauri::command]
pub(crate) async fn install_portable_update(
app: AppHandle,
expected_version: String,
) -> Result<(), String> {
#[cfg(target_os = "windows")]
{
let update = portable_update(&app)
.await?
.ok_or_else(|| "the selected update is no longer available".to_string())?;
if update.version != expected_version {
return Err(format!(
"available update changed from {expected_version} to {}",
update.version
));
}
let target = std::env::current_exe()
.map_err(|error| format!("failed to locate the running executable: {error}"))?;
ensure_target_writable(&target)
.map_err(|error| format!("the application directory is not writable: {error}"))?;
let bytes = update
.download(|_, _| {}, || {})
.await
.map_err(|error| format!("failed to download or verify the update: {error}"))?;
let payload = extract_windows_payload(&bytes)
.map_err(|error| format!("invalid Windows update archive: {error}"))?;
let staged = stage_payload(&target, &payload)
.map_err(|error| format!("failed to stage the update: {error}"))?;
let handshake = staged.with_extension("started");
let _ = fs::remove_file(&handshake);
let mut replacement = Command::new(&staged)
.arg("--apply-portable-update")
.arg("--update-target")
.arg(&target)
.arg("--update-wait-pid")
.arg(std::process::id().to_string())
.arg("--update-handshake")
.arg(&handshake)
.spawn()
.map_err(|error| format!("failed to start the update replacement process: {error}"))?;
wait_for_replacement_start(&mut replacement, &handshake).await?;
app.exit(0);
Ok(())
}
#[cfg(not(target_os = "windows"))]
{
let _ = (app, expected_version);
Err("portable updates are only supported on Windows".into())
}
}
#[cfg(target_os = "windows")]
async fn portable_update(app: &AppHandle) -> Result<Option<tauri_plugin_updater::Update>, String> {
let endpoint = PORTABLE_UPDATE_ENDPOINT
.parse()
.map_err(|error| format!("invalid portable update endpoint: {error}"))?;
let updater = app
.updater_builder()
.endpoints(vec![endpoint])
.map_err(|error| format!("failed to configure the updater: {error}"))?
.build()
.map_err(|error| format!("failed to initialize the updater: {error}"))?;
updater
.check()
.await
.map_err(|error| format!("failed to check for updates: {error}"))
}
#[cfg(target_os = "windows")]
async fn wait_for_replacement_start(
child: &mut std::process::Child,
handshake: &Path,
) -> Result<(), String> {
let wait = async {
loop {
if handshake.is_file() {
return Ok(());
}
if let Some(status) = child
.try_wait()
.map_err(|error| format!("failed to inspect replacement process: {error}"))?
{
return Err(format!(
"update replacement process exited before it was ready: {status}"
));
}
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
}
};
let result = match tokio::time::timeout(std::time::Duration::from_secs(5), wait).await {
Ok(result) => result,
Err(_) => {
let _ = child.kill();
let _ = child.wait();
return Err("update replacement process did not become ready".into());
}
};
if result.is_err() {
let _ = child.kill();
let _ = child.wait();
}
result
}
fn ensure_target_writable(target: &Path) -> std::io::Result<()> {
let parent = target
.parent()
.ok_or_else(|| std::io::Error::other("application executable has no parent directory"))?;
let probe = parent.join(format!(
".cursor-byok-update-write-test-{}",
std::process::id()
));
let mut file = OpenOptions::new()
.write(true)
.create_new(true)
.open(&probe)?;
file.write_all(b"test")?;
drop(file);
fs::remove_file(probe)
}
fn extract_windows_payload(bytes: &[u8]) -> Result<Vec<u8>, String> {
let mut archive = zip::ZipArchive::new(Cursor::new(bytes))
.map_err(|error| format!("failed to open ZIP: {error}"))?;
if archive.len() != 1 {
return Err("archive must contain exactly one file".into());
}
let mut entry = archive
.by_index(0)
.map_err(|error| format!("failed to read ZIP entry: {error}"))?;
let name = Path::new(entry.name())
.file_name()
.and_then(|name| name.to_str())
.ok_or_else(|| "archive entry has an invalid file name".to_string())?;
if name != WINDOWS_PAYLOAD_NAME || entry.is_dir() {
return Err(format!(
"expected {WINDOWS_PAYLOAD_NAME}, found {}",
entry.name()
));
}
let mut payload = Vec::with_capacity(entry.size() as usize);
entry
.read_to_end(&mut payload)
.map_err(|error| format!("failed to extract executable: {error}"))?;
if payload.len() < 2 || &payload[..2] != b"MZ" {
return Err("payload is not a Windows executable".into());
}
Ok(payload)
}
fn stage_payload(target: &Path, payload: &[u8]) -> std::io::Result<PathBuf> {
let directory = tempfile::Builder::new()
.prefix("cursor-byok-portable-update-")
.tempdir()?;
let name = target
.file_name()
.ok_or_else(|| std::io::Error::other("application executable has no file name"))?;
let path = directory.path().join(name);
fs::write(&path, payload)?;
let _ = directory.keep();
Ok(path)
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::Write;
fn update_zip(name: &str, payload: &[u8]) -> Vec<u8> {
let mut bytes = Cursor::new(Vec::new());
{
let mut archive = zip::ZipWriter::new(&mut bytes);
archive
.start_file(name, zip::write::SimpleFileOptions::default())
.unwrap();
archive.write_all(payload).unwrap();
archive.finish().unwrap();
}
bytes.into_inner()
}
#[test]
fn extracts_the_single_expected_windows_executable() {
let bytes = update_zip(WINDOWS_PAYLOAD_NAME, b"MZpayload");
assert_eq!(extract_windows_payload(&bytes).unwrap(), b"MZpayload");
}
#[test]
fn rejects_unexpected_or_non_executable_payloads() {
let wrong_name = update_zip("other.exe", b"MZpayload");
assert!(extract_windows_payload(&wrong_name).is_err());
let wrong_content = update_zip(WINDOWS_PAYLOAD_NAME, b"not an executable");
assert!(extract_windows_payload(&wrong_content).is_err());
}
}
@@ -0,0 +1,296 @@
use std::{
ffi::{OsStr, OsString},
fs, io,
path::{Path, PathBuf},
process::{Child, Command, ExitCode},
thread,
time::{Duration, Instant},
};
const APPLY_ARG: &str = "--apply-portable-update";
const TARGET_ARG: &str = "--update-target";
const PID_ARG: &str = "--update-wait-pid";
const HANDSHAKE_ARG: &str = "--update-handshake";
pub(super) const READY_ARG: &str = "--portable-update-ready";
const PROCESS_WAIT_TIMEOUT: Duration = Duration::from_secs(30);
const READY_WAIT_TIMEOUT: Duration = Duration::from_secs(30);
pub(super) struct ReplacementRequest {
target: PathBuf,
pid: u32,
handshake: PathBuf,
}
pub(super) fn request_from_args() -> Result<Option<ReplacementRequest>, String> {
let args = std::env::args_os().collect::<Vec<_>>();
if !args.iter().any(|arg| arg == APPLY_ARG) {
return Ok(None);
}
let target = PathBuf::from(
argument_value(&args, TARGET_ARG).ok_or_else(|| format!("{TARGET_ARG} is required"))?,
);
let pid = argument_value(&args, PID_ARG)
.ok_or_else(|| format!("{PID_ARG} is required"))?
.to_string_lossy()
.parse::<u32>()
.map_err(|error| format!("invalid {PID_ARG}: {error}"))?;
let handshake = PathBuf::from(
argument_value(&args, HANDSHAKE_ARG)
.ok_or_else(|| format!("{HANDSHAKE_ARG} is required"))?,
);
validate_target(&target).map_err(|error| error.to_string())?;
Ok(Some(ReplacementRequest {
target,
pid,
handshake,
}))
}
pub(super) fn ready_marker_from_args() -> Option<PathBuf> {
let args = std::env::args_os().collect::<Vec<_>>();
argument_value(&args, READY_ARG).map(PathBuf::from)
}
fn argument_value<'a>(args: &'a [OsString], name: &str) -> Option<&'a OsStr> {
args.iter()
.position(|arg| arg == name)
.and_then(|index| args.get(index + 1))
.map(OsString::as_os_str)
}
fn validate_target(target: &Path) -> io::Result<()> {
if !target.is_absolute() || !target.is_file() {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
"update target must be an existing absolute file",
));
}
let source_name = std::env::current_exe()?
.file_name()
.map(OsStr::to_os_string)
.ok_or_else(|| io::Error::new(io::ErrorKind::InvalidInput, "updater has no file name"))?;
if target.file_name() != Some(source_name.as_os_str()) {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
"update target file name does not match the updater",
));
}
Ok(())
}
pub(super) fn run(request: ReplacementRequest) -> ExitCode {
match run_inner(request) {
Ok(()) => ExitCode::SUCCESS,
Err(error) => {
eprintln!("portable update replacement failed: {error}");
ExitCode::FAILURE
}
}
}
fn run_inner(request: ReplacementRequest) -> io::Result<()> {
wait_for_process(request.pid, PROCESS_WAIT_TIMEOUT, &request.handshake)?;
remove_file_if_exists(&request.handshake)?;
let source = std::env::current_exe()?;
let backup = backup_path(&request.target);
let ready = source.with_extension("ready");
remove_file_if_exists(&ready)?;
if let Err(error) = install_staged(&source, &request.target, &backup) {
relaunch(&request.target);
return Err(error);
}
let mut child = match Command::new(&request.target)
.arg(READY_ARG)
.arg(&ready)
.spawn()
{
Ok(child) => child,
Err(error) => {
restore_backup(&request.target, &backup)?;
relaunch(&request.target);
return Err(error);
}
};
match wait_until_ready(&mut child, &ready, READY_WAIT_TIMEOUT) {
Ok(()) => {
let _ = fs::remove_file(&backup);
let _ = fs::remove_file(&ready);
Ok(())
}
Err(error) => {
let _ = child.kill();
let _ = child.wait();
restore_backup(&request.target, &backup)?;
relaunch(&request.target);
Err(error)
}
}
}
fn relaunch(target: &Path) {
if target.is_file() {
let _ = Command::new(target).spawn();
}
}
fn remove_file_if_exists(path: &Path) -> io::Result<()> {
match fs::remove_file(path) {
Ok(()) => Ok(()),
Err(error) if error.kind() == io::ErrorKind::NotFound => Ok(()),
Err(error) => Err(error),
}
}
fn backup_path(target: &Path) -> PathBuf {
path_with_suffix(target, ".old")
}
fn pending_path(target: &Path) -> PathBuf {
path_with_suffix(target, ".new")
}
fn path_with_suffix(target: &Path, suffix: &str) -> PathBuf {
let mut name = target.as_os_str().to_os_string();
name.push(suffix);
PathBuf::from(name)
}
fn install_staged(source: &Path, target: &Path, backup: &Path) -> io::Result<()> {
let pending = pending_path(target);
remove_file_if_exists(&pending)?;
fs::copy(source, &pending)?;
let result = activate_pending(&pending, target, backup);
if result.is_err() {
let _ = fs::remove_file(&pending);
}
result
}
fn activate_pending(pending: &Path, target: &Path, backup: &Path) -> io::Result<()> {
remove_file_if_exists(backup)?;
fs::rename(target, backup)?;
if let Err(error) = fs::rename(pending, target) {
if let Err(restore_error) = restore_backup(target, backup) {
return Err(io::Error::other(format!(
"failed to install update ({error}) and restore the original executable ({restore_error})"
)));
}
return Err(error);
}
Ok(())
}
fn restore_backup(target: &Path, backup: &Path) -> io::Result<()> {
remove_file_if_exists(target)?;
fs::rename(backup, target)
}
fn wait_until_ready(child: &mut Child, marker: &Path, timeout: Duration) -> io::Result<()> {
let deadline = Instant::now() + timeout;
loop {
if marker.is_file() {
return Ok(());
}
if let Some(status) = child.try_wait()? {
return Err(io::Error::other(format!(
"updated application exited before startup completed: {status}"
)));
}
if Instant::now() >= deadline {
return Err(io::Error::new(
io::ErrorKind::TimedOut,
"updated application did not report a successful startup",
));
}
thread::sleep(Duration::from_millis(100));
}
}
#[cfg(windows)]
fn wait_for_process(pid: u32, timeout: Duration, handshake: &Path) -> io::Result<()> {
use windows_sys::Win32::{
Foundation::{
CloseHandle, GetLastError, ERROR_INVALID_PARAMETER, WAIT_OBJECT_0, WAIT_TIMEOUT,
},
System::Threading::{OpenProcess, WaitForSingleObject},
};
const SYNCHRONIZE_ACCESS: u32 = 0x0010_0000;
let handle = unsafe { OpenProcess(SYNCHRONIZE_ACCESS, 0, pid) };
if handle.is_null() {
let error = unsafe { GetLastError() };
return if error == ERROR_INVALID_PARAMETER {
fs::write(handshake, b"started")
} else {
Err(io::Error::from_raw_os_error(error as i32))
};
}
if let Err(error) = fs::write(handshake, b"started") {
unsafe { CloseHandle(handle) };
return Err(error);
}
let milliseconds = timeout.as_millis().min(u32::MAX as u128) as u32;
let result = unsafe { WaitForSingleObject(handle, milliseconds) };
unsafe { CloseHandle(handle) };
match result {
WAIT_OBJECT_0 => Ok(()),
WAIT_TIMEOUT => Err(io::Error::new(
io::ErrorKind::TimedOut,
"running application did not exit before the update timeout",
)),
_ => Err(io::Error::last_os_error()),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn staged_file_can_be_restored() {
let directory = tempfile::tempdir().unwrap();
let source = directory.path().join("source.exe");
let target = directory.path().join("target.exe");
let backup = backup_path(&target);
fs::write(&source, b"new").unwrap();
fs::write(&target, b"old").unwrap();
install_staged(&source, &target, &backup).unwrap();
assert_eq!(fs::read(&target).unwrap(), b"new");
assert_eq!(fs::read(&backup).unwrap(), b"old");
restore_backup(&target, &backup).unwrap();
assert_eq!(fs::read(&target).unwrap(), b"old");
assert!(!backup.exists());
}
#[test]
fn missing_pending_file_restores_original_after_backup() {
let directory = tempfile::tempdir().unwrap();
let target = directory.path().join("target.exe");
let pending = pending_path(&target);
let backup = backup_path(&target);
fs::write(&target, b"old").unwrap();
assert!(activate_pending(&pending, &target, &backup).is_err());
assert_eq!(fs::read(&target).unwrap(), b"old");
assert!(!backup.exists());
}
#[test]
fn missing_source_preserves_original_file() {
let directory = tempfile::tempdir().unwrap();
let source = directory.path().join("missing.exe");
let target = directory.path().join("target.exe");
let backup = backup_path(&target);
fs::write(&target, b"old").unwrap();
assert!(install_staged(&source, &target, &backup).is_err());
assert_eq!(fs::read(&target).unwrap(), b"old");
assert!(!backup.exists());
assert!(!pending_path(&target).exists());
}
}
+2 -2
View File
@@ -1,7 +1,7 @@
{
"$schema": "https://schema.tauri.app/config/2",
"productName": "Cursor BYOK",
"version": "0.1.5-beta.1",
"version": "0.1.6",
"identifier": "dev.cursorbyok.desktop",
"build": {
"beforeDevCommand": "npm run dev",
@@ -12,7 +12,7 @@
"app": {
"windows": [],
"security": {
"csp": "default-src 'self'; img-src 'self' asset: http://asset.localhost data:; style-src 'self' 'unsafe-inline'; connect-src 'self' http://127.0.0.1:*",
"csp": "default-src 'self'; img-src 'self' asset: http://asset.localhost data: https:; style-src 'self' 'unsafe-inline'; connect-src 'self' http://127.0.0.1:*",
"dangerousDisableAssetCspModification": [
"style-src"
]
+1
View File
@@ -185,6 +185,7 @@ function createModel({ hash, order, name, type, url, modelId, endpoint = "/v1/re
model_hash: hash,
sort_order: order,
display_name: name,
group_name: null,
type,
base_url: url,
use_full_url: false,
@@ -32,6 +32,7 @@ type CursorModelCardsProps = {
onTestPluginModel: (model: PluginModelDescriptor) => void;
onPluginSettings: (model: PluginModelDescriptor) => void;
onReorder: (modelHashes: string[]) => void;
onGroupSettings: (group: CursorModelGroup) => void;
};
type ModelGridProps = Omit<CursorModelCardsProps, "grouping" | "pluginModels" | "onTestPluginModel" | "onPluginSettings"> & {
@@ -60,6 +61,7 @@ export function CursorModelCards(props: CursorModelCardsProps) {
key={group.key}
label={group.label}
icon={group.icon}
onSettings={props.grouping === "provider" ? () => props.onGroupSettings(group) : undefined}
>
{group.models.map((model) => <ModelListRow
key={model.model_hash}
@@ -107,25 +109,37 @@ function pluginGroups(models: PluginModelDescriptor[]) {
return groups;
}
function CollapsibleGroup({ label, icon, iconSrc, children }: {
function CollapsibleGroup({ label, icon, iconSrc, onSettings, children }: {
label: string;
icon?: IconifyIcon;
iconSrc?: string;
onSettings?: () => void;
children: ReactNode;
}) {
const [open, setOpen] = useState(true);
return <Card className={styles.groupCard}>
<button
type="button"
className={styles.groupToggle}
aria-expanded={open}
onClick={() => setOpen((current) => !current)}
>
{icon && <Icon icon={icon} size="1.1em" />}
{iconSrc && <Icon src={iconSrc} size="1.1em" />}
<span className={styles.groupLabel}>{label}</span>
<Icon icon={open ? chevronDownIcon : chevronRightIcon} size="1em" />
</button>
<div className={styles.groupHeader}>
<button
type="button"
className={styles.groupToggle}
aria-expanded={open}
onClick={() => setOpen((current) => !current)}
>
{icon && <Icon icon={icon} size="1.1em" />}
{iconSrc && <Icon src={iconSrc} size="1.1em" />}
<span className={styles.groupLabel}>{label}</span>
</button>
{onSettings && <button type="button" className={styles.groupSettings} onClick={onSettings}>{t("分组设置")}</button>}
<button
type="button"
className={styles.groupChevron}
tabIndex={-1}
aria-hidden="true"
onClick={() => setOpen((current) => !current)}
>
<Icon icon={open ? chevronDownIcon : chevronRightIcon} size="1em" />
</button>
</div>
{open && <div className={styles.modelList}>{children}</div>}
</Card>;
}
@@ -273,8 +287,9 @@ function ModelGrid({
}
function providerGroup(model: Model) {
const label = providerDomain(model.base_url);
return { key: label, label, icon: flatColorOrganizationIcon };
const key = providerDomain(model.base_url);
const label = model.group_name?.trim() || key;
return { key, label, icon: flatColorOrganizationIcon };
}
function providerDomain(baseUrl: string) {
@@ -24,6 +24,7 @@ export const emptyCursorModelDraft = (): CursorModelDraft => ({
model: {
sort_order: 0,
display_name: "",
group_name: null,
type: "openai",
base_url: "",
use_full_url: false,
@@ -73,8 +73,34 @@
.groupCard {
padding: 4px 12px 8px;
}
.groupHeader {
display: flex;
align-items: center;
gap: 8px;
}
.groupSettings {
padding: 4px 2px;
color: var(--vscode-descriptionForeground);
background: none;
border: none;
font-size: type.$font-size-xs;
white-space: nowrap;
cursor: pointer;
&:hover { color: var(--vscode-foreground); }
}
.groupChevron {
display: flex;
align-items: center;
padding: 4px 0;
color: var(--vscode-foreground);
background: none;
border: none;
cursor: pointer;
}
.groupToggle {
width: 100%;
flex: 1;
min-width: 0;
display: flex;
align-items: center;
gap: 8px;
@@ -2,13 +2,14 @@ import { useCallback, useEffect, useRef, useState } from "react";
import { useNavigate } from "react-router-dom";
import { api, configuredPluginModels, type Model, type ModelInput } from "../../shared/api";
import { CursorCaGate, CursorCaProvider, CursorModelGate, CursorModelProvider } from "./CursorGates";
import { CursorModelCards, cursorModelGroups, type CursorModelGrouping } from "./CursorModelCards";
import { CursorModelCards, cursorModelGroups, type CursorModelGroup, type CursorModelGrouping } from "./CursorModelCards";
import { CursorModelEditor, emptyCursorModelDraft, type CursorModelDraft } from "./CursorModelEditor";
import { CursorModelTestResult, type CursorModelTestState } from "./CursorModelTestResult";
import styles from "./CursorSettings.module.scss";
import { PageContent } from "../../shell/layout/PageContent";
import { LegacyModelImport } from "./LegacyModelImport";
import { ConfirmDialog } from "../../shared/ui/ConfirmDialog";
import { FormField, SecretTextInput, TextInput } from "../../shared/ui/FormControls";
import controls from "../../shared/ui/Controls.module.scss";
import { Icon } from "../../shared/ui/Icon";
import { Modal } from "../../shared/ui/Modal";
@@ -34,6 +35,11 @@ export function CursorSettingsPage() {
const [savingAndTesting, setSavingAndTesting] = useState(false);
const [batchTesting, setBatchTesting] = useState(false);
const [grouping, setGrouping] = useState<CursorModelGrouping>("flat");
const [settingsGroup, setSettingsGroup] = useState<CursorModelGroup | null>(null);
const [groupNameDraft, setGroupNameDraft] = useState("");
const [groupBaseUrlDraft, setGroupBaseUrlDraft] = useState("");
const [groupApiKeyDraft, setGroupApiKeyDraft] = useState("");
const [groupSettingsBusy, setGroupSettingsBusy] = useState(false);
const activeModelTests = useRef(new Map<string, { testId: string; controller: AbortController; cancelling: boolean }>());
const caReady = cursorHarness?.ca === "ready";
const pluginModels = configuredPluginModels(plugins);
@@ -208,6 +214,39 @@ export function CursorSettingsPage() {
}]);
if (created) message(t("模型已复制"));
};
const openGroupSettings = (group: CursorModelGroup) => {
setGroupNameDraft(group.models.find((model) => model.group_name?.trim())?.group_name?.trim() ?? "");
setGroupBaseUrlDraft(sharedValue(group.models.map((model) => model.base_url)) ?? "");
setGroupApiKeyDraft(sharedValue(group.models.map((model) => model.api_key)) ?? "");
setSettingsGroup(group);
};
const saveGroupSettings = async () => {
if (!settingsGroup) return;
const group_name = groupNameDraft.trim() || null;
const base_url = groupBaseUrlDraft.trim();
const api_key = groupApiKeyDraft.trim();
setGroupSettingsBusy(true);
try {
for (const model of settingsGroup.models) {
const input: ModelInput = {
...modelInput(model),
group_name,
...(base_url ? { base_url } : {}),
...(api_key ? { api_key } : {}),
};
if (input.group_name === (model.group_name ?? null)
&& input.base_url === model.base_url
&& input.api_key === model.api_key) continue;
await api.updateModel(model.model_hash, input);
}
await appStore.refresh();
setSettingsGroup(null);
} catch (cause) {
message(errorText(cause));
} finally {
setGroupSettingsBusy(false);
}
};
const reorderModels = useCallback(async (modelHashes: string[]) => {
if (!await appStore.reorderCursorModels(modelHashes)) {
message(appStore.getSnapshot().error || t("排序失败"));
@@ -228,6 +267,7 @@ export function CursorSettingsPage() {
onTestPluginModel={(model) => void testModel({ model_hash: model.id, display_name: model.displayName })}
onPluginSettings={() => navigate("/plugins")}
onReorder={reorderModels}
onGroupSettings={openGroupSettings}
/>;
const refreshCa = async () => {
@@ -274,6 +314,19 @@ export function CursorSettingsPage() {
<ConfirmDialog open={caCommand !== null} title={t("安装本地 CA")} cancelLabel={t("关闭")} confirmLabel={t("打开终端")} onCancel={() => setCaCommand(null)} onConfirm={openCaTerminal}>
<div className={styles.editor}><strong>{t("需要授权安装证书")}</strong><span>{t("安装命令已自动复制。点击“打开终端”,将命令粘贴到终端中执行,并按提示输入密码。")}</span><pre className={styles.command}>{caCommand}</pre></div>
</ConfirmDialog>
<Modal open={settingsGroup !== null} title={t("分组设置")} busy={groupSettingsBusy || cursorBusy} onClose={() => setSettingsGroup(null)} onSubmit={() => void saveGroupSettings()} submitLabel={t("保存")}>
{settingsGroup && <div className={styles.editor}>
<FormField label={t("分组名称")} hint={t("应用于该分组下的全部模型,并作为 Cursor 模型选择器中的徽章标签;清空则恢复显示服务器域名。")}>
<TextInput placeholder={settingsGroup.key} value={groupNameDraft} onChange={(event) => setGroupNameDraft(event.target.value)} />
</FormField>
<FormField label={t("服务器地址")} hint={t("修改后应用于该分组下的全部模型;留空保持各模型现有配置不变。")}>
<TextInput placeholder={t("留空保持不变")} value={groupBaseUrlDraft} onChange={(event) => setGroupBaseUrlDraft(event.target.value)} />
</FormField>
<FormField label="API Key" hint={t("修改后应用于该分组下的全部模型;留空保持各模型现有配置不变。")}>
<SecretTextInput placeholder={t("留空保持不变")} autoComplete="off" value={groupApiKeyDraft} onChange={(event) => setGroupApiKeyDraft(event.target.value)} />
</FormField>
</div>}
</Modal>
<ConfirmDialog open={deleting !== null} title={t("删除模型")} cancelLabel={t("取消")} confirmLabel={t("删除")} onCancel={() => setDeleting(null)} onConfirm={() => { if (deleting) void appStore.deleteModel(deleting.model_hash); setDeleting(null); }}><p>{t("确定删除这个模型吗?")}</p></ConfirmDialog>
</>;
}
@@ -283,6 +336,13 @@ function modelInput(model: Model): ModelInput {
return input;
}
/** 组内所有模型取值一致时返回该值,否则返回 null(表单留空表示保持不变)。 */
function sharedValue(values: string[]): string | null {
const [first, ...rest] = values;
if (first === undefined) return null;
return rest.every((value) => value === first) ? first : null;
}
function draftInput(draft: CursorModelDraft): ModelInput {
const model = {
...draft.model,
@@ -95,7 +95,8 @@ export function AppLifecycleSettingsCard() {
const nextVersion = await updateStore.check();
message(nextVersion ? t("发现新版本 {version}", { version: nextVersion }) : t("当前已是最新版本"));
} catch (cause) {
message(cause instanceof Error ? cause.message : String(cause));
const error = cause instanceof Error ? cause.message : String(cause);
message(t("检查更新失败:{error}", { error }));
}
};
@@ -103,7 +104,8 @@ export function AppLifecycleSettingsCard() {
try {
await updateStore.install();
} catch (cause) {
message(cause instanceof Error ? cause.message : String(cause));
const error = cause instanceof Error ? cause.message : String(cause);
message(t("安装更新失败:{error}", { error }));
}
};
File diff suppressed because it is too large Load Diff
+7
View File
@@ -79,6 +79,7 @@
"2f9daa828907b93f": "Delete",
"2fe5a8d0eee9f14c": "Invalid",
"303c30f301514250": "Search resources",
"3260348163d03b8e": "Leave blank to keep unchanged",
"32896fdaaaa4c106": "Account saved and the model catalog is synced.",
"346ff60e6c7c5181": "Reading…",
"36f33adaf0942634": "Confirm",
@@ -220,6 +221,7 @@
"91aaf184cfc17ffd": "Overview",
"91af6e57e7453fbe": "Add account",
"92156a483d4ba248": "Only request, response, and trace attachments are deleted; call summaries, metrics, and configuration are kept.",
"92e26b27d5ea8f0e": "Failed to check for updates: {error}",
"940a168911ade998": "Items per page",
"945fb1c67eca8493": "Installing the plugin runtime",
"946b3ffc02f026c0": "Delete this model?",
@@ -258,6 +260,7 @@
"a748cc074f78de00": "View details",
"a7617f42f898b2bf": "Use complete request URL",
"a8036485f9227f2c": "Drag to reorder",
"a80b53f8848e6d27": "Failed to install update: {error}",
"a98585871c5313ff": "Display name",
"ab9084a640fbb864": "Deselect all",
"abecab6701177721": "Launch at login enabled",
@@ -270,6 +273,7 @@
"b06325c5660f0c29": "Direct",
"b16c3b2ecedd6fe1": "Cursor integration is active. Add a model configuration to use a BYOK model.",
"b254ff315d861346": "Try initializing again",
"b2617bf9ae663752": "Group settings",
"b4411558b932266f": "Provider type",
"b4c9e08870d41aa2": "Initialize the plugin runtime first",
"b502b1d414664337": "Prompt: {tokens}",
@@ -290,6 +294,7 @@
"bb7efdcb6af6e805": "Default dark",
"bda62ce1d5e4ace9": "Tell us why",
"bda74b5674b6a57d": "Initialize plugins",
"be961dc60ab610da": "Applies to every model in this group when changed; leave blank to keep each model's current configuration.",
"bf57afd709694b55": "Overview time range",
"bfc01caf9fe0c841": "Cache hit rate {rate}",
"c0b3fbff51ccc40b": "Done",
@@ -322,6 +327,7 @@
"d60669bb26a22f5d": "Leave blank to use the default",
"d6b1f203680f5496": "Leave blank to use adaptive thinking",
"d766536c18e8e990": "Plugin runtime {version} is installed and ready to use.",
"d7e266bdc8064193": "Group name",
"d86fa42c3848c680": "Use system proxy",
"d8c47e9776cf1082": "Main menu",
"da521d1c1cbd36af": "Authorization is required to install the certificate",
@@ -361,6 +367,7 @@
"ea26b760e930a7ca": "Call observability",
"eb11e2df1d8ae387": "Provider URL",
"eb1be07f2ca6e506": "Estimated using Claude Opus 4.7 pricing.",
"eb4a3db23661fb52": "Applies to every model in this group and is used as the badge label in Cursor's model picker; clear it to fall back to the server domain.",
"eb77492c9f76a7e1": "The install command has been copied. Click “Open terminal”, paste it into the terminal, and enter your password when prompted.",
"eba54690937bc532": "Manage accounts",
"ed31fbb483ee1b0a": "Actions",
+7
View File
@@ -79,6 +79,7 @@
"2f9daa828907b93f": "删除",
"2fe5a8d0eee9f14c": "已失效",
"303c30f301514250": "搜索资源",
"3260348163d03b8e": "留空保持不变",
"32896fdaaaa4c106": "账号已保存,模型目录已同步。",
"346ff60e6c7c5181": "读取中…",
"36f33adaf0942634": "确认",
@@ -220,6 +221,7 @@
"91aaf184cfc17ffd": "数据概览",
"91af6e57e7453fbe": "添加账号",
"92156a483d4ba248": "仅删除请求、响应和追踪附件等详细内容,保留调用汇总、统计指标和配置。",
"92e26b27d5ea8f0e": "检查更新失败:{error}",
"940a168911ade998": "每页条数",
"945fb1c67eca8493": "正在安装插件运行时",
"946b3ffc02f026c0": "确定删除这个模型吗?",
@@ -258,6 +260,7 @@
"a748cc074f78de00": "查看详情",
"a7617f42f898b2bf": "使用完整请求地址",
"a8036485f9227f2c": "拖动排序",
"a80b53f8848e6d27": "安装更新失败:{error}",
"a98585871c5313ff": "显示名称",
"ab9084a640fbb864": "全不选",
"abecab6701177721": "已开启开机启动",
@@ -270,6 +273,7 @@
"b06325c5660f0c29": "直连",
"b16c3b2ecedd6fe1": "Cursor 接管已生效;添加模型配置后即可使用 BYOK 模型。",
"b254ff315d861346": "请重试初始化",
"b2617bf9ae663752": "分组设置",
"b4411558b932266f": "上游类型",
"b4c9e08870d41aa2": "需要先初始化插件运行时",
"b502b1d414664337": "提示词:{tokens}",
@@ -290,6 +294,7 @@
"bb7efdcb6af6e805": "默认暗色",
"bda62ce1d5e4ace9": "可以告诉我们原因",
"bda74b5674b6a57d": "初始化插件",
"be961dc60ab610da": "修改后应用于该分组下的全部模型;留空保持各模型现有配置不变。",
"bf57afd709694b55": "概览时间范围",
"bfc01caf9fe0c841": "缓存命中率 {rate}",
"c0b3fbff51ccc40b": "完成",
@@ -322,6 +327,7 @@
"d60669bb26a22f5d": "留空使用默认值",
"d6b1f203680f5496": "留空使用 adaptive thinking",
"d766536c18e8e990": "插件运行时 {version} 已安装,可以开始使用插件。",
"d7e266bdc8064193": "分组名称",
"d86fa42c3848c680": "使用系统代理",
"d8c47e9776cf1082": "主菜单",
"da521d1c1cbd36af": "需要授权安装证书",
@@ -361,6 +367,7 @@
"ea26b760e930a7ca": "调用观测",
"eb11e2df1d8ae387": "上游地址",
"eb1be07f2ca6e506": "按 Claude Opus 4.7 价格估算。",
"eb4a3db23661fb52": "应用于该分组下的全部模型,并作为 Cursor 模型选择器中的徽章标签;清空则恢复显示服务器域名。",
"eb77492c9f76a7e1": "安装命令已自动复制。点击“打开终端”,将命令粘贴到终端中执行,并按提示输入密码。",
"eba54690937bc532": "账号管理",
"ed31fbb483ee1b0a": "操作",
+2 -2
View File
@@ -7,6 +7,7 @@ export interface Model {
model_hash: string;
sort_order: number;
display_name: string;
group_name: string | null;
type: ModelType;
base_url: string;
use_full_url: boolean;
@@ -33,6 +34,7 @@ export interface Model {
export interface ModelInput {
sort_order: number;
display_name: string;
group_name: string | null;
type: ModelType;
base_url: string;
use_full_url: boolean;
@@ -233,9 +235,7 @@ export interface PluginModelDescriptor {
description: string | null;
icon: string;
providerType: string;
contextWindowTokens: number | null;
maxOutputTokens: number | null;
thinking: boolean;
images: boolean;
}
+34 -6
View File
@@ -1,5 +1,5 @@
import { getVersion, setDockVisibility } from "@tauri-apps/api/app";
import { isTauri } from "@tauri-apps/api/core";
import { invoke, isTauri } from "@tauri-apps/api/core";
import { disable, enable, isEnabled } from "@tauri-apps/plugin-autostart";
import { relaunch } from "@tauri-apps/plugin-process";
import { check, type Update } from "@tauri-apps/plugin-updater";
@@ -46,11 +46,39 @@ export async function writeDockIconVisibility(visible: boolean): Promise<void> {
}
}
export async function checkForUpdate(): Promise<Update | null> {
return check();
export type AppUpdate = {
version: string;
install(): Promise<void>;
close(): Promise<void>;
};
type PortableUpdateInfo = {
version: string;
};
export async function checkForUpdate(): Promise<AppUpdate | null> {
if (desktopPlatform() === "windows") {
const update = await invoke<PortableUpdateInfo | null>("check_portable_update");
if (!update) return null;
return {
version: update.version,
install: () => invoke("install_portable_update", { expectedVersion: update.version }),
close: async () => {},
};
}
const update: Update | null = await check();
if (!update) return null;
return {
version: update.version,
install: async () => {
await update.downloadAndInstall();
await relaunch();
},
close: () => update.close(),
};
}
export async function installUpdate(update: Update): Promise<void> {
await update.downloadAndInstall();
await relaunch();
export async function installUpdate(update: AppUpdate): Promise<void> {
await update.install();
}
+3 -3
View File
@@ -1,9 +1,9 @@
import { useSyncExternalStore } from "react";
import type { Update } from "@tauri-apps/plugin-updater";
import {
checkForUpdate,
hasNativeAppLifecycle,
installUpdate,
type AppUpdate,
} from "../native/appLifecycle";
export type UpdateSnapshot = {
@@ -17,7 +17,7 @@ let snapshot: UpdateSnapshot = {
checking: false,
installing: false,
};
let availableUpdate: Update | null = null;
let availableUpdate: AppUpdate | null = null;
let pendingCheck: Promise<string | null> | null = null;
const listeners = new Set<() => void>();
@@ -26,7 +26,7 @@ function update(patch: Partial<UpdateSnapshot>) {
listeners.forEach((listener) => listener());
}
async function replaceAvailableUpdate(next: Update | null) {
async function replaceAvailableUpdate(next: AppUpdate | null) {
const previous = availableUpdate;
availableUpdate = next;
update({ availableVersion: next?.version ?? null });
+2 -54
View File
@@ -1,9 +1,4 @@
以下是重构后完整目标版本
实现时,先创建所有目录和文件固化,每个文件头部都写好注释再实现
旧服务已被备份为server_backup,/Users/leokun/Documents/cursor-byok/server 目录已创建
行数均为目标估算,使用 `≈` 标记;不包含测试、生成代码和空行。
实现时可做略微调整,测试要求相对于目标文件旁边的独立文件,禁止码内测试
本文档目录 /Users/leokun/Documents/cursor-byok/cursor.md
## 完整目录
```text
@@ -656,54 +651,7 @@ store ─X→ cursor
model ─X→ cursor
```
## 当前代码迁移
```text
当前 目标
cursor/bidi_append.rs → api/cursor/bidi.rs
cursor/run_sse.rs → api/cursor/run_sse.rs
cursor/handlers.rs → api/cursor/handlers.rs
cursor/proxy.rs → api/cursor/proxy.rs
cursor/sessions.rs → cursor/transport/registry.rs
+ cursor/transport/handle.rs
+ cursor/transport/output.rs
cursor/inbox.rs → cursor/transport/inbox.rs
cursor/actor.rs → cursor/transport/
+ cursor/conversation/runtime.rs
+ cursor/conversation/delivery.rs
cursor/session.rs → cursor/conversation/runtime.rs
+ cursor/conversation/output.rs
+ cursor/checkpoint/
+ cursor/tools/
cursor/request/prepare.rs → cursor/compile/run.rs
cursor/request/context.rs → cursor/compile/context.rs
cursor/request/background.rs → cursor/compile/insert_messages.rs
cursor/request/runtime.rs → cursor/compile/break_messages.rs
cursor/request/images.rs → cursor/compile/images.rs
cursor/request/model.rs → cursor/compile/model.rs
cursor/interaction/mod.rs → cursor/protocol/events.rs
cursor/interaction/query.rs → cursor/tools/codec/query.rs
cursor/interaction/render.rs → cursor/tools/codec/render.rs
cursor/projection/decode.rs → cursor/checkpoint/messages/decode.rs
cursor/projection/encode.rs → cursor/checkpoint/messages/encode.rs
cursor/projection/tests.rs → cursor/checkpoint/messages/tests.rs
cursor/presentation.rs → cursor/checkpoint/steps.rs
run/runtime.rs RunRegistry → cursor/conversation/registry.rs
run/runtime.rs RunActor → run/engine.rs + run/handle.rs
run/port.rs → run/command.rs + run/event.rs + run/port.rs
store/revisions.rs → store/checkpoints.rs
```
## 最终核心
@@ -719,4 +667,4 @@ Bidi
→ Checkpoint
→ Transport
→ RunSSE
```
```
@@ -0,0 +1,4 @@
-- Custom provider-group display name shared by models with the same upstream host.
-- NULL means no custom name; the UI falls back to the base_url hostname and the
-- Cursor model picker badge falls back to the model type label.
ALTER TABLE model_configs ADD COLUMN group_name TEXT;
@@ -0,0 +1 @@
ALTER TABLE tool_round_calls ADD COLUMN argument_error TEXT;
@@ -8,6 +8,7 @@ import type { LlmRequest, ModelEvent } from "cursor-byok:provider";
import type { ResourceSnapshot } from "cursor-byok:resource";
import { codexDeviceOAuth } from "./oauth.ts";
import { parseOfficialModels } from "./models.ts";
import { buildResponsesBody } from "cursor-byok:protocol/openai-responses";
import { codexProvider, isQuotaError } from "./provider.ts";
import {
accountIdentity,
@@ -164,7 +165,10 @@ Deno.test("official model discovery excludes hidden models and puts the default
display_name: "GPT First",
supported_in_api: true,
visibility: "list",
supported_reasoning_efforts: ["low", "medium"],
supported_reasoning_levels: [
{ effort: "low", description: "Fast responses" },
{ effort: "medium", description: "Balanced" },
],
},
{ slug: "gpt-second", supported_in_api: true, visibility: "list" },
{ slug: "gpt-hidden", supported_in_api: true, visibility: "hidden" },
@@ -172,7 +176,7 @@ Deno.test("official model discovery excludes hidden models and puts the default
],
});
assertEquals(models.map((model) => model.id), ["gpt-second", "gpt-first"]);
assertEquals(models[1].capabilities, { thinking: true, images: true });
assertEquals(models[1].capabilities, { images: true });
assertEquals(models[1].privateData, { reasoningEfforts: ["low", "medium"] });
});
@@ -302,6 +306,43 @@ Deno.test("invoke streams normalized events from the Codex Responses API", async
]);
});
Deno.test("reasoning replay projects response items to valid input items", () => {
const replayRequest = request();
replayRequest.messages = [{
role: "assistant",
text: "",
thinking: "",
replayState: {
providerKind: "openai_responses",
value: {
items: [{
type: "reasoning",
id: "item-1",
status: "completed",
summary: [{ type: "summary_text", text: "why" }],
content: [],
encrypted_content: "opaque",
output_only: true,
}],
},
},
toolCalls: [],
}];
const body = buildResponsesBody({
url: "https://example.com/responses",
model: "gpt-test",
request: replayRequest,
});
assertEquals(body.input, [{
type: "reasoning",
id: "item-1",
summary: [{ type: "summary_text", text: "why" }],
content: [],
encrypted_content: "opaque",
}]);
});
Deno.test("invoke streams incremental tool calls and replays reasoning items", async () => {
const token = jwt({ "https://api.openai.com/auth": { chatgpt_account_id: "acct-1" } });
const draft = await credentialDraft({
+6 -7
View File
@@ -24,7 +24,11 @@ function positiveInteger(value: unknown): number | null {
}
function parseReasoningEfforts(model: Record<string, unknown>): string[] {
const source = model.supported_reasoning_efforts ??
const source = model.supported_reasoning_levels ??
model.supportedReasoningLevels ??
model.reasoning_levels ??
model.reasoningLevels ??
model.supported_reasoning_efforts ??
model.supportedReasoningEfforts ??
model.reasoning_efforts ??
model.reasoningEfforts;
@@ -61,10 +65,6 @@ export function parseOfficialModels(body: unknown): ModelDefinition[] {
seen.add(id);
const efforts = parseReasoningEfforts(model);
const description = text(model.description);
const contextWindowTokens = positiveInteger(
model.context_window_tokens ?? model.contextWindowTokens ?? model.context_window ??
model.contextWindow,
);
const maxOutputTokens = positiveInteger(
model.max_output_tokens ?? model.maxOutputTokens ?? model.max_completion_tokens ??
model.maxCompletionTokens,
@@ -74,9 +74,8 @@ export function parseOfficialModels(body: unknown): ModelDefinition[] {
displayName: text(model.display_name ?? model.displayName ?? model.title ?? model.name) ??
id,
...(description ? { description } : {}),
...(contextWindowTokens !== null ? { contextWindowTokens } : {}),
...(maxOutputTokens !== null ? { maxOutputTokens } : {}),
capabilities: { thinking: efforts.length > 0, images: true },
capabilities: { images: true },
privateData: { reasoningEfforts: efforts },
});
}
@@ -160,9 +160,8 @@ Deno.test("model discovery parses both language-models and standard list shapes"
});
assertEquals(richModels.map((model) => model.id), ["grok-4", "grok-3-mini"]);
assertEquals(richModels[0].displayName, "Grok 4");
assertEquals(richModels[0].contextWindowTokens, 256_000);
assertEquals(richModels[0].capabilities, { thinking: false, images: true });
assertEquals(richModels[1].capabilities, { thinking: false, images: false });
assertEquals(richModels[0].capabilities, { images: true });
assertEquals(richModels[1].capabilities, { images: false });
const plainModels = parseGrokModels({ data: [{ id: "grok-4-fast" }] });
assertEquals(plainModels.map((model) => model.id), ["grok-4-fast"]);
+2 -16
View File
@@ -9,12 +9,12 @@ export const FALLBACK_MODELS: ModelDefinition[] = [
{
id: "grok-4.6",
displayName: "Grok 4.6",
capabilities: { thinking: false, images: true },
capabilities: { images: true },
},
{
id: "grok-4.5",
displayName: "Grok 4.5",
capabilities: { thinking: false, images: true },
capabilities: { images: true },
},
];
@@ -28,15 +28,6 @@ function text(value: unknown): string | null {
return typeof value === "string" && value.trim() ? value.trim() : null;
}
function positiveInteger(value: unknown): number | null {
const parsed = typeof value === "number"
? value
: typeof value === "string"
? Number(value)
: NaN;
return Number.isFinite(parsed) && parsed > 0 ? Math.floor(parsed) : null;
}
function modalities(value: unknown): string[] {
return Array.isArray(value)
? value.flatMap((item) => (typeof item === "string" ? [item.toLowerCase()] : []))
@@ -66,15 +57,10 @@ export function parseGrokModels(body: unknown): ModelDefinition[] {
if (!id || seen.has(id)) continue;
seen.add(id);
const inputs = modalities(model?.input_modalities ?? model?.inputModalities);
const contextWindowTokens = positiveInteger(
model?.context_window ?? model?.contextWindow ?? model?.max_prompt_length,
);
models.push({
id,
displayName: displayName(id),
...(contextWindowTokens !== null ? { contextWindowTokens } : {}),
capabilities: {
thinking: false,
images: inputs.length === 0 || inputs.includes("image"),
},
});
+5 -1
View File
@@ -184,7 +184,11 @@ pub async fn append(
request: DecodedAppend,
parent: Option<TransportParent>,
) -> Result<ai::BidiAppendResponse> {
let handle = registry.get_or_create(&request.request_id).await?;
let replace_closing = request.model_id().is_some();
let handle = registry
.get_or_create_for_append(&request.request_id, replace_closing)
.await?;
let _admission = handle.admit()?;
if let Some(conversation_id) = request.conversation_id() {
handle.set_conversation_id(conversation_id)?;
}
+113 -28
View File
@@ -19,18 +19,26 @@ use crate::{
connect,
proto::{agent::v1 as agent, aiserver::v1 as ai},
},
services::{account, analytics, model_catalog, observability::CursorTraceRecorder, tab},
services::{account, analytics, knowledge, model_catalog, tab},
transport::{TransportParent, TransportRegistry},
},
Result,
};
pub fn router(registry: TransportRegistry) -> Result<Router> {
let proxy = CursorProxy::cursor(registry.store().clone())?;
Ok(router_with_proxy(registry, proxy))
pub fn router(
registry: TransportRegistry,
clients: crate::network::NetworkClients,
) -> Result<Router> {
let proxy = CursorProxy::cursor(clients);
let knowledge = knowledge::KnowledgeService::managed()?;
Ok(router_with_proxy(registry, proxy, knowledge))
}
fn router_with_proxy(registry: TransportRegistry, proxy: CursorProxy) -> Router {
fn router_with_proxy(
registry: TransportRegistry,
proxy: CursorProxy,
knowledge_service: knowledge::KnowledgeService,
) -> Router {
let web_cache = registry.web_cache().router();
Router::new()
.route("/__byok-api__/healthz", get(health))
@@ -69,6 +77,22 @@ fn router_with_proxy(registry: TransportRegistry, proxy: CursorProxy) -> Router
"/aiserver.v1.DashboardService/GetUsageLimitStatusAndActiveGrants",
post(account::usage_limit_status),
)
.route(
"/aiserver.v1.AiService/KnowledgeBaseAdd",
post(knowledge::add),
)
.route(
"/aiserver.v1.AiService/KnowledgeBaseList",
post(knowledge::list),
)
.route(
"/aiserver.v1.AiService/KnowledgeBaseUpdate",
post(knowledge::update),
)
.route(
"/aiserver.v1.AiService/KnowledgeBaseRemove",
post(knowledge::remove),
)
.route(
analytics::BOOTSTRAP_STATSIG_PATH,
post(analytics::bootstrap_statsig),
@@ -80,6 +104,7 @@ fn router_with_proxy(registry: TransportRegistry, proxy: CursorProxy) -> Router
.fallback(proxy::forward)
.method_not_allowed_fallback(proxy::forward)
.layer(Extension(proxy))
.layer(Extension(knowledge_service))
.with_state(registry)
.merge(web_cache)
}
@@ -96,16 +121,13 @@ async fn run_sse_handler(
let (parts, body) = buffered(request).await?;
let request: agent::BidiRequestId = connect::decode_unary(&body)?;
let route = registry.wait_route(&request.request_id).await;
let trace = CursorTraceRecorder::resume(registry.store().clone(), &request.request_id).await;
if let Some(trace) = &trace {
trace
.request(
"run_sse_request",
&body,
serde_json::json!({"request_id": request.request_id}),
)
.await;
}
let trace = registry.trace(&request.request_id);
trace.resume();
trace.request(
"run_sse_request",
body.clone(),
serde_json::json!({"request_id": request.request_id}),
);
match route {
crate::cursor::transport::TransportRoute::Local => {
run_sse::stream(&registry, &request.request_id).await
@@ -116,7 +138,14 @@ async fn run_sse_handler(
Request::from_parts(parts, Body::from(body)),
)
.await?;
Ok(run_sse::upstream(registry, request.request_id, generation, response, trace).await)
Ok(run_sse::upstream(
registry,
request.request_id,
generation,
response,
Some(trace),
)
.await)
}
}
}
@@ -132,6 +161,7 @@ async fn bidi_handler(
let first_model = decoded.model_id().map(str::to_owned);
let conversation_id = decoded.conversation_id().map(str::to_owned);
let trace_metadata = decoded.trace_metadata();
let trace = registry.trace(&decoded.request_id);
let local = if let Some(model_id) = decoded.model_id() {
// 插件模型 ID 只在本地有意义,永远不转发到 Cursor 官方上游。
if model_id.starts_with(crate::plugin::ADAPTER_ID_PREFIX)
@@ -156,14 +186,18 @@ async fn bidi_handler(
} else if registry.upstream(&decoded.request_id).await {
false
} else {
trace.resume();
trace.request(
"bidi_request",
body.clone(),
trace_outcome(trace_metadata, false, "missing_transport", None),
);
return Err(crate::Error::Protocol(
"first BidiAppend message must select a model".into(),
));
};
let trace = if first_model.is_some() {
CursorTraceRecorder::begin(
registry.store().clone(),
&decoded.request_id,
if first_model.is_some() {
trace.begin(
conversation_id.as_deref(),
if local {
"local_byok"
@@ -171,26 +205,61 @@ async fn bidi_handler(
"cursor_official"
},
first_model.as_deref(),
)
.await
);
} else {
CursorTraceRecorder::resume(registry.store().clone(), &decoded.request_id).await
};
if let Some(trace) = &trace {
trace.request("bidi_request", &body, trace_metadata).await;
trace.resume();
}
if !local {
if first_model.is_some() {
registry.mark_upstream(&decoded.request_id).await;
}
trace.request(
"bidi_request",
body.clone(),
trace_outcome(trace_metadata, true, "upstream", None),
);
return proxy::forward(
Extension(proxy),
Request::from_parts(parts, Body::from(body)),
)
.await;
}
let parent = parent_headers(&parts.headers)?;
bidi::append(&registry, decoded, parent).await?;
let parent = match parent_headers(&parts.headers) {
Ok(parent) => parent,
Err(error) => {
trace.request(
"bidi_request",
body,
trace_outcome(
trace_metadata,
false,
"invalid_parent",
Some(error.to_string()),
),
);
return Err(error);
}
};
match bidi::append(&registry, decoded, parent).await {
Ok(_) => trace.request(
"bidi_request",
body,
trace_outcome(trace_metadata, true, "local", None),
),
Err(error) => {
trace.request(
"bidi_request",
body,
trace_outcome(
trace_metadata,
false,
"command_rejected",
Some(error.to_string()),
),
);
return Err(error);
}
}
let mut response = Response::new(axum::body::Body::empty());
*response.status_mut() = StatusCode::OK;
response.headers_mut().insert(
@@ -200,6 +269,22 @@ async fn bidi_handler(
Ok(response)
}
fn trace_outcome(
mut metadata: serde_json::Value,
accepted: bool,
route_outcome: &str,
error: Option<String>,
) -> serde_json::Value {
if let Some(metadata) = metadata.as_object_mut() {
metadata.insert("accepted".into(), accepted.into());
metadata.insert("route_outcome".into(), route_outcome.into());
if let Some(error) = error {
metadata.insert("error".into(), error.into());
}
}
metadata
}
async fn buffered(request: Request<Body>) -> Result<(axum::http::request::Parts, Bytes)> {
let (parts, body) = request.into_parts();
let body = to_bytes(body, usize::MAX)
+6 -15
View File
@@ -14,8 +14,7 @@ pub const UPSTREAM_URL_HEADER: &str = "x-server-upstream-url";
#[derive(Clone)]
pub struct CursorProxy {
client: Option<reqwest::Client>,
store: Option<crate::store::Store>,
clients: crate::network::NetworkClients,
upstream: String,
}
@@ -47,23 +46,15 @@ impl BufferedResponse {
}
impl CursorProxy {
pub fn cursor(store: crate::store::Store) -> Result<Self> {
Ok(Self {
client: None,
store: Some(store),
pub fn cursor(clients: crate::network::NetworkClients) -> Self {
Self {
clients,
upstream: CURSOR_UPSTREAM.into(),
})
}
}
async fn client(&self) -> Result<reqwest::Client> {
match (&self.client, &self.store) {
(Some(client), _) => Ok(client.clone()),
(_, Some(store)) => Ok(crate::network::client_builder(store)
.await?
.redirect(reqwest::redirect::Policy::none())
.build()?),
_ => unreachable!("Cursor proxy always has a client or store"),
}
self.clients.cursor_client().await
}
}
+5 -5
View File
@@ -22,7 +22,7 @@ pub async fn stream(registry: &TransportRegistry, request_id: &str) -> Result<Re
let receiver = handle.subscribe();
let trace = handle.trace().cloned();
if let Some(trace) = &trace {
trace.response_started(StatusCode::OK.as_u16()).await;
trace.response_started(StatusCode::OK.as_u16());
}
let body_stream = local_body_stream(receiver, handle, trace);
let mut response = Response::new(Body::from_stream(body_stream));
@@ -133,7 +133,7 @@ pub async fn upstream(
) -> Response<Body> {
let (parts, body) = response.into_parts();
if let Some(trace) = &trace {
trace.response_started(parts.status.as_u16()).await;
trace.response_started(parts.status.as_u16());
}
let stream = async_stream::stream! {
let _guard = UpstreamRunGuard {
@@ -180,15 +180,15 @@ impl TraceStreamSink {
while let Some(event) = receiver.recv().await {
match event {
TraceStreamEvent::Chunk(chunk) => {
trace.response_chunk(source, &chunk).await;
trace.response_chunk(source, chunk);
}
TraceStreamEvent::Finish(error) => {
trace.finish(error.as_deref()).await;
trace.finish(error.as_deref());
return;
}
}
}
trace.finish(None).await;
trace.finish(None);
});
Self {
sender: Some(sender),
+3 -3
View File
@@ -1,7 +1,7 @@
//! Builds the top-level server router.
use crate::{cursor::transport::TransportRegistry, Result};
use crate::{cursor::transport::TransportRegistry, network::NetworkClients, Result};
pub fn router(registry: TransportRegistry) -> Result<axum::Router> {
super::cursor::router(registry)
pub fn router(registry: TransportRegistry, clients: NetworkClients) -> Result<axum::Router> {
super::cursor::router(registry, clients)
}
+12 -3
View File
@@ -44,9 +44,11 @@ impl App {
plugin_runtime.clone(),
config.app_version.clone(),
)?;
let clients = crate::network::NetworkClients::new(store.clone());
let provider = std::sync::Arc::new(ProviderRouter::new(
store.clone(),
plugins.clone(),
clients.clone(),
config.provider_request_timeout,
config.provider_stream_idle_timeout,
));
@@ -56,11 +58,18 @@ impl App {
compiler,
WebCache::managed()?,
plugins.clone(),
crate::config::managed_data_dir()?.join("rules"),
);
let control =
control::ControlService::new(store.clone(), provider, plugin_runtime, plugins)?;
let control = control::ControlService::new(
store.clone(),
provider,
plugin_runtime,
plugins,
clients.clone(),
config.app_version.clone(),
)?;
let harness = control.cursor_harness().clone();
let mut router = api::router(registry.clone())?;
let mut router = api::router(registry.clone(), clients)?;
router = match &config.console {
Some(ConsoleSource::Directory(directory)) => {
router.merge(control::web_router(control.clone(), directory))
-1
View File
@@ -46,7 +46,6 @@ pub struct ProviderConfig {
pub custom_headers: reqwest::header::HeaderMap,
pub max_output_tokens: Option<u64>,
pub request_timeout: Duration,
pub retry_count: u32,
pub allowed_body_fields: Option<std::collections::HashSet<String>>,
}
+1
View File
@@ -17,6 +17,7 @@ use super::ControlService;
// 开源项目广告为作者唯一收入来源,PR时请勿去除此广告
pub(super) const ADS_ENDPOINT: &str = "https://tab.leokun.cn/api/v1/ads?placement=menu";
// pub(super) const ADS_ENDPOINT: &str = "http://localhost:8080/api/v1/ads?placement=menu";
pub(super) const DEVICE_ID_HEADER: &str = "X-Cursor-Assistant-Device-ID";
pub(super) const OS_HEADER: &str = "X-Cursor-Assistant-OS";
pub(super) const APP_VERSION_HEADER: &str = "X-Cursor-Assistant-Version";
+14 -7
View File
@@ -41,6 +41,8 @@ pub struct ControlService {
provider: Arc<dyn Provider>,
plugin_runtime: PluginRuntime,
plugins: PluginRegistry,
clients: crate::network::NetworkClients,
app_version: String,
model_tests: Arc<Mutex<BTreeMap<String, CancellationToken>>>,
}
@@ -151,6 +153,8 @@ impl ControlService {
provider: Arc<dyn Provider>,
plugin_runtime: PluginRuntime,
plugins: PluginRegistry,
clients: crate::network::NetworkClients,
app_version: String,
) -> Result<Self> {
Ok(Self {
cursor_harness: CursorHarness::new(store.clone())?,
@@ -158,6 +162,8 @@ impl ControlService {
provider,
plugin_runtime,
plugins,
clients,
app_version,
model_tests: Arc::new(Mutex::new(BTreeMap::new())),
})
}
@@ -256,13 +262,13 @@ impl ControlService {
disabled_ad_ids: Option<&str>,
language: &str,
) -> Result<AdRuntime> {
let client = crate::network::client(&self.store).await?;
let client = self.clients.default_client().await?;
let installation_id = self.store.installation_id().await?;
let mut request = client
.get(ADS_ENDPOINT)
.header(DEVICE_ID_HEADER, installation_id)
.header(OS_HEADER, std::env::consts::OS)
.header(APP_VERSION_HEADER, env!("CARGO_PKG_VERSION"))
.header(APP_VERSION_HEADER, &self.app_version)
.header(LANGUAGE_HEADER, language)
.timeout(std::time::Duration::from_secs(60));
if let Some(disabled_ad_ids) = disabled_ad_ids.filter(|value| !value.is_empty()) {
@@ -281,7 +287,7 @@ impl ControlService {
}
pub(super) async fn dismiss_ad(&self, ad_id: &str, input: &AdDismissalInput) -> Result<()> {
let client = crate::network::client(&self.store).await?;
let client = self.clients.default_client().await?;
let installation_id = self.store.installation_id().await?;
let mut endpoint = Url::parse(ADS_ENDPOINT).map_err(|error| {
Error::Config(format!("advertisement endpoint is invalid: {error}"))
@@ -296,7 +302,7 @@ impl ControlService {
.post(endpoint)
.header(DEVICE_ID_HEADER, installation_id)
.header(OS_HEADER, std::env::consts::OS)
.header(APP_VERSION_HEADER, env!("CARGO_PKG_VERSION"))
.header(APP_VERSION_HEADER, &self.app_version)
.json(input)
.timeout(std::time::Duration::from_secs(5))
.send()
@@ -392,7 +398,6 @@ impl ControlService {
if model_hash.starts_with(crate::plugin::ADAPTER_ID_PREFIX) {
let descriptor = self.plugins.model_descriptor(model_hash).await?;
model.display_name = Some(descriptor.display_name);
model.context_window_tokens = descriptor.context_window_tokens;
model.max_output_tokens = Some(descriptor.max_output_tokens.unwrap_or(65_536));
} else {
let configured = self
@@ -511,7 +516,7 @@ impl ControlService {
}
pub async fn discover_models(&self, input: &ModelDiscoveryInput) -> Result<DiscoveredModels> {
let client = crate::network::client(&self.store).await?;
let client = self.clients.default_client().await?;
let base_url = crate::model::normalize_request_url(&input.base_url)?;
discover_models_from_endpoint(
&client,
@@ -678,7 +683,9 @@ impl ControlService {
}
pub async fn set_proxy_settings(&self, settings: ProxySettingsInput) -> Result<ProxySettings> {
self.store.set_proxy_settings(settings).await
let settings = self.store.set_proxy_settings(settings).await?;
self.clients.invalidate().await;
Ok(settings)
}
pub async fn tab_settings(&self) -> Result<TabSettings> {
+11 -13
View File
@@ -278,19 +278,17 @@ impl CheckpointBuilder {
),
});
if let Some(trace) = handle.trace() {
trace
.artifact(
"checkpoint",
"byok_server",
&checkpoint.encode_to_vec(),
serde_json::json!({
"root_message_count": checkpoint.root_prompt_messages_json.len(),
"turn_count": checkpoint.turns.len(),
"pending_tool_call_count": checkpoint.pending_tool_calls.len(),
"emit_status": if result.is_ok() { "sent" } else { "error" },
}),
)
.await;
trace.artifact(
"checkpoint",
"byok_server",
&checkpoint.encode_to_vec(),
serde_json::json!({
"root_message_count": checkpoint.root_prompt_messages_json.len(),
"turn_count": checkpoint.turns.len(),
"pending_tool_call_count": checkpoint.pending_tool_calls.len(),
"emit_status": if result.is_ok() { "sent" } else { "error" },
}),
);
}
result
}
@@ -102,6 +102,12 @@ pub fn decode_pending(value: &str) -> Result<RecoveredToolRound> {
.into_iter()
.enumerate()
.map(|(index, call)| {
let argument_error = wire
.pointer("/providerOptions/cursor/pendingToolExecutionContracts")
.and_then(|contracts| contracts.get(&call.call_id))
.and_then(|contract| contract.get("argumentError"))
.and_then(Value::as_str)
.map(str::to_string);
Ok(ToolCall {
index,
call_id: call.call_id,
@@ -109,6 +115,7 @@ pub fn decode_pending(value: &str) -> Result<RecoveredToolRound> {
name: call.name,
arguments_text: serde_json::to_string(&call.arguments)?,
arguments: call.arguments,
argument_error,
})
})
.collect::<Result<Vec<_>>>()?;
+20 -10
View File
@@ -71,6 +71,7 @@ pub fn staged_tool_round(
allowed_tools,
dynamic_tools,
started_at_ms,
tool_calls: Some(calls),
}),
)?)?)
}
@@ -96,6 +97,7 @@ pub fn staged_final(
allowed_tools,
dynamic_tools,
started_at_ms,
tool_calls: None,
}),
)?)?)
}
@@ -105,6 +107,7 @@ pub(super) struct PendingContext<'a> {
allowed_tools: &'a [String],
dynamic_tools: &'a HashSet<String>,
started_at_ms: u64,
tool_calls: Option<&'a [ToolCall]>,
}
pub(super) fn wire_message(
@@ -132,16 +135,23 @@ pub(super) fn wire_message(
calls
.iter()
.map(|call| {
(
call.call_id.clone(),
json!({
"toolCallId": call.call_id,
"outerToolName": call.name,
"toolIdentifier": tool_identifier(&call.name, pending.dynamic_tools),
"isDynamic": pending.dynamic_tools.contains(&call.name),
"allowedToolNames": pending.allowed_tools,
}),
)
let mut contract = json!({
"toolCallId": call.call_id,
"outerToolName": call.name,
"toolIdentifier": tool_identifier(&call.name, pending.dynamic_tools),
"isDynamic": pending.dynamic_tools.contains(&call.name),
"allowedToolNames": pending.allowed_tools,
});
if let Some(error) = pending
.tool_calls
.and_then(|calls| {
calls.iter().find(|candidate| candidate.call_id == call.call_id)
})
.and_then(|call| call.argument_error.as_deref())
{
contract["argumentError"] = Value::String(error.into());
}
(call.call_id.clone(), contract)
})
.collect(),
),
@@ -30,6 +30,7 @@ fn pending_tool_round_is_one_complete_assistant_message_and_round_trips() {
name: "Read".into(),
arguments_text: r#"{"path":"/a"}"#.into(),
arguments: json!({"path":"/a"}),
argument_error: Some("Read arguments are not valid JSON".into()),
},
ToolCall {
index: 1,
@@ -38,6 +39,7 @@ fn pending_tool_round_is_one_complete_assistant_message_and_round_trips() {
name: "Grep".into(),
arguments_text: r#"{"pattern":"x"}"#.into(),
arguments: json!({"pattern":"x"}),
argument_error: None,
},
];
let pending = staged_tool_round(
@@ -55,6 +57,10 @@ fn pending_tool_round_is_one_complete_assistant_message_and_round_trips() {
wire["providerOptions"]["cursor"]["pendingToolExecutionContracts"]["a"]["toolIdentifier"],
"READ"
);
assert_eq!(
wire["providerOptions"]["cursor"]["pendingToolExecutionContracts"]["a"]["argumentError"],
"Read arguments are not valid JSON"
);
assert_eq!(wire["role"], "assistant");
assert_eq!(
wire["providerOptions"]["cursor"]["pendingToolExecutionContracts"]
@@ -77,6 +83,10 @@ fn pending_tool_round_is_one_complete_assistant_message_and_round_trips() {
assert_eq!(recovered.assistant.replay_state, Some(replay_state));
assert_eq!(recovered.calls.len(), 2);
assert_eq!(recovered.calls[0].call_id, "a");
assert_eq!(
recovered.calls[0].argument_error.as_deref(),
Some("Read arguments are not valid JSON")
);
assert_eq!(recovered.calls[1].call_id, "b");
}
+27
View File
@@ -75,6 +75,11 @@ impl StepBuffer {
});
}
pub fn finish_model_attempt(&mut self) {
self.finish_text();
self.finish_thinking(Duration::ZERO);
}
pub fn discard_model_output(&mut self) {
self.text.clear();
self.thinking.clear();
@@ -101,6 +106,28 @@ impl StepBuffer {
mod tests {
use super::*;
#[test]
fn failed_attempt_output_is_retained_for_the_next_checkpoint() {
let mut buffer = StepBuffer::default();
buffer.text_delta("partial answer");
buffer.thinking_delta("partial reasoning");
buffer.finish_model_attempt();
let steps = buffer.take().steps;
assert_eq!(steps.len(), 2);
assert!(matches!(
&steps[0].message,
Some(pb::conversation_step::Message::AssistantMessage(message))
if message.text == "partial answer"
));
assert!(matches!(
&steps[1].message,
Some(pb::conversation_step::Message::ThinkingMessage(message))
if message.text == "partial reasoning"
));
}
#[test]
fn interrupted_model_output_is_not_persisted_as_checkpoint_steps() {
let mut buffer = StepBuffer::default();
+7 -1
View File
@@ -3,7 +3,7 @@ use prost::Message;
use crate::{
cursor::{checkpoint::PendingSteps, protocol::proto::agent::v1 as pb},
model::CanonicalMessage,
model::{estimate_context_tokens, project_messages, CanonicalMessage, PromptSpec},
store::{BlobEdge, BlobId},
Error, Result,
};
@@ -85,6 +85,12 @@ impl CheckpointBuilder {
.push(archive_id.as_bytes().to_vec());
}
self.base.self_summary_count = self.base.self_summary_count.saturating_add(1);
let projected = project_messages(messages)?;
let prompt = PromptSpec {
instructions: self.instructions.clone(),
tools: self.tool_definitions.clone(),
};
self.record_context_tokens(Some(estimate_context_tokens(&prompt, &projected)));
if let Some(details) = self.base.token_details.as_mut() {
details.breakdown = Some(crate::cursor::services::usage::breakdown(
details.used_tokens,
+78 -14
View File
@@ -12,7 +12,7 @@ use crate::{
protocol::proto::agent::v1 as pb, services::context_sync::RequestContextSynchronizer,
tools::runtime::McpRoute,
},
model::ToolDefinition,
model::{normalize_tool_name, ToolDefinition},
store::BlobId,
Error, Result,
};
@@ -146,6 +146,37 @@ async fn decode_part<T: Message + Default>(
.map_err(|error| Error::Protocol(format!("invalid {name} context Blob: {error}")))
}
/// 把本地 md 规则目录(rules 服务的存储)合并进请求上下文,
/// 使 BYOK 运行在 IDE 未携带这些规则时也能消费它们。
/// 与 IDE 已发规则按内容去重;读取失败只告警,不影响运行。
pub fn merge_local_rules(context: &mut pb::RequestContext, rules_dir: &Path) {
let records = match crate::cursor::services::knowledge::RuleStore::open(rules_dir.into())
.and_then(|store| store.list())
{
Ok(records) => records,
Err(error) => {
tracing::warn!(%error, "cannot read local rules; continuing without them");
return;
}
};
let existing = context
.rules
.iter()
.chain(context.non_file_rules.iter())
.map(|rule| rule.content.trim().to_owned())
.chain(context.cloud_rule.iter().map(|rule| rule.trim().to_owned()))
.collect::<HashSet<_>>();
for record in records {
if record.knowledge.trim().is_empty() || existing.contains(record.knowledge.trim()) {
continue;
}
context.non_file_rules.push(pb::CursorRule {
content: record.knowledge,
..Default::default()
});
}
}
pub fn request_context(request: &pb::AgentRunRequest) -> Option<&pb::RequestContext> {
let action = request.action.as_ref()?;
action
@@ -462,7 +493,7 @@ pub fn dynamic_mcp(
})?),
};
let parameters = normalize_mcp_parameters(&wire.name, parameters)?;
let name = model_tool_name(&wire.name);
let name = normalize_tool_name(&wire.name);
let definition = ToolDefinition {
name: name.clone(),
description: wire.description.clone(),
@@ -522,18 +553,6 @@ fn invalid_mcp_parameters(tool_name: &str) -> Error {
))
}
fn model_tool_name(name: &str) -> String {
name.chars()
.map(|character| {
if character.is_ascii_alphanumeric() || matches!(character, '_' | '-') {
character
} else {
'_'
}
})
.collect()
}
fn prost_value(value: &prost_types::Value) -> Value {
use prost_types::value::Kind;
match value.kind.as_ref() {
@@ -563,3 +582,48 @@ fn xml(value: &str) -> String {
.replace('<', "&lt;")
.replace('>', "&gt;")
}
#[cfg(test)]
mod tests {
use super::*;
fn rule(content: &str) -> pb::CursorRule {
pb::CursorRule {
content: content.into(),
..Default::default()
}
}
#[test]
fn merge_local_rules_appends_and_dedupes_by_content() {
let directory = tempfile::tempdir().unwrap();
std::fs::write(directory.path().join("a.md"), "shared rule").unwrap();
std::fs::write(directory.path().join("b.md"), "local only rule").unwrap();
std::fs::write(directory.path().join("c.md"), " \n").unwrap();
let mut context = pb::RequestContext {
non_file_rules: vec![rule(" shared rule ")],
..Default::default()
};
merge_local_rules(&mut context, directory.path());
let contents = context
.non_file_rules
.iter()
.map(|rule| rule.content.as_str())
.collect::<Vec<_>>();
assert_eq!(
contents,
[" shared rule ", "local only rule"],
"IDE-sent duplicate is kept once and blank local rules are skipped"
);
}
#[test]
fn merge_local_rules_survives_a_missing_directory() {
let directory = tempfile::tempdir().unwrap();
let mut context = pb::RequestContext::default();
merge_local_rules(&mut context, &directory.path().join("nested/rules"));
assert!(context.non_file_rules.is_empty());
}
}
+8 -4
View File
@@ -51,6 +51,7 @@ pub(crate) struct PrepareDependencies<'a> {
pub checkpoint: &'a CheckpointBuilder,
pub blob_sync: &'a BlobSynchronizer,
pub context_sync: &'a RequestContextSynchronizer,
pub local_rules_dir: Option<&'a std::path::Path>,
}
pub(crate) async fn prepare(
@@ -64,6 +65,7 @@ pub(crate) async fn prepare(
checkpoint,
blob_sync,
context_sync,
local_rules_dir,
} = dependencies;
checkpoint
.import_prefetched(&request.pre_fetched_blobs)
@@ -115,11 +117,13 @@ pub(crate) async fn prepare(
"selected_source": "root_prompt_messages_json",
});
let encoded = serde_json::to_vec(&summary)?;
trace
.artifact("history_projection", "byok_server", &encoded, summary)
.await;
trace.artifact("history_projection", "byok_server", &encoded, summary);
}
let request_context = context::hydrate(request, context_sync).await?;
let mut request_context = context::hydrate(request, context_sync).await?;
if let Some(rules_dir) = local_rules_dir {
context::merge_local_rules(&mut request_context, rules_dir);
}
let request_context = request_context;
let ActionProjection {
mode: mode_number,
mut turn_user,
+18 -2
View File
@@ -1,6 +1,19 @@
//! Defines commands accepted by a Conversation runtime.
use crate::cursor::protocol::proto::agent::v1 as pb;
use crate::{cursor::protocol::proto::agent::v1 as pb, Error};
#[derive(Debug)]
pub enum RunFinish {
TurnCompleted,
Transport(TransportFinish),
}
#[derive(Debug)]
pub enum TransportFinish {
Success,
Failed(Error),
Cancelled,
}
#[derive(Debug)]
pub enum TransportCommand {
@@ -8,6 +21,9 @@ pub enum TransportCommand {
seqno: i64,
message: Box<pb::AgentClientMessage>,
},
RunFinished {
generation: u64,
finish: RunFinish,
},
Disconnect,
Close,
}
+128 -37
View File
@@ -21,7 +21,7 @@ use crate::{
protocol::proto::agent::v1 as pb,
services::blob_sync::BlobSynchronizer,
tools::{
codec,
codec, compat,
runtime::CursorToolRuntime,
stream::ToolCallStream,
tool_call_result::{ToolCompletion, ToolResultReceiver},
@@ -34,7 +34,7 @@ use crate::{
Error, Result,
};
use super::{CompiledMessages, ConversationRegistry, MessageDelivery};
use super::{CompiledMessages, ConversationRegistry, MessageDelivery, RunFinish, TransportFinish};
use crate::cursor::transport::TransportHandle;
pub struct ConversationOutput {
@@ -110,7 +110,7 @@ impl ConversationOutput {
}
}
pub async fn run(mut self) -> Result<()> {
pub async fn run(mut self) -> Result<RunFinish> {
let result = self.run_inner().await;
if let Err(error) = &result {
if !self.superseded.is_cancelled() {
@@ -143,7 +143,7 @@ impl ConversationOutput {
result
}
async fn run_inner(&mut self) -> Result<()> {
async fn run_inner(&mut self) -> Result<RunFinish> {
if self.context.compacting {
self.handle.emit(&events::summary_started())?;
}
@@ -175,7 +175,7 @@ impl ConversationOutput {
if self.superseded.is_cancelled() {
worker.abort();
self.abort_execs().await;
return Ok(());
return Ok(RunFinish::Transport(TransportFinish::Cancelled));
}
let input = if let Ok(action) = self.runtime_actions.try_recv() {
Input::RuntimeAction(Some(Box::new(action)))
@@ -187,7 +187,7 @@ impl ConversationOutput {
_ = self.superseded.cancelled() => {
worker.abort();
self.abort_execs().await;
return Ok(());
return Ok(RunFinish::Transport(TransportFinish::Cancelled));
}
action = self.runtime_actions.recv() => Input::RuntimeAction(action.map(Box::new)),
event = self.core.events.recv() => Input::Event(event),
@@ -264,6 +264,32 @@ impl ConversationOutput {
streams.clear();
presentation.discard_model_output();
}
RunEvent::ModelAttemptFailed { attempt, message } => {
tracing::warn!(
run_id = %self.run.run_id(),
attempt,
%message,
"retrying model call from current checkpoint"
);
presentation.finish_model_attempt();
for call in calls.values_mut() {
if call.arguments.is_null() {
call.arguments = serde_json::from_str(&call.arguments_text)
.unwrap_or_else(|_| serde_json::json!({}));
}
let completion = compat::failure_with_message(
call,
format!("Model attempt failed before tool completion: {message}"),
);
self.handle
.emit(&codec::tool_completed(call, &completion))?;
presentation.tool_completed(&completion);
}
response_text.clear();
response_thinking.clear();
calls.clear();
streams.clear();
}
RunEvent::TextStart => {}
RunEvent::TextEnd => {
if !self.context.compacting {
@@ -312,6 +338,7 @@ impl ConversationOutput {
name: name.clone(),
arguments_text: String::new(),
arguments: serde_json::Value::Null,
argument_error: None,
};
self.emit_model_event(
crate::provider::ModelEvent::ToolCallStart {
@@ -335,15 +362,50 @@ impl ConversationOutput {
let stream = streams.get_mut(&index).ok_or_else(|| {
Error::Protocol(format!("missing Cursor tool stream: {index}"))
})?;
for message in stream.arguments_delta(call, &delta)? {
self.handle.emit(&message)?;
match stream.arguments_delta(call, &delta) {
Ok(messages) => {
for message in messages {
self.handle.emit(&message)?;
}
}
Err(Error::Protocol(message)) => {
tracing::warn!(
call_id = %call.call_id,
%message,
"ignoring invalid streaming tool arguments until completion"
);
}
Err(Error::Json(error)) => {
tracing::warn!(
call_id = %call.call_id,
%error,
"ignoring invalid streaming tool arguments until completion"
);
}
Err(error) => return Err(error),
}
}
RunEvent::ToolCallEnd { index } => {
let call = calls.get_mut(&index).ok_or_else(|| {
Error::Protocol(format!("unknown completed tool index: {index}"))
})?;
call.arguments = serde_json::from_str(&call.arguments_text)?;
// A tool call with no arguments streams no argument text.
// Treat empty text as an empty object, matching the model
// cycle, instead of failing the run on `from_str("")`.
call.arguments = if call.arguments_text.trim().is_empty() {
serde_json::json!({})
} else {
serde_json::from_str(&call.arguments_text)
.unwrap_or_else(|_| serde_json::json!({}))
};
}
RunEvent::UsageSnapshot(usage) => {
if !self.context.compacting {
if let Some(output_tokens) = usage.output_tokens {
self.handle.emit(&events::token_delta(output_tokens))?;
}
context_tokens = usage.context_input_tokens;
}
}
RunEvent::Usage(usage) => {
if !self.context.compacting {
@@ -352,10 +414,7 @@ impl ConversationOutput {
}
}
if !self.context.compacting {
context_tokens = usage
.input_tokens
.zip(usage.output_tokens)
.and_then(|(input, output)| input.checked_add(output));
context_tokens = usage.context_input_tokens;
}
match &mut turn_usage {
Some(total) => *total += usage,
@@ -424,21 +483,28 @@ impl ConversationOutput {
streams.clear();
}
if let CommitCause::RuntimeEvent { event_id } = &state.cause {
if let Some(injection_id) = event_id.strip_prefix("inject-context:") {
if let Some(pending) = self.pending_injections.remove(injection_id)
{
let delivered_at_ms = crate::cursor::tools::runtime::now_ms()
.min(i64::MAX as u64)
as i64;
self.handle.emit(&events::context_injection_delivered(
injection_id.to_owned(),
pending.delivery_batch_id.clone(),
delivered_at_ms,
))?;
if let Some(user_message) = pending.user_message {
self.handle
.emit(&events::user_message_appended(user_message))?;
}
// Injections key `pending_injections` by their raw
// injection id and commit under `inject-context:{id}`,
// while runtime user messages key it by (and commit
// under) the full `user-message:{id}` event id. Strip
// the injection prefix when present and otherwise use
// the event id verbatim so both are cleared and emit
// their delivered/appended events.
let injection_id = event_id
.strip_prefix("inject-context:")
.unwrap_or(event_id.as_str());
if let Some(pending) = self.pending_injections.remove(injection_id) {
let delivered_at_ms = crate::cursor::tools::runtime::now_ms()
.min(i64::MAX as u64)
as i64;
self.handle.emit(&events::context_injection_delivered(
injection_id.to_owned(),
pending.delivery_batch_id.clone(),
delivered_at_ms,
))?;
if let Some(user_message) = pending.user_message {
self.handle
.emit(&events::user_message_appended(user_message))?;
}
}
}
@@ -509,6 +575,7 @@ impl ConversationOutput {
.map_err(|_| Error::Protocol("checkpoint worker stopped".into()))?
{
Ok(checkpoint) => {
context_tokens = checkpoint_context_tokens(&checkpoint);
compaction_checkpoint = Some(checkpoint);
state.barrier.complete(Ok(()));
}
@@ -652,7 +719,7 @@ impl ConversationOutput {
if self.superseded.is_cancelled() {
worker.abort();
self.abort_execs().await;
return Ok(());
return Ok(RunFinish::Transport(TransportFinish::Cancelled));
}
return match outcome {
RunOutcome::Completed => {
@@ -668,8 +735,7 @@ impl ConversationOutput {
for _ in 0..3 {
self.checkpoint.publish(&self.handle, &checkpoint).await?;
}
finish_success(&self.handle);
return Ok(());
return Ok(RunFinish::TurnCompleted);
}
let checkpoints = final_checkpoint.take().ok_or_else(|| {
Error::Protocol("Completed without final state".into())
@@ -685,18 +751,19 @@ impl ConversationOutput {
ttft_breakdown: None,
message: Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(checkpoints.settled)),
})?;
finish_success(&self.handle);
Ok(())
Ok(RunFinish::TurnCompleted)
}
RunOutcome::Cancelled => {
worker.abort();
self.abort_execs().await;
finish_cancelled(&self.handle)
Ok(RunFinish::Transport(TransportFinish::Cancelled))
}
RunOutcome::Failed(failure) => {
worker.abort();
self.abort_execs().await;
finish_failed(&self.handle, &cursor_error(failure))
Ok(RunFinish::Transport(TransportFinish::Failed(cursor_error(
failure,
))))
}
};
}
@@ -1015,10 +1082,34 @@ pub(crate) fn finish_cancelled(handle: &TransportHandle) -> Result<()> {
Ok(())
}
fn checkpoint_context_tokens(checkpoint: &pb::ConversationStateStructure) -> Option<u64> {
checkpoint
.token_details
.as_ref()
.map(|details| u64::from(details.used_tokens))
}
#[cfg(test)]
mod tests {
use super::accept_tool_completion;
use crate::{run::CommandResult, Error};
use super::{accept_tool_completion, checkpoint_context_tokens};
use crate::{cursor::protocol::proto::agent::v1 as pb, run::CommandResult, Error};
#[test]
fn compacted_checkpoint_replaces_the_in_memory_context_usage() {
let compacted = pb::ConversationStateStructure {
token_details: Some(pb::ConversationTokenDetails {
used_tokens: 20_000,
..Default::default()
}),
..Default::default()
};
assert_eq!(checkpoint_context_tokens(&compacted), Some(20_000));
assert_eq!(
checkpoint_context_tokens(&pb::ConversationStateStructure::default()),
None
);
}
#[test]
fn closing_and_ended_runs_ignore_known_tool_completions() {
@@ -26,6 +26,8 @@ pub(crate) struct ConversationDependencies {
pub provider: Arc<dyn Provider>,
pub compiler: PromptCompiler,
pub web_cache: WebCache,
/// 本地 rules 服务的 md 存储目录;编译请求上下文时合并其中的规则。
pub local_rules_dir: Option<std::path::PathBuf>,
}
struct RegistryInner {
@@ -47,6 +49,7 @@ impl ConversationRegistry {
provider: Arc<dyn Provider>,
compiler: PromptCompiler,
web_cache: WebCache,
local_rules_dir: Option<std::path::PathBuf>,
) -> Self {
Self {
inner: Arc::new(RegistryInner {
@@ -58,6 +61,7 @@ impl ConversationRegistry {
provider,
compiler,
web_cache,
local_rules_dir,
},
}),
}
+300 -105
View File
@@ -11,25 +11,29 @@ use crate::{
protocol::proto::agent::v1 as pb,
services::{blob_sync::BlobSynchronizer, context_sync::RequestContextSynchronizer},
tools::{
codec,
codec, compat,
runtime::CursorToolRuntime,
tool_call_result::{tool_result_channel, ToolResultReceiver, ToolResultSender},
ClientToolEvent, ToolDispatcher,
},
transport::{OrderedInbox, TransportHandle},
},
run::{CommandResult, RunEngine, RunHandle},
run::{CommandResult, RunEngine, RunHandle, RunPhase},
};
use super::{
CompiledMessages, ConversationDependencies, ConversationOutput, ConversationOutputDependencies,
ConversationRegistry, MessageDelivery, TransportCommand,
ConversationRegistry, MessageDelivery, RunFinish, TransportCommand, TransportFinish,
};
pub struct ConversationRuntime;
const CONTINUATION_IDLE_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(2);
#[derive(Clone)]
struct RunGeneration {
id: u64,
request: pb::AgentRunRequest,
superseded: CancellationToken,
finished: CancellationToken,
run: Arc<parking_lot::Mutex<Option<RunHandle>>>,
@@ -41,6 +45,14 @@ struct RunGeneration {
struct FinishGeneration(CancellationToken);
struct TransportActorGuard(TransportHandle);
impl Drop for TransportActorGuard {
fn drop(&mut self) {
self.0.close_transport();
}
}
impl Drop for FinishGeneration {
fn drop(&mut self) {
self.0.cancel();
@@ -54,6 +66,7 @@ impl ConversationRuntime {
mut receiver: mpsc::Receiver<TransportCommand>,
) {
tokio::spawn(async move {
let _actor_guard = TransportActorGuard(handle.clone());
let dependencies = registry.dependencies().clone();
let blob_sync = BlobSynchronizer::new(
handle.request_id().into(),
@@ -65,19 +78,70 @@ impl ConversationRuntime {
let context_sync =
RequestContextSynchronizer::new(handle.clone(), dependencies.store.clone());
let mut current = None::<RunGeneration>;
let mut next_generation = 1_u64;
let mut pending_finish = None::<(u64, TransportFinish)>;
let mut draining = false;
let mut waiting_for_action = false;
loop {
let command = match receiver.recv().await {
Some(command) => command,
None => {
handle.mark_disconnected();
if let Some(generation) = current.as_ref() {
generation.superseded.cancel();
if let Some(run) = generation.run.lock().clone() {
run.cancel();
let command = if draining {
if !handle.admissions_drained() {
tokio::select! {
command = receiver.recv() => match command {
Some(command) => command,
None => {
finish_pending(&handle, &current, pending_finish.take());
break;
}
},
_ = handle.wait_admissions_drained() => continue,
}
} else {
handle.mark_draining();
match receiver.try_recv() {
Ok(command) => command,
Err(mpsc::error::TryRecvError::Empty)
| Err(mpsc::error::TryRecvError::Disconnected) => {
finish_pending(&handle, &current, pending_finish.take());
break;
}
}
super::finish_cancelled(&handle).ok();
break;
}
} else if waiting_for_action {
tokio::select! {
command = receiver.recv() => match command {
Some(command) => command,
None => {
handle.mark_disconnected();
super::finish_success(&handle);
break;
}
},
_ = tokio::time::sleep(CONTINUATION_IDLE_TIMEOUT) => {
let Some(generation) = current.as_ref() else {
super::finish_success(&handle);
break;
};
handle.begin_close();
pending_finish = Some((generation.id, TransportFinish::Success));
draining = true;
waiting_for_action = false;
continue;
}
}
} else {
match receiver.recv().await {
Some(command) => command,
None => {
handle.mark_disconnected();
if let Some(generation) = current.as_ref() {
generation.superseded.cancel();
if let Some(run) = generation.run.lock().clone() {
run.cancel();
}
}
super::finish_cancelled(&handle).ok();
break;
}
}
};
match command {
@@ -92,11 +156,41 @@ impl ConversationRuntime {
let _ = handle.emit(&codec::abort(id));
}
}
super::finish_cancelled(&handle).ok();
let turn_completed = waiting_for_action
|| current.as_ref().is_some_and(|generation| {
generation
.run
.lock()
.as_ref()
.is_none_or(|run| run.phase() != RunPhase::Running)
});
if turn_completed {
super::finish_success(&handle);
} else {
super::finish_cancelled(&handle).ok();
}
break;
}
TransportCommand::Close => {
break;
TransportCommand::RunFinished { generation, finish } => {
if !current
.as_ref()
.is_some_and(|current| current.id == generation)
{
continue;
}
match finish {
RunFinish::TurnCompleted => {
pending_finish = None;
draining = false;
waiting_for_action = true;
}
RunFinish::Transport(finish) => {
waiting_for_action = false;
handle.begin_close();
pending_finish = Some((generation, finish));
draining = true;
}
}
}
TransportCommand::Append { seqno, message } => {
for (_seqno, message) in inbox.push(seqno, *message) {
@@ -105,6 +199,12 @@ impl ConversationRuntime {
Some(pb::agent_client_message::Message::RunRequest(
request,
)) => {
waiting_for_action = false;
if draining {
handle.reopen();
draining = false;
pending_finish = None;
}
if let Some(conversation_id) =
request.conversation_id.as_deref()
{
@@ -117,60 +217,21 @@ impl ConversationRuntime {
"invalid Cursor conversation id"
);
let _ = super::finish_failed(&handle, &error);
let _ =
handle.command(TransportCommand::Close).await;
return;
}
}
let previous_finished =
if let Some(previous) = current.take() {
previous.superseded.cancel();
if let Some(run) = previous.run.lock().clone() {
run.cancel();
}
for id in previous
.tool_runtime
.interrupt_for_run_replacement()
.await
{
let _ = handle.emit(&codec::abort(id));
}
Some(previous.finished.clone())
} else {
None
};
let (results, result_receiver) = tool_result_channel();
let (runtime_actions, runtime_action_receiver) =
mpsc::unbounded_channel::<compile::RuntimeAction>();
let tool_runtime = tool_runtime_factory.next_run();
let tools = ToolDispatcher::with_results(
tool_runtime.clone(),
results.clone(),
dependencies.store.clone(),
dependencies.web_cache.clone(),
);
let generation = RunGeneration {
superseded: CancellationToken::new(),
finished: CancellationToken::new(),
run: Arc::new(parking_lot::Mutex::new(None)),
results,
runtime_actions,
tool_runtime,
tools,
};
current = Some(generation.clone());
spawn_run_request(
registry.clone(),
handle.clone(),
start_generation(
&registry,
&handle,
&dependencies,
&blob_sync,
&context_sync,
&tool_runtime_factory,
&mut current,
&mut next_generation,
request,
dependencies.clone(),
blob_sync.clone(),
context_sync.clone(),
generation,
previous_finished,
result_receiver,
runtime_action_receiver,
);
)
.await;
}
Some(pb::agent_client_message::Message::ExecClientMessage(
message,
@@ -262,17 +323,18 @@ impl ConversationRuntime {
.take_exec(throw.id)
.await
{
Some(pending) => generation.results.send_error(
crate::Error::Protocol(format!(
"Exec {} failed: {}",
pending.call.call_id, throw.error
)),
Some(pending) => generation.results.send(
compat::failure_with_message(
&pending.call,
format!(
"Exec {} failed: {}",
pending.call.call_id, throw.error
),
),
),
None => generation.results.send_error(
crate::Error::Protocol(format!(
"unknown ExecClientThrow id: {}",
throw.id
)),
None => tracing::warn!(
id = throw.id,
"ignoring failure for unknown tool execution"
),
}
}
@@ -330,27 +392,48 @@ impl ConversationRuntime {
// return an explicit Protocol Error rather than falling through silently.
Some(
pb::agent_client_message::Message::ConversationAction(
action,
conversation_action,
),
) => match action.action {
) => match conversation_action.action.clone() {
Some(
pb::conversation_action::Action::UserMessageAction(
action,
),
) => {
let Some(generation) = current.as_ref() else {
let delivered_to_active_run =
current.as_ref().is_some_and(|generation| {
generation.run.lock().as_ref().is_some_and(
|run| run.phase() == RunPhase::Running,
) && generation
.runtime_actions
.send(compile::RuntimeAction::UserMessage(
action.clone(),
))
.is_ok()
});
if delivered_to_active_run {
continue;
}
let Some(previous) = current.as_ref() else {
continue;
};
if generation
.runtime_actions
.send(compile::RuntimeAction::UserMessage(action))
.is_err()
{
generation.results.send_error(crate::Error::Protocol(
"UserMessageAction arrived without an active Run"
.into(),
));
}
let mut request = previous.request.clone();
request.action = Some(conversation_action);
request.conversation_state = None;
request.pre_fetched_blobs.clear();
waiting_for_action = false;
start_generation(
&registry,
&handle,
&dependencies,
&blob_sync,
&context_sync,
&tool_runtime_factory,
&mut current,
&mut next_generation,
request,
)
.await;
}
Some(pb::conversation_action::Action::CancelAction(_)) => {
if let Some(generation) = current.as_ref() {
@@ -427,6 +510,95 @@ impl ConversationRuntime {
}
}
#[allow(clippy::too_many_arguments)]
async fn start_generation(
registry: &ConversationRegistry,
handle: &TransportHandle,
dependencies: &ConversationDependencies,
blob_sync: &BlobSynchronizer,
context_sync: &RequestContextSynchronizer,
tool_runtime_factory: &CursorToolRuntime,
current: &mut Option<RunGeneration>,
next_generation: &mut u64,
request: pb::AgentRunRequest,
) {
let previous_finished = if let Some(previous) = current.take() {
previous.superseded.cancel();
if let Some(run) = previous.run.lock().clone() {
run.cancel();
}
for id in previous.tool_runtime.interrupt_for_run_replacement().await {
let _ = handle.emit(&codec::abort(id));
}
Some(previous.finished.clone())
} else {
None
};
let (results, result_receiver) = tool_result_channel();
let (runtime_actions, runtime_action_receiver) =
mpsc::unbounded_channel::<compile::RuntimeAction>();
let tool_runtime = tool_runtime_factory.next_run();
let tools = ToolDispatcher::with_results(
tool_runtime.clone(),
results.clone(),
dependencies.store.clone(),
dependencies.web_cache.clone(),
);
let generation = RunGeneration {
id: *next_generation,
request: request.clone(),
superseded: CancellationToken::new(),
finished: CancellationToken::new(),
run: Arc::new(parking_lot::Mutex::new(None)),
results,
runtime_actions,
tool_runtime,
tools,
};
*next_generation = next_generation.saturating_add(1);
*current = Some(generation.clone());
spawn_run_request(
registry.clone(),
handle.clone(),
request,
dependencies.clone(),
blob_sync.clone(),
context_sync.clone(),
generation,
previous_finished,
result_receiver,
runtime_action_receiver,
);
}
fn finish_pending(
handle: &TransportHandle,
current: &Option<RunGeneration>,
pending: Option<(u64, TransportFinish)>,
) {
let Some((generation, finish)) = pending else {
return;
};
if current
.as_ref()
.is_some_and(|current| current.id == generation)
{
finish_transport(handle, finish);
}
}
fn finish_transport(handle: &TransportHandle, finish: TransportFinish) {
match finish {
TransportFinish::Success => super::finish_success(handle),
TransportFinish::Failed(error) => {
let _ = super::finish_failed(handle, &error);
}
TransportFinish::Cancelled => {
let _ = super::finish_cancelled(handle);
}
}
}
#[allow(clippy::too_many_arguments)]
fn spawn_run_request(
registry: ConversationRegistry,
@@ -471,6 +643,7 @@ fn spawn_run_request(
checkpoint: &checkpoint,
blob_sync: &blob_sync,
context_sync: &context_sync,
local_rules_dir: dependencies.local_rules_dir.as_deref(),
},
) => prepared,
};
@@ -485,8 +658,12 @@ fn spawn_run_request(
%error,
"failed to prepare Cursor Run"
);
let _ = super::finish_failed(&handle, &error);
let _ = handle.command(TransportCommand::Close).await;
let _ = handle
.command(TransportCommand::RunFinished {
generation: generation.id,
finish: RunFinish::Transport(TransportFinish::Failed(error)),
})
.await;
return;
}
};
@@ -518,8 +695,12 @@ fn spawn_run_request(
{
CommandResult::Applied | CommandResult::Duplicate => {
if !generation.superseded.is_cancelled() {
super::finish_success(&handle);
let _ = handle.command(TransportCommand::Close).await;
let _ = handle
.command(TransportCommand::RunFinished {
generation: generation.id,
finish: RunFinish::Transport(TransportFinish::Success),
})
.await;
}
return;
}
@@ -543,8 +724,12 @@ fn spawn_run_request(
}
CommandResult::StaleTarget => {
if !generation.superseded.is_cancelled() {
super::finish_success(&handle);
let _ = handle.command(TransportCommand::Close).await;
let _ = handle
.command(TransportCommand::RunFinished {
generation: generation.id,
finish: RunFinish::Transport(TransportFinish::Success),
})
.await;
}
return;
}
@@ -603,16 +788,21 @@ fn spawn_run_request(
tool_runtime: generation.tool_runtime.clone(),
},
);
if let Err(error) = output.run().await {
if !generation.superseded.is_cancelled() {
tracing::error!(
request_id = handle.request_id(),
%error,
"Cursor session failed"
);
let _ = super::finish_failed(&handle, &error);
let finish = match output.run().await {
Ok(finish) => finish,
Err(error) => {
if generation.superseded.is_cancelled() {
RunFinish::Transport(TransportFinish::Cancelled)
} else {
tracing::error!(
request_id = handle.request_id(),
%error,
"Cursor session failed"
);
RunFinish::Transport(TransportFinish::Failed(error))
}
}
}
};
let _ = core_run.await;
registry.release(&conversation_id, &run_id).await;
if generation
@@ -624,7 +814,12 @@ fn spawn_run_request(
*generation.run.lock() = None;
}
if !generation.superseded.is_cancelled() {
let _ = handle.command(TransportCommand::Close).await;
let _ = handle
.command(TransportCommand::RunFinished {
generation: generation.id,
finish,
})
.await;
}
});
}
+73 -65
View File
@@ -20,6 +20,9 @@ use crate::{
type BlobSetSender = oneshot::Sender<Result<()>>;
const SET_TIMEOUT: Duration = Duration::from_secs(30 * 60);
const GET_TIMEOUT: Duration = Duration::from_secs(10 * 60);
#[derive(Clone)]
pub struct BlobSynchronizer {
inner: Arc<Inner>,
@@ -73,22 +76,20 @@ impl BlobSynchronizer {
let id = self.inner.store.put_blob(data, edges).await?;
let result = self.ensure_set(&id, data).await;
if let Some(trace) = self.inner.handle.trace() {
trace
.linked_blob(
"blob_set",
"byok_server",
&id,
serde_json::json!({
"byte_count": data.len(),
"status": if result.is_ok() { "acknowledged" } else { "error" },
"error": result.as_ref().err().map(ToString::to_string),
"edges": edges.iter().map(|edge| serde_json::json!({
"child_blob_id": edge.child.to_base64(),
"field_name": edge.field_name,
})).collect::<Vec<_>>(),
}),
)
.await;
trace.linked_blob(
"blob_set",
"byok_server",
&id,
serde_json::json!({
"byte_count": data.len(),
"status": if result.is_ok() { "acknowledged" } else { "error" },
"error": result.as_ref().err().map(ToString::to_string),
"edges": edges.iter().map(|edge| serde_json::json!({
"child_blob_id": edge.child.to_base64(),
"field_name": edge.field_name,
})).collect::<Vec<_>>(),
}),
);
}
result?;
Ok(id)
@@ -130,7 +131,7 @@ impl BlobSynchronizer {
let result = tokio::select! {
result = receiver => result.map_err(|_| Error::Protocol("KV SET response channel closed".into()))?,
_ = cancellation.cancelled() => Err(Error::Cancelled),
_ = tokio::time::sleep(Duration::from_secs(60)) => Err(Error::Protocol(format!("KV SET timed out: {}", blob_id.to_base64()))),
_ = tokio::time::sleep(SET_TIMEOUT) => Err(Error::Protocol(format!("KV SET timed out: {}", blob_id.to_base64()))),
};
if result.is_err() {
self.inner.set_requests.lock().await.remove(&id);
@@ -141,18 +142,16 @@ impl BlobSynchronizer {
pub async fn get(&self, blob_id: &BlobId) -> Result<Option<Vec<u8>>> {
if let Some(data) = self.inner.store.get_blob(blob_id).await? {
if let Some(trace) = self.inner.handle.trace() {
trace
.linked_blob(
"blob_get",
"byok_server",
blob_id,
serde_json::json!({
"byte_count": data.len(),
"source": "local_store",
"status": "found",
}),
)
.await;
trace.linked_blob(
"blob_get",
"byok_server",
blob_id,
serde_json::json!({
"byte_count": data.len(),
"source": "local_store",
"status": "found",
}),
);
}
return Ok(Some(data));
}
@@ -183,7 +182,7 @@ impl BlobSynchronizer {
let result = tokio::select! {
result = receiver => result.map_err(|_| Error::Protocol("KV GET response channel closed".into()))?,
_ = cancellation.cancelled() => Err(Error::Cancelled),
_ = tokio::time::sleep(Duration::from_secs(60)) => Err(Error::Protocol(format!("KV GET timed out: {}", blob_id.to_base64()))),
_ = tokio::time::sleep(GET_TIMEOUT) => Err(Error::Protocol(format!("KV GET timed out: {}", blob_id.to_base64()))),
};
if result.is_err() {
self.inner.get_requests.lock().await.remove(&id);
@@ -191,45 +190,39 @@ impl BlobSynchronizer {
if let Some(trace) = self.inner.handle.trace() {
match &result {
Ok(Some(data)) => {
trace
.linked_blob(
"blob_get",
"cursor_client",
blob_id,
serde_json::json!({
"byte_count": data.len(),
"source": "cursor_client",
"status": "found",
}),
)
.await;
trace.linked_blob(
"blob_get",
"cursor_client",
blob_id,
serde_json::json!({
"byte_count": data.len(),
"source": "cursor_client",
"status": "found",
}),
);
}
Ok(None) => {
trace
.artifact(
"blob_get",
"cursor_client",
&[],
serde_json::json!({
"blob_id": blob_id.to_base64(),
"status": "missing",
}),
)
.await;
trace.artifact(
"blob_get",
"cursor_client",
&[],
serde_json::json!({
"blob_id": blob_id.to_base64(),
"status": "missing",
}),
);
}
Err(error) => {
trace
.artifact(
"blob_get",
"cursor_client",
&[],
serde_json::json!({
"blob_id": blob_id.to_base64(),
"status": "error",
"error": error.to_string(),
}),
)
.await;
trace.artifact(
"blob_get",
"cursor_client",
&[],
serde_json::json!({
"blob_id": blob_id.to_base64(),
"status": "error",
"error": error.to_string(),
}),
);
}
}
}
@@ -318,3 +311,18 @@ impl BlobSynchronizer {
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn set_timeout_allows_slow_cursor_acknowledgements() {
assert_eq!(SET_TIMEOUT, Duration::from_secs(30 * 60));
}
#[test]
fn get_timeout_allows_slow_cursor_responses() {
assert_eq!(GET_TIMEOUT, Duration::from_secs(10 * 60));
}
}
+356
View File
@@ -0,0 +1,356 @@
//! Serves Cursor user rules: upstream-first with an offline markdown cache.
//!
//! 每个请求先回放离线日志再尝试上游;上游成功时把结果写穿到本地镜像,
//! 上游不可达时降级为本地 md 存储并记录日志等待回放。
mod store;
mod sync;
use std::sync::Arc;
use axum::{
body::{to_bytes, Body, Bytes},
extract::Extension,
http::{header, Request, Response},
};
use prost::Message;
use crate::{api::cursor::proxy, config, cursor::protocol::connect, Result};
pub(crate) use store::{RuleRecord, RuleStore};
#[derive(Clone, PartialEq, Message)]
pub(crate) struct KnowledgeBaseAddRequest {
#[prost(string, tag = "1")]
knowledge: String,
#[prost(string, tag = "2")]
title: String,
#[prost(string, tag = "3")]
git_origin: String,
#[prost(string, optional, tag = "4")]
composer_id: Option<String>,
}
#[derive(Clone, PartialEq, Message)]
pub(crate) struct KnowledgeBaseAddResponse {
#[prost(bool, tag = "1")]
success: bool,
#[prost(string, tag = "2")]
id: String,
}
#[derive(Clone, PartialEq, Message)]
pub(crate) struct KnowledgeBaseListRequest {
#[prost(int32, optional, tag = "1")]
limit: Option<i32>,
#[prost(string, optional, tag = "2")]
git_origin: Option<String>,
}
#[derive(Clone, PartialEq, Message)]
pub(crate) struct KnowledgeBaseListResponse {
#[prost(bool, tag = "1")]
success: bool,
#[prost(message, repeated, tag = "2")]
all_results: Vec<KnowledgeBaseListItem>,
}
#[derive(Clone, PartialEq, Message)]
pub(crate) struct KnowledgeBaseListItem {
#[prost(string, tag = "1")]
id: String,
#[prost(string, tag = "2")]
knowledge: String,
#[prost(string, tag = "3")]
title: String,
#[prost(string, tag = "4")]
created_at: String,
#[prost(bool, tag = "5")]
is_generated: bool,
}
#[derive(Clone, PartialEq, Message)]
pub(crate) struct KnowledgeBaseUpdateRequest {
#[prost(string, tag = "1")]
id: String,
#[prost(string, tag = "2")]
knowledge: String,
#[prost(string, tag = "3")]
title: String,
}
#[derive(Clone, PartialEq, Message)]
pub(crate) struct KnowledgeBaseUpdateResponse {
#[prost(bool, tag = "1")]
success: bool,
}
#[derive(Clone, PartialEq, Message)]
pub(crate) struct KnowledgeBaseRemoveRequest {
#[prost(string, tag = "1")]
id: String,
}
#[derive(Clone, PartialEq, Message)]
pub(crate) struct KnowledgeBaseRemoveResponse {
#[prost(bool, tag = "1")]
success: bool,
}
/// 规则存储与并发锁;经 axum Extension 注入四个 handler。
#[derive(Clone)]
pub struct KnowledgeService {
inner: Arc<Inner>,
}
struct Inner {
store: RuleStore,
lock: tokio::sync::Mutex<()>,
}
impl KnowledgeService {
pub fn managed() -> Result<Self> {
Self::with_root(config::managed_data_dir()?.join("rules"))
}
/// 指定存储根目录构造;managed() 与集成测试共用。
pub fn with_root(root: std::path::PathBuf) -> Result<Self> {
Ok(Self {
inner: Arc::new(Inner {
store: RuleStore::open(root)?,
lock: tokio::sync::Mutex::new(()),
}),
})
}
}
pub async fn add(
Extension(upstream): Extension<proxy::CursorProxy>,
Extension(service): Extension<KnowledgeService>,
request: Request<Body>,
) -> Result<Response<Body>> {
let (parts, body) = buffered(request).await?;
let message: KnowledgeBaseAddRequest = connect::decode_unary(&body)?;
let _guard = service.inner.lock.lock().await;
let store = &service.inner.store;
if sync::replay(&upstream, &parts.headers, store).await? {
match proxy::forward_buffered(&upstream, Request::from_parts(parts, Body::from(body))).await
{
Ok(response) if response.status.is_success() => {
if let Ok(reply) = connect::decode_unary::<KnowledgeBaseAddResponse>(&response.body)
{
if reply.success && !reply.id.is_empty() {
store.upsert(&RuleRecord {
id: reply.id,
knowledge: message.knowledge,
title: message.title,
created_at: now(),
is_generated: false,
git_origin: message.git_origin,
})?;
}
}
return Ok(response.into_response());
}
Ok(response) => {
tracing::warn!(status = %response.status, "rules upstream rejected add; storing locally");
}
Err(error) => {
tracing::warn!(%error, "rules upstream unavailable for add; storing locally");
}
}
}
let id = format!("{}{}", store::LOCAL_ID_PREFIX, uuid::Uuid::new_v4());
store.upsert(&RuleRecord {
id: id.clone(),
knowledge: message.knowledge,
title: message.title,
created_at: now(),
is_generated: false,
git_origin: message.git_origin,
})?;
store.record_add(&id)?;
proto(KnowledgeBaseAddResponse { success: true, id })
}
pub async fn list(
Extension(upstream): Extension<proxy::CursorProxy>,
Extension(service): Extension<KnowledgeService>,
request: Request<Body>,
) -> Result<Response<Body>> {
let (parts, body) = buffered(request).await?;
let message: KnowledgeBaseListRequest = connect::decode_unary(&body)?;
let _guard = service.inner.lock.lock().await;
let store = &service.inner.store;
let git_origin = message.git_origin.unwrap_or_default();
if sync::replay(&upstream, &parts.headers, store).await? {
match proxy::forward_buffered(&upstream, Request::from_parts(parts, Body::from(body))).await
{
Ok(response) if response.status.is_success() => {
if let Ok(reply) =
connect::decode_unary::<KnowledgeBaseListResponse>(&response.body)
{
// 带 git_origin 过滤的列表只是子集,整体覆盖会误删其他规则。
if reply.success && git_origin.is_empty() {
sync::mirror(store, reply.all_results)?;
}
}
return Ok(response.into_response());
}
Ok(response) => {
tracing::warn!(status = %response.status, "rules upstream rejected list; serving local cache");
}
Err(error) => {
tracing::warn!(%error, "rules upstream unavailable for list; serving local cache");
}
}
}
let mut records = store.list()?;
if !git_origin.is_empty() {
records.retain(|record| record.git_origin == git_origin);
}
if let Some(limit) = message.limit {
if limit >= 0 {
records.truncate(limit as usize);
}
}
proto(KnowledgeBaseListResponse {
success: true,
all_results: records
.into_iter()
.map(|record| KnowledgeBaseListItem {
id: record.id,
knowledge: record.knowledge,
title: record.title,
created_at: record.created_at,
is_generated: record.is_generated,
})
.collect(),
})
}
pub async fn update(
Extension(upstream): Extension<proxy::CursorProxy>,
Extension(service): Extension<KnowledgeService>,
request: Request<Body>,
) -> Result<Response<Body>> {
let (parts, body) = buffered(request).await?;
let message: KnowledgeBaseUpdateRequest = connect::decode_unary(&body)?;
let _guard = service.inner.lock.lock().await;
let store = &service.inner.store;
if sync::replay(&upstream, &parts.headers, store).await? {
match proxy::forward_buffered(&upstream, Request::from_parts(parts, Body::from(body))).await
{
Ok(response) if response.status.is_success() => {
if let Ok(reply) =
connect::decode_unary::<KnowledgeBaseUpdateResponse>(&response.body)
{
if reply.success {
let existing = store.get(&message.id)?;
store.upsert(&RuleRecord {
id: message.id,
knowledge: message.knowledge,
title: message.title,
created_at: existing
.as_ref()
.map_or_else(now, |record| record.created_at.clone()),
is_generated: existing
.as_ref()
.is_some_and(|record| record.is_generated),
git_origin: existing
.map(|record| record.git_origin)
.unwrap_or_default(),
})?;
}
}
return Ok(response.into_response());
}
Ok(response) => {
tracing::warn!(status = %response.status, "rules upstream rejected update; storing locally");
}
Err(error) => {
tracing::warn!(%error, "rules upstream unavailable for update; storing locally");
}
}
}
let Some(mut record) = store.get(&message.id)? else {
return proto(KnowledgeBaseUpdateResponse { success: false });
};
record.knowledge = message.knowledge;
record.title = message.title;
store.upsert(&record)?;
store.record_update(&message.id)?;
proto(KnowledgeBaseUpdateResponse { success: true })
}
pub async fn remove(
Extension(upstream): Extension<proxy::CursorProxy>,
Extension(service): Extension<KnowledgeService>,
request: Request<Body>,
) -> Result<Response<Body>> {
let (parts, body) = buffered(request).await?;
let message: KnowledgeBaseRemoveRequest = connect::decode_unary(&body)?;
let _guard = service.inner.lock.lock().await;
let store = &service.inner.store;
if sync::replay(&upstream, &parts.headers, store).await? {
match proxy::forward_buffered(&upstream, Request::from_parts(parts, Body::from(body))).await
{
Ok(response) if response.status.is_success() => {
if let Ok(reply) =
connect::decode_unary::<KnowledgeBaseRemoveResponse>(&response.body)
{
if reply.success {
store.remove(&message.id)?;
}
}
return Ok(response.into_response());
}
Ok(response) => {
tracing::warn!(status = %response.status, "rules upstream rejected remove; removing locally");
}
Err(error) => {
tracing::warn!(%error, "rules upstream unavailable for remove; removing locally");
}
}
}
store.remove(&message.id)?;
store.record_remove(&message.id)?;
proto(KnowledgeBaseRemoveResponse { success: true })
}
async fn buffered(request: Request<Body>) -> Result<(axum::http::request::Parts, Bytes)> {
let (parts, body) = request.into_parts();
let body = to_bytes(body, usize::MAX)
.await
.map_err(|error| crate::Error::Protocol(format!("cannot read request body: {error}")))?;
Ok((parts, body))
}
fn now() -> String {
chrono::Utc::now().to_rfc3339_opts(chrono::SecondsFormat::Millis, true)
}
fn proto(message: impl Message) -> Result<Response<Body>> {
let body = message.encode_to_vec();
let length = body.len();
let mut response = Response::new(Body::from(body));
response.headers_mut().insert(
header::CONTENT_TYPE,
axum::http::HeaderValue::from_static("application/proto"),
);
response.headers_mut().insert(
header::CONTENT_LENGTH,
length
.to_string()
.parse()
.expect("body length is always a valid header value"),
);
Ok(response)
}
@@ -0,0 +1,487 @@
//! Persists rules as markdown files with a JSON metadata sidecar.
use std::{
collections::BTreeMap,
path::{Path, PathBuf},
};
use serde::{Deserialize, Serialize};
use crate::{Error, Result};
const META_FILE: &str = "meta.json";
const RULE_EXTENSION: &str = "md";
pub const LOCAL_ID_PREFIX: &str = "local-";
/// 一条规则的完整视图:knowledge 来自 md 文件,其余字段来自 meta.json。
#[derive(Clone, Debug, PartialEq)]
pub struct RuleRecord {
pub id: String,
pub knowledge: String,
pub title: String,
pub created_at: String,
pub is_generated: bool,
pub git_origin: String,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum JournalOp {
Add,
Update,
Remove,
}
/// 离线期间未同步到上游的一次变更。
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct JournalEntry {
pub op: JournalOp,
pub id: String,
}
#[derive(Default, Serialize, Deserialize)]
struct Meta {
#[serde(default)]
rules: BTreeMap<String, RuleMeta>,
#[serde(default)]
journal: Vec<JournalEntry>,
}
#[derive(Clone, Default, Serialize, Deserialize)]
struct RuleMeta {
#[serde(default)]
title: String,
#[serde(default)]
created_at: String,
#[serde(default)]
is_generated: bool,
#[serde(default)]
git_origin: String,
}
/// md 文件为核心的规则存储;调用方需自行串行化并发访问。
pub struct RuleStore {
root: PathBuf,
}
impl RuleStore {
pub fn open(root: PathBuf) -> Result<Self> {
std::fs::create_dir_all(&root)?;
Ok(Self { root })
}
pub fn list(&self) -> Result<Vec<RuleRecord>> {
let meta = self.read_meta();
let mut records = Vec::new();
for entry in std::fs::read_dir(&self.root)? {
let path = entry?.path();
if path.extension().and_then(|value| value.to_str()) != Some(RULE_EXTENSION) {
continue;
}
let Some(id) = path.file_stem().and_then(|value| value.to_str()) else {
continue;
};
if validate_id(id).is_err() {
continue;
}
let knowledge = std::fs::read_to_string(&path)?;
records.push(assemble(id, knowledge, meta.rules.get(id), &path));
}
records.sort_by(|left, right| {
timestamp(&right.created_at)
.cmp(&timestamp(&left.created_at))
.then_with(|| left.id.cmp(&right.id))
});
Ok(records)
}
pub fn get(&self, id: &str) -> Result<Option<RuleRecord>> {
validate_id(id)?;
let path = self.rule_path(id);
let knowledge = match std::fs::read_to_string(&path) {
Ok(knowledge) => knowledge,
Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(None),
Err(error) => return Err(error.into()),
};
let meta = self.read_meta();
Ok(Some(assemble(id, knowledge, meta.rules.get(id), &path)))
}
pub fn upsert(&self, record: &RuleRecord) -> Result<()> {
validate_id(&record.id)?;
write_atomic(&self.rule_path(&record.id), record.knowledge.as_bytes())?;
let mut meta = self.read_meta();
meta.rules.insert(record.id.clone(), rule_meta(record));
self.write_meta(&meta)
}
pub fn remove(&self, id: &str) -> Result<()> {
validate_id(id)?;
remove_file_if_exists(&self.rule_path(id))?;
let mut meta = self.read_meta();
if meta.rules.remove(id).is_some() {
self.write_meta(&meta)?;
}
Ok(())
}
/// 离线新增的规则在上游落地后,把本地临时 id 换成上游分配的真实 id。
pub fn promote(&self, old_id: &str, new_id: &str) -> Result<()> {
validate_id(old_id)?;
validate_id(new_id)?;
let source = self.rule_path(old_id);
let target = self.rule_path(new_id);
#[cfg(windows)]
remove_file_if_exists(&target)?;
std::fs::rename(&source, &target)?;
let mut meta = self.read_meta();
if let Some(rule) = meta.rules.remove(old_id) {
meta.rules.insert(new_id.into(), rule);
}
for entry in &mut meta.journal {
if entry.id == old_id {
entry.id = new_id.into();
}
}
self.write_meta(&meta)
}
/// 用上游的完整列表覆盖本地镜像;仅应在日志为空(已全部回放)时调用。
pub fn replace_all(&self, records: &[RuleRecord]) -> Result<()> {
let mut meta = self.read_meta();
meta.rules.clear();
for record in records {
validate_id(&record.id)?;
write_atomic(&self.rule_path(&record.id), record.knowledge.as_bytes())?;
meta.rules.insert(record.id.clone(), rule_meta(record));
}
for entry in std::fs::read_dir(&self.root)? {
let path = entry?.path();
if path.extension().and_then(|value| value.to_str()) != Some(RULE_EXTENSION) {
continue;
}
let keep = path
.file_stem()
.and_then(|value| value.to_str())
.is_some_and(|id| meta.rules.contains_key(id));
if !keep {
remove_file_if_exists(&path)?;
}
}
self.write_meta(&meta)
}
pub fn journal_front(&self) -> Result<Option<JournalEntry>> {
Ok(self.read_meta().journal.first().cloned())
}
pub fn pop_journal(&self) -> Result<()> {
let mut meta = self.read_meta();
if !meta.journal.is_empty() {
meta.journal.remove(0);
self.write_meta(&meta)?;
}
Ok(())
}
pub fn record_add(&self, id: &str) -> Result<()> {
let mut meta = self.read_meta();
meta.journal.push(JournalEntry {
op: JournalOp::Add,
id: id.into(),
});
self.write_meta(&meta)
}
pub fn record_update(&self, id: &str) -> Result<()> {
let mut meta = self.read_meta();
if journal_contains(&meta.journal, id, JournalOp::Add) {
// 回放 add 时会读取最新内容,无需单独的 update 日志。
return Ok(());
}
let op = if id.starts_with(LOCAL_ID_PREFIX) {
// 本地临时 id 没有对应的 add 日志(如镜像覆盖后的残留),按新增回放。
JournalOp::Add
} else {
JournalOp::Update
};
if !journal_contains(&meta.journal, id, op) {
meta.journal.push(JournalEntry { op, id: id.into() });
self.write_meta(&meta)?;
}
Ok(())
}
pub fn record_remove(&self, id: &str) -> Result<()> {
let mut meta = self.read_meta();
let never_synced = journal_contains(&meta.journal, id, JournalOp::Add);
meta.journal.retain(|entry| entry.id != id);
if !never_synced && !id.starts_with(LOCAL_ID_PREFIX) {
meta.journal.push(JournalEntry {
op: JournalOp::Remove,
id: id.into(),
});
}
self.write_meta(&meta)
}
fn rule_path(&self, id: &str) -> PathBuf {
self.root.join(format!("{id}.{RULE_EXTENSION}"))
}
fn meta_path(&self) -> PathBuf {
self.root.join(META_FILE)
}
fn read_meta(&self) -> Meta {
match std::fs::read(self.meta_path()) {
Ok(bytes) => serde_json::from_slice(&bytes).unwrap_or_else(|error| {
tracing::warn!(%error, "rules meta.json is corrupt; starting from empty metadata");
Meta::default()
}),
Err(_) => Meta::default(),
}
}
fn write_meta(&self, meta: &Meta) -> Result<()> {
write_atomic(&self.meta_path(), &serde_json::to_vec_pretty(meta)?)
}
}
fn assemble(id: &str, knowledge: String, meta: Option<&RuleMeta>, path: &Path) -> RuleRecord {
match meta {
Some(meta) => RuleRecord {
id: id.into(),
knowledge,
title: meta.title.clone(),
created_at: meta.created_at.clone(),
is_generated: meta.is_generated,
git_origin: meta.git_origin.clone(),
},
// 用户手放的 md 文件没有元数据,用文件名当标题、修改时间当创建时间。
None => RuleRecord {
id: id.into(),
knowledge,
title: id.into(),
created_at: file_modified_at(path),
is_generated: false,
git_origin: String::new(),
},
}
}
fn rule_meta(record: &RuleRecord) -> RuleMeta {
RuleMeta {
title: record.title.clone(),
created_at: record.created_at.clone(),
is_generated: record.is_generated,
git_origin: record.git_origin.clone(),
}
}
fn journal_contains(journal: &[JournalEntry], id: &str, op: JournalOp) -> bool {
journal.iter().any(|entry| entry.id == id && entry.op == op)
}
fn timestamp(created_at: &str) -> i64 {
chrono::DateTime::parse_from_rfc3339(created_at)
.map(|time| time.timestamp_millis())
.unwrap_or(0)
}
fn file_modified_at(path: &Path) -> String {
let modified = std::fs::metadata(path)
.and_then(|meta| meta.modified())
.unwrap_or_else(|_| std::time::SystemTime::now());
chrono::DateTime::<chrono::Utc>::from(modified)
.to_rfc3339_opts(chrono::SecondsFormat::Millis, true)
}
fn validate_id(id: &str) -> Result<()> {
if id.is_empty()
|| id.len() > 128
|| !id
.bytes()
.all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_'))
{
return Err(Error::Protocol(format!("invalid rule id: {id:?}")));
}
Ok(())
}
fn remove_file_if_exists(path: &Path) -> Result<()> {
match std::fs::remove_file(path) {
Ok(()) => Ok(()),
Err(error) if error.kind() == std::io::ErrorKind::NotFound => Ok(()),
Err(error) => Err(error.into()),
}
}
fn write_atomic(path: &Path, bytes: &[u8]) -> Result<()> {
use std::io::Write;
let directory = path.parent().expect("rule path has a parent");
let temporary = directory.join(format!(".{}.tmp", uuid::Uuid::new_v4()));
let mut file = std::fs::File::create(&temporary)?;
file.write_all(bytes)?;
file.sync_all()?;
drop(file);
#[cfg(windows)]
remove_file_if_exists(path)?;
std::fs::rename(&temporary, path).inspect_err(|_| {
let _ = std::fs::remove_file(&temporary);
})?;
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
fn record(id: &str, knowledge: &str, created_at: &str) -> RuleRecord {
RuleRecord {
id: id.into(),
knowledge: knowledge.into(),
title: format!("title-{id}"),
created_at: created_at.into(),
is_generated: false,
git_origin: String::new(),
}
}
fn journal(store: &RuleStore) -> Vec<JournalEntry> {
store.read_meta().journal
}
#[test]
fn upserts_lists_and_removes_rules() {
let root = tempfile::tempdir().unwrap();
let store = RuleStore::open(root.path().join("rules")).unwrap();
store
.upsert(&record("100", "older", "2026-01-01T00:00:00.000Z"))
.unwrap();
store
.upsert(&record("200", "newer", "2026-02-01T00:00:00.000Z"))
.unwrap();
let listed = store.list().unwrap();
assert_eq!(
listed
.iter()
.map(|rule| rule.id.as_str())
.collect::<Vec<_>>(),
["200", "100"],
"list is sorted by created_at descending"
);
assert_eq!(listed[0].knowledge, "newer");
assert_eq!(listed[0].title, "title-200");
store.remove("200").unwrap();
assert!(store.get("200").unwrap().is_none());
assert_eq!(store.list().unwrap().len(), 1);
}
#[test]
fn rejects_path_traversal_ids() {
let root = tempfile::tempdir().unwrap();
let store = RuleStore::open(root.path().join("rules")).unwrap();
assert!(store.get("../escape").is_err());
assert!(store.get("a/b").is_err());
assert!(store.get("").is_err());
}
#[test]
fn compacts_offline_journal() {
let root = tempfile::tempdir().unwrap();
let store = RuleStore::open(root.path().join("rules")).unwrap();
// 离线新增后再更新:回放 add 即可携带最新内容,不产生 update 日志。
store
.upsert(&record("local-a", "v1", "2026-01-01T00:00:00.000Z"))
.unwrap();
store.record_add("local-a").unwrap();
store.record_update("local-a").unwrap();
assert_eq!(
journal(&store),
vec![JournalEntry {
op: JournalOp::Add,
id: "local-a".into()
}]
);
// 离线新增后又删除:上游从未见过它,日志清空。
store.record_remove("local-a").unwrap();
assert!(journal(&store).is_empty());
// 更新上游已有规则:多次更新合并为一条;删除后 update 日志被顶替。
store.record_update("42").unwrap();
store.record_update("42").unwrap();
assert_eq!(
journal(&store),
vec![JournalEntry {
op: JournalOp::Update,
id: "42".into()
}]
);
store.record_remove("42").unwrap();
assert_eq!(
journal(&store),
vec![JournalEntry {
op: JournalOp::Remove,
id: "42".into()
}]
);
}
#[test]
fn promote_renames_rule_and_journal_ids() {
let root = tempfile::tempdir().unwrap();
let store = RuleStore::open(root.path().join("rules")).unwrap();
store
.upsert(&record("local-a", "content", "2026-01-01T00:00:00.000Z"))
.unwrap();
store.record_add("local-a").unwrap();
store.promote("local-a", "17353272").unwrap();
assert!(store.get("local-a").unwrap().is_none());
let promoted = store.get("17353272").unwrap().unwrap();
assert_eq!(promoted.knowledge, "content");
assert_eq!(promoted.title, "title-local-a");
assert_eq!(journal(&store)[0].id, "17353272");
}
#[test]
fn replace_all_mirrors_upstream_state() {
let root = tempfile::tempdir().unwrap();
let store = RuleStore::open(root.path().join("rules")).unwrap();
store
.upsert(&record("stale", "gone soon", "2026-01-01T00:00:00.000Z"))
.unwrap();
store
.replace_all(&[record(
"17353272",
"from upstream",
"2026-02-01T00:00:00.000Z",
)])
.unwrap();
let listed = store.list().unwrap();
assert_eq!(listed.len(), 1);
assert_eq!(listed[0].id, "17353272");
assert_eq!(listed[0].knowledge, "from upstream");
assert!(store.get("stale").unwrap().is_none());
}
#[test]
fn lists_hand_written_markdown_without_metadata() {
let root = tempfile::tempdir().unwrap();
let store = RuleStore::open(root.path().join("rules")).unwrap();
std::fs::write(root.path().join("rules/manual_rule.md"), "hand written").unwrap();
let listed = store.list().unwrap();
assert_eq!(listed.len(), 1);
assert_eq!(listed[0].id, "manual_rule");
assert_eq!(listed[0].title, "manual_rule");
assert_eq!(listed[0].knowledge, "hand written");
}
}
@@ -0,0 +1,184 @@
//! Replays the offline journal to upstream and mirrors upstream list state.
use axum::{
body::{Body, Bytes},
http::{header, HeaderMap, HeaderValue, Method, Request},
};
use prost::Message;
use crate::{api::cursor::proxy, cursor::protocol::connect, Result};
use super::{
store::{JournalOp, RuleRecord, RuleStore},
KnowledgeBaseAddRequest, KnowledgeBaseAddResponse, KnowledgeBaseListItem,
KnowledgeBaseRemoveRequest, KnowledgeBaseRemoveResponse, KnowledgeBaseUpdateRequest,
KnowledgeBaseUpdateResponse,
};
const ADD_PATH: &str = "/aiserver.v1.AiService/KnowledgeBaseAdd";
const UPDATE_PATH: &str = "/aiserver.v1.AiService/KnowledgeBaseUpdate";
const REMOVE_PATH: &str = "/aiserver.v1.AiService/KnowledgeBaseRemove";
/// 逐条把离线日志推送到上游。返回 true 表示日志已清空(上游可用),
/// false 表示上游不可达,剩余日志保留、调用方应降级到本地。
pub async fn replay(
upstream: &proxy::CursorProxy,
headers: &HeaderMap,
store: &RuleStore,
) -> Result<bool> {
while let Some(entry) = store.journal_front()? {
let advanced = match entry.op {
JournalOp::Add => replay_add(upstream, headers, store, &entry.id).await?,
JournalOp::Update => replay_update(upstream, headers, store, &entry.id).await?,
JournalOp::Remove => replay_remove(upstream, headers, store, &entry.id).await?,
};
if !advanced {
return Ok(false);
}
}
Ok(true)
}
/// 用上游返回的完整列表覆盖本地镜像。仅应在日志已清空时调用。
pub fn mirror(store: &RuleStore, items: Vec<KnowledgeBaseListItem>) -> Result<()> {
let records = items
.into_iter()
.map(|item| RuleRecord {
id: item.id,
knowledge: item.knowledge,
title: item.title,
created_at: item.created_at,
is_generated: item.is_generated,
git_origin: String::new(),
})
.collect::<Vec<_>>();
store.replace_all(&records)
}
async fn replay_add(
upstream: &proxy::CursorProxy,
headers: &HeaderMap,
store: &RuleStore,
id: &str,
) -> Result<bool> {
let Some(record) = store.get(id)? else {
// 规则文件已不在(被手动删除等),日志作废。
store.pop_journal()?;
return Ok(true);
};
let message = KnowledgeBaseAddRequest {
knowledge: record.knowledge,
title: record.title,
git_origin: record.git_origin,
composer_id: None,
};
let Some(body) = send(upstream, headers, ADD_PATH, &message).await else {
return Ok(false);
};
let Ok(reply) = connect::decode_unary::<KnowledgeBaseAddResponse>(&body) else {
return Ok(false);
};
if !reply.success || reply.id.is_empty() {
tracing::warn!(
id,
"rules upstream declined replayed add; dropping journal entry"
);
store.pop_journal()?;
return Ok(true);
}
store.promote(id, &reply.id)?;
store.pop_journal()?;
tracing::info!(
local_id = id,
upstream_id = reply.id,
"replayed offline rule add to upstream"
);
Ok(true)
}
async fn replay_update(
upstream: &proxy::CursorProxy,
headers: &HeaderMap,
store: &RuleStore,
id: &str,
) -> Result<bool> {
let Some(record) = store.get(id)? else {
store.pop_journal()?;
return Ok(true);
};
let message = KnowledgeBaseUpdateRequest {
id: id.into(),
knowledge: record.knowledge,
title: record.title,
};
let Some(body) = send(upstream, headers, UPDATE_PATH, &message).await else {
return Ok(false);
};
let Ok(reply) = connect::decode_unary::<KnowledgeBaseUpdateResponse>(&body) else {
return Ok(false);
};
if !reply.success {
tracing::warn!(
id,
"rules upstream declined replayed update; dropping journal entry"
);
}
store.pop_journal()?;
Ok(true)
}
async fn replay_remove(
upstream: &proxy::CursorProxy,
headers: &HeaderMap,
store: &RuleStore,
id: &str,
) -> Result<bool> {
let message = KnowledgeBaseRemoveRequest { id: id.into() };
let Some(body) = send(upstream, headers, REMOVE_PATH, &message).await else {
return Ok(false);
};
let Ok(reply) = connect::decode_unary::<KnowledgeBaseRemoveResponse>(&body) else {
return Ok(false);
};
if !reply.success {
tracing::warn!(
id,
"rules upstream declined replayed remove; dropping journal entry"
);
}
store.pop_journal()?;
Ok(true)
}
/// 以当前请求的头为模板向上游发起一次 unary RPC。
/// 成功(2xx)返回响应体;不可达或被拒绝返回 None,由调用方保留日志。
async fn send(
upstream: &proxy::CursorProxy,
template: &HeaderMap,
path: &str,
message: &impl Message,
) -> Option<Bytes> {
let mut headers = template.clone();
// 模板里的上游 URL 头指向原始 RPC 路径,必须移除才能命中回放路径。
headers.remove(proxy::UPSTREAM_URL_HEADER);
headers.remove(header::CONTENT_LENGTH);
headers.insert(
header::CONTENT_TYPE,
HeaderValue::from_static("application/proto"),
);
let mut request = Request::new(Body::from(message.encode_to_vec()));
*request.method_mut() = Method::POST;
*request.uri_mut() = path.parse().expect("replay path is a valid URI");
*request.headers_mut() = headers;
match proxy::forward_buffered(upstream, request).await {
Ok(response) if response.status.is_success() => Some(response.body),
Ok(response) => {
tracing::warn!(path, status = %response.status, "rules journal replay rejected by upstream");
None
}
Err(error) => {
tracing::warn!(path, %error, "rules journal replay cannot reach upstream");
None
}
}
}
+1
View File
@@ -4,6 +4,7 @@ pub mod account;
pub mod analytics;
pub mod blob_sync;
pub mod context_sync;
pub mod knowledge;
pub mod model_catalog;
pub mod observability;
pub mod tab;
+21 -17
View File
@@ -10,7 +10,7 @@ use prost::Message;
use crate::{
api::cursor::proxy::{self, CursorProxy},
cursor::{protocol::proto::agent::v1 as agent, transport::TransportRegistry},
model::{format_token_count, parse_token_count, ModelConfig, ModelType},
model::{format_token_count, parse_token_count, ModelConfig},
plugin::PluginModelDescriptor,
Error, Result,
};
@@ -378,16 +378,25 @@ fn available_model(model: &ModelConfig) -> AvailableModel {
display_name: "Cursor".into(),
}),
model_picker_badges: vec![ModelPickerBadge {
label: match model.model_type {
ModelType::OpenAi => "OpenAI".into(),
ModelType::Anthropic => "Anthropic".into(),
},
label: model
.group_name
.clone()
.unwrap_or_else(|| provider_host(&model.base_url)),
variant: 1,
dismiss_on_selection: false,
}],
}
}
/// 徽章回退标签:base_url 的主机名。入库时已校验为带主机的 HTTP(S) URL,
/// 解析失败仅是理论分支,此时原样返回 base_url。
fn provider_host(base_url: &str) -> String {
reqwest::Url::parse(base_url.trim())
.ok()
.and_then(|url| url.host_str().map(str::to_lowercase))
.unwrap_or_else(|| base_url.trim().into())
}
fn model_parameters(
contexts: &[(String, String)],
thinking: bool,
@@ -571,14 +580,9 @@ fn available_plugin_model(model: &PluginModelDescriptor) -> AvailableModel {
let tooltip = TooltipData {
markdown_content: model.description.clone(),
};
let contexts = context_options(model.context_window_tokens);
let variants = model_variants(
&model.id,
&model.display_name,
&tooltip,
&contexts,
model.thinking,
);
// Effort 与上下文档位由宿主统一提供,与内置模型一致;插件不再声明这两项。
let contexts = context_options(None);
let variants = model_variants(&model.id, &model.display_name, &tooltip, &contexts, true);
let legacy_slugs = variants
.iter()
.filter_map(|variant| variant.legacy_slug.clone())
@@ -589,7 +593,7 @@ fn available_plugin_model(model: &PluginModelDescriptor) -> AvailableModel {
supports_agent: Some(true),
degradation_status: Some(0),
tooltip_data: Some(tooltip.clone()),
supports_thinking: Some(model.thinking),
supports_thinking: Some(true),
supports_images: Some(model.images),
supports_max_mode: Some(false),
client_display_name: Some(model.display_name.clone()),
@@ -601,7 +605,7 @@ fn available_plugin_model(model: &PluginModelDescriptor) -> AvailableModel {
inputbox_short_model_name: Some(model.display_name.clone()),
supports_sandboxing: Some(true),
supports_cmd_k: Some(false),
parameter_definitions: model_parameters(&contexts, model.thinking),
parameter_definitions: model_parameters(&contexts, true),
variants,
legacy_slugs,
named_model_section_index: Some(1),
@@ -611,7 +615,7 @@ fn available_plugin_model(model: &PluginModelDescriptor) -> AvailableModel {
display_name: model.provider_type.clone(),
}),
model_picker_badges: vec![ModelPickerBadge {
label: model.provider_type.clone(),
label: model.plugin_name.clone(),
variant: 1,
dismiss_on_selection: false,
}],
@@ -624,7 +628,7 @@ fn usable_plugin_model(model: &PluginModelDescriptor) -> agent::ModelDetails {
display_model_id: model.id.clone(),
display_name: model.display_name.clone(),
display_name_short: model.display_name.clone(),
thinking_details: model.thinking.then(agent::ThinkingDetails::default),
thinking_details: Some(agent::ThinkingDetails::default()),
..Default::default()
}
}
-225
View File
@@ -1,225 +0,0 @@
//! Records Cursor request traces and artifacts.
use std::{
sync::{
atomic::{AtomicBool, Ordering},
Arc,
},
time::{Duration, Instant},
};
use tokio::sync::Mutex;
use crate::store::{BlobId, BufferedCursorTraceChunk, Store};
#[derive(Clone)]
pub struct CursorTraceRecorder {
store: Store,
request_id: String,
chunks: Arc<Mutex<TraceChunkBuffer>>,
finished: Arc<AtomicBool>,
}
#[derive(Default)]
struct TraceChunkBuffer {
chunks: Vec<BufferedCursorTraceChunk>,
bytes: usize,
first_chunk_at: Option<Instant>,
generation: u64,
}
const MAX_BUFFERED_CHUNKS: usize = 32;
const MAX_BUFFERED_BYTES: usize = 256 * 1024;
const MAX_BUFFER_AGE: Duration = Duration::from_millis(50);
impl CursorTraceRecorder {
pub async fn begin(
store: Store,
request_id: &str,
conversation_id: Option<&str>,
route: &str,
model_id: Option<&str>,
) -> Option<Self> {
match store
.start_cursor_trace_if_detailed(request_id, conversation_id, route, model_id)
.await
{
Ok(true) => Some(Self {
store,
request_id: request_id.into(),
chunks: Arc::new(Mutex::new(TraceChunkBuffer::default())),
finished: Arc::new(AtomicBool::new(false)),
}),
Ok(false) => None,
Err(error) => {
tracing::warn!(request_id, %error, "failed to start Cursor trace");
None
}
}
}
pub async fn resume(store: Store, request_id: &str) -> Option<Self> {
match store.cursor_trace_exists(request_id).await {
Ok(true) => Some(Self {
store,
request_id: request_id.into(),
chunks: Arc::new(Mutex::new(TraceChunkBuffer::default())),
finished: Arc::new(AtomicBool::new(false)),
}),
Ok(false) => None,
Err(error) => {
tracing::warn!(request_id, %error, "failed to resume Cursor trace");
None
}
}
}
pub fn request_id(&self) -> &str {
&self.request_id
}
pub async fn request(&self, artifact_type: &str, data: &[u8], metadata: serde_json::Value) {
if let Err(error) = self
.store
.append_cursor_trace_artifact(
&self.request_id,
artifact_type,
"cursor_client",
data,
&metadata,
)
.await
{
tracing::warn!(request_id = self.request_id, %error, "failed to record Cursor request artifact");
return;
}
if let Err(error) = self
.store
.add_cursor_trace_request_bytes(&self.request_id, data.len())
.await
{
tracing::warn!(request_id = self.request_id, %error, "failed to update Cursor request trace size");
}
}
pub async fn artifact(
&self,
artifact_type: &str,
source: &str,
data: &[u8],
metadata: serde_json::Value,
) {
if let Err(error) = self
.store
.append_cursor_trace_artifact(&self.request_id, artifact_type, source, data, &metadata)
.await
{
tracing::warn!(request_id = self.request_id, %error, artifact_type, "failed to record Cursor trace artifact");
}
}
pub async fn linked_blob(
&self,
artifact_type: &str,
source: &str,
blob_id: &BlobId,
metadata: serde_json::Value,
) {
if let Err(error) = self
.store
.link_cursor_trace_artifact(&self.request_id, artifact_type, source, blob_id, &metadata)
.await
{
tracing::warn!(request_id = self.request_id, %error, artifact_type, "failed to link Cursor trace Blob");
}
}
pub async fn response_started(&self, status: u16) {
if let Err(error) = self
.store
.start_cursor_trace_response(&self.request_id, status)
.await
{
tracing::warn!(request_id = self.request_id, %error, "failed to start Cursor response trace");
}
}
pub async fn response_chunk(&self, source: &str, data: &[u8]) {
let mut buffer = self.chunks.lock().await;
if self.finished.load(Ordering::Acquire) {
return;
}
let schedule_flush = if buffer.chunks.is_empty() {
buffer.generation = buffer.generation.wrapping_add(1);
buffer.first_chunk_at = Some(Instant::now());
Some(buffer.generation)
} else {
None
};
buffer.bytes += data.len();
buffer
.chunks
.push(BufferedCursorTraceChunk::new(source, data));
let expired = buffer
.first_chunk_at
.is_some_and(|started| started.elapsed() >= MAX_BUFFER_AGE);
if buffer.chunks.len() >= MAX_BUFFERED_CHUNKS
|| buffer.bytes >= MAX_BUFFERED_BYTES
|| expired
{
if let Err(error) = self.flush_locked(&mut buffer).await {
tracing::warn!(request_id = self.request_id, %error, "failed to record Cursor response chunk");
}
}
drop(buffer);
if let Some(generation) = schedule_flush {
let recorder = self.clone();
tokio::spawn(async move {
tokio::time::sleep(MAX_BUFFER_AGE).await;
let mut buffer = recorder.chunks.lock().await;
if buffer.generation == generation {
if let Err(error) = recorder.flush_locked(&mut buffer).await {
tracing::warn!(request_id = recorder.request_id, %error, "failed to flush Cursor response chunks");
}
}
});
}
}
pub async fn finish(&self, error: Option<&str>) {
if self.finished.swap(true, Ordering::AcqRel) {
return;
}
let mut buffer = self.chunks.lock().await;
if let Err(store_error) = self.flush_locked(&mut buffer).await {
tracing::warn!(request_id = self.request_id, %store_error, "failed to flush Cursor response chunks");
}
drop(buffer);
if let Err(store_error) = self
.store
.finish_cursor_trace(&self.request_id, error)
.await
{
tracing::warn!(request_id = self.request_id, %store_error, "failed to finish Cursor trace");
}
}
async fn flush_locked(&self, buffer: &mut TraceChunkBuffer) -> crate::Result<()> {
if buffer.chunks.is_empty() {
return Ok(());
}
let chunks = std::mem::take(&mut buffer.chunks);
buffer.bytes = 0;
buffer.first_chunk_at = None;
if let Err(error) = self
.store
.add_cursor_trace_response_chunks(&self.request_id, &chunks)
.await
{
buffer.bytes = chunks.iter().map(|chunk| chunk.data.len()).sum();
buffer.first_chunk_at = Some(Instant::now());
buffer.chunks = chunks;
return Err(error);
}
Ok(())
}
}
@@ -0,0 +1,71 @@
use std::sync::{atomic::AtomicU8, Arc};
use bytes::Bytes;
use crate::store::BlobId;
pub(super) const TRACE_UNKNOWN: u8 = 0;
pub(super) const TRACE_ACTIVE: u8 = 1;
pub(super) const TRACE_DISABLED: u8 = 2;
pub(super) enum TraceEvent {
Begin {
request_id: String,
activation: Arc<AtomicU8>,
conversation_id: Option<String>,
route: String,
model_id: Option<String>,
},
Resume {
request_id: String,
activation: Arc<AtomicU8>,
},
Request {
request_id: String,
artifact_type: String,
data: Bytes,
metadata: serde_json::Value,
},
Artifact {
request_id: String,
artifact_type: String,
source: String,
data: Bytes,
metadata: serde_json::Value,
},
LinkedBlob {
request_id: String,
artifact_type: String,
source: String,
blob_id: BlobId,
metadata: serde_json::Value,
},
ResponseStarted {
request_id: String,
status: u16,
},
ResponseChunk {
request_id: String,
source: String,
data: Bytes,
},
Finish {
request_id: String,
error: Option<String>,
},
}
impl TraceEvent {
pub(super) fn request_id(&self) -> &str {
match self {
Self::Begin { request_id, .. }
| Self::Resume { request_id, .. }
| Self::Request { request_id, .. }
| Self::Artifact { request_id, .. }
| Self::LinkedBlob { request_id, .. }
| Self::ResponseStarted { request_id, .. }
| Self::ResponseChunk { request_id, .. }
| Self::Finish { request_id, .. } => request_id,
}
}
}
@@ -0,0 +1,157 @@
//! Records Cursor request traces without blocking request or runtime paths.
mod event;
mod worker;
use std::sync::{
atomic::{AtomicBool, AtomicU8, Ordering},
Arc,
};
use bytes::Bytes;
use tokio::sync::mpsc;
use crate::store::{BlobId, Store};
use event::{TraceEvent, TRACE_DISABLED, TRACE_UNKNOWN};
const TRACE_QUEUE_CAPACITY: usize = 512;
#[derive(Clone)]
pub struct CursorTraceService {
sender: mpsc::Sender<TraceEvent>,
}
impl CursorTraceService {
pub fn new(store: Store) -> Self {
let (sender, receiver) = mpsc::channel(TRACE_QUEUE_CAPACITY);
tokio::spawn(worker::run(store, receiver));
Self { sender }
}
pub fn recorder(&self, request_id: &str) -> CursorTraceRecorder {
CursorTraceRecorder {
request_id: Arc::from(request_id),
sender: self.sender.clone(),
finished: Arc::new(AtomicBool::new(false)),
activation: Arc::new(AtomicU8::new(TRACE_UNKNOWN)),
}
}
}
#[derive(Clone)]
pub struct CursorTraceRecorder {
request_id: Arc<str>,
sender: mpsc::Sender<TraceEvent>,
finished: Arc<AtomicBool>,
activation: Arc<AtomicU8>,
}
impl CursorTraceRecorder {
pub fn request_id(&self) -> &str {
&self.request_id
}
pub fn begin(&self, conversation_id: Option<&str>, route: &str, model_id: Option<&str>) {
self.send_control(TraceEvent::Begin {
request_id: self.request_id.to_string(),
activation: self.activation.clone(),
conversation_id: conversation_id.map(str::to_owned),
route: route.to_owned(),
model_id: model_id.map(str::to_owned),
});
}
pub fn resume(&self) {
self.send_control(TraceEvent::Resume {
request_id: self.request_id.to_string(),
activation: self.activation.clone(),
});
}
pub fn request(&self, artifact_type: &str, data: Bytes, metadata: serde_json::Value) {
self.send(TraceEvent::Request {
request_id: self.request_id.to_string(),
artifact_type: artifact_type.to_owned(),
data,
metadata,
});
}
pub fn artifact(
&self,
artifact_type: &str,
source: &str,
data: &[u8],
metadata: serde_json::Value,
) {
self.send(TraceEvent::Artifact {
request_id: self.request_id.to_string(),
artifact_type: artifact_type.to_owned(),
source: source.to_owned(),
data: Bytes::copy_from_slice(data),
metadata,
});
}
pub fn linked_blob(
&self,
artifact_type: &str,
source: &str,
blob_id: &BlobId,
metadata: serde_json::Value,
) {
self.send(TraceEvent::LinkedBlob {
request_id: self.request_id.to_string(),
artifact_type: artifact_type.to_owned(),
source: source.to_owned(),
blob_id: blob_id.clone(),
metadata,
});
}
pub fn response_started(&self, status: u16) {
self.send(TraceEvent::ResponseStarted {
request_id: self.request_id.to_string(),
status,
});
}
pub fn response_chunk(&self, source: &str, data: Bytes) {
if self.finished.load(Ordering::Acquire) {
return;
}
self.send(TraceEvent::ResponseChunk {
request_id: self.request_id.to_string(),
source: source.to_owned(),
data,
});
}
pub fn finish(&self, error: Option<&str>) {
if self.finished.swap(true, Ordering::AcqRel) {
return;
}
self.send_control(TraceEvent::Finish {
request_id: self.request_id.to_string(),
error: error.map(str::to_owned),
});
}
fn send(&self, event: TraceEvent) {
if self.activation.load(Ordering::Acquire) == TRACE_DISABLED {
return;
}
self.send_control(event);
}
fn send_control(&self, event: TraceEvent) {
if let Err(error) = self.sender.try_send(event) {
tracing::warn!(
request_id = %self.request_id,
%error,
"dropping Cursor trace event"
);
}
}
}
@@ -0,0 +1,349 @@
use std::{
collections::{BTreeMap, HashMap},
sync::atomic::Ordering,
time::Duration,
};
use tokio::sync::mpsc;
use crate::store::{BufferedCursorTraceChunk, Store};
use super::event::{TraceEvent, TRACE_ACTIVE, TRACE_DISABLED};
const MAX_BUFFERED_CHUNKS: usize = 32;
const MAX_BUFFERED_BYTES: usize = 256 * 1024;
const FLUSH_INTERVAL: Duration = Duration::from_millis(50);
#[derive(Clone, Copy)]
enum TraceState {
Active,
Disabled,
}
#[derive(Default)]
struct ResponseBuffer {
chunks: Vec<BufferedCursorTraceChunk>,
bytes: usize,
}
struct BufferedRequest {
artifact_type: String,
data: bytes::Bytes,
metadata: serde_json::Value,
}
#[derive(Default)]
struct RequestOrder {
next: i64,
pending: BTreeMap<i64, Vec<BufferedRequest>>,
}
pub(super) async fn run(store: Store, mut receiver: mpsc::Receiver<TraceEvent>) {
let mut states = HashMap::<String, TraceState>::new();
let mut buffers = HashMap::<String, ResponseBuffer>::new();
let mut request_orders = HashMap::<String, RequestOrder>::new();
let mut interval = tokio::time::interval(FLUSH_INTERVAL);
interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
loop {
tokio::select! {
event = receiver.recv() => {
let Some(event) = event else {
flush_all(&store, &mut buffers).await;
return;
};
process(&store, &mut states, &mut buffers, &mut request_orders, event).await;
}
_ = interval.tick() => flush_all(&store, &mut buffers).await,
}
}
}
async fn process(
store: &Store,
states: &mut HashMap<String, TraceState>,
buffers: &mut HashMap<String, ResponseBuffer>,
request_orders: &mut HashMap<String, RequestOrder>,
event: TraceEvent,
) {
let request_id = event.request_id().to_owned();
let finishes_trace = matches!(&event, TraceEvent::Finish { .. });
match event {
TraceEvent::Begin {
request_id,
activation,
conversation_id,
route,
model_id,
} => {
let state = match store
.start_cursor_trace_if_detailed(
&request_id,
conversation_id.as_deref(),
&route,
model_id.as_deref(),
)
.await
{
Ok(true) => TraceState::Active,
Ok(false) => TraceState::Disabled,
Err(error) => {
tracing::warn!(%request_id, %error, "failed to start Cursor trace");
TraceState::Disabled
}
};
activation.store(
match state {
TraceState::Active => TRACE_ACTIVE,
TraceState::Disabled => TRACE_DISABLED,
},
Ordering::Release,
);
states.insert(request_id, state);
return;
}
TraceEvent::Resume {
request_id,
activation,
} => {
let state = ensure_state(store, states, &request_id).await;
activation.store(
match state {
TraceState::Active => TRACE_ACTIVE,
TraceState::Disabled => TRACE_DISABLED,
},
Ordering::Release,
);
return;
}
_ => {}
}
if !matches!(
ensure_state(store, states, &request_id).await,
TraceState::Active
) {
if finishes_trace {
states.remove(&request_id);
buffers.remove(&request_id);
request_orders.remove(&request_id);
}
return;
}
let result = match event {
TraceEvent::Request {
artifact_type,
data,
metadata,
..
} => {
append_request(
store,
request_orders,
&request_id,
artifact_type,
data,
metadata,
)
.await
}
TraceEvent::Artifact {
artifact_type,
source,
data,
metadata,
..
} => {
store
.append_cursor_trace_artifact(
&request_id,
&artifact_type,
&source,
&data,
&metadata,
)
.await
}
TraceEvent::LinkedBlob {
artifact_type,
source,
blob_id,
metadata,
..
} => {
store
.link_cursor_trace_artifact(
&request_id,
&artifact_type,
&source,
&blob_id,
&metadata,
)
.await
}
TraceEvent::ResponseStarted { status, .. } => {
store.start_cursor_trace_response(&request_id, status).await
}
TraceEvent::ResponseChunk { source, data, .. } => {
let buffer = buffers.entry(request_id.clone()).or_default();
buffer.bytes += data.len();
buffer
.chunks
.push(BufferedCursorTraceChunk::new(&source, &data));
if buffer.chunks.len() >= MAX_BUFFERED_CHUNKS || buffer.bytes >= MAX_BUFFERED_BYTES {
flush_one(store, buffers, &request_id).await;
}
return;
}
TraceEvent::Finish { error, .. } => {
flush_request_order(store, request_orders, &request_id).await;
flush_one(store, buffers, &request_id).await;
store
.finish_cursor_trace(&request_id, error.as_deref())
.await
}
TraceEvent::Begin { .. } | TraceEvent::Resume { .. } => unreachable!(),
};
if let Err(error) = result {
tracing::warn!(%request_id, %error, "failed to record Cursor trace event");
}
if finishes_trace {
states.remove(&request_id);
buffers.remove(&request_id);
}
}
async fn append_request(
store: &Store,
request_orders: &mut HashMap<String, RequestOrder>,
request_id: &str,
artifact_type: String,
data: bytes::Bytes,
metadata: serde_json::Value,
) -> crate::Result<()> {
let append_seqno = metadata
.get("append_seqno")
.and_then(serde_json::Value::as_i64);
let ordered = artifact_type == "bidi_request"
&& metadata
.get("accepted")
.and_then(serde_json::Value::as_bool)
== Some(true)
&& metadata
.get("route_outcome")
.and_then(serde_json::Value::as_str)
== Some("local");
let Some(append_seqno) = append_seqno.filter(|_| ordered) else {
return store
.append_cursor_trace_request(
request_id,
&artifact_type,
"cursor_client",
&data,
&metadata,
)
.await;
};
let request = BufferedRequest {
artifact_type,
data,
metadata,
};
let order = request_orders.entry(request_id.to_owned()).or_default();
if append_seqno < order.next {
return store
.append_cursor_trace_request(
request_id,
&request.artifact_type,
"cursor_client",
&request.data,
&request.metadata,
)
.await;
}
order.pending.entry(append_seqno).or_default().push(request);
while let Some(requests) = order.pending.remove(&order.next) {
for request in requests {
store
.append_cursor_trace_request(
request_id,
&request.artifact_type,
"cursor_client",
&request.data,
&request.metadata,
)
.await?;
}
order.next = order.next.saturating_add(1);
}
Ok(())
}
async fn flush_request_order(
store: &Store,
request_orders: &mut HashMap<String, RequestOrder>,
request_id: &str,
) {
let Some(order) = request_orders.remove(request_id) else {
return;
};
for requests in order.pending.into_values() {
for request in requests {
if let Err(error) = store
.append_cursor_trace_request(
request_id,
&request.artifact_type,
"cursor_client",
&request.data,
&request.metadata,
)
.await
{
tracing::warn!(%request_id, %error, "failed to flush ordered Cursor request trace");
}
}
}
}
async fn ensure_state(
store: &Store,
states: &mut HashMap<String, TraceState>,
request_id: &str,
) -> TraceState {
if let Some(state) = states.get(request_id).copied() {
return state;
}
let state = match store.cursor_trace_exists(request_id).await {
Ok(true) => TraceState::Active,
Ok(false) => TraceState::Disabled,
Err(error) => {
tracing::warn!(%request_id, %error, "failed to resume Cursor trace");
TraceState::Disabled
}
};
states.insert(request_id.to_owned(), state);
state
}
async fn flush_one(store: &Store, buffers: &mut HashMap<String, ResponseBuffer>, request_id: &str) {
let Some(mut buffer) = buffers.remove(request_id) else {
return;
};
if let Err(error) = store
.add_cursor_trace_response_chunks(request_id, &buffer.chunks)
.await
{
tracing::warn!(%request_id, %error, "failed to flush Cursor response chunks");
buffer.bytes = buffer.chunks.iter().map(|chunk| chunk.data.len()).sum();
buffers.insert(request_id.to_owned(), buffer);
}
}
async fn flush_all(store: &Store, buffers: &mut HashMap<String, ResponseBuffer>) {
let request_ids = buffers.keys().cloned().collect::<Vec<_>>();
for request_id in request_ids {
flush_one(store, buffers, &request_id).await;
}
}
+1 -1
View File
@@ -12,7 +12,7 @@ use crate::{
};
pub use query::tool_query;
pub(crate) use render::{create_plan_partial, edit_content_delta, edit_path_partial};
pub(crate) use render::{create_plan_partial, edit_content_delta, edit_path_partial, task_partial};
pub use render::{dynamic_mcp_placeholder, render_dynamic_mcp, tool_completed};
pub use request::{abort, mcp_request, mcp_state_request, request};
pub(crate) use request::{edit_read_request, json_object_to_prost, mcp_meta_request};
+64 -1
View File
@@ -54,6 +54,49 @@ pub(crate) fn edit_content_delta(call: &ToolCall, content: String) -> pb::AgentS
)))
}
pub(crate) fn task_partial(
call: &ToolCall,
description: &str,
prompt: &str,
subagent: &str,
model: &str,
resume: &str,
environment: &str,
) -> pb::AgentServerMessage {
server_interaction(pb::interaction_update::Message::PartialToolCall(
pb::PartialToolCallUpdate {
call_id: call.call_id.clone(),
tool_call: Some(pb::ToolCall {
hook_additional_contexts: Vec::new(),
tool_call_id: Some(call.call_id.clone()),
started_at_ms: None,
completed_at_ms: None,
tool: Some(pb::tool_call::Tool::TaskToolCall(pb::TaskToolCall {
args: Some(pb::TaskArgs {
description: description.into(),
prompt: prompt.into(),
subagent_type: Some(subagent_type(subagent)),
model: (!model.is_empty()).then(|| model.into()),
resume: (!resume.is_empty()).then(|| resume.into()),
agent_id: None,
attachments: Vec::new(),
mode: 0,
responding_to_message_ids: Vec::new(),
environment: execution_environment(
(!environment.is_empty()).then_some(environment),
),
machine: None,
}),
result: None,
..Default::default()
})),
}),
args_text_delta: String::new(),
model_call_id: call.model_call_id.clone(),
},
))
}
pub(crate) fn create_plan_partial(
call: &ToolCall,
name: &str,
@@ -167,7 +210,7 @@ pub fn tool_completed(call: &ToolCall, completion: &ToolCompletion) -> pb::Agent
pub fn tool_placeholder(name: &str, call_id: &str) -> Result<pb::ToolCall> {
use pb::tool_call::Tool;
let tool = match normalized(name).as_str() {
"shell" => Tool::ShellToolCall(pb::ShellToolCall::default()),
"shell" | "bash" => Tool::ShellToolCall(pb::ShellToolCall::default()),
"delete" => Tool::DeleteToolCall(pb::DeleteToolCall::default()),
"glob" => Tool::GlobToolCall(pb::GlobToolCall::default()),
"grep" => Tool::GrepToolCall(pb::GrepToolCall::default()),
@@ -527,3 +570,23 @@ fn now_ms() -> u64 {
.unwrap_or_default()
.as_millis() as u64
}
#[cfg(test)]
mod tests {
use super::tool_placeholder;
use crate::cursor::protocol::proto::agent::v1 as pb;
#[test]
fn bash_renders_as_a_shell_placeholder() {
// The dispatcher treats `bash`/`Bash` as a Shell alias, so the streaming
// placeholder must too; otherwise a `Bash` tool call aborts the turn with
// `unsupported tool: bash` before it ever runs.
for name in ["shell", "Shell", "bash", "Bash"] {
let tool = tool_placeholder(name, "call-1").unwrap().tool;
assert!(
matches!(tool, Some(pb::tool_call::Tool::ShellToolCall(_))),
"{name} should render as a Shell tool"
);
}
}
}
+36 -1
View File
@@ -35,7 +35,7 @@ pub fn request(id: u32, call: &ToolCall, context: &ExecContext) -> Result<pb::Ag
.map(|v| v as i32)
};
let message = match normalize(&call.name).as_str() {
"shell" => {
"shell" | "bash" => {
let command = string("command")?;
let (simple_commands, parsing_result) = shell_command_metadata(&command);
Message::ShellStreamArgs(pb::ShellArgs {
@@ -520,3 +520,38 @@ fn prost_value(value: &Value) -> prost_types::Value {
};
ProstValue { kind: Some(kind) }
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::request;
use crate::cursor::protocol::proto::agent::v1 as pb;
use crate::cursor::tools::runtime::ExecContext;
use crate::model::ToolCall;
#[test]
fn bash_is_encoded_as_a_shell_exec_request() {
// The dispatcher routes `bash`/`Bash` to the shell executor, so the
// request codec must encode it as a Shell stream instead of erroring
// with `tool bash is not executed through ExecServerMessage`.
let call = ToolCall {
index: 0,
call_id: "call-1".into(),
model_call_id: "model-1".into(),
name: "Bash".into(),
arguments_text: String::new(),
arguments: json!({ "command": "ls -la" }),
argument_error: None,
};
let message = request(1, &call, &ExecContext::default()).unwrap();
let Some(pb::agent_server_message::Message::ExecServerMessage(exec)) = message.message
else {
panic!("expected an ExecServerMessage");
};
let Some(pb::exec_server_message::Message::ShellStreamArgs(args)) = exec.message else {
panic!("expected ShellStreamArgs");
};
assert_eq!(args.command, "ls -la");
}
}
+37 -19
View File
@@ -3,7 +3,7 @@ use crate::{
cursor::{
protocol::{events, proto::agent::v1 as pb},
tools::{
edit,
compat, edit,
runtime::{CursorToolRuntime, ExecStage, PendingExec},
tool_call_result::{self as result, ToolCompletion},
},
@@ -34,16 +34,15 @@ pub async fn client_event(
let call = match pending.exec_call(message.id).await {
Some(call) => call,
None if pending.completed_call(message.id).await.is_some() => {
return Err(Error::Protocol(format!(
"duplicate terminal ExecClientMessage id: {}",
message.id
)))
tracing::warn!(id = message.id, "ignoring duplicate terminal tool response");
return Ok(ClientExecEvent::Pending);
}
None => {
return Err(Error::Protocol(format!(
"unknown ExecClientMessage id: {}",
message.id
)))
tracing::warn!(
id = message.id,
"ignoring response for unknown tool execution"
);
return Ok(ClientExecEvent::Pending);
}
};
let Some(wire_result) = &message.message else {
@@ -144,7 +143,8 @@ pub async fn stream_closed(id: u32, pending: &CursorToolRuntime) -> Result<Optio
return Ok(None);
};
let error = "Cursor Exec stream closed before returning a terminal result";
if entry.call.name.eq_ignore_ascii_case("Shell") {
if entry.call.name.eq_ignore_ascii_case("Shell") || entry.call.name.eq_ignore_ascii_case("Bash")
{
let command = entry
.call
.arguments
@@ -173,9 +173,22 @@ pub async fn stream_closed(id: u32, pending: &CursorToolRuntime) -> Result<Optio
}
let rendered = match &entry.stage {
ExecStage::DynamicMcp(definition) => {
super::render_dynamic_mcp(&entry.call, definition, false)
Ok(super::render_dynamic_mcp(&entry.call, definition, false))
}
_ => super::render_tool_call(&entry.call, false)?,
_ => super::render_tool_call(&entry.call, false),
};
let rendered = match rendered {
Ok(rendered) => rendered,
Err(Error::Protocol(message)) => {
return Ok(Some(compat::failure_with_message(&entry.call, message)));
}
Err(Error::Json(error)) => {
return Ok(Some(compat::failure_with_message(
&entry.call,
error.to_string(),
)));
}
Err(error) => return Err(error),
};
Ok(Some(ToolCompletion::from_rendered(
&entry.call,
@@ -211,10 +224,10 @@ async fn advance_edit(
pb::exec_client_message::Message::ReadResult(result)
| pb::exec_client_message::Message::RedactedReadResult(result) => result,
_ => {
return Err(Error::Protocol(format!(
"expected ReadResult for edit tool {}",
entry.call.name
)))
let message = format!("expected ReadResult for edit tool {}", entry.call.name);
return Ok(ClientExecEvent::Completed(Box::new(
compat::failure_with_message(&entry.call, message),
)));
}
};
let write = match edit::after_read(&entry.call, read) {
@@ -259,9 +272,14 @@ fn completed(
pending: PendingExec,
result: pb::exec_client_message::Message,
) -> Result<ClientExecEvent> {
Ok(ClientExecEvent::Completed(Box::new(result::from_exec(
pending, &result,
)?)))
let call = pending.call.clone();
let completion = match result::from_exec(pending, &result) {
Ok(completion) => completion,
Err(Error::Protocol(message)) => compat::failure_with_message(&call, message),
Err(Error::Json(error)) => compat::failure_with_message(&call, error.to_string()),
Err(error) => return Err(error),
};
Ok(ClientExecEvent::Completed(Box::new(completion)))
}
fn shell_exit_result(
+4 -1
View File
@@ -49,7 +49,10 @@ pub(crate) fn render(call: &ToolCall, completed: bool) -> pb::ToolCall {
}
pub(crate) fn failure(call: &ToolCall) -> ToolCompletion {
let error = failure_message(&call.name);
failure_with_message(call, failure_message(&call.name))
}
pub(crate) fn failure_with_message(call: &ToolCall, error: String) -> ToolCompletion {
let arguments = call
.arguments
.as_object()
+50
View File
@@ -183,6 +183,9 @@ fn edit_notebook(call: &ToolCall, before: &str) -> std::result::Result<String, S
.unwrap_or_default();
let old =
normalize_newlines(&string(call, "old_string").map_err(|error| error.to_string())?);
if old.is_empty() {
return Err("old_string must not be empty".into());
}
let occurrences = source.match_indices(&old).count();
let edited = match occurrences {
0 => return Err("old_string was not found in the notebook cell".into()),
@@ -241,3 +244,50 @@ fn normalized(value: &str) -> String {
.flat_map(char::to_lowercase)
.collect()
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::edit_notebook;
use crate::model::ToolCall;
fn notebook_call(old_string: &str) -> ToolCall {
ToolCall {
index: 0,
call_id: "call".into(),
model_call_id: "model".into(),
name: "EditNotebook".into(),
arguments_text: String::new(),
arguments: json!({
"target_notebook": "/notebook.ipynb",
"cell_idx": 0,
"old_string": old_string,
"new_string": "replacement",
}),
argument_error: None,
}
}
fn single_cell_notebook() -> String {
json!({
"cells": [{"cell_type": "code", "source": ["print('hi')\n"]}],
})
.to_string()
}
#[test]
fn edit_notebook_rejects_empty_old_string() {
// StrReplace rejects an empty old_string; EditNotebook must do the same
// instead of prepending new_string (empty cell) or reporting a
// misleading "not unique" error (non-empty cell).
let error = edit_notebook(&notebook_call(""), &single_cell_notebook()).unwrap_err();
assert_eq!(error, "old_string must not be empty");
}
#[test]
fn edit_notebook_replaces_a_unique_old_string() {
let edited = edit_notebook(&notebook_call("hi"), &single_cell_notebook()).unwrap();
assert!(edited.contains("print('replacement')"));
}
}
+91 -41
View File
@@ -103,10 +103,20 @@ impl ToolDispatcher {
}
let message_index = first_tool_index + position;
let publish_started = !state.started.contains(&call.call_id);
if let Some(error) = &call.argument_error {
dispatched.push(validation_failure(call, error.clone()));
continue;
}
let edit_path = if dynamic_mcp.contains_key(&call.name) {
None
} else {
edit::execution_path(call)?
match edit::execution_path(call) {
Ok(path) => path,
Err(error) => {
dispatched.push(recover_validation_failure(call, error)?);
continue;
}
}
};
if let Some(path) = edit_path {
let next = self.edit_schedule.lock().await.start_or_defer(
@@ -121,22 +131,28 @@ impl ToolDispatcher {
let Some(next) = next else {
continue;
};
dispatched.push(
self.start(
let started = self
.start(
&next.call,
next.message_index,
next.publish_started,
dynamic_mcp,
&next.context,
)
.await?,
);
.await;
dispatched.push(match started {
Ok(started) => started,
Err(error) => recover_validation_failure(&next.call, error)?,
});
continue;
}
dispatched.push(
self.start(call, message_index, publish_started, dynamic_mcp, context)
.await?,
);
let started = self
.start(call, message_index, publish_started, dynamic_mcp, context)
.await;
dispatched.push(match started {
Ok(started) => started,
Err(error) => recover_validation_failure(call, error)?,
});
}
Ok(dispatched)
}
@@ -146,15 +162,19 @@ impl ToolDispatcher {
let Some(next) = next else {
return Ok(None);
};
self.start(
&next.call,
next.message_index,
next.publish_started,
&BTreeMap::new(),
&next.context,
)
.await
.map(Some)
match self
.start(
&next.call,
next.message_index,
next.publish_started,
&BTreeMap::new(),
&next.context,
)
.await
{
Ok(started) => Ok(Some(started)),
Err(error) => recover_validation_failure(&next.call, error).map(Some),
}
}
pub async fn interrupt_for_message(&self) -> Vec<u32> {
@@ -203,34 +223,64 @@ impl ToolDispatcher {
let pending = match self.runtime.take_interaction(response.id).await {
Some(pending) => pending,
None if self.runtime.completed_call(response.id).await.is_some() => {
return Err(Error::Protocol(format!(
"duplicate terminal InteractionResponse id: {}",
response.id
)));
tracing::warn!(
id = response.id,
"ignoring duplicate terminal interaction response"
);
return Ok(ClientToolEvent::Pending);
}
None => {
return Err(Error::Protocol(format!(
"unknown InteractionResponse id: {}",
response.id
)));
tracing::warn!(
id = response.id,
"ignoring response for unknown interaction"
);
return Ok(ClientToolEvent::Pending);
}
};
Ok(
match tool_call_dispatch::resume_interaction(
&self.results,
&self.search,
&self.fetch,
pending,
response,
)
.await?
{
tool_call_dispatch::InteractionContinuation::Completed(completion) => {
ClientToolEvent::Completed(completion)
}
tool_call_dispatch::InteractionContinuation::Pending => ClientToolEvent::Pending,
},
let call = pending.call.clone();
let continuation = match tool_call_dispatch::resume_interaction(
&self.results,
&self.search,
&self.fetch,
pending,
response,
)
.await
{
Ok(continuation) => continuation,
Err(Error::Protocol(message)) => {
return Ok(ClientToolEvent::Completed(Box::new(
compat::failure_with_message(&call, message),
)));
}
Err(Error::Json(error)) => {
return Ok(ClientToolEvent::Completed(Box::new(
compat::failure_with_message(&call, error.to_string()),
)));
}
Err(error) => return Err(error),
};
Ok(match continuation {
tool_call_dispatch::InteractionContinuation::Completed(completion) => {
ClientToolEvent::Completed(completion)
}
tool_call_dispatch::InteractionContinuation::Pending => ClientToolEvent::Pending,
})
}
}
fn validation_failure(call: &ToolCall, message: String) -> DispatchedTool {
DispatchedTool {
messages: Vec::new(),
completion: Some(compat::failure_with_message(call, message)),
}
}
fn recover_validation_failure(call: &ToolCall, error: Error) -> Result<DispatchedTool> {
match error {
Error::Protocol(message) => Ok(validation_failure(call, message)),
Error::Json(error) => Ok(validation_failure(call, error.to_string())),
error => Err(error),
}
}
+102
View File
@@ -20,6 +20,7 @@ enum Presentation {
DynamicMcp(pb::McpToolDefinition),
Edit(EditProjection),
CreatePlan(CreatePlanProjection),
Task(TaskProjection),
}
struct EditProjection {
@@ -38,6 +39,17 @@ struct CreatePlanProjection {
overview: String,
}
#[derive(Default)]
struct TaskProjection {
fields: JsonStringFields,
description: String,
prompt: String,
subagent_type: String,
model: String,
resume: String,
environment: String,
}
impl ToolCallStream {
pub fn new(name: &str, dynamic_mcp: Option<&pb::McpToolDefinition>) -> Self {
let presentation = match dynamic_mcp {
@@ -49,6 +61,7 @@ impl ToolCallStream {
Presentation::Edit(EditProjection::new("target_notebook", "new_string"))
}
"createplan" => Presentation::CreatePlan(CreatePlanProjection::default()),
"task" => Presentation::Task(TaskProjection::default()),
_ => Presentation::Plain,
},
};
@@ -73,10 +86,48 @@ impl ToolCallStream {
Ok(messages)
}
Presentation::CreatePlan(plan) => plan.project(call, raw_delta),
Presentation::Task(task) => task.project(call, raw_delta),
}
}
}
impl TaskProjection {
fn project(&mut self, call: &ToolCall, raw_delta: &str) -> Result<Vec<pb::AgentServerMessage>> {
let mut description_completed = false;
for event in self.fields.push(raw_delta)? {
match event {
StringFieldEvent::Delta { name, text } => match name.as_str() {
"description" => self.description.push_str(&text),
"prompt" => self.prompt.push_str(&text),
"subagent_type" => self.subagent_type.push_str(&text),
"model" => self.model.push_str(&text),
"resume" => self.resume.push_str(&text),
"environment" => self.environment.push_str(&text),
_ => {}
},
StringFieldEvent::End { name } if name == "description" => {
description_completed = true
}
_ => {}
}
}
Ok(description_completed
.then(|| {
interaction::task_partial(
call,
&self.description,
&self.prompt,
&self.subagent_type,
&self.model,
&self.resume,
&self.environment,
)
})
.into_iter()
.collect())
}
}
impl CreatePlanProjection {
fn project(&mut self, call: &ToolCall, raw_delta: &str) -> Result<Vec<pb::AgentServerMessage>> {
let mut completed_field = false;
@@ -184,3 +235,54 @@ fn normalized(value: &str) -> String {
.flat_map(char::to_lowercase)
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
fn task_call(arguments_text: &str) -> ToolCall {
ToolCall {
index: 0,
call_id: "task-1".into(),
model_call_id: "model-1".into(),
name: "Task".into(),
arguments_text: arguments_text.into(),
arguments: serde_json::Value::Null,
argument_error: None,
}
}
#[test]
fn task_description_projects_a_visible_partial_card_before_execution() {
let mut stream = ToolCallStream::new("Task", None);
let first = r#"{"description":"Review K10"#;
assert!(stream
.arguments_delta(&task_call(first), first)
.unwrap()
.is_empty());
let closing = "\",";
let messages = stream
.arguments_delta(&task_call(&format!("{first}{closing}")), closing)
.unwrap();
assert_eq!(messages.len(), 1);
let Some(pb::agent_server_message::Message::InteractionUpdate(update)) =
messages[0].message.as_ref()
else {
panic!("expected interaction update")
};
let Some(pb::interaction_update::Message::PartialToolCall(partial)) =
update.message.as_ref()
else {
panic!("expected partial tool call")
};
let tool_call = partial.tool_call.as_ref().expect("expected tool call");
assert_eq!(tool_call.started_at_ms, None);
let Some(pb::tool_call::Tool::TaskToolCall(task)) = tool_call.tool.as_ref() else {
panic!("expected Task tool call")
};
let args = task.args.as_ref().expect("expected partial Task args");
assert_eq!(args.description, "Review K10");
assert_eq!(args.prompt, "");
}
}
@@ -406,6 +406,7 @@ mod tests {
name: "WebFetch".into(),
arguments_text: r#"{"url":"https://example.com"}"#.into(),
arguments: json!({"url": "https://example.com"}),
argument_error: None,
},
started_at_ms: 1,
}
+43 -7
View File
@@ -15,7 +15,7 @@ use crate::{
Error, Result,
};
use super::OutputHub;
use super::{OutputHub, TransportAdmission, TransportLifecycle};
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct TransportParent {
@@ -30,7 +30,8 @@ pub struct TransportHandle {
output: Arc<OutputHub>,
conversation_id: Arc<OnceLock<String>>,
parent: Arc<OnceLock<TransportParent>>,
trace: Option<CursorTraceRecorder>,
trace: CursorTraceRecorder,
lifecycle: TransportLifecycle,
disconnect: CancellationToken,
}
@@ -39,7 +40,7 @@ impl TransportHandle {
request_id: String,
commands: mpsc::Sender<TransportCommand>,
output: Arc<OutputHub>,
trace: Option<CursorTraceRecorder>,
trace: CursorTraceRecorder,
) -> Self {
Self {
request_id,
@@ -48,6 +49,7 @@ impl TransportHandle {
conversation_id: Arc::new(OnceLock::new()),
parent: Arc::new(OnceLock::new()),
trace,
lifecycle: TransportLifecycle::new(),
disconnect: CancellationToken::new(),
}
}
@@ -127,12 +129,46 @@ impl TransportHandle {
self.output.close()
}
pub(crate) async fn wait_closed(&self) {
self.output.wait_closed().await;
pub(crate) fn trace(&self) -> Option<&CursorTraceRecorder> {
Some(&self.trace)
}
pub(crate) fn trace(&self) -> Option<&CursorTraceRecorder> {
self.trace.as_ref()
pub(crate) fn accepting_appends(&self) -> bool {
self.lifecycle.is_open()
}
pub(crate) fn admit(&self) -> Result<TransportAdmission> {
self.lifecycle
.admit()
.ok_or_else(|| Error::RunNotFound(self.request_id.clone()))
}
pub(crate) fn begin_close(&self) {
self.lifecycle.begin_close();
}
pub(crate) fn admissions_drained(&self) -> bool {
self.lifecycle.admissions_drained()
}
pub(crate) async fn wait_admissions_drained(&self) {
self.lifecycle.wait_admissions_drained().await;
}
pub(crate) fn mark_draining(&self) {
self.lifecycle.mark_draining();
}
pub(crate) fn reopen(&self) {
self.lifecycle.reopen();
}
pub(crate) fn close_transport(&self) {
self.lifecycle.close();
}
pub(crate) async fn wait_transport_closed(&self) {
self.lifecycle.wait_closed().await;
}
pub(crate) fn disconnect_token(&self) -> CancellationToken {
+170
View File
@@ -0,0 +1,170 @@
//! Coordinates append admission with transport shutdown.
use std::sync::Arc;
use tokio::sync::Notify;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum TransportState {
Open,
Closing,
Draining,
Closed,
}
#[derive(Clone)]
pub(crate) struct TransportLifecycle {
inner: Arc<LifecycleInner>,
}
struct LifecycleInner {
state: parking_lot::Mutex<LifecycleState>,
admissions_drained: Notify,
closed: Notify,
}
struct LifecycleState {
state: TransportState,
admissions: usize,
}
pub(crate) struct TransportAdmission {
inner: Arc<LifecycleInner>,
}
impl TransportLifecycle {
pub(crate) fn new() -> Self {
Self {
inner: Arc::new(LifecycleInner {
state: parking_lot::Mutex::new(LifecycleState {
state: TransportState::Open,
admissions: 0,
}),
admissions_drained: Notify::new(),
closed: Notify::new(),
}),
}
}
pub(crate) fn is_open(&self) -> bool {
self.inner.state.lock().state == TransportState::Open
}
pub(crate) fn admit(&self) -> Option<TransportAdmission> {
let mut lifecycle = self.inner.state.lock();
if lifecycle.state != TransportState::Open {
return None;
}
lifecycle.admissions += 1;
Some(TransportAdmission {
inner: self.inner.clone(),
})
}
pub(crate) fn begin_close(&self) {
let mut lifecycle = self.inner.state.lock();
if lifecycle.state != TransportState::Open {
return;
}
lifecycle.state = TransportState::Closing;
let drained = lifecycle.admissions == 0;
drop(lifecycle);
if drained {
self.inner.admissions_drained.notify_waiters();
}
}
pub(crate) fn admissions_drained(&self) -> bool {
self.inner.state.lock().admissions == 0
}
pub(crate) async fn wait_admissions_drained(&self) {
loop {
let notified = self.inner.admissions_drained.notified();
if self.inner.state.lock().admissions == 0 {
return;
}
notified.await;
}
}
pub(crate) fn mark_draining(&self) {
let mut lifecycle = self.inner.state.lock();
if lifecycle.state == TransportState::Closing && lifecycle.admissions == 0 {
lifecycle.state = TransportState::Draining;
}
}
pub(crate) fn reopen(&self) {
let mut lifecycle = self.inner.state.lock();
if matches!(
lifecycle.state,
TransportState::Closing | TransportState::Draining
) {
lifecycle.state = TransportState::Open;
}
}
pub(crate) fn close(&self) {
let mut lifecycle = self.inner.state.lock();
if lifecycle.state == TransportState::Closed {
return;
}
lifecycle.state = TransportState::Closed;
drop(lifecycle);
self.inner.closed.notify_waiters();
}
pub(crate) async fn wait_closed(&self) {
loop {
let notified = self.inner.closed.notified();
if self.inner.state.lock().state == TransportState::Closed {
return;
}
notified.await;
}
}
}
impl Drop for TransportAdmission {
fn drop(&mut self) {
let mut lifecycle = self.inner.state.lock();
lifecycle.admissions = lifecycle.admissions.saturating_sub(1);
let drained = lifecycle.state == TransportState::Closing && lifecycle.admissions == 0;
drop(lifecycle);
if drained {
self.inner.admissions_drained.notify_waiters();
}
}
}
#[cfg(test)]
mod tests {
use super::{TransportLifecycle, TransportState};
#[tokio::test]
async fn closing_waits_for_existing_admissions() {
let lifecycle = TransportLifecycle::new();
let admission = lifecycle.admit().unwrap();
lifecycle.begin_close();
assert!(lifecycle.admit().is_none());
drop(admission);
lifecycle.wait_admissions_drained().await;
lifecycle.mark_draining();
assert_eq!(lifecycle.inner.state.lock().state, TransportState::Draining);
}
#[tokio::test]
async fn an_admitted_continuation_reopens_the_transport() {
let lifecycle = TransportLifecycle::new();
let admission = lifecycle.admit().unwrap();
lifecycle.begin_close();
drop(admission);
lifecycle.wait_admissions_drained().await;
lifecycle.mark_draining();
lifecycle.reopen();
assert_eq!(lifecycle.inner.state.lock().state, TransportState::Open);
assert!(lifecycle.admit().is_some());
}
}
+2
View File
@@ -2,10 +2,12 @@
mod handle;
mod inbox;
mod lifecycle;
mod output;
mod registry;
pub use handle::*;
pub use inbox::*;
pub(crate) use lifecycle::*;
pub use output::*;
pub use registry::*;
+1 -13
View File
@@ -1,12 +1,11 @@
//! Buffers, replays, broadcasts, and atomically closes downstream output.
use bytes::Bytes;
use tokio::sync::{mpsc, Notify};
use tokio::sync::mpsc;
#[derive(Default)]
pub struct OutputHub {
state: parking_lot::Mutex<OutputState>,
closed: Notify,
}
#[derive(Default)]
@@ -49,17 +48,6 @@ impl OutputHub {
state.closed = true;
state.subscribers.clear();
drop(state);
self.closed.notify_waiters();
true
}
pub async fn wait_closed(&self) {
loop {
let notified = self.closed.notified();
if self.state.lock().closed {
return;
}
notified.await;
}
}
}
+105 -21
View File
@@ -1,13 +1,19 @@
//! Maps request IDs to active transport handles.
use std::{collections::HashMap, sync::Arc};
use std::{
collections::HashMap,
sync::{
atomic::{AtomicU64, Ordering},
Arc,
},
};
use tokio::sync::{mpsc, Mutex, Notify};
use crate::{
cursor::{
conversation::ConversationRegistry, prompting::PromptCompiler,
services::observability::CursorTraceRecorder,
services::observability::CursorTraceService,
},
plugin::PluginRegistry,
provider::Provider,
@@ -24,15 +30,23 @@ pub struct TransportRegistry {
}
struct RegistryInner {
local: Mutex<HashMap<String, TransportHandle>>,
local: Mutex<HashMap<String, LocalTransport>>,
next_local_generation: AtomicU64,
upstream: Mutex<HashMap<String, u64>>,
route_changed: Notify,
store: Store,
traces: CursorTraceService,
web_cache: WebCache,
plugins: Option<PluginRegistry>,
conversations: ConversationRegistry,
}
#[derive(Clone)]
struct LocalTransport {
generation: u64,
handle: TransportHandle,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum TransportRoute {
Local,
@@ -50,7 +64,24 @@ impl TransportRegistry {
compiler: PromptCompiler,
web_cache: WebCache,
) -> Self {
Self::build(store, provider, compiler, web_cache, None)
Self::build(store, provider, compiler, web_cache, None, None)
}
/// 附带本地 rules 目录的构造;编译请求上下文时会合并该目录下的 md 规则。
pub fn with_local_rules(
store: Store,
provider: Arc<dyn Provider>,
compiler: PromptCompiler,
local_rules_dir: std::path::PathBuf,
) -> Self {
Self::build(
store,
provider,
compiler,
WebCache::default(),
None,
Some(local_rules_dir),
)
}
pub fn with_plugins(
@@ -59,8 +90,16 @@ impl TransportRegistry {
compiler: PromptCompiler,
web_cache: WebCache,
plugins: PluginRegistry,
local_rules_dir: std::path::PathBuf,
) -> Self {
Self::build(store, provider, compiler, web_cache, Some(plugins))
Self::build(
store,
provider,
compiler,
web_cache,
Some(plugins),
Some(local_rules_dir),
)
}
fn build(
@@ -69,17 +108,21 @@ impl TransportRegistry {
compiler: PromptCompiler,
web_cache: WebCache,
plugins: Option<PluginRegistry>,
local_rules_dir: Option<std::path::PathBuf>,
) -> Self {
Self {
inner: Arc::new(RegistryInner {
local: Mutex::new(HashMap::new()),
next_local_generation: AtomicU64::new(1),
upstream: Mutex::new(HashMap::new()),
route_changed: Notify::new(),
traces: CursorTraceService::new(store.clone()),
conversations: ConversationRegistry::new(
store.clone(),
provider,
compiler,
web_cache.clone(),
local_rules_dir,
),
store,
web_cache,
@@ -92,6 +135,13 @@ impl TransportRegistry {
&self.inner.store
}
pub fn trace(
&self,
request_id: &str,
) -> crate::cursor::services::observability::CursorTraceRecorder {
self.inner.traces.recorder(request_id)
}
pub fn web_cache(&self) -> &WebCache {
&self.inner.web_cache
}
@@ -105,18 +155,37 @@ impl TransportRegistry {
}
pub async fn get_or_create(&self, request_id: &str) -> Result<TransportHandle> {
if let Some(handle) = self.inner.local.lock().await.get(request_id).cloned() {
return Ok(handle);
self.get_or_create_for_append(request_id, false).await
}
pub(crate) async fn get_or_create_for_append(
&self,
request_id: &str,
replace_closing: bool,
) -> Result<TransportHandle> {
let mut local = self.inner.local.lock().await;
if let Some(transport) = local.get(request_id) {
if transport.handle.accepting_appends() || !replace_closing {
return Ok(transport.handle.clone());
}
}
local.remove(request_id);
let (commands, receiver) = mpsc::channel(128);
let output = Arc::new(OutputHub::default());
let trace = CursorTraceRecorder::resume(self.inner.store.clone(), request_id).await;
let handle = TransportHandle::new(request_id.into(), commands, output.clone(), trace);
let mut local = self.inner.local.lock().await;
if let Some(existing) = local.get(request_id).cloned() {
return Ok(existing);
}
local.insert(request_id.into(), handle.clone());
let trace = self.inner.traces.recorder(request_id);
trace.resume();
let handle = TransportHandle::new(request_id.into(), commands, output, trace);
let generation = self
.inner
.next_local_generation
.fetch_add(1, Ordering::Relaxed);
local.insert(
request_id.into(),
LocalTransport {
generation,
handle: handle.clone(),
},
);
drop(local);
self.inner.route_changed.notify_waiters();
self.inner
@@ -125,17 +194,29 @@ impl TransportRegistry {
let registry = Arc::downgrade(&self.inner);
let request_id = request_id.to_string();
let lifecycle = handle.clone();
tokio::spawn(async move {
output.wait_closed().await;
lifecycle.wait_transport_closed().await;
if let Some(registry) = registry.upgrade() {
registry.local.lock().await.remove(&request_id);
let mut local = registry.local.lock().await;
if local
.get(&request_id)
.is_some_and(|transport| transport.generation == generation)
{
local.remove(&request_id);
}
}
});
Ok(handle)
}
pub async fn local(&self, request_id: &str) -> Option<TransportHandle> {
self.inner.local.lock().await.get(request_id).cloned()
self.inner
.local
.lock()
.await
.get(request_id)
.map(|transport| transport.handle.clone())
}
pub async fn mark_upstream(&self, request_id: &str) {
@@ -179,10 +260,13 @@ impl TransportRegistry {
self.inner.conversations.shutdown().await;
let handles = std::mem::take(&mut *self.inner.local.lock().await);
self.inner.upstream.lock().await.clear();
for handle in handles.into_values() {
handle.disconnect().await;
let _ =
tokio::time::timeout(std::time::Duration::from_secs(2), handle.wait_closed()).await;
for transport in handles.into_values() {
transport.handle.disconnect().await;
let _ = tokio::time::timeout(
std::time::Duration::from_secs(2),
transport.handle.wait_transport_closed(),
)
.await;
}
}
}
+4
View File
@@ -161,6 +161,10 @@ fn is_local_path(path: &str) -> bool {
| "/aiserver.v1.DashboardService/GetUserProfile"
| "/aiserver.v1.DashboardService/GetCurrentPeriodUsage"
| "/aiserver.v1.DashboardService/GetUsageLimitStatusAndActiveGrants"
| "/aiserver.v1.AiService/KnowledgeBaseAdd"
| "/aiserver.v1.AiService/KnowledgeBaseList"
| "/aiserver.v1.AiService/KnowledgeBaseUpdate"
| "/aiserver.v1.AiService/KnowledgeBaseRemove"
| "/aiserver.v1.AnalyticsService/BootstrapStatsig"
| "/auth/full_stripe_profile"
)
+11
View File
@@ -87,6 +87,9 @@ pub struct ModelConfigInput {
#[serde(default)]
pub sort_order: i64,
pub display_name: String,
/// 供应商分组的自定义显示名;同一 base_url 主机下的模型共享。
#[serde(default)]
pub group_name: Option<String>,
#[serde(rename = "type")]
pub model_type: ModelType,
pub base_url: String,
@@ -124,6 +127,7 @@ pub struct ModelConfig {
pub model_hash: String,
pub sort_order: i64,
pub display_name: String,
pub group_name: Option<String>,
#[serde(rename = "type")]
pub model_type: ModelType,
pub base_url: String,
@@ -204,6 +208,12 @@ impl ModelConfig {
pub fn normalize_model_input(input: &ModelConfigInput) -> Result<ModelConfigInput> {
let display_name = required(&input.display_name, "model display name")?;
let group_name = input
.group_name
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(String::from);
let base_url = normalize_request_url(&input.base_url)?;
let api_key = required(&input.api_key, "model API key")?;
let tooltip_data = required(&input.tooltip_data, "model tooltip")?;
@@ -230,6 +240,7 @@ pub fn normalize_model_input(input: &ModelConfigInput) -> Result<ModelConfigInpu
let normalized = ModelConfigInput {
sort_order: input.sort_order.max(0),
display_name,
group_name,
model_type: input.model_type,
base_url,
use_full_url: input.use_full_url,
+3 -26
View File
@@ -6,11 +6,10 @@ mod usage {
use serde::{Deserialize, Serialize};
use super::ProviderType;
#[derive(Clone, Copy, Debug, Default, Serialize, Deserialize, PartialEq, Eq)]
pub struct Usage {
pub input_tokens: Option<u64>,
pub context_input_tokens: Option<u64>,
pub output_tokens: Option<u64>,
pub total_tokens: Option<u64>,
pub cache_read_tokens: Option<u64>,
@@ -18,24 +17,10 @@ mod usage {
pub reasoning_tokens: Option<u64>,
}
impl Usage {
/// Returns the provider-visible input context without counting cached tokens twice.
pub(crate) fn context_input_tokens(self, provider: ProviderType) -> Option<u64> {
let input = self.input_tokens?;
match provider {
ProviderType::OpenAiChat | ProviderType::OpenAiResponses | ProviderType::Plugin => {
Some(input)
}
ProviderType::Anthropic => input
.checked_add(self.cache_read_tokens.unwrap_or_default())?
.checked_add(self.cache_write_tokens.unwrap_or_default()),
}
}
}
impl AddAssign for Usage {
fn add_assign(&mut self, rhs: Self) {
self.input_tokens = sum(self.input_tokens, rhs.input_tokens);
self.context_input_tokens = sum(self.context_input_tokens, rhs.context_input_tokens);
self.output_tokens = sum(self.output_tokens, rhs.output_tokens);
self.total_tokens = sum(self.total_tokens, rhs.total_tokens);
self.cache_read_tokens = sum(self.cache_read_tokens, rhs.cache_read_tokens);
@@ -53,7 +38,7 @@ pub use usage::*;
mod llm_call {
use serde::Serialize;
use super::{ProviderType, Usage};
use super::ProviderType;
#[derive(Clone, Debug)]
pub struct NewLlmCall {
@@ -75,14 +60,6 @@ mod llm_call {
pub detailed: bool,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) struct LlmCallUsageAnchor {
pub request_type: ProviderType,
pub usage: Usage,
pub message_count: usize,
pub tool_count: usize,
}
#[derive(Clone, Debug, Serialize)]
pub struct LlmCallSummary {
pub call_id: String,
+55 -6
View File
@@ -6,8 +6,8 @@ use serde::{Deserialize, Serialize};
use crate::{Error, Result};
use super::{
CanonicalMessage, ContentPart, MessageContent, ProviderReplayState, Role, ToolCallContent,
ToolResultContent,
normalize_tool_name, CanonicalMessage, ContentPart, MessageContent, ProviderReplayState, Role,
ToolCallContent, ToolResultContent,
};
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
@@ -91,7 +91,7 @@ fn project_tool_round(
"tool round repeats provider replay state".into(),
));
}
calls.extend(part_calls.iter().cloned());
calls.extend(part_calls.iter().map(normalized_tool_call));
cursor += 1;
while cursor < messages.len() {
@@ -139,7 +139,7 @@ fn project_tool_round(
.map(|(message_id, result)| ProjectedMessage {
message_id,
role: Role::Tool,
content: ProjectedContent::ToolResult(result),
content: ProjectedContent::ToolResult(normalized_tool_result(&result)),
}),
);
Ok(Some((output, cursor)))
@@ -158,9 +158,11 @@ fn project_message(message: &CanonicalMessage) -> ProjectedMessage {
text: text.clone(),
thinking: thinking.clone(),
replay_state: replay_state.clone(),
calls: tool_calls.clone(),
calls: tool_calls.iter().map(normalized_tool_call).collect(),
},
MessageContent::ToolResult(result) => ProjectedContent::ToolResult(result.clone()),
MessageContent::ToolResult(result) => {
ProjectedContent::ToolResult(normalized_tool_result(result))
}
};
ProjectedMessage {
message_id: message.message_id.clone(),
@@ -168,3 +170,50 @@ fn project_message(message: &CanonicalMessage) -> ProjectedMessage {
content,
}
}
fn normalized_tool_call(call: &ToolCallContent) -> ToolCallContent {
let mut call = call.clone();
call.name = normalize_tool_name(&call.name);
call
}
fn normalized_tool_result(result: &ToolResultContent) -> ToolResultContent {
let mut result = result.clone();
result.name = normalize_tool_name(&result.name);
result
}
#[cfg(test)]
mod tests {
use super::*;
use crate::model::Origin;
use serde_json::json;
#[test]
fn tool_names_are_normalized_before_provider_dispatch() {
let messages = [CanonicalMessage {
message_id: "assistant-1".into(),
role: Role::Assistant,
origin: Origin::Assistant,
content: MessageContent::Assistant {
text: String::new(),
thinking: String::new(),
tool_round_id: None,
replay_state: None,
tool_calls: vec![ToolCallContent {
index: 0,
call_id: "call-1".into(),
name: "multi_tool_use.parallel".into(),
arguments: json!({}),
}],
},
runtime_event_id: None,
}];
let projected = project_messages(&messages).unwrap();
let ProjectedContent::Assistant { calls, .. } = &projected[0].content else {
panic!("expected assistant projection");
};
assert_eq!(calls[0].name, "multi_tool_use_parallel");
}
}
+229 -1
View File
@@ -1,4 +1,99 @@
//! Estimates and records model token usage.
//! Estimates provider-visible context size and formats configured token counts.
use super::{ContentPart, ProjectedContent, ProjectedMessage, PromptSpec};
const TOKENS_PER_MESSAGE_OVERHEAD: u64 = 8;
const TOKENS_PER_TOOL_CALL_OVERHEAD: u64 = 6;
const TOKENS_PER_IMAGE: u64 = 1_024;
pub(crate) fn estimate_context_tokens(prompt: &PromptSpec, messages: &[ProjectedMessage]) -> u64 {
let instructions = estimate_text_tokens(&prompt.instructions);
let tools = prompt.tools.iter().fold(0_u64, |total, tool| {
total.saturating_add(estimate_json_tokens(tool))
});
instructions
.saturating_add(tools)
.saturating_add(estimate_projected_messages_tokens(messages))
}
pub(crate) fn estimate_projected_messages_tokens(messages: &[ProjectedMessage]) -> u64 {
messages.iter().fold(0_u64, |total, message| {
total.saturating_add(estimate_message_tokens(message))
})
}
fn estimate_message_tokens(message: &ProjectedMessage) -> u64 {
let content = match &message.content {
ProjectedContent::Parts(parts) => estimate_parts_tokens(parts),
ProjectedContent::Assistant {
text,
thinking,
replay_state: _,
calls,
} => {
let calls = calls.iter().fold(0_u64, |total, call| {
total
.saturating_add(TOKENS_PER_TOOL_CALL_OVERHEAD)
.saturating_add(estimate_text_tokens(&call.call_id))
.saturating_add(estimate_text_tokens(&call.name))
.saturating_add(estimate_json_tokens(&call.arguments))
});
estimate_text_tokens(text)
.saturating_add(estimate_text_tokens(thinking))
.saturating_add(calls)
}
ProjectedContent::ToolResult(result) => {
let content = if result.provider_parts.is_empty() {
estimate_text_tokens(&result.content).saturating_add(
result
.image
.as_ref()
.map(|image| {
TOKENS_PER_IMAGE.saturating_add(estimate_text_tokens(&image.mime_type))
})
.unwrap_or_default(),
)
} else {
estimate_parts_tokens(&result.provider_parts)
};
estimate_text_tokens(&result.call_id)
.saturating_add(estimate_text_tokens(&result.name))
.saturating_add(content)
}
};
TOKENS_PER_MESSAGE_OVERHEAD.saturating_add(content)
}
fn estimate_parts_tokens(parts: &[ContentPart]) -> u64 {
parts.iter().fold(0_u64, |total, part| {
let tokens = match part {
ContentPart::Text { text } => estimate_text_tokens(text),
ContentPart::Image { mime_type, .. } => {
TOKENS_PER_IMAGE.saturating_add(estimate_text_tokens(mime_type))
}
};
total.saturating_add(tokens)
})
}
fn estimate_json_tokens(value: &impl serde::Serialize) -> u64 {
serde_json::to_string(value)
.map(|value| estimate_text_tokens(&value))
.unwrap_or_default()
}
fn estimate_text_tokens(text: &str) -> u64 {
let text = text.trim();
if text.is_empty() {
return 0;
}
let characters = text.chars().count() as u64;
characters
.div_ceil(4)
.saturating_add(text.bytes().filter(|byte| *byte == b'\n').count() as u64)
.max(1)
}
pub(crate) fn parse_token_count(value: &str) -> Option<u64> {
let value = value.trim().to_ascii_lowercase();
let (number, multiplier) = match value.chars().last()? {
@@ -18,3 +113,136 @@ pub(crate) fn format_token_count(tokens: u64) -> String {
tokens.to_string()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::model::{
ProviderReplayState, Role, ToolCallContent, ToolDefinition, ToolResultContent,
};
fn prompt() -> PromptSpec {
PromptSpec {
instructions: "system instructions".into(),
tools: vec![ToolDefinition {
name: "Read".into(),
description: "Read a file".into(),
parameters: serde_json::json!({"type": "object", "properties": {"path": {"type": "string"}}}),
}],
}
}
#[test]
fn context_estimate_grows_with_provider_visible_text_and_tools() {
let short = vec![ProjectedMessage {
message_id: "short".into(),
role: Role::User,
content: ProjectedContent::Parts(vec![ContentPart::Text {
text: "hello".into(),
}]),
}];
let long = vec![ProjectedMessage {
message_id: "long".into(),
role: Role::User,
content: ProjectedContent::Parts(vec![ContentPart::Text {
text: "x".repeat(40_000),
}]),
}];
let without_tools = PromptSpec {
instructions: prompt().instructions,
tools: Vec::new(),
};
assert!(
estimate_context_tokens(&prompt(), &short)
> estimate_context_tokens(&without_tools, &short)
);
assert!(
estimate_context_tokens(&prompt(), &long) > estimate_context_tokens(&prompt(), &short)
);
}
#[test]
fn context_estimate_counts_tool_calls_results_and_images() {
let assistant = ProjectedMessage {
message_id: "assistant".into(),
role: Role::Assistant,
content: ProjectedContent::Assistant {
text: String::new(),
thinking: "reasoning".into(),
replay_state: None,
calls: vec![ToolCallContent {
index: 0,
call_id: "call-1".into(),
name: "Read".into(),
arguments: serde_json::json!({"path": "/tmp/file"}),
}],
},
};
let text_result = ProjectedMessage {
message_id: "result-text".into(),
role: Role::Tool,
content: ProjectedContent::ToolResult(ToolResultContent {
call_id: "call-1".into(),
name: "Read".into(),
content: "file contents".into(),
is_error: false,
image: None,
provider_parts: Vec::new(),
}),
};
let image_result = ProjectedMessage {
message_id: "result-image".into(),
role: Role::Tool,
content: ProjectedContent::ToolResult(ToolResultContent {
call_id: "call-1".into(),
name: "Read".into(),
content: "file contents".into(),
is_error: false,
image: None,
provider_parts: vec![ContentPart::Image {
mime_type: "image/png".into(),
data: vec![0; 32],
}],
}),
};
let base = estimate_context_tokens(&prompt(), &[]);
let with_call = estimate_context_tokens(&prompt(), std::slice::from_ref(&assistant));
let with_text = estimate_context_tokens(&prompt(), &[assistant.clone(), text_result]);
let with_image = estimate_context_tokens(&prompt(), &[assistant, image_result]);
assert!(with_call > base);
assert!(with_text > with_call);
assert!(with_image > with_text);
}
#[test]
fn assistant_replay_state_does_not_duplicate_thinking_or_count_signature() {
let assistant = |replay_state| ProjectedMessage {
message_id: "assistant".into(),
role: Role::Assistant,
content: ProjectedContent::Assistant {
text: "answer".into(),
thinking: "reasoning".repeat(1_000),
replay_state,
calls: Vec::new(),
},
};
let without_replay = assistant(None);
let with_replay = assistant(Some(ProviderReplayState {
provider_kind: "anthropic".into(),
value: serde_json::json!({
"blocks": [{
"type": "thinking",
"thinking": "reasoning".repeat(1_000),
"signature": "s".repeat(282_100)
}]
}),
}));
assert_eq!(
estimate_projected_messages_tokens(&[without_replay]),
estimate_projected_messages_tokens(&[with_replay])
);
}
}
+19
View File
@@ -4,6 +4,24 @@ use serde_json::Value;
use super::ProviderReplayState;
pub fn normalize_tool_name(name: &str) -> String {
let normalized = name
.chars()
.map(|character| {
if character.is_ascii_alphanumeric() || matches!(character, '_' | '-') {
character
} else {
'_'
}
})
.collect::<String>();
if normalized.is_empty() {
"_".into()
} else {
normalized
}
}
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
pub struct ToolDefinition {
pub name: String,
@@ -19,6 +37,7 @@ pub struct ToolCall {
pub name: String,
pub arguments_text: String,
pub arguments: Value,
pub argument_error: Option<String>,
}
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
+87 -2
View File
@@ -1,8 +1,93 @@
//! Provides shared network client and transport configuration.
//! Outbound HTTP clients configured from persisted application proxy settings.
//! Owns reusable outbound HTTP clients configured from persisted proxy settings.
use std::{sync::Arc, time::Duration};
use tokio::sync::RwLock;
use crate::{store::Store, Result};
#[derive(Clone)]
pub struct NetworkClients {
store: Store,
cache: Arc<RwLock<ClientCache>>,
}
#[derive(Default)]
struct ClientCache {
default: Option<reqwest::Client>,
cursor: Option<reqwest::Client>,
provider: Option<(Duration, reqwest::Client)>,
}
impl NetworkClients {
pub fn new(store: Store) -> Self {
Self {
store,
cache: Arc::new(RwLock::new(ClientCache::default())),
}
}
pub async fn default_client(&self) -> Result<reqwest::Client> {
if let Some(client) = self.cache.read().await.default.clone() {
return Ok(client);
}
let mut cache = self.cache.write().await;
if let Some(client) = cache.default.clone() {
return Ok(client);
}
let client = client_builder(&self.store).await?.build()?;
cache.default = Some(client.clone());
Ok(client)
}
pub async fn cursor_client(&self) -> Result<reqwest::Client> {
if let Some(client) = self.cache.read().await.cursor.clone() {
return Ok(client);
}
let mut cache = self.cache.write().await;
if let Some(client) = cache.cursor.clone() {
return Ok(client);
}
let client = client_builder(&self.store)
.await?
.redirect(reqwest::redirect::Policy::none())
.build()?;
cache.cursor = Some(client.clone());
Ok(client)
}
pub async fn provider_client(&self, timeout: Duration) -> Result<reqwest::Client> {
if let Some((_, client)) = self
.cache
.read()
.await
.provider
.as_ref()
.filter(|(cached_timeout, _)| *cached_timeout == timeout)
{
return Ok(client.clone());
}
let mut cache = self.cache.write().await;
if let Some((_, client)) = cache
.provider
.as_ref()
.filter(|(cached_timeout, _)| *cached_timeout == timeout)
{
return Ok(client.clone());
}
let client = client_builder(&self.store)
.await?
.timeout(timeout)
.build()?;
cache.provider = Some((timeout, client.clone()));
Ok(client)
}
pub async fn invalidate(&self) {
*self.cache.write().await = ClientCache::default();
}
}
pub async fn client_builder(store: &Store) -> Result<reqwest::ClientBuilder> {
let settings = store.proxy_settings_secret().await?;
// Use the platform TLS stack for compatibility with provider gateways that
-4
View File
@@ -107,9 +107,7 @@ pub struct PluginModelDescriptor {
pub description: Option<String>,
pub icon: String,
pub provider_type: String,
pub context_window_tokens: Option<u64>,
pub max_output_tokens: Option<u64>,
pub thinking: bool,
pub images: bool,
}
@@ -209,9 +207,7 @@ impl PluginModelDescriptor {
description: model.description.clone(),
icon: icon.to_owned(),
provider_type: provider.provider_type.clone(),
context_window_tokens: model.context_window_tokens,
max_output_tokens: model.max_output_tokens,
thinking: model.thinking,
images: model.images,
}
}
-1
View File
@@ -20,7 +20,6 @@ pub use descriptor::{
};
pub use registry::{ImportResponse, OAuthBeginResponse, OAuthPollResponse, PluginRegistry};
pub use runtime::{PluginRuntime, PluginRuntimePhase, PluginRuntimeState, PluginRuntimeStatus};
pub(crate) use wire::llm_request as plugin_llm_request;
/// Windows 下阻止 Deno 子进程弹出控制台窗口(CREATE_NO_WINDOW)。
#[cfg(windows)]
+7 -3
View File
@@ -20,8 +20,11 @@ use super::{
worker::{PluginWorker, WorkerStreamItem},
};
use crate::{
model::ModelInvocation, provider::ModelEvent, provider::ProviderStream, store::Store, Error,
Result,
model::ModelInvocation,
provider::ProviderStream,
provider::{CallRecorder, ModelEvent},
store::Store,
Error, Result,
};
const OAUTH_SLOW_DOWN_STEP_MS: i64 = 5_000;
@@ -203,6 +206,7 @@ impl PluginRegistry {
&self,
invocation: ModelInvocation,
cancellation: CancellationToken,
recorder: CallRecorder,
) -> ProviderStream {
let registry = self.clone();
Box::pin(try_stream! {
@@ -232,7 +236,7 @@ impl PluginRegistry {
"request": request,
});
let worker = registry.worker(&entry, &executable).await;
let mut items = worker.invoke_streaming("provider.invoke", params, cancellation.clone()).await?;
let mut items = worker.invoke_streaming("provider.invoke", params, cancellation.clone(), Some(recorder)).await?;
yield ModelEvent::Start { model_call_id: invocation.call_id.clone() };
while let Some(item) = items.recv().await {
match item {
-2
View File
@@ -2,7 +2,6 @@ import type { JsonValue, PluginContext } from "./plugin.ts";
import type { ResourceSnapshot } from "./resource.ts";
export type ModelCapabilities = {
thinking?: boolean;
images?: boolean;
};
@@ -10,7 +9,6 @@ export type ModelDefinition = {
id: string;
displayName: string;
description?: string;
contextWindowTokens?: number;
maxOutputTokens?: number;
capabilities?: ModelCapabilities;
/** 之后的调用原样传回;永远不会展示给用户。 */
@@ -53,7 +53,17 @@ function replayItems(value: JsonValue): JsonValue[] {
if (!Array.isArray(items)) {
throw new Error("OpenAI Responses replay state is missing items");
}
return items;
return items.map((item) => {
const source = record(item);
if (source?.type !== "reasoning") {
throw new Error("OpenAI Responses replay state contains a non-reasoning item");
}
const projected: Record<string, JsonValue> = { type: "reasoning" };
for (const field of ["id", "summary", "content", "encrypted_content"] as const) {
if (field in source) projected[field] = source[field] as JsonValue;
}
return projected;
});
}
export function buildResponsesBody(call: OpenAiResponsesCall): Record<string, JsonValue> {
+3 -12
View File
@@ -135,12 +135,8 @@ pub struct StoredModel {
#[serde(default)]
pub description: Option<String>,
#[serde(default)]
pub context_window_tokens: Option<u64>,
#[serde(default)]
pub max_output_tokens: Option<u64>,
#[serde(default)]
pub thinking: bool,
#[serde(default)]
pub images: bool,
#[serde(default)]
pub private_data: serde_json::Value,
@@ -179,13 +175,9 @@ impl StoredModel {
.get("description")
.and_then(serde_json::Value::as_str)
.map(str::to_owned),
context_window_tokens: object
.get("contextWindowTokens")
.and_then(serde_json::Value::as_u64),
max_output_tokens: object
.get("maxOutputTokens")
.and_then(serde_json::Value::as_u64),
thinking: capability("thinking"),
images: capability("images"),
private_data: object
.get("privateData")
@@ -200,9 +192,8 @@ impl StoredModel {
"id": self.id,
"displayName": self.display_name,
"description": self.description,
"contextWindowTokens": self.context_window_tokens,
"maxOutputTokens": self.max_output_tokens,
"capabilities": { "thinking": self.thinking, "images": self.images },
"capabilities": { "images": self.images },
"privateData": self.private_data,
})
}
@@ -450,7 +441,7 @@ mod tests {
let model = StoredModel::from_definition(&serde_json::json!({
"id": "gpt-test",
"displayName": "GPT Test",
"capabilities": {"thinking": true},
"capabilities": {"images": true},
"privateData": {"reasoningEfforts": ["low"]},
}))
.unwrap();
@@ -460,7 +451,7 @@ mod tests {
.unwrap();
let models = store.models("dev.example", "codex").await.unwrap();
assert_eq!(models.len(), 1);
assert!(models[0].thinking);
assert!(models[0].images);
assert_eq!(models[0].private_data["reasoningEfforts"][0], "low");
}
}
+4 -1
View File
@@ -147,8 +147,10 @@ pub fn model_event(value: &serde_json::Value) -> Result<ModelEvent> {
.get("usage")
.ok_or_else(|| Error::Protocol("plugin usage event requires usage".into()))?;
let tokens = |name: &str| usage.get(name).and_then(serde_json::Value::as_u64);
let input_tokens = tokens("inputTokens");
ModelEvent::Usage(Usage {
input_tokens: tokens("inputTokens"),
input_tokens,
context_input_tokens: input_tokens,
output_tokens: tokens("outputTokens"),
total_tokens: tokens("totalTokens"),
cache_read_tokens: tokens("cacheReadTokens"),
@@ -285,6 +287,7 @@ mod tests {
usage,
ModelEvent::Usage(Usage {
input_tokens: Some(10),
context_input_tokens: Some(10),
output_tokens: Some(2),
total_tokens: None,
cache_read_tokens: Some(4),
+269 -24
View File
@@ -3,7 +3,10 @@ use std::{
collections::{HashMap, HashSet},
path::PathBuf,
process::Stdio,
sync::Arc,
sync::{
atomic::{AtomicBool, Ordering},
Arc,
},
time::Duration,
};
@@ -19,7 +22,7 @@ use super::{
definition::{file_url, PluginDefinitionLoader},
protocol::{HostMessage, WorkerMessage},
};
use crate::{store::Store, Error, Result};
use crate::{provider::CallRecorder, store::Store, Error, Result};
const INVOCATION_TIMEOUT: Duration = Duration::from_secs(10 * 60);
const MAX_NETWORK_RESPONSE_BYTES: u64 = 16 * 1024 * 1024;
@@ -56,12 +59,29 @@ struct WorkerProcess {
stdin: Arc<Mutex<ChildStdin>>,
}
struct InvocationState {
cancellation: CancellationToken,
recorder: Option<CallRecorder>,
recorder_claimed: AtomicBool,
}
impl InvocationState {
fn claim_recorder(&self) -> Option<CallRecorder> {
self.recorder.as_ref().and_then(|recorder| {
self.recorder_claimed
.compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire)
.ok()
.map(|_| recorder.clone())
})
}
}
#[derive(Clone)]
struct HostContext {
plugin_id: String,
network_hosts: Arc<HashSet<String>>,
store: Store,
cancellations: Arc<Mutex<HashMap<String, CancellationToken>>>,
invocations: Arc<Mutex<HashMap<String, Arc<InvocationState>>>>,
streams: Arc<Mutex<HashMap<String, StreamLines>>>,
}
@@ -87,7 +107,7 @@ impl PluginWorker {
.collect(),
),
store,
cancellations: Arc::new(Mutex::new(HashMap::new())),
invocations: Arc::new(Mutex::new(HashMap::new())),
streams: Arc::new(Mutex::new(HashMap::new())),
},
plugin_id,
@@ -108,7 +128,9 @@ impl PluginWorker {
params: serde_json::Value,
cancellation: CancellationToken,
) -> Result<serde_json::Value> {
let mut items = self.invoke_streaming(method, params, cancellation).await?;
let mut items = self
.invoke_streaming(method, params, cancellation, None)
.await?;
let result = tokio::time::timeout(INVOCATION_TIMEOUT, async {
while let Some(item) = items.recv().await {
if let WorkerStreamItem::Result(result) = item {
@@ -137,15 +159,18 @@ impl PluginWorker {
method: &str,
params: serde_json::Value,
cancellation: CancellationToken,
recorder: Option<CallRecorder>,
) -> Result<mpsc::UnboundedReceiver<WorkerStreamItem>> {
let id = uuid::Uuid::new_v4().to_string();
let request_cancellation = CancellationToken::new();
self.inner
.host
.cancellations
.lock()
.await
.insert(id.clone(), request_cancellation.clone());
self.inner.host.invocations.lock().await.insert(
id.clone(),
Arc::new(InvocationState {
cancellation: request_cancellation.clone(),
recorder,
recorder_claimed: AtomicBool::new(false),
}),
);
let (sender, receiver) = mpsc::unbounded_channel();
self.inner
.pending
@@ -182,10 +207,10 @@ impl PluginWorker {
}
let _ = sender.send(WorkerStreamItem::Result(Err(Error::Cancelled)));
inner.pending.lock().await.remove(&request_id);
inner.host.cancellations.lock().await.remove(&request_id);
inner.host.invocations.lock().await.remove(&request_id);
}
_ = sender.closed() => {
inner.host.cancellations.lock().await.remove(&request_id);
inner.host.invocations.lock().await.remove(&request_id);
}
}
});
@@ -201,7 +226,7 @@ impl PluginWorker {
async fn cleanup(&self, id: &str) {
self.inner.pending.lock().await.remove(id);
self.inner.host.cancellations.lock().await.remove(id);
self.inner.host.invocations.lock().await.remove(id);
}
async fn stdin(&self) -> Result<Arc<Mutex<ChildStdin>>> {
@@ -393,6 +418,28 @@ async fn fail_pending(pending: &Pending, message: &str) {
}
}
fn recorded_network_request(
params: &serde_json::Value,
) -> Result<(serde_json::Value, serde_json::Value)> {
let mut recorded_headers = serde_json::Map::new();
if let Some(headers) = params.get("headers").and_then(serde_json::Value::as_object) {
for (name, value) in headers {
let value = value.as_str().ok_or_else(|| {
Error::Config(format!("plugin HTTP header '{name}' must be a string"))
})?;
if !crate::model::is_sensitive_header(name) {
recorded_headers.insert(name.clone(), value.into());
}
}
}
let body = params
.get("body")
.and_then(serde_json::Value::as_str)
.map(|body| serde_json::from_str(body).unwrap_or_else(|_| body.into()))
.unwrap_or(serde_json::Value::Null);
Ok((serde_json::Value::Object(recorded_headers), body))
}
impl HostContext {
async fn call(
&self,
@@ -421,7 +468,11 @@ impl HostContext {
&self,
request_id: &str,
params: &serde_json::Value,
) -> Result<(reqwest::RequestBuilder, CancellationToken)> {
) -> Result<(
reqwest::RequestBuilder,
CancellationToken,
Option<CallRecorder>,
)> {
let raw_url = required_string(params, "url")?;
let url = url::Url::parse(raw_url)
.map_err(|error| Error::Config(format!("invalid plugin network URL: {error}")))?;
@@ -463,14 +514,17 @@ impl HostContext {
if let Some(body) = params.get("body").and_then(serde_json::Value::as_str) {
request = request.body(body.to_owned());
}
let cancellation = self
.cancellations
.lock()
.await
.get(request_id)
.cloned()
let invocation = self.invocations.lock().await.get(request_id).cloned();
let cancellation = invocation
.as_ref()
.map(|state| state.cancellation.clone())
.unwrap_or_default();
Ok((request, cancellation))
let recorder = invocation.and_then(|state| state.claim_recorder());
if let Some(recorder) = &recorder {
let (headers, body) = recorded_network_request(params)?;
recorder.request(headers, &body).await?;
}
Ok((request, cancellation, recorder))
}
async fn fetch(
@@ -478,13 +532,16 @@ impl HostContext {
request_id: &str,
params: serde_json::Value,
) -> Result<serde_json::Value> {
let (request, cancellation) = self.request(request_id, &params).await?;
let (request, cancellation, recorder) = self.request(request_id, &params).await?;
let request = request.timeout(Duration::from_secs(60));
let response = tokio::select! {
_ = cancellation.cancelled() => return Err(Error::Cancelled),
response = request.send() => response?,
};
let status = response.status().as_u16();
if let Some(recorder) = &recorder {
recorder.response_headers(status).await?;
}
if response
.content_length()
.is_some_and(|size| size > MAX_NETWORK_RESPONSE_BYTES)
@@ -503,6 +560,9 @@ impl HostContext {
"plugin network response is larger than allowed".into(),
));
}
if let Some(recorder) = &recorder {
recorder.response_chunk(&body).await?;
}
Ok(
serde_json::json!({ "status": status, "headers": headers, "body": String::from_utf8_lossy(&body) }),
)
@@ -514,12 +574,15 @@ impl HostContext {
request_id: &str,
params: serde_json::Value,
) -> Result<serde_json::Value> {
let (request, cancellation) = self.request(request_id, &params).await?;
let (request, cancellation, recorder) = self.request(request_id, &params).await?;
let response = tokio::select! {
_ = cancellation.cancelled() => return Err(Error::Cancelled),
response = request.send() => response?,
};
let status = response.status().as_u16();
if let Some(recorder) = &recorder {
recorder.response_headers(status).await?;
}
let headers = header_map(&response);
let (sender, receiver) = mpsc::channel::<Result<String>>(256);
tokio::spawn(async move {
@@ -552,6 +615,12 @@ impl HostContext {
.await;
return;
}
if let Some(recorder) = &recorder {
if let Err(error) = recorder.response_chunk(&chunk).await {
let _ = sender.send(Err(error)).await;
return;
}
}
buffered.extend_from_slice(&chunk);
while let Some(position) = buffered.iter().position(|byte| *byte == b'\n') {
let mut line = buffered.drain(..=position).collect::<Vec<u8>>();
@@ -645,3 +714,179 @@ fn required_string<'a>(params: &'a serde_json::Value, key: &str) -> Result<&'a s
.and_then(serde_json::Value::as_str)
.ok_or_else(|| Error::Protocol(format!("plugin host call requires string '{key}'")))
}
#[cfg(test)]
mod tests {
use crate::{
model::{NewLlmCall, ProviderType},
provider::{CallRecorder, FinishReason},
store::Store,
};
use super::*;
async fn recorder(detailed: bool, call_id: &str) -> (tempfile::TempDir, Store, CallRecorder) {
let directory = tempfile::tempdir().unwrap();
let store = Store::connect(&format!(
"sqlite://{}",
directory.path().join("test.db").display()
))
.await
.unwrap();
store.set_detailed_logging(detailed).await.unwrap();
let recorder = CallRecorder::start(
store.clone(),
NewLlmCall {
call_id: call_id.into(),
run_id: "run".into(),
conversation_id: "conversation".into(),
provider_call_index: 0,
model_hash: "plugin:test/provider/model".into(),
provider_type: ProviderType::Plugin,
provider_url: "plugin://test/provider".into(),
request_type: ProviderType::Plugin,
request_url: "plugin://test/provider".into(),
model_id: "model".into(),
display_name: "Model".into(),
reasoning_effort: None,
fast: false,
message_count: 1,
tool_count: 0,
detailed: false,
},
)
.await
.unwrap();
(directory, store, recorder)
}
fn network_params() -> serde_json::Value {
serde_json::json!({
"url": "https://example.com/v1/responses",
"method": "POST",
"headers": {
"Authorization": "Bearer secret",
"X-Api-Key": "secret-key",
"Cookie": "session=secret",
"content-type": "application/json",
"x-client-request-id": "request-1"
},
"body": "{\"model\":\"test\",\"stream\":true}"
})
}
async fn host_with_recorder(store: Store, recorder: CallRecorder) -> HostContext {
let invocations = Arc::new(Mutex::new(HashMap::new()));
invocations.lock().await.insert(
"invocation".into(),
Arc::new(InvocationState {
cancellation: CancellationToken::new(),
recorder: Some(recorder),
recorder_claimed: AtomicBool::new(false),
}),
);
HostContext {
plugin_id: "test".into(),
network_hosts: Arc::new(HashSet::from(["example.com".into()])),
store,
invocations,
streams: Arc::new(Mutex::new(HashMap::new())),
}
}
#[test]
fn recorded_plugin_request_omits_sensitive_headers() {
let (headers, body) = recorded_network_request(&network_params()).unwrap();
assert_eq!(
headers,
serde_json::json!({
"content-type": "application/json",
"x-client-request-id": "request-1"
})
);
assert_eq!(body, serde_json::json!({ "model": "test", "stream": true }));
}
#[tokio::test]
async fn detailed_plugin_network_recording_persists_request_and_raw_response() {
let (_directory, store, recorder) = recorder(true, "detailed-plugin").await;
let host = host_with_recorder(store.clone(), recorder.clone()).await;
let params = network_params();
let (_, _, first_recorder) = host.request("invocation", &params).await.unwrap();
let (_, _, second_recorder) = host.request("invocation", &params).await.unwrap();
let (_, body) = recorded_network_request(&params).unwrap();
assert!(first_recorder.is_some());
assert!(second_recorder.is_none());
recorder.response_headers(200).await.unwrap();
recorder
.response_chunk(b"data: {\"type\":\"response.created\"}\n\n")
.await
.unwrap();
recorder.response_chunk(b"data: [DONE]\n\n").await.unwrap();
recorder.completed(FinishReason::Stop).await.unwrap();
let request = store
.llm_call_request("detailed-plugin")
.await
.unwrap()
.unwrap();
assert_eq!(
request.headers,
serde_json::json!({
"content-type": "application/json",
"x-client-request-id": "request-1"
})
);
assert_eq!(request.body, body);
let chunks = store.llm_call_chunks("detailed-plugin").await.unwrap();
let expected_response = "data: {\"type\":\"response.created\"}\n\ndata: [DONE]\n\n";
assert_eq!(chunks.len(), 2);
assert_eq!(
chunks
.iter()
.map(|chunk| chunk.data.as_str())
.collect::<String>(),
expected_response
);
let summary = store.llm_call("detailed-plugin").await.unwrap().unwrap();
assert_eq!(summary.http_status, Some(200));
assert_eq!(summary.stream_event_count, 2);
assert_eq!(summary.response_bytes, expected_response.len() as i64);
assert!(summary.detailed);
}
#[tokio::test]
async fn standard_plugin_network_recording_keeps_metrics_without_payloads() {
let (_directory, store, recorder) = recorder(false, "standard-plugin").await;
let host = host_with_recorder(store.clone(), recorder.clone()).await;
let params = network_params();
let (_, body) = recorded_network_request(&params).unwrap();
let request_bytes = serde_json::to_string(&body).unwrap().len() as i64;
let response = b"data: [DONE]\n\n";
let (_, _, observed) = host.request("invocation", &params).await.unwrap();
assert!(observed.is_some());
recorder.response_headers(204).await.unwrap();
recorder.response_chunk(response).await.unwrap();
recorder.completed(FinishReason::Stop).await.unwrap();
assert!(store
.llm_call_request("standard-plugin")
.await
.unwrap()
.is_none());
assert!(store
.llm_call_chunks("standard-plugin")
.await
.unwrap()
.is_empty());
let summary = store.llm_call("standard-plugin").await.unwrap().unwrap();
assert_eq!(summary.http_status, Some(204));
assert_eq!(summary.request_bytes, Some(request_bytes));
assert_eq!(summary.response_bytes, response.len() as i64);
assert_eq!(summary.stream_event_count, 1);
assert!(!summary.detailed);
}
}
+69 -10
View File
@@ -12,9 +12,9 @@ use crate::{
};
use super::{
attempt::{send_once, Attempt},
map_sse_error, merge_extra_params, provider_event_error,
recorder::recorded_headers,
retry::{send_with_retry, Attempt, RetryPolicy},
CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream,
};
@@ -90,17 +90,14 @@ impl Provider for AnthropicProvider {
if let Some(recorder) = &recorder {
recorder.request(request_headers.clone(), &body).await?;
}
let attempt = send_with_retry(
let attempt = send_once(
"Anthropic",
|| client.post(&config.request_url)
.header("x-api-key", &config.api_key).header("anthropic-version", "2023-06-01")
.headers(config.custom_headers.clone())
.json(&body),
RetryPolicy { retries: config.retry_count, ..RetryPolicy::default() },
&cancellation,
recorder.as_ref(),
request_headers,
&body,
).await?;
let Attempt::Response(response) = attempt else { return };
yield ModelEvent::Start { model_call_id: call_id };
@@ -321,10 +318,19 @@ fn apply_model(body: &mut Value, model: &crate::model::ModelSpec) -> Result<()>
fn merge_usage(total: &mut Usage, update: Usage) {
merge_usage_field(&mut total.input_tokens, update.input_tokens);
merge_usage_field(&mut total.context_input_tokens, update.context_input_tokens);
merge_usage_field(&mut total.output_tokens, update.output_tokens);
merge_usage_field(&mut total.total_tokens, update.total_tokens);
merge_usage_field(&mut total.cache_read_tokens, update.cache_read_tokens);
merge_usage_field(&mut total.cache_write_tokens, update.cache_write_tokens);
merge_usage_field(&mut total.reasoning_tokens, update.reasoning_tokens);
if let Some(request_tokens) = total
.context_input_tokens
.zip(total.output_tokens)
.and_then(|(input, output)| input.checked_add(output))
{
total.total_tokens = Some(request_tokens);
}
}
fn merge_usage_field(total: &mut Option<u64>, update: Option<u64>) {
@@ -475,14 +481,67 @@ fn required_u64(value: &Value, name: &str) -> Result<u64> {
}
fn anthropic_usage(value: &Value) -> Usage {
let input_tokens = value.get("input_tokens").and_then(Value::as_u64);
let cache_read_tokens = value.get("cache_read_input_tokens").and_then(Value::as_u64);
let cache_write_tokens = value
.get("cache_creation_input_tokens")
.and_then(Value::as_u64);
let context_input_tokens = input_tokens.map(|input| {
input
.saturating_add(cache_read_tokens.unwrap_or_default())
.saturating_add(cache_write_tokens.unwrap_or_default())
});
Usage {
input_tokens: value.get("input_tokens").and_then(Value::as_u64),
input_tokens,
context_input_tokens,
output_tokens: value.get("output_tokens").and_then(Value::as_u64),
total_tokens: value.get("total_tokens").and_then(Value::as_u64),
cache_read_tokens: value.get("cache_read_input_tokens").and_then(Value::as_u64),
cache_write_tokens: value
.get("cache_creation_input_tokens")
.and_then(Value::as_u64),
cache_read_tokens,
cache_write_tokens,
reasoning_tokens: None,
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn cached_tokens_are_included_once_in_anthropic_context_input() {
let usage = anthropic_usage(&serde_json::json!({
"input_tokens": 10,
"output_tokens": 5,
"cache_read_input_tokens": 20,
"cache_creation_input_tokens": 30
}));
assert_eq!(usage.input_tokens, Some(10));
assert_eq!(usage.context_input_tokens, Some(60));
assert_eq!(usage.output_tokens, Some(5));
}
#[test]
fn streamed_anthropic_usage_includes_cached_input_in_total() {
let mut usage = Usage::default();
merge_usage(
&mut usage,
anthropic_usage(&serde_json::json!({
"input_tokens": 414,
"output_tokens": 0,
"cache_read_input_tokens": 100_352,
"cache_creation_input_tokens": 0
})),
);
merge_usage(
&mut usage,
anthropic_usage(&serde_json::json!({"output_tokens": 191})),
);
assert_eq!(usage.input_tokens, Some(414));
assert_eq!(usage.cache_read_tokens, Some(100_352));
assert_eq!(usage.cache_write_tokens, Some(0));
assert_eq!(usage.context_input_tokens, Some(100_766));
assert_eq!(usage.output_tokens, Some(191));
assert_eq!(usage.total_tokens, Some(100_957));
}
}
+100
View File
@@ -0,0 +1,100 @@
//! Sends one Provider HTTP attempt without applying retry policy.
use tokio_util::sync::CancellationToken;
use crate::{Error, Result};
use super::CallRecorder;
#[derive(Debug)]
pub(crate) enum Attempt {
Response(reqwest::Response),
Cancelled,
}
pub(crate) async fn send_once<F>(
label: &str,
build: F,
cancellation: &CancellationToken,
recorder: Option<&CallRecorder>,
) -> Result<Attempt>
where
F: FnOnce() -> reqwest::RequestBuilder,
{
let response = tokio::select! {
_ = cancellation.cancelled() => return Ok(Attempt::Cancelled),
response = build().send() => response,
}?;
if let Some(recorder) = recorder {
recorder
.response_headers(response.status().as_u16())
.await?;
}
if response.status().is_success() {
return Ok(Attempt::Response(response));
}
let status = response.status();
let bytes = tokio::select! {
_ = cancellation.cancelled() => return Ok(Attempt::Cancelled),
bytes = response.bytes() => bytes,
}?;
Err(Error::Provider(format!(
"{label} {status}: {}",
String::from_utf8_lossy(&bytes)
)))
}
#[cfg(test)]
mod tests {
use super::*;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
async fn server(response: &'static [u8]) -> String {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut request = [0_u8; 1024];
let _ = socket.read(&mut request).await;
socket.write_all(response).await.unwrap();
});
format!("http://{address}")
}
#[tokio::test]
async fn non_success_status_is_one_failed_attempt() {
let url =
server(b"HTTP/1.1 503 Service Unavailable\r\nContent-Length: 4\r\n\r\ndown").await;
let client = reqwest::Client::new();
let error = send_once("test", || client.get(&url), &CancellationToken::new(), None)
.await
.unwrap_err();
assert!(
matches!(error, Error::Provider(message) if message.contains("503") && message.contains("down"))
);
}
#[tokio::test]
async fn response_body_transport_failure_is_one_failed_attempt() {
let url =
server(b"HTTP/1.1 500 Internal Server Error\r\nContent-Length: 100\r\n\r\nshort").await;
let client = reqwest::Client::new();
let error = send_once("test", || client.get(&url), &CancellationToken::new(), None)
.await
.unwrap_err();
assert!(matches!(error, Error::Http(_)));
}
#[tokio::test]
async fn request_transport_failure_is_one_failed_attempt() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
drop(listener);
let url = format!("http://{address}");
let client = reqwest::Client::new();
let error = send_once("test", || client.get(&url), &CancellationToken::new(), None)
.await
.unwrap_err();
assert!(matches!(error, Error::Http(_)));
}
}
+1 -1
View File
@@ -1,11 +1,11 @@
//! Defines the provider interface and exports provider implementations.
mod anthropic;
mod attempt;
mod event;
mod normalize;
mod openai_chat;
mod openai_responses;
mod recorder;
mod retry;
mod router;
use std::pin::Pin;
+7 -8
View File
@@ -17,10 +17,10 @@ use crate::{
};
use super::{
apply_body_allowlist, apply_openai_prompt_cache_key, map_sse_error, merge_extra_params,
provider_event_error,
apply_body_allowlist, apply_openai_prompt_cache_key,
attempt::{send_once, Attempt},
map_sse_error, merge_extra_params, provider_event_error,
recorder::recorded_headers,
retry::{send_with_retry, Attempt, RetryPolicy},
CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream,
};
@@ -92,15 +92,12 @@ impl Provider for OpenAiChatProvider {
if let Some(recorder) = &recorder {
recorder.request(request_headers.clone(), &body).await?;
}
let attempt = send_with_retry(
let attempt = send_once(
"OpenAI Chat",
|| client.post(&config.request_url)
.bearer_auth(&config.api_key).headers(config.custom_headers.clone()).json(&body),
RetryPolicy { retries: config.retry_count, ..RetryPolicy::default() },
&cancellation,
recorder.as_ref(),
request_headers,
&body,
).await?;
let Attempt::Response(response) = attempt else { return };
yield ModelEvent::Start { model_call_id: call_id };
@@ -421,8 +418,10 @@ fn merge_chat_fragment(target: &mut String, fragment: &str) {
}
pub(crate) fn openai_usage(value: &Value) -> Usage {
let input_tokens = value.get("prompt_tokens").and_then(Value::as_u64);
Usage {
input_tokens: value.get("prompt_tokens").and_then(Value::as_u64),
input_tokens,
context_input_tokens: input_tokens,
output_tokens: value.get("completion_tokens").and_then(Value::as_u64),
total_tokens: value.get("total_tokens").and_then(Value::as_u64),
cache_read_tokens: value

Some files were not shown because too many files have changed in this diff Show More