mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-05 12:13:05 +08:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
75b7ea9cc8 | ||
|
|
0fd5e9d6f2 | ||
|
|
ddaa61c827 | ||
|
|
919c1d8032 | ||
|
|
669f129dcd | ||
|
|
f22c7b6680 | ||
|
|
5cdf642dd1 | ||
|
|
2dad593263 | ||
|
|
76417e005b | ||
|
|
e768980dad | ||
|
|
2c63bd845a | ||
|
|
6e74637c69 | ||
|
|
d004139526 | ||
|
|
8c6c415a84 | ||
|
|
d83e14af9a | ||
|
|
29fde7d7c7 | ||
|
|
ee2592c469 | ||
|
|
49c1fb6378 | ||
|
|
788868f8b9 | ||
|
|
4c3fe230ce | ||
|
|
ac14245d19 | ||
|
|
84addec26a | ||
|
|
5de547041c | ||
|
|
8bd0d70add | ||
|
|
3ac4402a86 | ||
|
|
e535a98945 | ||
|
|
45e694fd63 | ||
|
|
9120b90be7 | ||
|
|
e673a034df | ||
|
|
6ac666da4f | ||
|
|
9e2d22418b | ||
|
|
97ee138de8 |
@@ -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/);
|
||||
});
|
||||
@@ -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
@@ -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]]
|
||||
|
||||
Generated
+2
-2
@@ -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",
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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"] }
|
||||
|
||||
@@ -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"]
|
||||
@@ -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!());
|
||||
|
||||
@@ -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,
|
||||
})
|
||||
}
|
||||
@@ -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());
|
||||
}
|
||||
}
|
||||
@@ -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 }));
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
]
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -221,6 +221,7 @@
|
||||
"91aaf184cfc17ffd": "数据概览",
|
||||
"91af6e57e7453fbe": "添加账号",
|
||||
"92156a483d4ba248": "仅删除请求、响应和追踪附件等详细内容,保留调用汇总、统计指标和配置。",
|
||||
"92e26b27d5ea8f0e": "检查更新失败:{error}",
|
||||
"940a168911ade998": "每页条数",
|
||||
"945fb1c67eca8493": "正在安装插件运行时",
|
||||
"946b3ffc02f026c0": "确定删除这个模型吗?",
|
||||
@@ -259,6 +260,7 @@
|
||||
"a748cc074f78de00": "查看详情",
|
||||
"a7617f42f898b2bf": "使用完整请求地址",
|
||||
"a8036485f9227f2c": "拖动排序",
|
||||
"a80b53f8848e6d27": "安装更新失败:{error}",
|
||||
"a98585871c5313ff": "显示名称",
|
||||
"ab9084a640fbb864": "全不选",
|
||||
"abecab6701177721": "已开启开机启动",
|
||||
|
||||
@@ -235,9 +235,7 @@ export interface PluginModelDescriptor {
|
||||
description: string | null;
|
||||
icon: string;
|
||||
providerType: string;
|
||||
contextWindowTokens: number | null;
|
||||
maxOutputTokens: number | null;
|
||||
thinking: boolean;
|
||||
images: boolean;
|
||||
}
|
||||
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
|
||||
@@ -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 });
|
||||
|
||||
@@ -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({
|
||||
|
||||
@@ -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"]);
|
||||
|
||||
@@ -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"),
|
||||
},
|
||||
});
|
||||
|
||||
@@ -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)?;
|
||||
}
|
||||
|
||||
@@ -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(®istry, &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(®istry, 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(®istry, 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)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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
@@ -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))
|
||||
|
||||
@@ -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>>,
|
||||
}
|
||||
|
||||
|
||||
@@ -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";
|
||||
|
||||
@@ -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> {
|
||||
|
||||
@@ -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<_>>>()?;
|
||||
|
||||
@@ -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");
|
||||
}
|
||||
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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() {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
@@ -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() {
|
||||
|
||||
@@ -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, ¤t, 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, ¤t, 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(
|
||||
®istry,
|
||||
&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(
|
||||
®istry,
|
||||
&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;
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
@@ -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));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
@@ -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};
|
||||
|
||||
@@ -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"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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(¬ebook_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(¬ebook_call("hi"), &single_cell_notebook()).unwrap();
|
||||
assert!(edited.contains("print('replacement')"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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,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,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;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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])
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)]
|
||||
|
||||
@@ -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,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> {
|
||||
|
||||
@@ -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");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
@@ -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, ¶ms).await?;
|
||||
let (request, cancellation, recorder) = self.request(request_id, ¶ms).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, ¶ms).await?;
|
||||
let (request, cancellation, recorder) = self.request(request_id, ¶ms).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", ¶ms).await.unwrap();
|
||||
let (_, _, second_recorder) = host.request("invocation", ¶ms).await.unwrap();
|
||||
let (_, body) = recorded_network_request(¶ms).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(¶ms).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", ¶ms).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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,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;
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,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,
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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
@@ -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
@@ -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, ¤t_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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -7,6 +7,7 @@ mod event;
|
||||
mod handle;
|
||||
mod messages;
|
||||
mod model_cycle;
|
||||
mod model_retry;
|
||||
mod port;
|
||||
mod tool_round;
|
||||
|
||||
|
||||
+130
-16
@@ -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");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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));
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
Reference in New Issue
Block a user