Compare commits

...
32 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
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
311 changed files with 6084 additions and 49160 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"
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",
"version": "0.1.6",
"lockfileVersion": 3,
"requires": true,
"packages": {
"": {
"name": "cursor-byok-desktop",
"version": "0.1.5",
"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",
"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"
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",
"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"
]
@@ -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 }));
}
};
+48 -20
View File
@@ -606,7 +606,7 @@
"refs": [
{
"file": "features/settings/AppLifecycleSettingsCard.tsx",
"line": 156,
"line": 158,
"column": 27
}
]
@@ -916,7 +916,7 @@
"refs": [
{
"file": "features/settings/AppLifecycleSettingsCard.tsx",
"line": 152,
"line": 154,
"column": 13
}
]
@@ -1673,7 +1673,7 @@
"refs": [
{
"file": "features/settings/AppLifecycleSettingsCard.tsx",
"line": 126,
"line": 128,
"column": 17
}
]
@@ -1725,7 +1725,7 @@
"refs": [
{
"file": "features/settings/AppLifecycleSettingsCard.tsx",
"line": 151,
"line": 153,
"column": 13
}
]
@@ -2353,7 +2353,7 @@
"refs": [
{
"file": "shared/api.ts",
"line": 488,
"line": 486,
"column": 43
}
]
@@ -2365,7 +2365,7 @@
"refs": [
{
"file": "features/settings/AppLifecycleSettingsCard.tsx",
"line": 149,
"line": 151,
"column": 18
}
]
@@ -2451,12 +2451,12 @@
"refs": [
{
"file": "features/settings/AppLifecycleSettingsCard.tsx",
"line": 137,
"line": 139,
"column": 18
},
{
"file": "features/settings/AppLifecycleSettingsCard.tsx",
"line": 143,
"line": 145,
"column": 16
}
]
@@ -2680,7 +2680,7 @@
"refs": [
{
"file": "shared/api.ts",
"line": 483,
"line": 481,
"column": 43
}
]
@@ -2795,7 +2795,7 @@
"refs": [
{
"file": "features/settings/AppLifecycleSettingsCard.tsx",
"line": 160,
"line": 162,
"column": 37
}
]
@@ -2834,7 +2834,7 @@
"refs": [
{
"file": "shared/api.ts",
"line": 420,
"line": 418,
"column": 21
}
]
@@ -2900,12 +2900,12 @@
"refs": [
{
"file": "features/settings/AppLifecycleSettingsCard.tsx",
"line": 113,
"line": 115,
"column": 18
},
{
"file": "features/settings/AppLifecycleSettingsCard.tsx",
"line": 119,
"line": 121,
"column": 16
}
]
@@ -3194,6 +3194,20 @@
}
]
},
"92e26b27d5ea8f0e": {
"source": "检查更新失败:{error}",
"kind": "template",
"placeholders": [
"error"
],
"refs": [
{
"file": "features/settings/AppLifecycleSettingsCard.tsx",
"line": 99,
"column": 15
}
]
},
"940a168911ade998": {
"source": "每页条数",
"kind": "text",
@@ -3520,7 +3534,7 @@
"refs": [
{
"file": "features/settings/AppLifecycleSettingsCard.tsx",
"line": 110,
"line": 112,
"column": 29
}
]
@@ -3532,12 +3546,12 @@
"refs": [
{
"file": "features/settings/AppLifecycleSettingsCard.tsx",
"line": 125,
"line": 127,
"column": 18
},
{
"file": "features/settings/AppLifecycleSettingsCard.tsx",
"line": 131,
"line": 133,
"column": 16
}
]
@@ -3749,6 +3763,20 @@
}
]
},
"a80b53f8848e6d27": {
"source": "安装更新失败:{error}",
"kind": "template",
"placeholders": [
"error"
],
"refs": [
{
"file": "features/settings/AppLifecycleSettingsCard.tsx",
"line": 108,
"column": 15
}
]
},
"a98585871c5313ff": {
"source": "显示名称",
"kind": "text",
@@ -3814,7 +3842,7 @@
"refs": [
{
"file": "features/settings/AppLifecycleSettingsCard.tsx",
"line": 156,
"line": 158,
"column": 39
}
]
@@ -3854,7 +3882,7 @@
"refs": [
{
"file": "features/settings/AppLifecycleSettingsCard.tsx",
"line": 138,
"line": 140,
"column": 17
}
]
@@ -4577,7 +4605,7 @@
"refs": [
{
"file": "features/settings/AppLifecycleSettingsCard.tsx",
"line": 114,
"line": 116,
"column": 17
}
]
@@ -5522,7 +5550,7 @@
},
{
"file": "features/settings/AppLifecycleSettingsCard.tsx",
"line": 160,
"line": 162,
"column": 25
}
]
+2
View File
@@ -221,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?",
@@ -259,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",
+2
View File
@@ -221,6 +221,7 @@
"91aaf184cfc17ffd": "数据概览",
"91af6e57e7453fbe": "添加账号",
"92156a483d4ba248": "仅删除请求、响应和追踪附件等详细内容,保留调用汇总、统计指标和配置。",
"92e26b27d5ea8f0e": "检查更新失败:{error}",
"940a168911ade998": "每页条数",
"945fb1c67eca8493": "正在安装插件运行时",
"946b3ffc02f026c0": "确定删除这个模型吗?",
@@ -259,6 +260,7 @@
"a748cc074f78de00": "查看详情",
"a7617f42f898b2bf": "使用完整请求地址",
"a8036485f9227f2c": "拖动排序",
"a80b53f8848e6d27": "安装更新失败:{error}",
"a98585871c5313ff": "显示名称",
"ab9084a640fbb864": "全不选",
"abecab6701177721": "已开启开机启动",
-2
View File
@@ -235,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 @@
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)?;
}
+89 -28
View File
@@ -19,16 +19,17 @@ use crate::{
connect,
proto::{agent::v1 as agent, aiserver::v1 as ai},
},
services::{
account, analytics, knowledge, 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())?;
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))
}
@@ -120,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
@@ -140,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)
}
}
}
@@ -156,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)
@@ -180,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"
@@ -195,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(
@@ -224,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)
}
+11 -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,
));
@@ -58,10 +60,16 @@ impl App {
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,
+2 -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,
};
@@ -493,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(),
@@ -553,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() {
+1 -3
View File
@@ -117,9 +117,7 @@ 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 mut request_context = context::hydrate(request, context_sync).await?;
if let Some(rules_dir) = local_rules_dir {
+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() {
+299 -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,
@@ -486,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;
}
};
@@ -519,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;
}
@@ -544,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;
}
@@ -604,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
@@ -625,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));
}
}
+6 -11
View File
@@ -580,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())
@@ -598,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()),
@@ -610,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),
@@ -633,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;
}
}
}
+76 -19
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,
@@ -99,8 +113,10 @@ impl TransportRegistry {
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,
@@ -119,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
}
@@ -132,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
@@ -152,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) {
@@ -206,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;
}
}
}
+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
+74 -9
View File
@@ -14,10 +14,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,
};
@@ -89,15 +89,12 @@ impl Provider for OpenAiResponsesProvider {
if let Some(recorder) = &recorder {
recorder.request(request_headers.clone(), &body).await?;
}
let attempt = send_with_retry(
let attempt = send_once(
"OpenAI Responses",
|| 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 };
@@ -423,7 +420,12 @@ fn responses_input(messages: &[ProjectedMessage]) -> Result<Vec<Value>> {
.ok_or_else(|| {
Error::Protocol("OpenAI Responses replay state is missing items".into())
})?;
input.extend(items.iter().cloned());
input.extend(
items
.iter()
.map(response_reasoning_input)
.collect::<Result<Vec<_>>>()?,
);
}
push_responses_text(&mut input, &message.role, text);
for call in calls {
@@ -440,6 +442,23 @@ fn responses_input(messages: &[ProjectedMessage]) -> Result<Vec<Value>> {
Ok(input)
}
fn response_reasoning_input(item: &Value) -> Result<Value> {
let source = item
.as_object()
.filter(|object| object.get("type").and_then(Value::as_str) == Some("reasoning"))
.ok_or_else(|| {
Error::Protocol("OpenAI Responses replay state contains a non-reasoning item".into())
})?;
let mut projected = Map::new();
projected.insert("type".into(), json!("reasoning"));
for field in ["id", "summary", "content", "encrypted_content"] {
if let Some(value) = source.get(field) {
projected.insert(field.into(), value.clone());
}
}
Ok(Value::Object(projected))
}
fn push_responses_parts(input: &mut Vec<Value>, role: &Role, parts: &[ContentPart]) -> Result<()> {
let text_type = if *role == Role::Assistant {
"output_text"
@@ -505,8 +524,10 @@ fn required_u64(value: &Value, name: &str) -> Result<u64> {
}
fn responses_usage(value: &Value) -> Usage {
let input_tokens = value.get("input_tokens").and_then(Value::as_u64);
Usage {
input_tokens: value.get("input_tokens").and_then(Value::as_u64),
input_tokens,
context_input_tokens: 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
@@ -518,3 +539,47 @@ fn responses_usage(value: &Value) -> Usage {
.and_then(Value::as_u64),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::model::ProviderReplayState;
#[test]
fn reasoning_replay_projects_response_items_to_valid_input_items() {
let messages = [ProjectedMessage {
message_id: "assistant-1".into(),
role: Role::Assistant,
content: ProjectedContent::Assistant {
text: String::new(),
thinking: String::new(),
replay_state: Some(ProviderReplayState {
provider_kind: "openai_responses".into(),
value: json!({
"items": [{
"type": "reasoning",
"id": "item-1",
"status": "completed",
"summary": [{"type": "summary_text", "text": "why"}],
"content": [],
"encrypted_content": "opaque",
"output_only": true
}]
}),
}),
calls: Vec::new(),
},
}];
assert_eq!(
responses_input(&messages).unwrap(),
vec![json!({
"type": "reasoning",
"id": "item-1",
"summary": [{"type": "summary_text", "text": "why"}],
"content": [],
"encrypted_content": "opaque"
})]
);
}
}
+1 -31
View File
@@ -1,7 +1,7 @@
//! Records provider requests, responses, usage, and timing.
use std::{
sync::{
atomic::{AtomicBool, AtomicI64, AtomicU32, AtomicU64, Ordering},
atomic::{AtomicBool, AtomicI64, AtomicU64, Ordering},
Arc,
},
time::Instant,
@@ -71,7 +71,6 @@ struct Inner {
base_call: NewLlmCall,
detailed: bool,
attempt: Mutex<AttemptState>,
next_attempt: AtomicU32,
next_generation: AtomicU64,
finished: AtomicBool,
}
@@ -120,7 +119,6 @@ impl CallRecorder {
base_call: call.clone(),
detailed: call.detailed,
attempt: Mutex::new(AttemptState::new(call.call_id.clone())),
next_attempt: AtomicU32::new(0),
next_generation: AtomicU64::new(0),
finished: AtomicBool::new(false),
}),
@@ -284,34 +282,6 @@ impl CallRecorder {
self.finish("cancelled", None, None, None).await
}
pub async fn retry(
&self,
error: &crate::Error,
headers: serde_json::Value,
body: &serde_json::Value,
) -> Result<()> {
self.failed(error).await?;
let attempt_number = self.inner.next_attempt.fetch_add(1, Ordering::Relaxed) + 1;
let mut call = self.inner.base_call.clone();
call.call_id = format!("{}:retry-{attempt_number}", self.inner.base_call.call_id);
{
let mut attempt = self.inner.attempt.lock().await;
*attempt = AttemptState::new(call.call_id.clone());
self.inner.finished.store(false, Ordering::Release);
if let Err(error) = self.inner.store.start_llm_call(&call).await {
self.inner.finished.store(true, Ordering::Release);
return Err(error);
}
}
if let Err(error) = self.request(headers, body).await {
self.failed(&error).await?;
return Err(error);
}
Ok(())
}
async fn finish(
&self,
status: &str,
-84
View File
@@ -1,84 +0,0 @@
//! Applies provider retry and backoff behavior.
use std::time::Duration;
use tokio_util::sync::CancellationToken;
use crate::{Error, Result};
use super::CallRecorder;
#[derive(Clone, Copy, Debug)]
pub(crate) struct RetryPolicy {
pub retries: u32,
pub delay: Duration,
}
impl Default for RetryPolicy {
fn default() -> Self {
Self {
retries: 5,
delay: Duration::from_secs(5),
}
}
}
#[derive(Debug)]
pub(crate) enum Attempt {
Response(reqwest::Response),
Cancelled,
}
pub(crate) async fn send_with_retry<F>(
label: &str,
build: F,
policy: RetryPolicy,
cancellation: &CancellationToken,
recorder: Option<&CallRecorder>,
request_headers: serde_json::Value,
request_body: &serde_json::Value,
) -> Result<Attempt>
where
F: Fn() -> reqwest::RequestBuilder,
{
for attempt in 0..=policy.retries {
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 = response.bytes().await?;
let error = Error::Provider(format!(
"{label} {status}: {}",
String::from_utf8_lossy(&bytes)
));
if attempt == policy.retries {
return Err(error);
}
tracing::warn!(
provider = label,
status = status.as_u16(),
attempt = attempt + 1,
retries = policy.retries,
delay_ms = policy.delay.as_millis(),
"provider returned a non-success status, retrying"
);
if let Some(recorder) = recorder {
recorder
.retry(&error, request_headers.clone(), request_body)
.await?;
}
tokio::select! {
_ = cancellation.cancelled() => return Ok(Attempt::Cancelled),
_ = tokio::time::sleep(policy.delay) => {}
}
}
unreachable!("the retry loop returns on the final attempt")
}
+9 -9
View File
@@ -18,11 +18,10 @@ use super::{
OpenAiChatProvider, OpenAiResponsesProvider, Provider, ProviderStream,
};
const BUILTIN_PROVIDER_RETRIES: u32 = 5;
pub struct ProviderRouter {
store: Store,
plugins: PluginRegistry,
clients: crate::network::NetworkClients,
request_timeout: Duration,
stream_idle_timeout: Duration,
}
@@ -31,12 +30,14 @@ impl ProviderRouter {
pub fn new(
store: Store,
plugins: PluginRegistry,
clients: crate::network::NetworkClients,
request_timeout: Duration,
stream_idle_timeout: Duration,
) -> Self {
Self {
store,
plugins,
clients,
request_timeout,
stream_idle_timeout,
}
@@ -51,6 +52,7 @@ impl Provider for ProviderRouter {
) -> ProviderStream {
let store = self.store.clone();
let plugins = self.plugins.clone();
let clients = self.clients.clone();
let request_timeout = self.request_timeout;
let stream_idle_timeout = self.stream_idle_timeout;
Box::pin(try_stream! {
@@ -64,17 +66,14 @@ impl Provider for ProviderRouter {
let plan = plugins.plan_model(&selected).await?;
let recorder = start_recorder(&store, &invocation, &selected, &plan.model.display_name, ProviderType::Plugin, &plan.request_url, &plan.model.model_id).await?;
let guard = recorder.cancel_on_drop();
recorder.request(serde_json::json!({}), &crate::plugin::plugin_llm_request(&invocation)?).await?;
let mut routed = invocation.clone();
routed.request.model.display_name = Some(plan.model.display_name.clone());
if let Some(tokens) = plan.model.context_window_tokens {
routed.request.model.context_window_tokens.get_or_insert(tokens);
}
if let Some(tokens) = plan.model.max_output_tokens {
routed.request.model.max_output_tokens.get_or_insert(tokens);
}
let provider: Arc<dyn Provider> = Arc::new(NormalizedProvider::new(Arc::new(PluginModelProvider {
registry: plugins.clone(),
recorder: recorder.clone(),
})));
(recorder, guard, provider.stream(routed, cancellation.clone()))
} else {
@@ -94,10 +93,9 @@ impl Provider for ProviderRouter {
custom_headers: if model.custom_headers_enabled { custom_headers(&model.custom_headers)? } else { reqwest::header::HeaderMap::new() },
max_output_tokens: model.max_output_tokens(),
request_timeout,
retry_count: BUILTIN_PROVIDER_RETRIES,
allowed_body_fields: None,
};
let client = crate::network::client_builder(&store).await?.timeout(request_timeout).build()?;
let client = clients.provider_client(request_timeout).await?;
let provider = build_observed(&config, recorder.clone(), client)?;
(recorder, guard, provider.stream(routed, cancellation.clone()))
};
@@ -233,6 +231,7 @@ async fn finish_stream(recorder: &CallRecorder, cancellation: &CancellationToken
/// 插件模型的 Provider 实现;对路由与规范化层完全等同于内置 Provider。
struct PluginModelProvider {
registry: PluginRegistry,
recorder: CallRecorder,
}
impl Provider for PluginModelProvider {
@@ -241,7 +240,8 @@ impl Provider for PluginModelProvider {
invocation: ModelInvocation,
cancellation: CancellationToken,
) -> ProviderStream {
self.registry.stream_model(invocation, cancellation)
self.registry
.stream_model(invocation, cancellation, self.recorder.clone())
}
}
+176 -82
View File
@@ -1,67 +1,78 @@
//! Decides when to compact context and builds a stable fallback summary.
//! Decides when to compact provider-visible context and builds a stable fallback summary.
use std::collections::HashSet;
use crate::model::{
CanonicalMessage, LlmCallUsageAnchor, PreparedRun, ProjectedMessage, RunAction,
use crate::{
model::{
estimate_context_tokens, estimate_projected_messages_tokens, CanonicalMessage, PreparedRun,
ProjectedMessage,
},
store::ContextUsageAnchor,
};
const FALLBACK_CHARS: usize = 12_000;
pub(super) const RESERVE_TOKENS: u64 = 10_000;
pub(super) const OUTPUT_TOKENS: u64 = 4_096;
pub(super) const INSTRUCTIONS: &str = "Summarize the conversation for the next model turn. Preserve goals, constraints, decisions, files, commands, errors, results, and unfinished work. Do not call tools. Return only the concise durable summary.";
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(super) struct ContextUsageAnchor {
input_tokens: u64,
message_count: usize,
tool_count: usize,
pub(super) fn input_budget(prepared: &PreparedRun) -> Option<u64> {
prepared
.model
.context_window_tokens
.map(|window| window.saturating_sub(RESERVE_TOKENS))
}
impl ContextUsageAnchor {
pub(super) fn from_llm_call(anchor: LlmCallUsageAnchor) -> Option<Self> {
Some(Self {
input_tokens: anchor.usage.context_input_tokens(anchor.request_type)?,
message_count: anchor.message_count,
tool_count: anchor.tool_count,
pub(super) fn estimated_tokens(
prepared: &PreparedRun,
projected_messages: &[ProjectedMessage],
anchor: Option<ContextUsageAnchor>,
) -> u64 {
anchor
.filter(|anchor| anchor.message_count <= projected_messages.len())
.map(|anchor| {
anchor
.context_input_tokens
.saturating_add(estimate_projected_messages_tokens(
&projected_messages[anchor.message_count..],
))
})
}
.unwrap_or_else(|| estimate_context_tokens(&prepared.prompt, projected_messages))
}
pub(super) fn compaction_estimate(
prepared: &PreparedRun,
projected_messages: &[ProjectedMessage],
anchor: Option<ContextUsageAnchor>,
) -> Option<u64> {
let budget = input_budget(prepared)?;
let estimated = estimated_tokens(prepared, projected_messages, anchor);
(estimated > budget).then_some(estimated)
}
#[cfg(test)]
pub(super) fn should_compact(
prepared: &PreparedRun,
messages: &[CanonicalMessage],
projected_messages: &[ProjectedMessage],
anchor: Option<ContextUsageAnchor>,
) -> bool {
if prepared.action != RunAction::Start {
return false;
}
let Some(context_window) = prepared.model.context_window_tokens else {
return false;
compaction_estimate(prepared, projected_messages, anchor).is_some()
}
pub(super) fn validate_compacted(
prepared: &PreparedRun,
projected_messages: &[ProjectedMessage],
) -> std::result::Result<u64, String> {
let estimated = estimate_context_tokens(&prepared.prompt, projected_messages);
let Some(budget) = input_budget(prepared) else {
return Ok(estimated);
};
if context_window == 0 || messages.len() <= prepared.initial_messages.len() {
return false;
if estimated <= budget {
return Ok(estimated);
}
let estimated_input = anchor
.filter(|anchor| {
anchor.message_count <= projected_messages.len()
&& anchor.tool_count == prepared.prompt.tools.len()
})
.map(|anchor| {
anchor
.input_tokens
.saturating_add(estimate_serialized_tokens(
&serde_json::to_string(&projected_messages[anchor.message_count..])
.unwrap_or_default(),
))
})
.unwrap_or_else(|| {
estimate_serialized_tokens(
&serde_json::to_string(&(&prepared.prompt, messages)).unwrap_or_default(),
)
});
estimated_input > context_window
Err(format!(
"context overflow after compaction: estimated input {estimated} tokens exceeds budget {budget} tokens"
))
}
pub(super) fn partition(
@@ -100,28 +111,18 @@ pub(super) fn fallback_summary(messages: &[CanonicalMessage]) -> String {
)
}
fn estimate_serialized_tokens(serialized: &str) -> u64 {
serialized
.chars()
.fold(0_u64, |units, character| {
units.saturating_add(if character.is_ascii() { 273 } else { 550 })
})
.div_ceil(1_000)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::model::{
project_messages, CheckpointId, ConversationId, ModelSpec, Origin, PromptSpec, Role, RunId,
RunKind,
project_messages, CheckpointId, ConversationId, ModelSpec, Origin, PromptSpec, Role,
RunAction, RunId, RunKind,
};
#[test]
fn automatic_compaction_starts_only_after_the_context_window_is_exceeded() {
fn prepared(context_window_tokens: u64) -> PreparedRun {
let mut model = ModelSpec::new("model");
model.context_window_tokens = Some(200_000);
let prepared = PreparedRun {
model.context_window_tokens = Some(context_window_tokens);
PreparedRun {
run_id: RunId::new("run"),
cursor_request_id: None,
conversation_id: ConversationId::new("conversation"),
@@ -134,40 +135,133 @@ mod tests {
initial_messages: Vec::new(),
action: RunAction::Start,
base_checkpoint_id: CheckpointId(1),
};
}
}
#[test]
fn automatic_compaction_uses_fixed_reserve_for_every_action() {
let messages = vec![CanonicalMessage::text(
"user",
Role::User,
Origin::Runtime,
"hello",
"x".repeat(40_000),
)];
let projected = project_messages(&messages).unwrap();
let tail_tokens = estimate_serialized_tokens(&serde_json::to_string(&projected).unwrap());
let anchor = |estimated_input| {
Some(ContextUsageAnchor {
input_tokens: estimated_input - tail_tokens,
message_count: 0,
tool_count: 0,
})
};
let estimated = estimate_context_tokens(&prepared(1).prompt, &projected);
let mut prepared = prepared(estimated + RESERVE_TOKENS);
assert!(!should_compact(&prepared, &projected, None));
prepared.model.context_window_tokens = Some(estimated + RESERVE_TOKENS - 1);
assert!(should_compact(&prepared, &projected, None));
prepared.action = RunAction::Resume {
pending_tool_round: None,
};
assert!(should_compact(&prepared, &projected, None));
}
#[test]
fn provider_usage_anchor_only_estimates_messages_added_after_last_request() {
let messages = vec![
CanonicalMessage::text("old", Role::User, Origin::Runtime, "x".repeat(400_000)),
CanonicalMessage::text("new", Role::User, Origin::Runtime, "short follow-up"),
];
let projected = project_messages(&messages).unwrap();
let anchor = ContextUsageAnchor {
context_input_tokens: 103_904,
message_count: 1,
};
let expected = 103_904 + estimate_projected_messages_tokens(&projected[1..]);
assert_eq!(
estimated_tokens(&prepared(200_000), &projected, Some(anchor)),
expected
);
assert!(!should_compact(
&prepared,
&messages,
&prepared(200_000),
&projected,
anchor(199_999)
));
assert!(!should_compact(
&prepared,
&messages,
&projected,
anchor(200_000)
));
assert!(should_compact(
&prepared,
&messages,
&projected,
anchor(200_001)
Some(anchor)
));
}
#[test]
fn provider_usage_anchor_triggers_after_new_messages_cross_budget() {
let messages = vec![
CanonicalMessage::text("old", Role::User, Origin::Runtime, "old"),
CanonicalMessage::text("new", Role::User, Origin::Runtime, "x".repeat(80_000)),
];
let projected = project_messages(&messages).unwrap();
assert!(should_compact(
&prepared(200_000),
&projected,
Some(ContextUsageAnchor {
context_input_tokens: 180_000,
message_count: 1,
})
));
}
#[test]
fn missing_anchor_uses_full_fallback() {
let messages = vec![CanonicalMessage::text(
"user",
Role::User,
Origin::Runtime,
"x".repeat(40_000),
)];
let projected = project_messages(&messages).unwrap();
let prepared = prepared(200_000);
assert_eq!(
estimated_tokens(&prepared, &projected, None),
estimate_context_tokens(&prepared.prompt, &projected)
);
}
#[test]
fn invalid_anchor_message_count_uses_full_fallback() {
let messages = vec![CanonicalMessage::text(
"user",
Role::User,
Origin::Runtime,
"x".repeat(40_000),
)];
let projected = project_messages(&messages).unwrap();
let expected = estimate_context_tokens(&prepared(200_000).prompt, &projected);
assert_eq!(
estimated_tokens(
&prepared(200_000),
&projected,
Some(ContextUsageAnchor {
context_input_tokens: 1,
message_count: 2,
})
),
expected
);
}
#[test]
fn compacted_history_is_validated_against_the_same_budget() {
let messages = vec![CanonicalMessage::text(
"user",
Role::User,
Origin::Runtime,
"x".repeat(40_000),
)];
let projected = project_messages(&messages).unwrap();
let estimated = estimate_context_tokens(&prepared(1).prompt, &projected);
assert_eq!(
validate_compacted(&prepared(estimated + RESERVE_TOKENS), &projected),
Ok(estimated)
);
assert!(
validate_compacted(&prepared(estimated + RESERVE_TOKENS - 1), &projected)
.unwrap_err()
.contains("context overflow after compaction")
);
}
}
+290 -132
View File
@@ -10,12 +10,14 @@ use crate::{
ToolRoundId, Usage,
},
provider::Provider,
store::{RunStatus, Store},
store::{ContextUsageAnchor, RunStatus, Store},
};
use super::{
consume_model_cycle, CommitBarrier, CommitCause, MessagesCommitted, ModelCycleFailure,
RunCommand, RunEvent, RunFailure, RunOutcome, RunPort,
consume_model_cycle,
model_retry::{should_retry, MODEL_RETRY_DELAY},
CommitBarrier, CommitCause, MessagesCommitted, RunCommand, RunEvent, RunFailure, RunOutcome,
RunPort,
};
pub struct RunEngine {
@@ -90,6 +92,14 @@ impl RunEngine {
cancellation: &CancellationToken,
) -> (RunOutcome, Option<Usage>) {
let mut usage = None;
let mut context_usage_anchor = match self
.store
.latest_context_usage(prepared.conversation_id.as_str())
.await
{
Ok(anchor) => anchor,
Err(error) => return (RunOutcome::Failed(error.into()), usage),
};
tracing::info!(
checkpoint_id = checkpoint.0,
"Run claimed conversation ownership"
@@ -161,7 +171,6 @@ impl RunEngine {
};
}
let mut auto_compacted = prepared.action == RunAction::Compact;
'model: loop {
if cancellation.is_cancelled() {
return (RunOutcome::Cancelled, usage);
@@ -170,37 +179,32 @@ impl RunEngine {
Ok(messages) => messages,
Err(error) => return (RunOutcome::Failed(error.into()), usage),
};
let context_anchor = if !auto_compacted && prepared.action == RunAction::Start {
match self
.store
.latest_llm_call_usage_anchor(
&prepared.conversation_id,
&prepared.model.model_id,
)
.await
{
Ok(anchor) => {
anchor.and_then(super::compaction::ContextUsageAnchor::from_llm_call)
}
Err(error) => return (RunOutcome::Failed(error.into()), usage),
}
} else {
None
};
let history = match crate::model::project_messages(&messages) {
Ok(history) => history,
Err(error) => return (RunOutcome::Failed(error.into()), usage),
};
if !auto_compacted
&& super::compaction::should_compact(prepared, &messages, &history, context_anchor)
{
auto_compacted = true;
let compaction_estimate = (prepared.action != RunAction::Compact)
.then(|| {
super::compaction::compaction_estimate(prepared, &history, context_usage_anchor)
})
.flatten();
if let Some(estimated_tokens) = compaction_estimate {
if emit(
client,
RunEvent::UsageSnapshot(context_usage_snapshot(estimated_tokens)),
)
.await
.is_err()
{
return (client_failure(), usage);
}
match self
.auto_compact(prepared, checkpoint, &messages, client, cancellation)
.await
{
Ok((next_checkpoint, compaction_usage)) => {
checkpoint = next_checkpoint;
context_usage_anchor = None;
if let Some(compaction_usage) = compaction_usage {
accumulate_usage(&mut usage, compaction_usage);
}
@@ -227,122 +231,236 @@ impl RunEngine {
model: prepared.model.clone(),
history,
};
let invocation = crate::model::ModelInvocation {
call_id: format!("{}:{provider_call_index}", prepared.run_id),
run_id: prepared.run_id.to_string(),
conversation_id: prepared.conversation_id.to_string(),
provider_call_index,
request,
};
let cycle_cancellation = cancellation.child_token();
let cycle_events = client.events.clone();
let cycle = consume_model_cycle(
self.provider.stream(invocation, cycle_cancellation.clone()),
&cycle_events,
&cycle_cancellation,
);
tokio::pin!(cycle);
let mut retries = 0_u32;
let mut pending_insertions = Vec::new();
let cycle = loop {
tokio::select! {
biased;
command = client.commands.recv() => {
let interruption = match command {
Some(RunCommand::InsertMessages(insertion)) => {
pending_insertions.push(insertion);
continue;
let cycle = 'attempt: loop {
let call_id = if retries == 0 {
format!("{}:{provider_call_index}", prepared.run_id)
} else {
format!("{}:{provider_call_index}:retry-{retries}", prepared.run_id)
};
let invocation = crate::model::ModelInvocation {
call_id,
run_id: prepared.run_id.to_string(),
conversation_id: prepared.conversation_id.to_string(),
provider_call_index,
request: request.clone(),
};
let cycle_cancellation = cancellation.child_token();
let cycle_events = client.events.clone();
let cycle = consume_model_cycle(
self.provider.stream(invocation, cycle_cancellation.clone()),
&cycle_events,
&cycle_cancellation,
);
tokio::pin!(cycle);
let cycle = loop {
tokio::select! {
biased;
command = client.commands.recv() => {
let interruption = match command {
Some(RunCommand::InsertMessages(insertion)) => {
pending_insertions.push(insertion);
continue;
}
Some(RunCommand::BreakMessages(messages)) => messages,
Some(RunCommand::Cancel) => {
cycle_cancellation.cancel();
let _ = cycle.await;
let _ = emit(client, RunEvent::CycleInterrupted).await;
return (RunOutcome::Cancelled, usage);
}
Some(RunCommand::ToolResult(_)) => {
cycle_cancellation.cancel();
let _ = cycle.await;
let _ = emit(client, RunEvent::CycleInterrupted).await;
return (
RunOutcome::Failed(RunFailure::Protocol(
"received a tool result while the model was running".into(),
)),
usage,
);
}
None => {
cycle_cancellation.cancel();
let _ = cycle.await;
let _ = emit(client, RunEvent::CycleInterrupted).await;
return (client_failure(), usage);
}
};
cycle_cancellation.cancel();
let interrupted = cycle.await;
match interrupted {
Ok(cycle) => {
if let Some(cycle_usage) = cycle.usage {
update_context_usage_anchor(
&mut context_usage_anchor,
cycle_usage,
request.history.len(),
);
accumulate_usage(&mut usage, cycle_usage);
}
}
Err(failure) => {
if let Some(cycle_usage) = failure.usage {
update_context_usage_anchor(
&mut context_usage_anchor,
cycle_usage,
request.history.len(),
);
accumulate_usage(&mut usage, cycle_usage);
}
}
}
Some(RunCommand::BreakMessages(messages)) => messages,
Some(RunCommand::Cancel) => {
cycle_cancellation.cancel();
let _ = cycle.await;
let _ = emit(client, RunEvent::CycleInterrupted).await;
return (RunOutcome::Cancelled, usage);
}
Some(RunCommand::ToolResult(_)) => {
cycle_cancellation.cancel();
let _ = cycle.await;
let _ = emit(client, RunEvent::CycleInterrupted).await;
return (
RunOutcome::Failed(RunFailure::Protocol(
"received a tool result while the model was running".into(),
)),
usage,
);
}
None => {
cycle_cancellation.cancel();
let _ = cycle.await;
let _ = emit(client, RunEvent::CycleInterrupted).await;
if emit(client, RunEvent::CycleInterrupted).await.is_err() {
return (client_failure(), usage);
}
};
cycle_cancellation.cancel();
let interrupted = cycle.await;
match interrupted {
Ok(cycle) => {
if let Some(cycle_usage) = cycle.usage {
accumulate_usage(&mut usage, cycle_usage);
}
}
Err(failure) => {
if let Some(cycle_usage) = failure.usage {
accumulate_usage(&mut usage, cycle_usage);
}
}
checkpoint = match super::messages::append_batches(
&self.store,
prepared,
client,
cancellation,
checkpoint,
std::mem::take(&mut pending_insertions),
)
.await
{
Ok((checkpoint, _)) => checkpoint,
Err(outcome) => return (outcome, usage),
};
checkpoint = match super::messages::append_batches(
&self.store,
prepared,
client,
cancellation,
checkpoint,
vec![interruption],
)
.await
{
Ok((checkpoint, _)) => checkpoint,
Err(outcome) => return (outcome, usage),
};
continue 'model;
},
result = &mut cycle => break result,
}
};
match cycle {
Ok(cycle) => break 'attempt cycle,
Err(cycle_failure) => {
if let Some(cycle_usage) = cycle_failure.usage {
update_context_usage_anchor(
&mut context_usage_anchor,
cycle_usage,
request.history.len(),
);
accumulate_usage(&mut usage, cycle_usage);
}
if emit(client, RunEvent::CycleInterrupted).await.is_err() {
if cancellation.is_cancelled() {
let _ = emit(client, RunEvent::CycleInterrupted).await;
return (RunOutcome::Cancelled, usage);
}
if !should_retry(&cycle_failure, retries) {
return (RunOutcome::Failed(cycle_failure.failure), usage);
}
retries += 1;
let message = failure_message(&cycle_failure.failure);
tracing::warn!(
provider_call_index,
retries,
max_retries = super::model_retry::MAX_MODEL_RETRIES,
delay_ms = MODEL_RETRY_DELAY.as_millis() as u64,
%message,
checkpoint_id = checkpoint.0,
"model attempt failed; retrying from current checkpoint"
);
if emit(
client,
RunEvent::ModelAttemptFailed {
attempt: retries,
message,
},
)
.await
.is_err()
{
return (client_failure(), usage);
}
checkpoint = match super::messages::append_batches(
&self.store,
prepared,
client,
cancellation,
checkpoint,
std::mem::take(&mut pending_insertions),
)
.await
{
Ok((checkpoint, _)) => checkpoint,
Err(outcome) => return (outcome, usage),
};
checkpoint = match super::messages::append_batches(
&self.store,
prepared,
client,
cancellation,
checkpoint,
vec![interruption],
)
.await
{
Ok((checkpoint, _)) => checkpoint,
Err(outcome) => return (outcome, usage),
};
continue 'model;
},
result = &mut cycle => break result,
}
};
let cycle = match cycle {
Ok(cycle) => cycle,
Err(ModelCycleFailure {
failure,
usage: cycle_usage,
..
}) => {
if let Some(cycle_usage) = cycle_usage {
accumulate_usage(&mut usage, cycle_usage);
let delay = tokio::time::sleep(MODEL_RETRY_DELAY);
tokio::pin!(delay);
loop {
tokio::select! {
biased;
command = client.commands.recv() => {
let interruption = match command {
Some(RunCommand::InsertMessages(insertion)) => {
pending_insertions.push(insertion);
continue;
}
Some(RunCommand::BreakMessages(messages)) => messages,
Some(RunCommand::Cancel) => {
let _ = emit(client, RunEvent::CycleInterrupted).await;
return (RunOutcome::Cancelled, usage);
}
Some(RunCommand::ToolResult(_)) => {
return (
RunOutcome::Failed(RunFailure::Protocol(
"received a tool result while waiting to retry the model".into(),
)),
usage,
);
}
None => return (client_failure(), usage),
};
if emit(client, RunEvent::CycleInterrupted).await.is_err() {
return (client_failure(), usage);
}
checkpoint = match super::messages::append_batches(
&self.store,
prepared,
client,
cancellation,
checkpoint,
std::mem::take(&mut pending_insertions),
)
.await
{
Ok((checkpoint, _)) => checkpoint,
Err(outcome) => return (outcome, usage),
};
checkpoint = match super::messages::append_batches(
&self.store,
prepared,
client,
cancellation,
checkpoint,
vec![interruption],
)
.await
{
Ok((checkpoint, _)) => checkpoint,
Err(outcome) => return (outcome, usage),
};
continue 'model;
}
_ = cancellation.cancelled() => {
let _ = emit(client, RunEvent::CycleInterrupted).await;
return (RunOutcome::Cancelled, usage);
}
_ = &mut delay => break,
}
}
}
if cancellation.is_cancelled() {
let _ = emit(client, RunEvent::CycleInterrupted).await;
return (RunOutcome::Cancelled, usage);
}
return (RunOutcome::Failed(failure), usage);
}
};
if let Some(cycle_usage) = cycle.usage {
update_context_usage_anchor(
&mut context_usage_anchor,
cycle_usage,
request.history.len(),
);
accumulate_usage(&mut usage, cycle_usage);
}
@@ -583,7 +701,15 @@ impl RunEngine {
let (compactable, retained_request_context) =
super::compaction::partition(messages, &current_ids);
if compactable.is_empty() {
return Ok((checkpoint, None));
let projected = crate::model::project_messages(messages)
.map_err(|error| RunOutcome::Failed(error.into()))?;
let message = super::compaction::validate_compacted(prepared, &projected)
.err()
.unwrap_or_else(|| {
"context overflow after compaction: no conversation history can be compacted"
.into()
});
return Err(RunOutcome::Failed(RunFailure::Protocol(message)));
}
emit(client, RunEvent::AutoCompactionStarted)
@@ -689,7 +815,7 @@ impl RunEngine {
)
}
};
let event_id = format!("summary:auto:{}", prepared.run_id);
let event_id = format!("summary:auto:{}:{provider_call_index}", prepared.run_id);
let summary_message = CanonicalMessage {
message_id: format!("runtime:{event_id}"),
role: Role::User,
@@ -704,6 +830,10 @@ impl RunEngine {
let mut replacement = retained_request_context.into_iter().collect::<Vec<_>>();
replacement.push(summary_message);
replacement.extend(prepared.initial_messages.iter().cloned());
let projected_replacement = crate::model::project_messages(&replacement)
.map_err(|error| RunOutcome::Failed(error.into()))?;
super::compaction::validate_compacted(prepared, &projected_replacement)
.map_err(|message| RunOutcome::Failed(RunFailure::Protocol(message)))?;
let mut checkpoint = self
.store
.replace_checkpoint(
@@ -730,6 +860,9 @@ impl RunEngine {
emit(client, RunEvent::AutoCompactionCompleted)
.await
.map_err(|_| client_failure())?;
emit(client, RunEvent::UsageSnapshot(context_usage_snapshot(0)))
.await
.map_err(|_| client_failure())?;
checkpoint = super::messages::append_batches(
&self.store,
prepared,
@@ -790,6 +923,31 @@ async fn hydrate_tool_images(
Ok(())
}
fn context_usage_snapshot(tokens: u64) -> Usage {
Usage {
input_tokens: Some(tokens),
context_input_tokens: Some(tokens),
output_tokens: Some(0),
total_tokens: Some(tokens),
cache_read_tokens: Some(0),
cache_write_tokens: Some(0),
reasoning_tokens: Some(0),
}
}
fn update_context_usage_anchor(
anchor: &mut Option<ContextUsageAnchor>,
usage: Usage,
message_count: usize,
) {
if let Some(context_input_tokens) = usage.context_input_tokens {
*anchor = Some(ContextUsageAnchor {
context_input_tokens,
message_count,
});
}
}
fn accumulate_usage(total: &mut Option<Usage>, usage: Usage) {
match total {
Some(total) => *total += usage,
+5
View File
@@ -104,6 +104,10 @@ pub enum RunEvent {
AutoCompactionStarted,
AutoCompactionCompleted,
CycleInterrupted,
ModelAttemptFailed {
attempt: u32,
message: String,
},
TextStart,
TextDelta(String),
TextEnd,
@@ -125,6 +129,7 @@ pub enum RunEvent {
ToolCallEnd {
index: usize,
},
UsageSnapshot(Usage),
Usage(Usage),
ExecuteToolRound {
round_id: ToolRoundId,
+1
View File
@@ -7,6 +7,7 @@ mod event;
mod handle;
mod messages;
mod model_cycle;
mod model_retry;
mod port;
mod tool_round;
+130 -16
View File
@@ -6,7 +6,7 @@ use tokio::sync::mpsc;
use tokio_util::sync::CancellationToken;
use crate::{
model::{ProviderReplayState, ToolCall, Usage},
model::{normalize_tool_name, ProviderReplayState, ToolCall, Usage},
provider::{FinishReason, ModelEvent, ProviderStream},
};
@@ -29,6 +29,7 @@ pub struct ModelCycleFailure {
pub partial_text: String,
pub partial_reasoning: String,
pub usage: Option<Usage>,
pub retryable: bool,
}
struct OpenTool {
@@ -164,6 +165,7 @@ pub async fn consume_model_cycle(
call_id,
name,
} => {
let name = normalize_tool_name(&name);
let Some(model_call_id) = model_call_id.as_ref() else {
return Err(failure(
RunFailure::Protocol("provider emitted content before Start".into()),
@@ -186,6 +188,7 @@ pub async fn consume_model_cycle(
name: name.clone(),
arguments_text: String::new(),
arguments: serde_json::Value::Null,
argument_error: None,
},
ended: false,
});
@@ -218,13 +221,26 @@ pub async fn consume_model_cycle(
serde_json::from_str(&tool.call.arguments_text)
};
match arguments {
Ok(arguments) => {
Ok(arguments) if arguments.is_object() => {
tool.call.arguments = arguments;
tool.ended = true;
send(client, RunEvent::ToolCallEnd { index }).await
}
Err(_) => Err("provider ended a tool call with invalid JSON arguments"),
Ok(_) => {
tool.call.arguments = serde_json::json!({});
tool.call.argument_error = Some(format!(
"{} arguments must be a JSON object",
tool.call.name
));
}
Err(error) => {
tool.call.arguments = serde_json::json!({});
tool.call.argument_error = Some(format!(
"{} arguments are not valid JSON: {error}",
tool.call.name
));
}
}
tool.ended = true;
send(client, RunEvent::ToolCallEnd { index }).await
}
Some(_) => Err("provider emitted duplicate ToolCallEnd"),
None => Err("provider ended an unknown tool index"),
@@ -239,6 +255,13 @@ pub async fn consume_model_cycle(
ModelEvent::Usage(value) => {
if usage.replace(value).is_some() {
Err("provider emitted duplicate Usage")
} else if send(client, RunEvent::Usage(value)).await.is_err() {
return Err(failure(
RunFailure::Client("client event channel closed".into()),
text,
reasoning,
usage,
));
} else {
Ok(())
}
@@ -280,7 +303,7 @@ pub async fn consume_model_cycle(
.map(|tool| tool.call)
.collect::<Vec<_>>();
if finish_reason == FinishReason::Length {
return Err(failure(
return Err(terminal_failure(
RunFailure::Provider("model stopped before completing the response".into()),
text,
reasoning,
@@ -296,16 +319,6 @@ pub async fn consume_model_cycle(
usage,
));
}
if let Some(usage) = usage {
if send(client, RunEvent::Usage(usage)).await.is_err() {
return Err(failure(
RunFailure::Client("client event channel closed".into()),
text,
reasoning,
Some(usage),
));
}
}
let model_call_id = model_call_id.ok_or_else(|| {
failure(
RunFailure::Protocol("provider completed without Start".into()),
@@ -360,10 +373,111 @@ fn failure(
partial_reasoning: String,
usage: Option<Usage>,
) -> ModelCycleFailure {
let retryable = matches!(failure, RunFailure::Protocol(_) | RunFailure::Provider(_));
ModelCycleFailure {
failure,
partial_text,
partial_reasoning,
usage,
retryable,
}
}
fn terminal_failure(
failure: RunFailure,
partial_text: String,
partial_reasoning: String,
usage: Option<Usage>,
) -> ModelCycleFailure {
ModelCycleFailure {
failure,
partial_text,
partial_reasoning,
usage,
retryable: false,
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{
model::Usage,
provider::{FinishReason, ModelEvent},
};
use tokio_stream::wrappers::ReceiverStream;
#[tokio::test]
async fn provider_tool_names_are_normalized_when_received() {
let events = vec![
Ok(ModelEvent::Start {
model_call_id: "call".into(),
}),
Ok(ModelEvent::ToolCallStart {
index: 0,
call_id: "tool-call".into(),
name: "multi_tool_use.parallel".into(),
}),
Ok(ModelEvent::ToolCallEnd { index: 0 }),
Ok(ModelEvent::Done(FinishReason::ToolUse)),
];
let stream = Box::pin(tokio_stream::iter(events));
let (event_tx, mut event_rx) = tokio::sync::mpsc::channel(4);
let result = consume_model_cycle(stream, &event_tx, &CancellationToken::new())
.await
.unwrap();
assert_eq!(result.calls[0].name, "multi_tool_use_parallel");
assert!(matches!(
event_rx.recv().await,
Some(RunEvent::ToolCallStart { name, .. }) if name == "multi_tool_use_parallel"
));
}
#[tokio::test]
async fn usage_is_forwarded_before_the_provider_call_finishes() {
let (provider_tx, provider_rx) = tokio::sync::mpsc::channel(4);
let stream = Box::pin(ReceiverStream::new(provider_rx));
let (event_tx, mut event_rx) = tokio::sync::mpsc::channel(4);
let cancellation = CancellationToken::new();
let cycle_cancellation = cancellation.clone();
let cycle = tokio::spawn(async move {
consume_model_cycle(stream, &event_tx, &cycle_cancellation).await
});
let usage = Usage {
input_tokens: Some(100),
context_input_tokens: Some(100),
output_tokens: Some(20),
total_tokens: Some(120),
..Default::default()
};
provider_tx
.send(Ok(ModelEvent::Start {
model_call_id: "call".into(),
}))
.await
.unwrap();
provider_tx
.send(Ok(ModelEvent::Usage(usage)))
.await
.unwrap();
let event = tokio::time::timeout(std::time::Duration::from_secs(1), event_rx.recv())
.await
.unwrap()
.unwrap();
assert!(matches!(event, RunEvent::Usage(value) if value == usage));
assert!(!cycle.is_finished(), "usage must arrive before Done");
provider_tx
.send(Ok(ModelEvent::Done(FinishReason::Stop)))
.await
.unwrap();
drop(provider_tx);
let result = cycle.await.unwrap().unwrap();
assert_eq!(result.usage, Some(usage));
assert!(event_rx.try_recv().is_err(), "usage must be forwarded once");
}
}
+42
View File
@@ -0,0 +1,42 @@
//! Defines retry policy for one logical model call.
use std::time::Duration;
use super::ModelCycleFailure;
pub(super) const MAX_MODEL_RETRIES: u32 = 8;
pub(super) const MODEL_RETRY_DELAY: Duration = Duration::from_secs(5);
pub(super) fn should_retry(failure: &ModelCycleFailure, retries: u32) -> bool {
failure.retryable && retries < MAX_MODEL_RETRIES
}
#[cfg(test)]
mod tests {
use super::*;
use crate::run::{ModelCycleFailure, RunFailure};
fn failure(retryable: bool) -> ModelCycleFailure {
ModelCycleFailure {
failure: RunFailure::Provider("failed".into()),
partial_text: String::new(),
partial_reasoning: String::new(),
usage: None,
retryable,
}
}
#[test]
fn permits_eight_retries_after_the_initial_attempt() {
let retryable = failure(true);
for retries in 0..MAX_MODEL_RETRIES {
assert!(should_retry(&retryable, retries));
}
assert!(!should_retry(&retryable, MAX_MODEL_RETRIES));
}
#[test]
fn terminal_failures_never_retry() {
assert!(!should_retry(&failure(false), 0));
}
}
+34 -17
View File
@@ -62,6 +62,40 @@ impl Store {
.await?)
}
pub async fn append_cursor_trace_request(
&self,
request_id: &str,
artifact_type: &str,
source: &str,
data: &[u8],
metadata: &serde_json::Value,
) -> Result<()> {
let metadata_json = serde_json::to_string(metadata)?;
let blob_id = BlobId::digest(data);
let _write = self.writes.lock().await;
let mut tx = self.pool.begin_with("BEGIN IMMEDIATE").await?;
Self::put_blob_tx(&mut tx, &blob_id, data, &[]).await?;
Self::link_cursor_trace_artifact_tx(
&mut tx,
request_id,
artifact_type,
source,
&blob_id,
&metadata_json,
)
.await?;
sqlx::query(
"UPDATE cursor_run_traces
SET request_bytes = request_bytes + ? WHERE request_id = ?",
)
.bind(as_i64(data.len()))
.bind(request_id)
.execute(&mut *tx)
.await?;
tx.commit().await?;
Ok(())
}
pub async fn append_cursor_trace_artifact(
&self,
request_id: &str,
@@ -144,23 +178,6 @@ impl Store {
Ok(())
}
pub async fn add_cursor_trace_request_bytes(
&self,
request_id: &str,
bytes: usize,
) -> Result<()> {
let _write = self.writes.lock().await;
sqlx::query(
"UPDATE cursor_run_traces
SET request_bytes = request_bytes + ? WHERE request_id = ?",
)
.bind(as_i64(bytes))
.bind(request_id)
.execute(&self.pool)
.await?;
Ok(())
}
pub async fn start_cursor_trace_response(&self, request_id: &str, status: u16) -> Result<()> {
let now = now_ms();
let _write = self.writes.lock().await;
+100 -41
View File
@@ -1,18 +1,19 @@
//! Persists provider call payloads, timing, and usage.
use std::str::FromStr;
use sqlx::Row;
use crate::{
model::{
ConversationId, LlmCallRequest, LlmCallResponseChunk, LlmCallSummary, LlmCallUsageAnchor,
NewLlmCall, ProviderType, Usage,
},
model::{LlmCallRequest, LlmCallResponseChunk, LlmCallSummary, NewLlmCall, Usage},
Result,
};
use super::{now_ms, Store};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) struct ContextUsageAnchor {
pub(crate) context_input_tokens: u64,
pub(crate) message_count: usize,
}
#[derive(Clone, Debug)]
pub(crate) struct BufferedLlmChunk {
pub(crate) seq: i64,
@@ -271,6 +272,37 @@ impl Store {
Ok(())
}
pub(crate) async fn latest_context_usage(
&self,
conversation_id: &str,
) -> Result<Option<ContextUsageAnchor>> {
let row = sqlx::query(
"SELECT usage_json, message_count FROM llm_calls
WHERE conversation_id = ?
AND json_extract(usage_json, '$.context_input_tokens') IS NOT NULL
ORDER BY created_at_ms DESC, rowid DESC
LIMIT 1",
)
.bind(conversation_id)
.fetch_optional(&self.pool)
.await?;
let Some(row) = row else {
return Ok(None);
};
let usage: Usage = serde_json::from_str(row.try_get("usage_json")?)?;
let Some(context_input_tokens) = usage.context_input_tokens else {
return Ok(None);
};
let message_count = row.try_get::<i64, _>("message_count")?;
let Ok(message_count) = usize::try_from(message_count) else {
return Ok(None);
};
Ok(Some(ContextUsageAnchor {
context_input_tokens,
message_count,
}))
}
pub async fn llm_calls(&self, limit: i64) -> Result<Vec<LlmCallSummary>> {
let rows = sqlx::query("SELECT * FROM llm_calls ORDER BY created_at_ms DESC LIMIT ?")
.bind(limit.clamp(1, 500))
@@ -288,41 +320,6 @@ impl Store {
.transpose()
}
pub(crate) async fn latest_llm_call_usage_anchor(
&self,
conversation_id: &ConversationId,
model_hash: &str,
) -> Result<Option<LlmCallUsageAnchor>> {
let row = sqlx::query(
r#"SELECT request_type, usage_json, message_count, tool_count
FROM llm_calls
WHERE conversation_id = ?
AND model_hash = ?
AND status = 'completed'
AND input_tokens IS NOT NULL
AND usage_json IS NOT NULL
ORDER BY rowid DESC
LIMIT 1"#,
)
.bind(conversation_id.as_str())
.bind(model_hash)
.fetch_optional(&self.pool)
.await?;
row.map(|row| {
let message_count =
usize::try_from(row.try_get::<i64, _>("message_count")?).unwrap_or(usize::MAX);
let tool_count =
usize::try_from(row.try_get::<i64, _>("tool_count")?).unwrap_or(usize::MAX);
Ok(LlmCallUsageAnchor {
request_type: ProviderType::from_str(row.try_get("request_type")?)?,
usage: serde_json::from_str(row.try_get("usage_json")?)?,
message_count,
tool_count,
})
})
.transpose()
}
pub async fn llm_call_request(&self, call_id: &str) -> Result<Option<LlmCallRequest>> {
let row = sqlx::query(
"SELECT headers_json, body_json, byte_count FROM llm_call_requests WHERE call_id = ?",
@@ -416,6 +413,7 @@ fn summary_from_row(row: sqlx::sqlite::SqliteRow) -> Result<LlmCallSummary> {
#[cfg(test)]
mod tests {
use super::*;
use crate::model::ProviderType;
/// 插件模型不在 model_configs 中,调用记录必须照常落库并可按其稳定 ID 筛选。
#[tokio::test]
@@ -460,4 +458,65 @@ mod tests {
assert_eq!(overview.metrics.llm_calls, 1);
assert_eq!(overview.metrics.successful_calls, 1);
}
#[tokio::test]
async fn latest_context_usage_follows_conversation_chronology() {
let directory = tempfile::tempdir().unwrap();
let store = Store::connect(&format!(
"sqlite://{}",
directory.path().join("test.db").display()
))
.await
.unwrap();
for (call_id, model_id, context_input_tokens, message_count) in [
("call-a-1", "model-a", 100_u64, 3_usize),
("call-b", "model-b", 200_u64, 5_usize),
("call-a-2", "model-a", 300_u64, 7_usize),
] {
store
.start_llm_call(&NewLlmCall {
call_id: call_id.into(),
run_id: format!("run-{call_id}"),
conversation_id: "conversation".into(),
provider_call_index: 0,
model_hash: model_id.into(),
provider_type: ProviderType::Plugin,
provider_url: "plugin://test".into(),
request_type: ProviderType::Plugin,
request_url: "plugin://test".into(),
model_id: model_id.into(),
display_name: model_id.into(),
reasoning_effort: None,
fast: false,
message_count,
tool_count: 0,
detailed: false,
})
.await
.unwrap();
store
.record_llm_usage(
call_id,
Usage {
input_tokens: Some(context_input_tokens),
context_input_tokens: Some(context_input_tokens),
output_tokens: Some(10),
total_tokens: Some(context_input_tokens + 10),
..Default::default()
},
)
.await
.unwrap();
}
assert_eq!(
store.latest_context_usage("conversation").await.unwrap(),
Some(ContextUsageAnchor {
context_input_tokens: 300,
message_count: 7,
})
);
assert_eq!(store.latest_context_usage("other").await.unwrap(), None);
}
}
+12 -1
View File
@@ -445,9 +445,20 @@ mod tests {
.await
.unwrap();
let argument_error_column_exists: i64 = sqlx::query_scalar(
"SELECT EXISTS(
SELECT 1 FROM pragma_table_info('tool_round_calls')
WHERE name = 'argument_error'
)",
)
.fetch_one(&pool)
.await
.unwrap();
assert_eq!(checksum_after, checksum_before);
assert_eq!(versions, vec![1, 2, 3, 4, 5, 6, 7, 8]);
assert_eq!(versions, vec![1, 2, 3, 4, 5, 6, 7, 8, 9]);
assert_eq!(checkpoint_table_exists, 1);
assert_eq!(argument_error_column_exists, 1);
}
}

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