mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-07 06:04:53 +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
|
mkdir -p legacy-update
|
||||||
tar -czf "legacy-update/cursor-byok-${VERSION}-linux-amd64.tar.gz" -C target/release cursor-byok-desktop
|
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'
|
if: matrix.platform == 'windows-x86_64'
|
||||||
shell: pwsh
|
shell: pwsh
|
||||||
env:
|
env:
|
||||||
VERSION: ${{ needs.prepare.outputs.version }}
|
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: |
|
run: |
|
||||||
New-Item -ItemType Directory -Force legacy-update | Out-Null
|
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
|
- name: Package legacy macOS updater asset
|
||||||
if: contains(matrix.platform, 'macos')
|
if: contains(matrix.platform, 'macos')
|
||||||
@@ -213,6 +226,20 @@ jobs:
|
|||||||
--output legacy-update/update.json \
|
--output legacy-update/update.json \
|
||||||
--notes "Cursor BYOK v${VERSION}"
|
--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
|
- name: Normalize Tauri updater download URLs
|
||||||
env:
|
env:
|
||||||
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||||
|
|||||||
Generated
+5
-1
@@ -1172,10 +1172,11 @@ checksum = "52560adf09603e58c9a7ee1fe1dcb95a16927b17c127f0ac02d6e768a0e25bc1"
|
|||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "cursor-byok-desktop"
|
name = "cursor-byok-desktop"
|
||||||
version = "0.1.5"
|
version = "0.1.6"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"axum",
|
"axum",
|
||||||
"cursor-server",
|
"cursor-server",
|
||||||
|
"libc",
|
||||||
"rfd",
|
"rfd",
|
||||||
"serde",
|
"serde",
|
||||||
"serde_json",
|
"serde_json",
|
||||||
@@ -1187,12 +1188,15 @@ dependencies = [
|
|||||||
"tauri-plugin-process",
|
"tauri-plugin-process",
|
||||||
"tauri-plugin-single-instance",
|
"tauri-plugin-single-instance",
|
||||||
"tauri-plugin-updater",
|
"tauri-plugin-updater",
|
||||||
|
"tempfile",
|
||||||
"tokio",
|
"tokio",
|
||||||
"tokio-util",
|
"tokio-util",
|
||||||
"tracing",
|
"tracing",
|
||||||
"tracing-appender",
|
"tracing-appender",
|
||||||
"tracing-subscriber",
|
"tracing-subscriber",
|
||||||
"url",
|
"url",
|
||||||
|
"windows-sys 0.61.2",
|
||||||
|
"zip",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
|
|||||||
Generated
+2
-2
@@ -1,12 +1,12 @@
|
|||||||
{
|
{
|
||||||
"name": "cursor-byok-desktop",
|
"name": "cursor-byok-desktop",
|
||||||
"version": "0.1.5",
|
"version": "0.1.6",
|
||||||
"lockfileVersion": 3,
|
"lockfileVersion": 3,
|
||||||
"requires": true,
|
"requires": true,
|
||||||
"packages": {
|
"packages": {
|
||||||
"": {
|
"": {
|
||||||
"name": "cursor-byok-desktop",
|
"name": "cursor-byok-desktop",
|
||||||
"version": "0.1.5",
|
"version": "0.1.6",
|
||||||
"license": "MIT",
|
"license": "MIT",
|
||||||
"dependencies": {
|
"dependencies": {
|
||||||
"@floating-ui/dom": "^1.8.0",
|
"@floating-ui/dom": "^1.8.0",
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
{
|
{
|
||||||
"name": "cursor-byok-desktop",
|
"name": "cursor-byok-desktop",
|
||||||
"version": "0.1.5",
|
"version": "0.1.6",
|
||||||
"description": "Cursor BYOK desktop management application",
|
"description": "Cursor BYOK desktop management application",
|
||||||
"type": "module",
|
"type": "module",
|
||||||
"scripts": {
|
"scripts": {
|
||||||
@@ -8,7 +8,7 @@
|
|||||||
"dev": "vite",
|
"dev": "vite",
|
||||||
"typecheck": "tsc --noEmit",
|
"typecheck": "tsc --noEmit",
|
||||||
"typecheck:node": "tsc --noEmit -p tsconfig.node.json",
|
"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": "vite build",
|
||||||
"build:demo": "npm run typecheck && npm run typecheck:node && vite build --config vite.demo.config.ts",
|
"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",
|
"check": "npm run typecheck && npm run typecheck:node && npm run build",
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
[package]
|
[package]
|
||||||
name = "cursor-byok-desktop"
|
name = "cursor-byok-desktop"
|
||||||
version = "0.1.5"
|
version = "0.1.6"
|
||||||
edition = "2021"
|
edition = "2021"
|
||||||
publish = false
|
publish = false
|
||||||
|
|
||||||
@@ -14,6 +14,7 @@ tauri-build = { version = "2", features = [] }
|
|||||||
[dependencies]
|
[dependencies]
|
||||||
axum = "0.8"
|
axum = "0.8"
|
||||||
cursor-server = { path = "../../../server" }
|
cursor-server = { path = "../../../server" }
|
||||||
|
libc = "0.2"
|
||||||
rfd = "0.15"
|
rfd = "0.15"
|
||||||
serde = { version = "1", features = ["derive"] }
|
serde = { version = "1", features = ["derive"] }
|
||||||
serde_json = "1"
|
serde_json = "1"
|
||||||
@@ -24,9 +25,14 @@ tauri-plugin-opener = "2"
|
|||||||
tauri-plugin-autostart = "2"
|
tauri-plugin-autostart = "2"
|
||||||
tauri-plugin-process = "2"
|
tauri-plugin-process = "2"
|
||||||
tauri-plugin-updater = "2"
|
tauri-plugin-updater = "2"
|
||||||
|
tempfile = "3"
|
||||||
tokio = { version = "1", features = ["time"] }
|
tokio = { version = "1", features = ["time"] }
|
||||||
tokio-util = "0.7"
|
tokio-util = "0.7"
|
||||||
tracing = "0.1"
|
tracing = "0.1"
|
||||||
tracing-appender = "0.2"
|
tracing-appender = "0.2"
|
||||||
tracing-subscriber = { version = "0.3", features = ["env-filter"] }
|
tracing-subscriber = { version = "0.3", features = ["env-filter"] }
|
||||||
url = "2"
|
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() {
|
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))
|
tauri_build::try_build(tauri_build::Attributes::new().app_manifest(manifest))
|
||||||
.expect("failed to build Tauri application")
|
.expect("failed to build Tauri application")
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -18,6 +18,8 @@
|
|||||||
"core:window:allow-close",
|
"core:window:allow-close",
|
||||||
"core:app:allow-set-dock-visibility",
|
"core:app:allow-set-dock-visibility",
|
||||||
"allow-open-terminal-with-command",
|
"allow-open-terminal-with-command",
|
||||||
|
"allow-check-portable-update",
|
||||||
|
"allow-install-portable-update",
|
||||||
"clipboard-manager:allow-write-text",
|
"clipboard-manager:allow-write-text",
|
||||||
"autostart:default",
|
"autostart:default",
|
||||||
"process:allow-restart",
|
"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;
|
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!(
|
tracing::info!(
|
||||||
version = env!("CARGO_PKG_VERSION"),
|
version = env!("CARGO_PKG_VERSION"),
|
||||||
os = std::env::consts::OS,
|
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 started_by_autostart = std::env::args_os().any(|arg| arg == AUTOSTART_ARG);
|
||||||
|
|
||||||
let app = tauri::Builder::default()
|
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, _| {
|
.plugin(tauri_plugin_single_instance::init(|app, args, _| {
|
||||||
if !args.iter().any(|arg| arg == AUTOSTART_ARG) {
|
if !args.iter().any(|arg| arg == AUTOSTART_ARG) {
|
||||||
tray::show_main_window(app);
|
tray::show_main_window(app);
|
||||||
@@ -228,6 +249,7 @@ pub fn run() -> ExitCode {
|
|||||||
window.set_focus()?;
|
window.set_focus()?;
|
||||||
}
|
}
|
||||||
tray::create(app)?;
|
tray::create(app)?;
|
||||||
|
crate::update::signal_ready_if_requested()?;
|
||||||
Ok(())
|
Ok(())
|
||||||
})
|
})
|
||||||
.build(tauri::generate_context!());
|
.build(tauri::generate_context!());
|
||||||
|
|||||||
@@ -1,7 +1,14 @@
|
|||||||
mod desktop;
|
mod desktop;
|
||||||
#[cfg(not(dev))]
|
#[cfg(not(dev))]
|
||||||
mod frontend;
|
mod frontend;
|
||||||
|
mod resource_limits;
|
||||||
mod startup;
|
mod startup;
|
||||||
mod tray;
|
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",
|
"$schema": "https://schema.tauri.app/config/2",
|
||||||
"productName": "Cursor BYOK",
|
"productName": "Cursor BYOK",
|
||||||
"version": "0.1.5",
|
"version": "0.1.6",
|
||||||
"identifier": "dev.cursorbyok.desktop",
|
"identifier": "dev.cursorbyok.desktop",
|
||||||
"build": {
|
"build": {
|
||||||
"beforeDevCommand": "npm run dev",
|
"beforeDevCommand": "npm run dev",
|
||||||
@@ -12,7 +12,7 @@
|
|||||||
"app": {
|
"app": {
|
||||||
"windows": [],
|
"windows": [],
|
||||||
"security": {
|
"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": [
|
"dangerousDisableAssetCspModification": [
|
||||||
"style-src"
|
"style-src"
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -95,7 +95,8 @@ export function AppLifecycleSettingsCard() {
|
|||||||
const nextVersion = await updateStore.check();
|
const nextVersion = await updateStore.check();
|
||||||
message(nextVersion ? t("发现新版本 {version}", { version: nextVersion }) : t("当前已是最新版本"));
|
message(nextVersion ? t("发现新版本 {version}", { version: nextVersion }) : t("当前已是最新版本"));
|
||||||
} catch (cause) {
|
} 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 {
|
try {
|
||||||
await updateStore.install();
|
await updateStore.install();
|
||||||
} catch (cause) {
|
} 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": [
|
"refs": [
|
||||||
{
|
{
|
||||||
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
||||||
"line": 156,
|
"line": 158,
|
||||||
"column": 27
|
"column": 27
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
@@ -916,7 +916,7 @@
|
|||||||
"refs": [
|
"refs": [
|
||||||
{
|
{
|
||||||
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
||||||
"line": 152,
|
"line": 154,
|
||||||
"column": 13
|
"column": 13
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
@@ -1673,7 +1673,7 @@
|
|||||||
"refs": [
|
"refs": [
|
||||||
{
|
{
|
||||||
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
||||||
"line": 126,
|
"line": 128,
|
||||||
"column": 17
|
"column": 17
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
@@ -1725,7 +1725,7 @@
|
|||||||
"refs": [
|
"refs": [
|
||||||
{
|
{
|
||||||
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
||||||
"line": 151,
|
"line": 153,
|
||||||
"column": 13
|
"column": 13
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
@@ -2353,7 +2353,7 @@
|
|||||||
"refs": [
|
"refs": [
|
||||||
{
|
{
|
||||||
"file": "shared/api.ts",
|
"file": "shared/api.ts",
|
||||||
"line": 488,
|
"line": 486,
|
||||||
"column": 43
|
"column": 43
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
@@ -2365,7 +2365,7 @@
|
|||||||
"refs": [
|
"refs": [
|
||||||
{
|
{
|
||||||
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
||||||
"line": 149,
|
"line": 151,
|
||||||
"column": 18
|
"column": 18
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
@@ -2451,12 +2451,12 @@
|
|||||||
"refs": [
|
"refs": [
|
||||||
{
|
{
|
||||||
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
||||||
"line": 137,
|
"line": 139,
|
||||||
"column": 18
|
"column": 18
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
||||||
"line": 143,
|
"line": 145,
|
||||||
"column": 16
|
"column": 16
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
@@ -2680,7 +2680,7 @@
|
|||||||
"refs": [
|
"refs": [
|
||||||
{
|
{
|
||||||
"file": "shared/api.ts",
|
"file": "shared/api.ts",
|
||||||
"line": 483,
|
"line": 481,
|
||||||
"column": 43
|
"column": 43
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
@@ -2795,7 +2795,7 @@
|
|||||||
"refs": [
|
"refs": [
|
||||||
{
|
{
|
||||||
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
||||||
"line": 160,
|
"line": 162,
|
||||||
"column": 37
|
"column": 37
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
@@ -2834,7 +2834,7 @@
|
|||||||
"refs": [
|
"refs": [
|
||||||
{
|
{
|
||||||
"file": "shared/api.ts",
|
"file": "shared/api.ts",
|
||||||
"line": 420,
|
"line": 418,
|
||||||
"column": 21
|
"column": 21
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
@@ -2900,12 +2900,12 @@
|
|||||||
"refs": [
|
"refs": [
|
||||||
{
|
{
|
||||||
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
||||||
"line": 113,
|
"line": 115,
|
||||||
"column": 18
|
"column": 18
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
||||||
"line": 119,
|
"line": 121,
|
||||||
"column": 16
|
"column": 16
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
@@ -3194,6 +3194,20 @@
|
|||||||
}
|
}
|
||||||
]
|
]
|
||||||
},
|
},
|
||||||
|
"92e26b27d5ea8f0e": {
|
||||||
|
"source": "检查更新失败:{error}",
|
||||||
|
"kind": "template",
|
||||||
|
"placeholders": [
|
||||||
|
"error"
|
||||||
|
],
|
||||||
|
"refs": [
|
||||||
|
{
|
||||||
|
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
||||||
|
"line": 99,
|
||||||
|
"column": 15
|
||||||
|
}
|
||||||
|
]
|
||||||
|
},
|
||||||
"940a168911ade998": {
|
"940a168911ade998": {
|
||||||
"source": "每页条数",
|
"source": "每页条数",
|
||||||
"kind": "text",
|
"kind": "text",
|
||||||
@@ -3520,7 +3534,7 @@
|
|||||||
"refs": [
|
"refs": [
|
||||||
{
|
{
|
||||||
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
||||||
"line": 110,
|
"line": 112,
|
||||||
"column": 29
|
"column": 29
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
@@ -3532,12 +3546,12 @@
|
|||||||
"refs": [
|
"refs": [
|
||||||
{
|
{
|
||||||
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
||||||
"line": 125,
|
"line": 127,
|
||||||
"column": 18
|
"column": 18
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
||||||
"line": 131,
|
"line": 133,
|
||||||
"column": 16
|
"column": 16
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
@@ -3749,6 +3763,20 @@
|
|||||||
}
|
}
|
||||||
]
|
]
|
||||||
},
|
},
|
||||||
|
"a80b53f8848e6d27": {
|
||||||
|
"source": "安装更新失败:{error}",
|
||||||
|
"kind": "template",
|
||||||
|
"placeholders": [
|
||||||
|
"error"
|
||||||
|
],
|
||||||
|
"refs": [
|
||||||
|
{
|
||||||
|
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
||||||
|
"line": 108,
|
||||||
|
"column": 15
|
||||||
|
}
|
||||||
|
]
|
||||||
|
},
|
||||||
"a98585871c5313ff": {
|
"a98585871c5313ff": {
|
||||||
"source": "显示名称",
|
"source": "显示名称",
|
||||||
"kind": "text",
|
"kind": "text",
|
||||||
@@ -3814,7 +3842,7 @@
|
|||||||
"refs": [
|
"refs": [
|
||||||
{
|
{
|
||||||
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
||||||
"line": 156,
|
"line": 158,
|
||||||
"column": 39
|
"column": 39
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
@@ -3854,7 +3882,7 @@
|
|||||||
"refs": [
|
"refs": [
|
||||||
{
|
{
|
||||||
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
||||||
"line": 138,
|
"line": 140,
|
||||||
"column": 17
|
"column": 17
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
@@ -4577,7 +4605,7 @@
|
|||||||
"refs": [
|
"refs": [
|
||||||
{
|
{
|
||||||
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
||||||
"line": 114,
|
"line": 116,
|
||||||
"column": 17
|
"column": 17
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
@@ -5522,7 +5550,7 @@
|
|||||||
},
|
},
|
||||||
{
|
{
|
||||||
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
"file": "features/settings/AppLifecycleSettingsCard.tsx",
|
||||||
"line": 160,
|
"line": 162,
|
||||||
"column": 25
|
"column": 25
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -221,6 +221,7 @@
|
|||||||
"91aaf184cfc17ffd": "Overview",
|
"91aaf184cfc17ffd": "Overview",
|
||||||
"91af6e57e7453fbe": "Add account",
|
"91af6e57e7453fbe": "Add account",
|
||||||
"92156a483d4ba248": "Only request, response, and trace attachments are deleted; call summaries, metrics, and configuration are kept.",
|
"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",
|
"940a168911ade998": "Items per page",
|
||||||
"945fb1c67eca8493": "Installing the plugin runtime",
|
"945fb1c67eca8493": "Installing the plugin runtime",
|
||||||
"946b3ffc02f026c0": "Delete this model?",
|
"946b3ffc02f026c0": "Delete this model?",
|
||||||
@@ -259,6 +260,7 @@
|
|||||||
"a748cc074f78de00": "View details",
|
"a748cc074f78de00": "View details",
|
||||||
"a7617f42f898b2bf": "Use complete request URL",
|
"a7617f42f898b2bf": "Use complete request URL",
|
||||||
"a8036485f9227f2c": "Drag to reorder",
|
"a8036485f9227f2c": "Drag to reorder",
|
||||||
|
"a80b53f8848e6d27": "Failed to install update: {error}",
|
||||||
"a98585871c5313ff": "Display name",
|
"a98585871c5313ff": "Display name",
|
||||||
"ab9084a640fbb864": "Deselect all",
|
"ab9084a640fbb864": "Deselect all",
|
||||||
"abecab6701177721": "Launch at login enabled",
|
"abecab6701177721": "Launch at login enabled",
|
||||||
|
|||||||
@@ -221,6 +221,7 @@
|
|||||||
"91aaf184cfc17ffd": "数据概览",
|
"91aaf184cfc17ffd": "数据概览",
|
||||||
"91af6e57e7453fbe": "添加账号",
|
"91af6e57e7453fbe": "添加账号",
|
||||||
"92156a483d4ba248": "仅删除请求、响应和追踪附件等详细内容,保留调用汇总、统计指标和配置。",
|
"92156a483d4ba248": "仅删除请求、响应和追踪附件等详细内容,保留调用汇总、统计指标和配置。",
|
||||||
|
"92e26b27d5ea8f0e": "检查更新失败:{error}",
|
||||||
"940a168911ade998": "每页条数",
|
"940a168911ade998": "每页条数",
|
||||||
"945fb1c67eca8493": "正在安装插件运行时",
|
"945fb1c67eca8493": "正在安装插件运行时",
|
||||||
"946b3ffc02f026c0": "确定删除这个模型吗?",
|
"946b3ffc02f026c0": "确定删除这个模型吗?",
|
||||||
@@ -259,6 +260,7 @@
|
|||||||
"a748cc074f78de00": "查看详情",
|
"a748cc074f78de00": "查看详情",
|
||||||
"a7617f42f898b2bf": "使用完整请求地址",
|
"a7617f42f898b2bf": "使用完整请求地址",
|
||||||
"a8036485f9227f2c": "拖动排序",
|
"a8036485f9227f2c": "拖动排序",
|
||||||
|
"a80b53f8848e6d27": "安装更新失败:{error}",
|
||||||
"a98585871c5313ff": "显示名称",
|
"a98585871c5313ff": "显示名称",
|
||||||
"ab9084a640fbb864": "全不选",
|
"ab9084a640fbb864": "全不选",
|
||||||
"abecab6701177721": "已开启开机启动",
|
"abecab6701177721": "已开启开机启动",
|
||||||
|
|||||||
@@ -235,9 +235,7 @@ export interface PluginModelDescriptor {
|
|||||||
description: string | null;
|
description: string | null;
|
||||||
icon: string;
|
icon: string;
|
||||||
providerType: string;
|
providerType: string;
|
||||||
contextWindowTokens: number | null;
|
|
||||||
maxOutputTokens: number | null;
|
maxOutputTokens: number | null;
|
||||||
thinking: boolean;
|
|
||||||
images: boolean;
|
images: boolean;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -1,5 +1,5 @@
|
|||||||
import { getVersion, setDockVisibility } from "@tauri-apps/api/app";
|
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 { disable, enable, isEnabled } from "@tauri-apps/plugin-autostart";
|
||||||
import { relaunch } from "@tauri-apps/plugin-process";
|
import { relaunch } from "@tauri-apps/plugin-process";
|
||||||
import { check, type Update } from "@tauri-apps/plugin-updater";
|
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> {
|
export type AppUpdate = {
|
||||||
return check();
|
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> {
|
export async function installUpdate(update: AppUpdate): Promise<void> {
|
||||||
await update.downloadAndInstall();
|
await update.install();
|
||||||
await relaunch();
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,9 +1,9 @@
|
|||||||
import { useSyncExternalStore } from "react";
|
import { useSyncExternalStore } from "react";
|
||||||
import type { Update } from "@tauri-apps/plugin-updater";
|
|
||||||
import {
|
import {
|
||||||
checkForUpdate,
|
checkForUpdate,
|
||||||
hasNativeAppLifecycle,
|
hasNativeAppLifecycle,
|
||||||
installUpdate,
|
installUpdate,
|
||||||
|
type AppUpdate,
|
||||||
} from "../native/appLifecycle";
|
} from "../native/appLifecycle";
|
||||||
|
|
||||||
export type UpdateSnapshot = {
|
export type UpdateSnapshot = {
|
||||||
@@ -17,7 +17,7 @@ let snapshot: UpdateSnapshot = {
|
|||||||
checking: false,
|
checking: false,
|
||||||
installing: false,
|
installing: false,
|
||||||
};
|
};
|
||||||
let availableUpdate: Update | null = null;
|
let availableUpdate: AppUpdate | null = null;
|
||||||
let pendingCheck: Promise<string | null> | null = null;
|
let pendingCheck: Promise<string | null> | null = null;
|
||||||
const listeners = new Set<() => void>();
|
const listeners = new Set<() => void>();
|
||||||
|
|
||||||
@@ -26,7 +26,7 @@ function update(patch: Partial<UpdateSnapshot>) {
|
|||||||
listeners.forEach((listener) => listener());
|
listeners.forEach((listener) => listener());
|
||||||
}
|
}
|
||||||
|
|
||||||
async function replaceAvailableUpdate(next: Update | null) {
|
async function replaceAvailableUpdate(next: AppUpdate | null) {
|
||||||
const previous = availableUpdate;
|
const previous = availableUpdate;
|
||||||
availableUpdate = next;
|
availableUpdate = next;
|
||||||
update({ availableVersion: next?.version ?? null });
|
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
|
```text
|
||||||
@@ -656,54 +651,7 @@ store ─X→ cursor
|
|||||||
model ─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
|
→ Checkpoint
|
||||||
→ Transport
|
→ Transport
|
||||||
→ RunSSE
|
→ 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 type { ResourceSnapshot } from "cursor-byok:resource";
|
||||||
import { codexDeviceOAuth } from "./oauth.ts";
|
import { codexDeviceOAuth } from "./oauth.ts";
|
||||||
import { parseOfficialModels } from "./models.ts";
|
import { parseOfficialModels } from "./models.ts";
|
||||||
|
import { buildResponsesBody } from "cursor-byok:protocol/openai-responses";
|
||||||
import { codexProvider, isQuotaError } from "./provider.ts";
|
import { codexProvider, isQuotaError } from "./provider.ts";
|
||||||
import {
|
import {
|
||||||
accountIdentity,
|
accountIdentity,
|
||||||
@@ -164,7 +165,10 @@ Deno.test("official model discovery excludes hidden models and puts the default
|
|||||||
display_name: "GPT First",
|
display_name: "GPT First",
|
||||||
supported_in_api: true,
|
supported_in_api: true,
|
||||||
visibility: "list",
|
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-second", supported_in_api: true, visibility: "list" },
|
||||||
{ slug: "gpt-hidden", supported_in_api: true, visibility: "hidden" },
|
{ 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.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"] });
|
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 () => {
|
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 token = jwt({ "https://api.openai.com/auth": { chatgpt_account_id: "acct-1" } });
|
||||||
const draft = await credentialDraft({
|
const draft = await credentialDraft({
|
||||||
|
|||||||
@@ -24,7 +24,11 @@ function positiveInteger(value: unknown): number | null {
|
|||||||
}
|
}
|
||||||
|
|
||||||
function parseReasoningEfforts(model: Record<string, unknown>): string[] {
|
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.supportedReasoningEfforts ??
|
||||||
model.reasoning_efforts ??
|
model.reasoning_efforts ??
|
||||||
model.reasoningEfforts;
|
model.reasoningEfforts;
|
||||||
@@ -61,10 +65,6 @@ export function parseOfficialModels(body: unknown): ModelDefinition[] {
|
|||||||
seen.add(id);
|
seen.add(id);
|
||||||
const efforts = parseReasoningEfforts(model);
|
const efforts = parseReasoningEfforts(model);
|
||||||
const description = text(model.description);
|
const description = text(model.description);
|
||||||
const contextWindowTokens = positiveInteger(
|
|
||||||
model.context_window_tokens ?? model.contextWindowTokens ?? model.context_window ??
|
|
||||||
model.contextWindow,
|
|
||||||
);
|
|
||||||
const maxOutputTokens = positiveInteger(
|
const maxOutputTokens = positiveInteger(
|
||||||
model.max_output_tokens ?? model.maxOutputTokens ?? model.max_completion_tokens ??
|
model.max_output_tokens ?? model.maxOutputTokens ?? model.max_completion_tokens ??
|
||||||
model.maxCompletionTokens,
|
model.maxCompletionTokens,
|
||||||
@@ -74,9 +74,8 @@ export function parseOfficialModels(body: unknown): ModelDefinition[] {
|
|||||||
displayName: text(model.display_name ?? model.displayName ?? model.title ?? model.name) ??
|
displayName: text(model.display_name ?? model.displayName ?? model.title ?? model.name) ??
|
||||||
id,
|
id,
|
||||||
...(description ? { description } : {}),
|
...(description ? { description } : {}),
|
||||||
...(contextWindowTokens !== null ? { contextWindowTokens } : {}),
|
|
||||||
...(maxOutputTokens !== null ? { maxOutputTokens } : {}),
|
...(maxOutputTokens !== null ? { maxOutputTokens } : {}),
|
||||||
capabilities: { thinking: efforts.length > 0, images: true },
|
capabilities: { images: true },
|
||||||
privateData: { reasoningEfforts: efforts },
|
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.map((model) => model.id), ["grok-4", "grok-3-mini"]);
|
||||||
assertEquals(richModels[0].displayName, "Grok 4");
|
assertEquals(richModels[0].displayName, "Grok 4");
|
||||||
assertEquals(richModels[0].contextWindowTokens, 256_000);
|
assertEquals(richModels[0].capabilities, { images: true });
|
||||||
assertEquals(richModels[0].capabilities, { thinking: false, images: true });
|
assertEquals(richModels[1].capabilities, { images: false });
|
||||||
assertEquals(richModels[1].capabilities, { thinking: false, images: false });
|
|
||||||
|
|
||||||
const plainModels = parseGrokModels({ data: [{ id: "grok-4-fast" }] });
|
const plainModels = parseGrokModels({ data: [{ id: "grok-4-fast" }] });
|
||||||
assertEquals(plainModels.map((model) => model.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",
|
id: "grok-4.6",
|
||||||
displayName: "Grok 4.6",
|
displayName: "Grok 4.6",
|
||||||
capabilities: { thinking: false, images: true },
|
capabilities: { images: true },
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
id: "grok-4.5",
|
id: "grok-4.5",
|
||||||
displayName: "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;
|
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[] {
|
function modalities(value: unknown): string[] {
|
||||||
return Array.isArray(value)
|
return Array.isArray(value)
|
||||||
? value.flatMap((item) => (typeof item === "string" ? [item.toLowerCase()] : []))
|
? value.flatMap((item) => (typeof item === "string" ? [item.toLowerCase()] : []))
|
||||||
@@ -66,15 +57,10 @@ export function parseGrokModels(body: unknown): ModelDefinition[] {
|
|||||||
if (!id || seen.has(id)) continue;
|
if (!id || seen.has(id)) continue;
|
||||||
seen.add(id);
|
seen.add(id);
|
||||||
const inputs = modalities(model?.input_modalities ?? model?.inputModalities);
|
const inputs = modalities(model?.input_modalities ?? model?.inputModalities);
|
||||||
const contextWindowTokens = positiveInteger(
|
|
||||||
model?.context_window ?? model?.contextWindow ?? model?.max_prompt_length,
|
|
||||||
);
|
|
||||||
models.push({
|
models.push({
|
||||||
id,
|
id,
|
||||||
displayName: displayName(id),
|
displayName: displayName(id),
|
||||||
...(contextWindowTokens !== null ? { contextWindowTokens } : {}),
|
|
||||||
capabilities: {
|
capabilities: {
|
||||||
thinking: false,
|
|
||||||
images: inputs.length === 0 || inputs.includes("image"),
|
images: inputs.length === 0 || inputs.includes("image"),
|
||||||
},
|
},
|
||||||
});
|
});
|
||||||
|
|||||||
@@ -184,7 +184,11 @@ pub async fn append(
|
|||||||
request: DecodedAppend,
|
request: DecodedAppend,
|
||||||
parent: Option<TransportParent>,
|
parent: Option<TransportParent>,
|
||||||
) -> Result<ai::BidiAppendResponse> {
|
) -> 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() {
|
if let Some(conversation_id) = request.conversation_id() {
|
||||||
handle.set_conversation_id(conversation_id)?;
|
handle.set_conversation_id(conversation_id)?;
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -19,16 +19,17 @@ use crate::{
|
|||||||
connect,
|
connect,
|
||||||
proto::{agent::v1 as agent, aiserver::v1 as ai},
|
proto::{agent::v1 as agent, aiserver::v1 as ai},
|
||||||
},
|
},
|
||||||
services::{
|
services::{account, analytics, knowledge, model_catalog, tab},
|
||||||
account, analytics, knowledge, model_catalog, observability::CursorTraceRecorder, tab,
|
|
||||||
},
|
|
||||||
transport::{TransportParent, TransportRegistry},
|
transport::{TransportParent, TransportRegistry},
|
||||||
},
|
},
|
||||||
Result,
|
Result,
|
||||||
};
|
};
|
||||||
|
|
||||||
pub fn router(registry: TransportRegistry) -> Result<Router> {
|
pub fn router(
|
||||||
let proxy = CursorProxy::cursor(registry.store().clone())?;
|
registry: TransportRegistry,
|
||||||
|
clients: crate::network::NetworkClients,
|
||||||
|
) -> Result<Router> {
|
||||||
|
let proxy = CursorProxy::cursor(clients);
|
||||||
let knowledge = knowledge::KnowledgeService::managed()?;
|
let knowledge = knowledge::KnowledgeService::managed()?;
|
||||||
Ok(router_with_proxy(registry, proxy, knowledge))
|
Ok(router_with_proxy(registry, proxy, knowledge))
|
||||||
}
|
}
|
||||||
@@ -120,16 +121,13 @@ async fn run_sse_handler(
|
|||||||
let (parts, body) = buffered(request).await?;
|
let (parts, body) = buffered(request).await?;
|
||||||
let request: agent::BidiRequestId = connect::decode_unary(&body)?;
|
let request: agent::BidiRequestId = connect::decode_unary(&body)?;
|
||||||
let route = registry.wait_route(&request.request_id).await;
|
let route = registry.wait_route(&request.request_id).await;
|
||||||
let trace = CursorTraceRecorder::resume(registry.store().clone(), &request.request_id).await;
|
let trace = registry.trace(&request.request_id);
|
||||||
if let Some(trace) = &trace {
|
trace.resume();
|
||||||
trace
|
trace.request(
|
||||||
.request(
|
"run_sse_request",
|
||||||
"run_sse_request",
|
body.clone(),
|
||||||
&body,
|
serde_json::json!({"request_id": request.request_id}),
|
||||||
serde_json::json!({"request_id": request.request_id}),
|
);
|
||||||
)
|
|
||||||
.await;
|
|
||||||
}
|
|
||||||
match route {
|
match route {
|
||||||
crate::cursor::transport::TransportRoute::Local => {
|
crate::cursor::transport::TransportRoute::Local => {
|
||||||
run_sse::stream(®istry, &request.request_id).await
|
run_sse::stream(®istry, &request.request_id).await
|
||||||
@@ -140,7 +138,14 @@ async fn run_sse_handler(
|
|||||||
Request::from_parts(parts, Body::from(body)),
|
Request::from_parts(parts, Body::from(body)),
|
||||||
)
|
)
|
||||||
.await?;
|
.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 first_model = decoded.model_id().map(str::to_owned);
|
||||||
let conversation_id = decoded.conversation_id().map(str::to_owned);
|
let conversation_id = decoded.conversation_id().map(str::to_owned);
|
||||||
let trace_metadata = decoded.trace_metadata();
|
let trace_metadata = decoded.trace_metadata();
|
||||||
|
let trace = registry.trace(&decoded.request_id);
|
||||||
let local = if let Some(model_id) = decoded.model_id() {
|
let local = if let Some(model_id) = decoded.model_id() {
|
||||||
// 插件模型 ID 只在本地有意义,永远不转发到 Cursor 官方上游。
|
// 插件模型 ID 只在本地有意义,永远不转发到 Cursor 官方上游。
|
||||||
if model_id.starts_with(crate::plugin::ADAPTER_ID_PREFIX)
|
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 {
|
} else if registry.upstream(&decoded.request_id).await {
|
||||||
false
|
false
|
||||||
} else {
|
} else {
|
||||||
|
trace.resume();
|
||||||
|
trace.request(
|
||||||
|
"bidi_request",
|
||||||
|
body.clone(),
|
||||||
|
trace_outcome(trace_metadata, false, "missing_transport", None),
|
||||||
|
);
|
||||||
return Err(crate::Error::Protocol(
|
return Err(crate::Error::Protocol(
|
||||||
"first BidiAppend message must select a model".into(),
|
"first BidiAppend message must select a model".into(),
|
||||||
));
|
));
|
||||||
};
|
};
|
||||||
let trace = if first_model.is_some() {
|
if first_model.is_some() {
|
||||||
CursorTraceRecorder::begin(
|
trace.begin(
|
||||||
registry.store().clone(),
|
|
||||||
&decoded.request_id,
|
|
||||||
conversation_id.as_deref(),
|
conversation_id.as_deref(),
|
||||||
if local {
|
if local {
|
||||||
"local_byok"
|
"local_byok"
|
||||||
@@ -195,26 +205,61 @@ async fn bidi_handler(
|
|||||||
"cursor_official"
|
"cursor_official"
|
||||||
},
|
},
|
||||||
first_model.as_deref(),
|
first_model.as_deref(),
|
||||||
)
|
);
|
||||||
.await
|
|
||||||
} else {
|
} else {
|
||||||
CursorTraceRecorder::resume(registry.store().clone(), &decoded.request_id).await
|
trace.resume();
|
||||||
};
|
|
||||||
if let Some(trace) = &trace {
|
|
||||||
trace.request("bidi_request", &body, trace_metadata).await;
|
|
||||||
}
|
}
|
||||||
if !local {
|
if !local {
|
||||||
if first_model.is_some() {
|
if first_model.is_some() {
|
||||||
registry.mark_upstream(&decoded.request_id).await;
|
registry.mark_upstream(&decoded.request_id).await;
|
||||||
}
|
}
|
||||||
|
trace.request(
|
||||||
|
"bidi_request",
|
||||||
|
body.clone(),
|
||||||
|
trace_outcome(trace_metadata, true, "upstream", None),
|
||||||
|
);
|
||||||
return proxy::forward(
|
return proxy::forward(
|
||||||
Extension(proxy),
|
Extension(proxy),
|
||||||
Request::from_parts(parts, Body::from(body)),
|
Request::from_parts(parts, Body::from(body)),
|
||||||
)
|
)
|
||||||
.await;
|
.await;
|
||||||
}
|
}
|
||||||
let parent = parent_headers(&parts.headers)?;
|
let parent = match parent_headers(&parts.headers) {
|
||||||
bidi::append(®istry, decoded, parent).await?;
|
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());
|
let mut response = Response::new(axum::body::Body::empty());
|
||||||
*response.status_mut() = StatusCode::OK;
|
*response.status_mut() = StatusCode::OK;
|
||||||
response.headers_mut().insert(
|
response.headers_mut().insert(
|
||||||
@@ -224,6 +269,22 @@ async fn bidi_handler(
|
|||||||
Ok(response)
|
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)> {
|
async fn buffered(request: Request<Body>) -> Result<(axum::http::request::Parts, Bytes)> {
|
||||||
let (parts, body) = request.into_parts();
|
let (parts, body) = request.into_parts();
|
||||||
let body = to_bytes(body, usize::MAX)
|
let body = to_bytes(body, usize::MAX)
|
||||||
|
|||||||
@@ -14,8 +14,7 @@ pub const UPSTREAM_URL_HEADER: &str = "x-server-upstream-url";
|
|||||||
|
|
||||||
#[derive(Clone)]
|
#[derive(Clone)]
|
||||||
pub struct CursorProxy {
|
pub struct CursorProxy {
|
||||||
client: Option<reqwest::Client>,
|
clients: crate::network::NetworkClients,
|
||||||
store: Option<crate::store::Store>,
|
|
||||||
upstream: String,
|
upstream: String,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -47,23 +46,15 @@ impl BufferedResponse {
|
|||||||
}
|
}
|
||||||
|
|
||||||
impl CursorProxy {
|
impl CursorProxy {
|
||||||
pub fn cursor(store: crate::store::Store) -> Result<Self> {
|
pub fn cursor(clients: crate::network::NetworkClients) -> Self {
|
||||||
Ok(Self {
|
Self {
|
||||||
client: None,
|
clients,
|
||||||
store: Some(store),
|
|
||||||
upstream: CURSOR_UPSTREAM.into(),
|
upstream: CURSOR_UPSTREAM.into(),
|
||||||
})
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn client(&self) -> Result<reqwest::Client> {
|
async fn client(&self) -> Result<reqwest::Client> {
|
||||||
match (&self.client, &self.store) {
|
self.clients.cursor_client().await
|
||||||
(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"),
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -22,7 +22,7 @@ pub async fn stream(registry: &TransportRegistry, request_id: &str) -> Result<Re
|
|||||||
let receiver = handle.subscribe();
|
let receiver = handle.subscribe();
|
||||||
let trace = handle.trace().cloned();
|
let trace = handle.trace().cloned();
|
||||||
if let Some(trace) = &trace {
|
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 body_stream = local_body_stream(receiver, handle, trace);
|
||||||
let mut response = Response::new(Body::from_stream(body_stream));
|
let mut response = Response::new(Body::from_stream(body_stream));
|
||||||
@@ -133,7 +133,7 @@ pub async fn upstream(
|
|||||||
) -> Response<Body> {
|
) -> Response<Body> {
|
||||||
let (parts, body) = response.into_parts();
|
let (parts, body) = response.into_parts();
|
||||||
if let Some(trace) = &trace {
|
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 stream = async_stream::stream! {
|
||||||
let _guard = UpstreamRunGuard {
|
let _guard = UpstreamRunGuard {
|
||||||
@@ -180,15 +180,15 @@ impl TraceStreamSink {
|
|||||||
while let Some(event) = receiver.recv().await {
|
while let Some(event) = receiver.recv().await {
|
||||||
match event {
|
match event {
|
||||||
TraceStreamEvent::Chunk(chunk) => {
|
TraceStreamEvent::Chunk(chunk) => {
|
||||||
trace.response_chunk(source, &chunk).await;
|
trace.response_chunk(source, chunk);
|
||||||
}
|
}
|
||||||
TraceStreamEvent::Finish(error) => {
|
TraceStreamEvent::Finish(error) => {
|
||||||
trace.finish(error.as_deref()).await;
|
trace.finish(error.as_deref());
|
||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
trace.finish(None).await;
|
trace.finish(None);
|
||||||
});
|
});
|
||||||
Self {
|
Self {
|
||||||
sender: Some(sender),
|
sender: Some(sender),
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
//! Builds the top-level server router.
|
//! 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> {
|
pub fn router(registry: TransportRegistry, clients: NetworkClients) -> Result<axum::Router> {
|
||||||
super::cursor::router(registry)
|
super::cursor::router(registry, clients)
|
||||||
}
|
}
|
||||||
|
|||||||
+11
-3
@@ -44,9 +44,11 @@ impl App {
|
|||||||
plugin_runtime.clone(),
|
plugin_runtime.clone(),
|
||||||
config.app_version.clone(),
|
config.app_version.clone(),
|
||||||
)?;
|
)?;
|
||||||
|
let clients = crate::network::NetworkClients::new(store.clone());
|
||||||
let provider = std::sync::Arc::new(ProviderRouter::new(
|
let provider = std::sync::Arc::new(ProviderRouter::new(
|
||||||
store.clone(),
|
store.clone(),
|
||||||
plugins.clone(),
|
plugins.clone(),
|
||||||
|
clients.clone(),
|
||||||
config.provider_request_timeout,
|
config.provider_request_timeout,
|
||||||
config.provider_stream_idle_timeout,
|
config.provider_stream_idle_timeout,
|
||||||
));
|
));
|
||||||
@@ -58,10 +60,16 @@ impl App {
|
|||||||
plugins.clone(),
|
plugins.clone(),
|
||||||
crate::config::managed_data_dir()?.join("rules"),
|
crate::config::managed_data_dir()?.join("rules"),
|
||||||
);
|
);
|
||||||
let control =
|
let control = control::ControlService::new(
|
||||||
control::ControlService::new(store.clone(), provider, plugin_runtime, plugins)?;
|
store.clone(),
|
||||||
|
provider,
|
||||||
|
plugin_runtime,
|
||||||
|
plugins,
|
||||||
|
clients.clone(),
|
||||||
|
config.app_version.clone(),
|
||||||
|
)?;
|
||||||
let harness = control.cursor_harness().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 {
|
router = match &config.console {
|
||||||
Some(ConsoleSource::Directory(directory)) => {
|
Some(ConsoleSource::Directory(directory)) => {
|
||||||
router.merge(control::web_router(control.clone(), directory))
|
router.merge(control::web_router(control.clone(), directory))
|
||||||
|
|||||||
@@ -46,7 +46,6 @@ pub struct ProviderConfig {
|
|||||||
pub custom_headers: reqwest::header::HeaderMap,
|
pub custom_headers: reqwest::header::HeaderMap,
|
||||||
pub max_output_tokens: Option<u64>,
|
pub max_output_tokens: Option<u64>,
|
||||||
pub request_timeout: Duration,
|
pub request_timeout: Duration,
|
||||||
pub retry_count: u32,
|
|
||||||
pub allowed_body_fields: Option<std::collections::HashSet<String>>,
|
pub allowed_body_fields: Option<std::collections::HashSet<String>>,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -17,6 +17,7 @@ use super::ControlService;
|
|||||||
// 开源项目广告为作者唯一收入来源,PR时请勿去除此广告
|
// 开源项目广告为作者唯一收入来源,PR时请勿去除此广告
|
||||||
|
|
||||||
pub(super) const ADS_ENDPOINT: &str = "https://tab.leokun.cn/api/v1/ads?placement=menu";
|
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 DEVICE_ID_HEADER: &str = "X-Cursor-Assistant-Device-ID";
|
||||||
pub(super) const OS_HEADER: &str = "X-Cursor-Assistant-OS";
|
pub(super) const OS_HEADER: &str = "X-Cursor-Assistant-OS";
|
||||||
pub(super) const APP_VERSION_HEADER: &str = "X-Cursor-Assistant-Version";
|
pub(super) const APP_VERSION_HEADER: &str = "X-Cursor-Assistant-Version";
|
||||||
|
|||||||
@@ -41,6 +41,8 @@ pub struct ControlService {
|
|||||||
provider: Arc<dyn Provider>,
|
provider: Arc<dyn Provider>,
|
||||||
plugin_runtime: PluginRuntime,
|
plugin_runtime: PluginRuntime,
|
||||||
plugins: PluginRegistry,
|
plugins: PluginRegistry,
|
||||||
|
clients: crate::network::NetworkClients,
|
||||||
|
app_version: String,
|
||||||
model_tests: Arc<Mutex<BTreeMap<String, CancellationToken>>>,
|
model_tests: Arc<Mutex<BTreeMap<String, CancellationToken>>>,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -151,6 +153,8 @@ impl ControlService {
|
|||||||
provider: Arc<dyn Provider>,
|
provider: Arc<dyn Provider>,
|
||||||
plugin_runtime: PluginRuntime,
|
plugin_runtime: PluginRuntime,
|
||||||
plugins: PluginRegistry,
|
plugins: PluginRegistry,
|
||||||
|
clients: crate::network::NetworkClients,
|
||||||
|
app_version: String,
|
||||||
) -> Result<Self> {
|
) -> Result<Self> {
|
||||||
Ok(Self {
|
Ok(Self {
|
||||||
cursor_harness: CursorHarness::new(store.clone())?,
|
cursor_harness: CursorHarness::new(store.clone())?,
|
||||||
@@ -158,6 +162,8 @@ impl ControlService {
|
|||||||
provider,
|
provider,
|
||||||
plugin_runtime,
|
plugin_runtime,
|
||||||
plugins,
|
plugins,
|
||||||
|
clients,
|
||||||
|
app_version,
|
||||||
model_tests: Arc::new(Mutex::new(BTreeMap::new())),
|
model_tests: Arc::new(Mutex::new(BTreeMap::new())),
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
@@ -256,13 +262,13 @@ impl ControlService {
|
|||||||
disabled_ad_ids: Option<&str>,
|
disabled_ad_ids: Option<&str>,
|
||||||
language: &str,
|
language: &str,
|
||||||
) -> Result<AdRuntime> {
|
) -> 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 installation_id = self.store.installation_id().await?;
|
||||||
let mut request = client
|
let mut request = client
|
||||||
.get(ADS_ENDPOINT)
|
.get(ADS_ENDPOINT)
|
||||||
.header(DEVICE_ID_HEADER, installation_id)
|
.header(DEVICE_ID_HEADER, installation_id)
|
||||||
.header(OS_HEADER, std::env::consts::OS)
|
.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)
|
.header(LANGUAGE_HEADER, language)
|
||||||
.timeout(std::time::Duration::from_secs(60));
|
.timeout(std::time::Duration::from_secs(60));
|
||||||
if let Some(disabled_ad_ids) = disabled_ad_ids.filter(|value| !value.is_empty()) {
|
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<()> {
|
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 installation_id = self.store.installation_id().await?;
|
||||||
let mut endpoint = Url::parse(ADS_ENDPOINT).map_err(|error| {
|
let mut endpoint = Url::parse(ADS_ENDPOINT).map_err(|error| {
|
||||||
Error::Config(format!("advertisement endpoint is invalid: {error}"))
|
Error::Config(format!("advertisement endpoint is invalid: {error}"))
|
||||||
@@ -296,7 +302,7 @@ impl ControlService {
|
|||||||
.post(endpoint)
|
.post(endpoint)
|
||||||
.header(DEVICE_ID_HEADER, installation_id)
|
.header(DEVICE_ID_HEADER, installation_id)
|
||||||
.header(OS_HEADER, std::env::consts::OS)
|
.header(OS_HEADER, std::env::consts::OS)
|
||||||
.header(APP_VERSION_HEADER, env!("CARGO_PKG_VERSION"))
|
.header(APP_VERSION_HEADER, &self.app_version)
|
||||||
.json(input)
|
.json(input)
|
||||||
.timeout(std::time::Duration::from_secs(5))
|
.timeout(std::time::Duration::from_secs(5))
|
||||||
.send()
|
.send()
|
||||||
@@ -392,7 +398,6 @@ impl ControlService {
|
|||||||
if model_hash.starts_with(crate::plugin::ADAPTER_ID_PREFIX) {
|
if model_hash.starts_with(crate::plugin::ADAPTER_ID_PREFIX) {
|
||||||
let descriptor = self.plugins.model_descriptor(model_hash).await?;
|
let descriptor = self.plugins.model_descriptor(model_hash).await?;
|
||||||
model.display_name = Some(descriptor.display_name);
|
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));
|
model.max_output_tokens = Some(descriptor.max_output_tokens.unwrap_or(65_536));
|
||||||
} else {
|
} else {
|
||||||
let configured = self
|
let configured = self
|
||||||
@@ -511,7 +516,7 @@ impl ControlService {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub async fn discover_models(&self, input: &ModelDiscoveryInput) -> Result<DiscoveredModels> {
|
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)?;
|
let base_url = crate::model::normalize_request_url(&input.base_url)?;
|
||||||
discover_models_from_endpoint(
|
discover_models_from_endpoint(
|
||||||
&client,
|
&client,
|
||||||
@@ -678,7 +683,9 @@ impl ControlService {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub async fn set_proxy_settings(&self, settings: ProxySettingsInput) -> Result<ProxySettings> {
|
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> {
|
pub async fn tab_settings(&self) -> Result<TabSettings> {
|
||||||
|
|||||||
@@ -278,19 +278,17 @@ impl CheckpointBuilder {
|
|||||||
),
|
),
|
||||||
});
|
});
|
||||||
if let Some(trace) = handle.trace() {
|
if let Some(trace) = handle.trace() {
|
||||||
trace
|
trace.artifact(
|
||||||
.artifact(
|
"checkpoint",
|
||||||
"checkpoint",
|
"byok_server",
|
||||||
"byok_server",
|
&checkpoint.encode_to_vec(),
|
||||||
&checkpoint.encode_to_vec(),
|
serde_json::json!({
|
||||||
serde_json::json!({
|
"root_message_count": checkpoint.root_prompt_messages_json.len(),
|
||||||
"root_message_count": checkpoint.root_prompt_messages_json.len(),
|
"turn_count": checkpoint.turns.len(),
|
||||||
"turn_count": checkpoint.turns.len(),
|
"pending_tool_call_count": checkpoint.pending_tool_calls.len(),
|
||||||
"pending_tool_call_count": checkpoint.pending_tool_calls.len(),
|
"emit_status": if result.is_ok() { "sent" } else { "error" },
|
||||||
"emit_status": if result.is_ok() { "sent" } else { "error" },
|
}),
|
||||||
}),
|
);
|
||||||
)
|
|
||||||
.await;
|
|
||||||
}
|
}
|
||||||
result
|
result
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -102,6 +102,12 @@ pub fn decode_pending(value: &str) -> Result<RecoveredToolRound> {
|
|||||||
.into_iter()
|
.into_iter()
|
||||||
.enumerate()
|
.enumerate()
|
||||||
.map(|(index, call)| {
|
.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 {
|
Ok(ToolCall {
|
||||||
index,
|
index,
|
||||||
call_id: call.call_id,
|
call_id: call.call_id,
|
||||||
@@ -109,6 +115,7 @@ pub fn decode_pending(value: &str) -> Result<RecoveredToolRound> {
|
|||||||
name: call.name,
|
name: call.name,
|
||||||
arguments_text: serde_json::to_string(&call.arguments)?,
|
arguments_text: serde_json::to_string(&call.arguments)?,
|
||||||
arguments: call.arguments,
|
arguments: call.arguments,
|
||||||
|
argument_error,
|
||||||
})
|
})
|
||||||
})
|
})
|
||||||
.collect::<Result<Vec<_>>>()?;
|
.collect::<Result<Vec<_>>>()?;
|
||||||
|
|||||||
@@ -71,6 +71,7 @@ pub fn staged_tool_round(
|
|||||||
allowed_tools,
|
allowed_tools,
|
||||||
dynamic_tools,
|
dynamic_tools,
|
||||||
started_at_ms,
|
started_at_ms,
|
||||||
|
tool_calls: Some(calls),
|
||||||
}),
|
}),
|
||||||
)?)?)
|
)?)?)
|
||||||
}
|
}
|
||||||
@@ -96,6 +97,7 @@ pub fn staged_final(
|
|||||||
allowed_tools,
|
allowed_tools,
|
||||||
dynamic_tools,
|
dynamic_tools,
|
||||||
started_at_ms,
|
started_at_ms,
|
||||||
|
tool_calls: None,
|
||||||
}),
|
}),
|
||||||
)?)?)
|
)?)?)
|
||||||
}
|
}
|
||||||
@@ -105,6 +107,7 @@ pub(super) struct PendingContext<'a> {
|
|||||||
allowed_tools: &'a [String],
|
allowed_tools: &'a [String],
|
||||||
dynamic_tools: &'a HashSet<String>,
|
dynamic_tools: &'a HashSet<String>,
|
||||||
started_at_ms: u64,
|
started_at_ms: u64,
|
||||||
|
tool_calls: Option<&'a [ToolCall]>,
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(super) fn wire_message(
|
pub(super) fn wire_message(
|
||||||
@@ -132,16 +135,23 @@ pub(super) fn wire_message(
|
|||||||
calls
|
calls
|
||||||
.iter()
|
.iter()
|
||||||
.map(|call| {
|
.map(|call| {
|
||||||
(
|
let mut contract = json!({
|
||||||
call.call_id.clone(),
|
"toolCallId": call.call_id,
|
||||||
json!({
|
"outerToolName": call.name,
|
||||||
"toolCallId": call.call_id,
|
"toolIdentifier": tool_identifier(&call.name, pending.dynamic_tools),
|
||||||
"outerToolName": call.name,
|
"isDynamic": pending.dynamic_tools.contains(&call.name),
|
||||||
"toolIdentifier": tool_identifier(&call.name, pending.dynamic_tools),
|
"allowedToolNames": pending.allowed_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(),
|
.collect(),
|
||||||
),
|
),
|
||||||
|
|||||||
@@ -30,6 +30,7 @@ fn pending_tool_round_is_one_complete_assistant_message_and_round_trips() {
|
|||||||
name: "Read".into(),
|
name: "Read".into(),
|
||||||
arguments_text: r#"{"path":"/a"}"#.into(),
|
arguments_text: r#"{"path":"/a"}"#.into(),
|
||||||
arguments: json!({"path":"/a"}),
|
arguments: json!({"path":"/a"}),
|
||||||
|
argument_error: Some("Read arguments are not valid JSON".into()),
|
||||||
},
|
},
|
||||||
ToolCall {
|
ToolCall {
|
||||||
index: 1,
|
index: 1,
|
||||||
@@ -38,6 +39,7 @@ fn pending_tool_round_is_one_complete_assistant_message_and_round_trips() {
|
|||||||
name: "Grep".into(),
|
name: "Grep".into(),
|
||||||
arguments_text: r#"{"pattern":"x"}"#.into(),
|
arguments_text: r#"{"pattern":"x"}"#.into(),
|
||||||
arguments: json!({"pattern":"x"}),
|
arguments: json!({"pattern":"x"}),
|
||||||
|
argument_error: None,
|
||||||
},
|
},
|
||||||
];
|
];
|
||||||
let pending = staged_tool_round(
|
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"],
|
wire["providerOptions"]["cursor"]["pendingToolExecutionContracts"]["a"]["toolIdentifier"],
|
||||||
"READ"
|
"READ"
|
||||||
);
|
);
|
||||||
|
assert_eq!(
|
||||||
|
wire["providerOptions"]["cursor"]["pendingToolExecutionContracts"]["a"]["argumentError"],
|
||||||
|
"Read arguments are not valid JSON"
|
||||||
|
);
|
||||||
assert_eq!(wire["role"], "assistant");
|
assert_eq!(wire["role"], "assistant");
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
wire["providerOptions"]["cursor"]["pendingToolExecutionContracts"]
|
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.assistant.replay_state, Some(replay_state));
|
||||||
assert_eq!(recovered.calls.len(), 2);
|
assert_eq!(recovered.calls.len(), 2);
|
||||||
assert_eq!(recovered.calls[0].call_id, "a");
|
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");
|
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) {
|
pub fn discard_model_output(&mut self) {
|
||||||
self.text.clear();
|
self.text.clear();
|
||||||
self.thinking.clear();
|
self.thinking.clear();
|
||||||
@@ -101,6 +106,28 @@ impl StepBuffer {
|
|||||||
mod tests {
|
mod tests {
|
||||||
use super::*;
|
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]
|
#[test]
|
||||||
fn interrupted_model_output_is_not_persisted_as_checkpoint_steps() {
|
fn interrupted_model_output_is_not_persisted_as_checkpoint_steps() {
|
||||||
let mut buffer = StepBuffer::default();
|
let mut buffer = StepBuffer::default();
|
||||||
|
|||||||
@@ -3,7 +3,7 @@ use prost::Message;
|
|||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
cursor::{checkpoint::PendingSteps, protocol::proto::agent::v1 as pb},
|
cursor::{checkpoint::PendingSteps, protocol::proto::agent::v1 as pb},
|
||||||
model::CanonicalMessage,
|
model::{estimate_context_tokens, project_messages, CanonicalMessage, PromptSpec},
|
||||||
store::{BlobEdge, BlobId},
|
store::{BlobEdge, BlobId},
|
||||||
Error, Result,
|
Error, Result,
|
||||||
};
|
};
|
||||||
@@ -85,6 +85,12 @@ impl CheckpointBuilder {
|
|||||||
.push(archive_id.as_bytes().to_vec());
|
.push(archive_id.as_bytes().to_vec());
|
||||||
}
|
}
|
||||||
self.base.self_summary_count = self.base.self_summary_count.saturating_add(1);
|
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() {
|
if let Some(details) = self.base.token_details.as_mut() {
|
||||||
details.breakdown = Some(crate::cursor::services::usage::breakdown(
|
details.breakdown = Some(crate::cursor::services::usage::breakdown(
|
||||||
details.used_tokens,
|
details.used_tokens,
|
||||||
|
|||||||
@@ -12,7 +12,7 @@ use crate::{
|
|||||||
protocol::proto::agent::v1 as pb, services::context_sync::RequestContextSynchronizer,
|
protocol::proto::agent::v1 as pb, services::context_sync::RequestContextSynchronizer,
|
||||||
tools::runtime::McpRoute,
|
tools::runtime::McpRoute,
|
||||||
},
|
},
|
||||||
model::ToolDefinition,
|
model::{normalize_tool_name, ToolDefinition},
|
||||||
store::BlobId,
|
store::BlobId,
|
||||||
Error, Result,
|
Error, Result,
|
||||||
};
|
};
|
||||||
@@ -493,7 +493,7 @@ pub fn dynamic_mcp(
|
|||||||
})?),
|
})?),
|
||||||
};
|
};
|
||||||
let parameters = normalize_mcp_parameters(&wire.name, parameters)?;
|
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 {
|
let definition = ToolDefinition {
|
||||||
name: name.clone(),
|
name: name.clone(),
|
||||||
description: wire.description.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 {
|
fn prost_value(value: &prost_types::Value) -> Value {
|
||||||
use prost_types::value::Kind;
|
use prost_types::value::Kind;
|
||||||
match value.kind.as_ref() {
|
match value.kind.as_ref() {
|
||||||
|
|||||||
@@ -117,9 +117,7 @@ pub(crate) async fn prepare(
|
|||||||
"selected_source": "root_prompt_messages_json",
|
"selected_source": "root_prompt_messages_json",
|
||||||
});
|
});
|
||||||
let encoded = serde_json::to_vec(&summary)?;
|
let encoded = serde_json::to_vec(&summary)?;
|
||||||
trace
|
trace.artifact("history_projection", "byok_server", &encoded, summary);
|
||||||
.artifact("history_projection", "byok_server", &encoded, summary)
|
|
||||||
.await;
|
|
||||||
}
|
}
|
||||||
let mut request_context = context::hydrate(request, context_sync).await?;
|
let mut request_context = context::hydrate(request, context_sync).await?;
|
||||||
if let Some(rules_dir) = local_rules_dir {
|
if let Some(rules_dir) = local_rules_dir {
|
||||||
|
|||||||
@@ -1,6 +1,19 @@
|
|||||||
//! Defines commands accepted by a Conversation runtime.
|
//! 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)]
|
#[derive(Debug)]
|
||||||
pub enum TransportCommand {
|
pub enum TransportCommand {
|
||||||
@@ -8,6 +21,9 @@ pub enum TransportCommand {
|
|||||||
seqno: i64,
|
seqno: i64,
|
||||||
message: Box<pb::AgentClientMessage>,
|
message: Box<pb::AgentClientMessage>,
|
||||||
},
|
},
|
||||||
|
RunFinished {
|
||||||
|
generation: u64,
|
||||||
|
finish: RunFinish,
|
||||||
|
},
|
||||||
Disconnect,
|
Disconnect,
|
||||||
Close,
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -21,7 +21,7 @@ use crate::{
|
|||||||
protocol::proto::agent::v1 as pb,
|
protocol::proto::agent::v1 as pb,
|
||||||
services::blob_sync::BlobSynchronizer,
|
services::blob_sync::BlobSynchronizer,
|
||||||
tools::{
|
tools::{
|
||||||
codec,
|
codec, compat,
|
||||||
runtime::CursorToolRuntime,
|
runtime::CursorToolRuntime,
|
||||||
stream::ToolCallStream,
|
stream::ToolCallStream,
|
||||||
tool_call_result::{ToolCompletion, ToolResultReceiver},
|
tool_call_result::{ToolCompletion, ToolResultReceiver},
|
||||||
@@ -34,7 +34,7 @@ use crate::{
|
|||||||
Error, Result,
|
Error, Result,
|
||||||
};
|
};
|
||||||
|
|
||||||
use super::{CompiledMessages, ConversationRegistry, MessageDelivery};
|
use super::{CompiledMessages, ConversationRegistry, MessageDelivery, RunFinish, TransportFinish};
|
||||||
use crate::cursor::transport::TransportHandle;
|
use crate::cursor::transport::TransportHandle;
|
||||||
|
|
||||||
pub struct ConversationOutput {
|
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;
|
let result = self.run_inner().await;
|
||||||
if let Err(error) = &result {
|
if let Err(error) = &result {
|
||||||
if !self.superseded.is_cancelled() {
|
if !self.superseded.is_cancelled() {
|
||||||
@@ -143,7 +143,7 @@ impl ConversationOutput {
|
|||||||
result
|
result
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn run_inner(&mut self) -> Result<()> {
|
async fn run_inner(&mut self) -> Result<RunFinish> {
|
||||||
if self.context.compacting {
|
if self.context.compacting {
|
||||||
self.handle.emit(&events::summary_started())?;
|
self.handle.emit(&events::summary_started())?;
|
||||||
}
|
}
|
||||||
@@ -175,7 +175,7 @@ impl ConversationOutput {
|
|||||||
if self.superseded.is_cancelled() {
|
if self.superseded.is_cancelled() {
|
||||||
worker.abort();
|
worker.abort();
|
||||||
self.abort_execs().await;
|
self.abort_execs().await;
|
||||||
return Ok(());
|
return Ok(RunFinish::Transport(TransportFinish::Cancelled));
|
||||||
}
|
}
|
||||||
let input = if let Ok(action) = self.runtime_actions.try_recv() {
|
let input = if let Ok(action) = self.runtime_actions.try_recv() {
|
||||||
Input::RuntimeAction(Some(Box::new(action)))
|
Input::RuntimeAction(Some(Box::new(action)))
|
||||||
@@ -187,7 +187,7 @@ impl ConversationOutput {
|
|||||||
_ = self.superseded.cancelled() => {
|
_ = self.superseded.cancelled() => {
|
||||||
worker.abort();
|
worker.abort();
|
||||||
self.abort_execs().await;
|
self.abort_execs().await;
|
||||||
return Ok(());
|
return Ok(RunFinish::Transport(TransportFinish::Cancelled));
|
||||||
}
|
}
|
||||||
action = self.runtime_actions.recv() => Input::RuntimeAction(action.map(Box::new)),
|
action = self.runtime_actions.recv() => Input::RuntimeAction(action.map(Box::new)),
|
||||||
event = self.core.events.recv() => Input::Event(event),
|
event = self.core.events.recv() => Input::Event(event),
|
||||||
@@ -264,6 +264,32 @@ impl ConversationOutput {
|
|||||||
streams.clear();
|
streams.clear();
|
||||||
presentation.discard_model_output();
|
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::TextStart => {}
|
||||||
RunEvent::TextEnd => {
|
RunEvent::TextEnd => {
|
||||||
if !self.context.compacting {
|
if !self.context.compacting {
|
||||||
@@ -312,6 +338,7 @@ impl ConversationOutput {
|
|||||||
name: name.clone(),
|
name: name.clone(),
|
||||||
arguments_text: String::new(),
|
arguments_text: String::new(),
|
||||||
arguments: serde_json::Value::Null,
|
arguments: serde_json::Value::Null,
|
||||||
|
argument_error: None,
|
||||||
};
|
};
|
||||||
self.emit_model_event(
|
self.emit_model_event(
|
||||||
crate::provider::ModelEvent::ToolCallStart {
|
crate::provider::ModelEvent::ToolCallStart {
|
||||||
@@ -335,15 +362,50 @@ impl ConversationOutput {
|
|||||||
let stream = streams.get_mut(&index).ok_or_else(|| {
|
let stream = streams.get_mut(&index).ok_or_else(|| {
|
||||||
Error::Protocol(format!("missing Cursor tool stream: {index}"))
|
Error::Protocol(format!("missing Cursor tool stream: {index}"))
|
||||||
})?;
|
})?;
|
||||||
for message in stream.arguments_delta(call, &delta)? {
|
match stream.arguments_delta(call, &delta) {
|
||||||
self.handle.emit(&message)?;
|
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 } => {
|
RunEvent::ToolCallEnd { index } => {
|
||||||
let call = calls.get_mut(&index).ok_or_else(|| {
|
let call = calls.get_mut(&index).ok_or_else(|| {
|
||||||
Error::Protocol(format!("unknown completed tool index: {index}"))
|
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) => {
|
RunEvent::Usage(usage) => {
|
||||||
if !self.context.compacting {
|
if !self.context.compacting {
|
||||||
@@ -352,10 +414,7 @@ impl ConversationOutput {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
if !self.context.compacting {
|
if !self.context.compacting {
|
||||||
context_tokens = usage
|
context_tokens = usage.context_input_tokens;
|
||||||
.input_tokens
|
|
||||||
.zip(usage.output_tokens)
|
|
||||||
.and_then(|(input, output)| input.checked_add(output));
|
|
||||||
}
|
}
|
||||||
match &mut turn_usage {
|
match &mut turn_usage {
|
||||||
Some(total) => *total += usage,
|
Some(total) => *total += usage,
|
||||||
@@ -424,21 +483,28 @@ impl ConversationOutput {
|
|||||||
streams.clear();
|
streams.clear();
|
||||||
}
|
}
|
||||||
if let CommitCause::RuntimeEvent { event_id } = &state.cause {
|
if let CommitCause::RuntimeEvent { event_id } = &state.cause {
|
||||||
if let Some(injection_id) = event_id.strip_prefix("inject-context:") {
|
// Injections key `pending_injections` by their raw
|
||||||
if let Some(pending) = self.pending_injections.remove(injection_id)
|
// injection id and commit under `inject-context:{id}`,
|
||||||
{
|
// while runtime user messages key it by (and commit
|
||||||
let delivered_at_ms = crate::cursor::tools::runtime::now_ms()
|
// under) the full `user-message:{id}` event id. Strip
|
||||||
.min(i64::MAX as u64)
|
// the injection prefix when present and otherwise use
|
||||||
as i64;
|
// the event id verbatim so both are cleared and emit
|
||||||
self.handle.emit(&events::context_injection_delivered(
|
// their delivered/appended events.
|
||||||
injection_id.to_owned(),
|
let injection_id = event_id
|
||||||
pending.delivery_batch_id.clone(),
|
.strip_prefix("inject-context:")
|
||||||
delivered_at_ms,
|
.unwrap_or(event_id.as_str());
|
||||||
))?;
|
if let Some(pending) = self.pending_injections.remove(injection_id) {
|
||||||
if let Some(user_message) = pending.user_message {
|
let delivered_at_ms = crate::cursor::tools::runtime::now_ms()
|
||||||
self.handle
|
.min(i64::MAX as u64)
|
||||||
.emit(&events::user_message_appended(user_message))?;
|
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()))?
|
.map_err(|_| Error::Protocol("checkpoint worker stopped".into()))?
|
||||||
{
|
{
|
||||||
Ok(checkpoint) => {
|
Ok(checkpoint) => {
|
||||||
|
context_tokens = checkpoint_context_tokens(&checkpoint);
|
||||||
compaction_checkpoint = Some(checkpoint);
|
compaction_checkpoint = Some(checkpoint);
|
||||||
state.barrier.complete(Ok(()));
|
state.barrier.complete(Ok(()));
|
||||||
}
|
}
|
||||||
@@ -652,7 +719,7 @@ impl ConversationOutput {
|
|||||||
if self.superseded.is_cancelled() {
|
if self.superseded.is_cancelled() {
|
||||||
worker.abort();
|
worker.abort();
|
||||||
self.abort_execs().await;
|
self.abort_execs().await;
|
||||||
return Ok(());
|
return Ok(RunFinish::Transport(TransportFinish::Cancelled));
|
||||||
}
|
}
|
||||||
return match outcome {
|
return match outcome {
|
||||||
RunOutcome::Completed => {
|
RunOutcome::Completed => {
|
||||||
@@ -668,8 +735,7 @@ impl ConversationOutput {
|
|||||||
for _ in 0..3 {
|
for _ in 0..3 {
|
||||||
self.checkpoint.publish(&self.handle, &checkpoint).await?;
|
self.checkpoint.publish(&self.handle, &checkpoint).await?;
|
||||||
}
|
}
|
||||||
finish_success(&self.handle);
|
return Ok(RunFinish::TurnCompleted);
|
||||||
return Ok(());
|
|
||||||
}
|
}
|
||||||
let checkpoints = final_checkpoint.take().ok_or_else(|| {
|
let checkpoints = final_checkpoint.take().ok_or_else(|| {
|
||||||
Error::Protocol("Completed without final state".into())
|
Error::Protocol("Completed without final state".into())
|
||||||
@@ -685,18 +751,19 @@ impl ConversationOutput {
|
|||||||
ttft_breakdown: None,
|
ttft_breakdown: None,
|
||||||
message: Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(checkpoints.settled)),
|
message: Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(checkpoints.settled)),
|
||||||
})?;
|
})?;
|
||||||
finish_success(&self.handle);
|
Ok(RunFinish::TurnCompleted)
|
||||||
Ok(())
|
|
||||||
}
|
}
|
||||||
RunOutcome::Cancelled => {
|
RunOutcome::Cancelled => {
|
||||||
worker.abort();
|
worker.abort();
|
||||||
self.abort_execs().await;
|
self.abort_execs().await;
|
||||||
finish_cancelled(&self.handle)
|
Ok(RunFinish::Transport(TransportFinish::Cancelled))
|
||||||
}
|
}
|
||||||
RunOutcome::Failed(failure) => {
|
RunOutcome::Failed(failure) => {
|
||||||
worker.abort();
|
worker.abort();
|
||||||
self.abort_execs().await;
|
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(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn checkpoint_context_tokens(checkpoint: &pb::ConversationStateStructure) -> Option<u64> {
|
||||||
|
checkpoint
|
||||||
|
.token_details
|
||||||
|
.as_ref()
|
||||||
|
.map(|details| u64::from(details.used_tokens))
|
||||||
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
use super::accept_tool_completion;
|
use super::{accept_tool_completion, checkpoint_context_tokens};
|
||||||
use crate::{run::CommandResult, Error};
|
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]
|
#[test]
|
||||||
fn closing_and_ended_runs_ignore_known_tool_completions() {
|
fn closing_and_ended_runs_ignore_known_tool_completions() {
|
||||||
|
|||||||
@@ -11,25 +11,29 @@ use crate::{
|
|||||||
protocol::proto::agent::v1 as pb,
|
protocol::proto::agent::v1 as pb,
|
||||||
services::{blob_sync::BlobSynchronizer, context_sync::RequestContextSynchronizer},
|
services::{blob_sync::BlobSynchronizer, context_sync::RequestContextSynchronizer},
|
||||||
tools::{
|
tools::{
|
||||||
codec,
|
codec, compat,
|
||||||
runtime::CursorToolRuntime,
|
runtime::CursorToolRuntime,
|
||||||
tool_call_result::{tool_result_channel, ToolResultReceiver, ToolResultSender},
|
tool_call_result::{tool_result_channel, ToolResultReceiver, ToolResultSender},
|
||||||
ClientToolEvent, ToolDispatcher,
|
ClientToolEvent, ToolDispatcher,
|
||||||
},
|
},
|
||||||
transport::{OrderedInbox, TransportHandle},
|
transport::{OrderedInbox, TransportHandle},
|
||||||
},
|
},
|
||||||
run::{CommandResult, RunEngine, RunHandle},
|
run::{CommandResult, RunEngine, RunHandle, RunPhase},
|
||||||
};
|
};
|
||||||
|
|
||||||
use super::{
|
use super::{
|
||||||
CompiledMessages, ConversationDependencies, ConversationOutput, ConversationOutputDependencies,
|
CompiledMessages, ConversationDependencies, ConversationOutput, ConversationOutputDependencies,
|
||||||
ConversationRegistry, MessageDelivery, TransportCommand,
|
ConversationRegistry, MessageDelivery, RunFinish, TransportCommand, TransportFinish,
|
||||||
};
|
};
|
||||||
|
|
||||||
pub struct ConversationRuntime;
|
pub struct ConversationRuntime;
|
||||||
|
|
||||||
|
const CONTINUATION_IDLE_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(2);
|
||||||
|
|
||||||
#[derive(Clone)]
|
#[derive(Clone)]
|
||||||
struct RunGeneration {
|
struct RunGeneration {
|
||||||
|
id: u64,
|
||||||
|
request: pb::AgentRunRequest,
|
||||||
superseded: CancellationToken,
|
superseded: CancellationToken,
|
||||||
finished: CancellationToken,
|
finished: CancellationToken,
|
||||||
run: Arc<parking_lot::Mutex<Option<RunHandle>>>,
|
run: Arc<parking_lot::Mutex<Option<RunHandle>>>,
|
||||||
@@ -41,6 +45,14 @@ struct RunGeneration {
|
|||||||
|
|
||||||
struct FinishGeneration(CancellationToken);
|
struct FinishGeneration(CancellationToken);
|
||||||
|
|
||||||
|
struct TransportActorGuard(TransportHandle);
|
||||||
|
|
||||||
|
impl Drop for TransportActorGuard {
|
||||||
|
fn drop(&mut self) {
|
||||||
|
self.0.close_transport();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
impl Drop for FinishGeneration {
|
impl Drop for FinishGeneration {
|
||||||
fn drop(&mut self) {
|
fn drop(&mut self) {
|
||||||
self.0.cancel();
|
self.0.cancel();
|
||||||
@@ -54,6 +66,7 @@ impl ConversationRuntime {
|
|||||||
mut receiver: mpsc::Receiver<TransportCommand>,
|
mut receiver: mpsc::Receiver<TransportCommand>,
|
||||||
) {
|
) {
|
||||||
tokio::spawn(async move {
|
tokio::spawn(async move {
|
||||||
|
let _actor_guard = TransportActorGuard(handle.clone());
|
||||||
let dependencies = registry.dependencies().clone();
|
let dependencies = registry.dependencies().clone();
|
||||||
let blob_sync = BlobSynchronizer::new(
|
let blob_sync = BlobSynchronizer::new(
|
||||||
handle.request_id().into(),
|
handle.request_id().into(),
|
||||||
@@ -65,19 +78,70 @@ impl ConversationRuntime {
|
|||||||
let context_sync =
|
let context_sync =
|
||||||
RequestContextSynchronizer::new(handle.clone(), dependencies.store.clone());
|
RequestContextSynchronizer::new(handle.clone(), dependencies.store.clone());
|
||||||
let mut current = None::<RunGeneration>;
|
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 {
|
loop {
|
||||||
let command = match receiver.recv().await {
|
let command = if draining {
|
||||||
Some(command) => command,
|
if !handle.admissions_drained() {
|
||||||
None => {
|
tokio::select! {
|
||||||
handle.mark_disconnected();
|
command = receiver.recv() => match command {
|
||||||
if let Some(generation) = current.as_ref() {
|
Some(command) => command,
|
||||||
generation.superseded.cancel();
|
None => {
|
||||||
if let Some(run) = generation.run.lock().clone() {
|
finish_pending(&handle, ¤t, pending_finish.take());
|
||||||
run.cancel();
|
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 {
|
match command {
|
||||||
@@ -92,11 +156,41 @@ impl ConversationRuntime {
|
|||||||
let _ = handle.emit(&codec::abort(id));
|
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;
|
break;
|
||||||
}
|
}
|
||||||
TransportCommand::Close => {
|
TransportCommand::RunFinished { generation, finish } => {
|
||||||
break;
|
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 } => {
|
TransportCommand::Append { seqno, message } => {
|
||||||
for (_seqno, message) in inbox.push(seqno, *message) {
|
for (_seqno, message) in inbox.push(seqno, *message) {
|
||||||
@@ -105,6 +199,12 @@ impl ConversationRuntime {
|
|||||||
Some(pb::agent_client_message::Message::RunRequest(
|
Some(pb::agent_client_message::Message::RunRequest(
|
||||||
request,
|
request,
|
||||||
)) => {
|
)) => {
|
||||||
|
waiting_for_action = false;
|
||||||
|
if draining {
|
||||||
|
handle.reopen();
|
||||||
|
draining = false;
|
||||||
|
pending_finish = None;
|
||||||
|
}
|
||||||
if let Some(conversation_id) =
|
if let Some(conversation_id) =
|
||||||
request.conversation_id.as_deref()
|
request.conversation_id.as_deref()
|
||||||
{
|
{
|
||||||
@@ -117,60 +217,21 @@ impl ConversationRuntime {
|
|||||||
"invalid Cursor conversation id"
|
"invalid Cursor conversation id"
|
||||||
);
|
);
|
||||||
let _ = super::finish_failed(&handle, &error);
|
let _ = super::finish_failed(&handle, &error);
|
||||||
let _ =
|
|
||||||
handle.command(TransportCommand::Close).await;
|
|
||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
let previous_finished =
|
start_generation(
|
||||||
if let Some(previous) = current.take() {
|
®istry,
|
||||||
previous.superseded.cancel();
|
&handle,
|
||||||
if let Some(run) = previous.run.lock().clone() {
|
&dependencies,
|
||||||
run.cancel();
|
&blob_sync,
|
||||||
}
|
&context_sync,
|
||||||
for id in previous
|
&tool_runtime_factory,
|
||||||
.tool_runtime
|
&mut current,
|
||||||
.interrupt_for_run_replacement()
|
&mut next_generation,
|
||||||
.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(),
|
|
||||||
request,
|
request,
|
||||||
dependencies.clone(),
|
)
|
||||||
blob_sync.clone(),
|
.await;
|
||||||
context_sync.clone(),
|
|
||||||
generation,
|
|
||||||
previous_finished,
|
|
||||||
result_receiver,
|
|
||||||
runtime_action_receiver,
|
|
||||||
);
|
|
||||||
}
|
}
|
||||||
Some(pb::agent_client_message::Message::ExecClientMessage(
|
Some(pb::agent_client_message::Message::ExecClientMessage(
|
||||||
message,
|
message,
|
||||||
@@ -262,17 +323,18 @@ impl ConversationRuntime {
|
|||||||
.take_exec(throw.id)
|
.take_exec(throw.id)
|
||||||
.await
|
.await
|
||||||
{
|
{
|
||||||
Some(pending) => generation.results.send_error(
|
Some(pending) => generation.results.send(
|
||||||
crate::Error::Protocol(format!(
|
compat::failure_with_message(
|
||||||
"Exec {} failed: {}",
|
&pending.call,
|
||||||
pending.call.call_id, throw.error
|
format!(
|
||||||
)),
|
"Exec {} failed: {}",
|
||||||
|
pending.call.call_id, throw.error
|
||||||
|
),
|
||||||
|
),
|
||||||
),
|
),
|
||||||
None => generation.results.send_error(
|
None => tracing::warn!(
|
||||||
crate::Error::Protocol(format!(
|
id = throw.id,
|
||||||
"unknown ExecClientThrow id: {}",
|
"ignoring failure for unknown tool execution"
|
||||||
throw.id
|
|
||||||
)),
|
|
||||||
),
|
),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -330,27 +392,48 @@ impl ConversationRuntime {
|
|||||||
// return an explicit Protocol Error rather than falling through silently.
|
// return an explicit Protocol Error rather than falling through silently.
|
||||||
Some(
|
Some(
|
||||||
pb::agent_client_message::Message::ConversationAction(
|
pb::agent_client_message::Message::ConversationAction(
|
||||||
action,
|
conversation_action,
|
||||||
),
|
),
|
||||||
) => match action.action {
|
) => match conversation_action.action.clone() {
|
||||||
Some(
|
Some(
|
||||||
pb::conversation_action::Action::UserMessageAction(
|
pb::conversation_action::Action::UserMessageAction(
|
||||||
action,
|
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;
|
continue;
|
||||||
};
|
};
|
||||||
if generation
|
let mut request = previous.request.clone();
|
||||||
.runtime_actions
|
request.action = Some(conversation_action);
|
||||||
.send(compile::RuntimeAction::UserMessage(action))
|
request.conversation_state = None;
|
||||||
.is_err()
|
request.pre_fetched_blobs.clear();
|
||||||
{
|
waiting_for_action = false;
|
||||||
generation.results.send_error(crate::Error::Protocol(
|
start_generation(
|
||||||
"UserMessageAction arrived without an active Run"
|
®istry,
|
||||||
.into(),
|
&handle,
|
||||||
));
|
&dependencies,
|
||||||
}
|
&blob_sync,
|
||||||
|
&context_sync,
|
||||||
|
&tool_runtime_factory,
|
||||||
|
&mut current,
|
||||||
|
&mut next_generation,
|
||||||
|
request,
|
||||||
|
)
|
||||||
|
.await;
|
||||||
}
|
}
|
||||||
Some(pb::conversation_action::Action::CancelAction(_)) => {
|
Some(pb::conversation_action::Action::CancelAction(_)) => {
|
||||||
if let Some(generation) = current.as_ref() {
|
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)]
|
#[allow(clippy::too_many_arguments)]
|
||||||
fn spawn_run_request(
|
fn spawn_run_request(
|
||||||
registry: ConversationRegistry,
|
registry: ConversationRegistry,
|
||||||
@@ -486,8 +658,12 @@ fn spawn_run_request(
|
|||||||
%error,
|
%error,
|
||||||
"failed to prepare Cursor Run"
|
"failed to prepare Cursor Run"
|
||||||
);
|
);
|
||||||
let _ = super::finish_failed(&handle, &error);
|
let _ = handle
|
||||||
let _ = handle.command(TransportCommand::Close).await;
|
.command(TransportCommand::RunFinished {
|
||||||
|
generation: generation.id,
|
||||||
|
finish: RunFinish::Transport(TransportFinish::Failed(error)),
|
||||||
|
})
|
||||||
|
.await;
|
||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
@@ -519,8 +695,12 @@ fn spawn_run_request(
|
|||||||
{
|
{
|
||||||
CommandResult::Applied | CommandResult::Duplicate => {
|
CommandResult::Applied | CommandResult::Duplicate => {
|
||||||
if !generation.superseded.is_cancelled() {
|
if !generation.superseded.is_cancelled() {
|
||||||
super::finish_success(&handle);
|
let _ = handle
|
||||||
let _ = handle.command(TransportCommand::Close).await;
|
.command(TransportCommand::RunFinished {
|
||||||
|
generation: generation.id,
|
||||||
|
finish: RunFinish::Transport(TransportFinish::Success),
|
||||||
|
})
|
||||||
|
.await;
|
||||||
}
|
}
|
||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
@@ -544,8 +724,12 @@ fn spawn_run_request(
|
|||||||
}
|
}
|
||||||
CommandResult::StaleTarget => {
|
CommandResult::StaleTarget => {
|
||||||
if !generation.superseded.is_cancelled() {
|
if !generation.superseded.is_cancelled() {
|
||||||
super::finish_success(&handle);
|
let _ = handle
|
||||||
let _ = handle.command(TransportCommand::Close).await;
|
.command(TransportCommand::RunFinished {
|
||||||
|
generation: generation.id,
|
||||||
|
finish: RunFinish::Transport(TransportFinish::Success),
|
||||||
|
})
|
||||||
|
.await;
|
||||||
}
|
}
|
||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
@@ -604,16 +788,21 @@ fn spawn_run_request(
|
|||||||
tool_runtime: generation.tool_runtime.clone(),
|
tool_runtime: generation.tool_runtime.clone(),
|
||||||
},
|
},
|
||||||
);
|
);
|
||||||
if let Err(error) = output.run().await {
|
let finish = match output.run().await {
|
||||||
if !generation.superseded.is_cancelled() {
|
Ok(finish) => finish,
|
||||||
tracing::error!(
|
Err(error) => {
|
||||||
request_id = handle.request_id(),
|
if generation.superseded.is_cancelled() {
|
||||||
%error,
|
RunFinish::Transport(TransportFinish::Cancelled)
|
||||||
"Cursor session failed"
|
} else {
|
||||||
);
|
tracing::error!(
|
||||||
let _ = super::finish_failed(&handle, &error);
|
request_id = handle.request_id(),
|
||||||
|
%error,
|
||||||
|
"Cursor session failed"
|
||||||
|
);
|
||||||
|
RunFinish::Transport(TransportFinish::Failed(error))
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
};
|
||||||
let _ = core_run.await;
|
let _ = core_run.await;
|
||||||
registry.release(&conversation_id, &run_id).await;
|
registry.release(&conversation_id, &run_id).await;
|
||||||
if generation
|
if generation
|
||||||
@@ -625,7 +814,12 @@ fn spawn_run_request(
|
|||||||
*generation.run.lock() = None;
|
*generation.run.lock() = None;
|
||||||
}
|
}
|
||||||
if !generation.superseded.is_cancelled() {
|
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<()>>;
|
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)]
|
#[derive(Clone)]
|
||||||
pub struct BlobSynchronizer {
|
pub struct BlobSynchronizer {
|
||||||
inner: Arc<Inner>,
|
inner: Arc<Inner>,
|
||||||
@@ -73,22 +76,20 @@ impl BlobSynchronizer {
|
|||||||
let id = self.inner.store.put_blob(data, edges).await?;
|
let id = self.inner.store.put_blob(data, edges).await?;
|
||||||
let result = self.ensure_set(&id, data).await;
|
let result = self.ensure_set(&id, data).await;
|
||||||
if let Some(trace) = self.inner.handle.trace() {
|
if let Some(trace) = self.inner.handle.trace() {
|
||||||
trace
|
trace.linked_blob(
|
||||||
.linked_blob(
|
"blob_set",
|
||||||
"blob_set",
|
"byok_server",
|
||||||
"byok_server",
|
&id,
|
||||||
&id,
|
serde_json::json!({
|
||||||
serde_json::json!({
|
"byte_count": data.len(),
|
||||||
"byte_count": data.len(),
|
"status": if result.is_ok() { "acknowledged" } else { "error" },
|
||||||
"status": if result.is_ok() { "acknowledged" } else { "error" },
|
"error": result.as_ref().err().map(ToString::to_string),
|
||||||
"error": result.as_ref().err().map(ToString::to_string),
|
"edges": edges.iter().map(|edge| serde_json::json!({
|
||||||
"edges": edges.iter().map(|edge| serde_json::json!({
|
"child_blob_id": edge.child.to_base64(),
|
||||||
"child_blob_id": edge.child.to_base64(),
|
"field_name": edge.field_name,
|
||||||
"field_name": edge.field_name,
|
})).collect::<Vec<_>>(),
|
||||||
})).collect::<Vec<_>>(),
|
}),
|
||||||
}),
|
);
|
||||||
)
|
|
||||||
.await;
|
|
||||||
}
|
}
|
||||||
result?;
|
result?;
|
||||||
Ok(id)
|
Ok(id)
|
||||||
@@ -130,7 +131,7 @@ impl BlobSynchronizer {
|
|||||||
let result = tokio::select! {
|
let result = tokio::select! {
|
||||||
result = receiver => result.map_err(|_| Error::Protocol("KV SET response channel closed".into()))?,
|
result = receiver => result.map_err(|_| Error::Protocol("KV SET response channel closed".into()))?,
|
||||||
_ = cancellation.cancelled() => Err(Error::Cancelled),
|
_ = 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() {
|
if result.is_err() {
|
||||||
self.inner.set_requests.lock().await.remove(&id);
|
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>>> {
|
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(data) = self.inner.store.get_blob(blob_id).await? {
|
||||||
if let Some(trace) = self.inner.handle.trace() {
|
if let Some(trace) = self.inner.handle.trace() {
|
||||||
trace
|
trace.linked_blob(
|
||||||
.linked_blob(
|
"blob_get",
|
||||||
"blob_get",
|
"byok_server",
|
||||||
"byok_server",
|
blob_id,
|
||||||
blob_id,
|
serde_json::json!({
|
||||||
serde_json::json!({
|
"byte_count": data.len(),
|
||||||
"byte_count": data.len(),
|
"source": "local_store",
|
||||||
"source": "local_store",
|
"status": "found",
|
||||||
"status": "found",
|
}),
|
||||||
}),
|
);
|
||||||
)
|
|
||||||
.await;
|
|
||||||
}
|
}
|
||||||
return Ok(Some(data));
|
return Ok(Some(data));
|
||||||
}
|
}
|
||||||
@@ -183,7 +182,7 @@ impl BlobSynchronizer {
|
|||||||
let result = tokio::select! {
|
let result = tokio::select! {
|
||||||
result = receiver => result.map_err(|_| Error::Protocol("KV GET response channel closed".into()))?,
|
result = receiver => result.map_err(|_| Error::Protocol("KV GET response channel closed".into()))?,
|
||||||
_ = cancellation.cancelled() => Err(Error::Cancelled),
|
_ = 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() {
|
if result.is_err() {
|
||||||
self.inner.get_requests.lock().await.remove(&id);
|
self.inner.get_requests.lock().await.remove(&id);
|
||||||
@@ -191,45 +190,39 @@ impl BlobSynchronizer {
|
|||||||
if let Some(trace) = self.inner.handle.trace() {
|
if let Some(trace) = self.inner.handle.trace() {
|
||||||
match &result {
|
match &result {
|
||||||
Ok(Some(data)) => {
|
Ok(Some(data)) => {
|
||||||
trace
|
trace.linked_blob(
|
||||||
.linked_blob(
|
"blob_get",
|
||||||
"blob_get",
|
"cursor_client",
|
||||||
"cursor_client",
|
blob_id,
|
||||||
blob_id,
|
serde_json::json!({
|
||||||
serde_json::json!({
|
"byte_count": data.len(),
|
||||||
"byte_count": data.len(),
|
"source": "cursor_client",
|
||||||
"source": "cursor_client",
|
"status": "found",
|
||||||
"status": "found",
|
}),
|
||||||
}),
|
);
|
||||||
)
|
|
||||||
.await;
|
|
||||||
}
|
}
|
||||||
Ok(None) => {
|
Ok(None) => {
|
||||||
trace
|
trace.artifact(
|
||||||
.artifact(
|
"blob_get",
|
||||||
"blob_get",
|
"cursor_client",
|
||||||
"cursor_client",
|
&[],
|
||||||
&[],
|
serde_json::json!({
|
||||||
serde_json::json!({
|
"blob_id": blob_id.to_base64(),
|
||||||
"blob_id": blob_id.to_base64(),
|
"status": "missing",
|
||||||
"status": "missing",
|
}),
|
||||||
}),
|
);
|
||||||
)
|
|
||||||
.await;
|
|
||||||
}
|
}
|
||||||
Err(error) => {
|
Err(error) => {
|
||||||
trace
|
trace.artifact(
|
||||||
.artifact(
|
"blob_get",
|
||||||
"blob_get",
|
"cursor_client",
|
||||||
"cursor_client",
|
&[],
|
||||||
&[],
|
serde_json::json!({
|
||||||
serde_json::json!({
|
"blob_id": blob_id.to_base64(),
|
||||||
"blob_id": blob_id.to_base64(),
|
"status": "error",
|
||||||
"status": "error",
|
"error": error.to_string(),
|
||||||
"error": error.to_string(),
|
}),
|
||||||
}),
|
);
|
||||||
)
|
|
||||||
.await;
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -318,3 +311,18 @@ impl BlobSynchronizer {
|
|||||||
Ok(())
|
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 {
|
let tooltip = TooltipData {
|
||||||
markdown_content: model.description.clone(),
|
markdown_content: model.description.clone(),
|
||||||
};
|
};
|
||||||
let contexts = context_options(model.context_window_tokens);
|
// Effort 与上下文档位由宿主统一提供,与内置模型一致;插件不再声明这两项。
|
||||||
let variants = model_variants(
|
let contexts = context_options(None);
|
||||||
&model.id,
|
let variants = model_variants(&model.id, &model.display_name, &tooltip, &contexts, true);
|
||||||
&model.display_name,
|
|
||||||
&tooltip,
|
|
||||||
&contexts,
|
|
||||||
model.thinking,
|
|
||||||
);
|
|
||||||
let legacy_slugs = variants
|
let legacy_slugs = variants
|
||||||
.iter()
|
.iter()
|
||||||
.filter_map(|variant| variant.legacy_slug.clone())
|
.filter_map(|variant| variant.legacy_slug.clone())
|
||||||
@@ -598,7 +593,7 @@ fn available_plugin_model(model: &PluginModelDescriptor) -> AvailableModel {
|
|||||||
supports_agent: Some(true),
|
supports_agent: Some(true),
|
||||||
degradation_status: Some(0),
|
degradation_status: Some(0),
|
||||||
tooltip_data: Some(tooltip.clone()),
|
tooltip_data: Some(tooltip.clone()),
|
||||||
supports_thinking: Some(model.thinking),
|
supports_thinking: Some(true),
|
||||||
supports_images: Some(model.images),
|
supports_images: Some(model.images),
|
||||||
supports_max_mode: Some(false),
|
supports_max_mode: Some(false),
|
||||||
client_display_name: Some(model.display_name.clone()),
|
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()),
|
inputbox_short_model_name: Some(model.display_name.clone()),
|
||||||
supports_sandboxing: Some(true),
|
supports_sandboxing: Some(true),
|
||||||
supports_cmd_k: Some(false),
|
supports_cmd_k: Some(false),
|
||||||
parameter_definitions: model_parameters(&contexts, model.thinking),
|
parameter_definitions: model_parameters(&contexts, true),
|
||||||
variants,
|
variants,
|
||||||
legacy_slugs,
|
legacy_slugs,
|
||||||
named_model_section_index: Some(1),
|
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_model_id: model.id.clone(),
|
||||||
display_name: model.display_name.clone(),
|
display_name: model.display_name.clone(),
|
||||||
display_name_short: 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()
|
..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 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 render::{dynamic_mcp_placeholder, render_dynamic_mcp, tool_completed};
|
||||||
pub use request::{abort, mcp_request, mcp_state_request, request};
|
pub use request::{abort, mcp_request, mcp_state_request, request};
|
||||||
pub(crate) use request::{edit_read_request, json_object_to_prost, mcp_meta_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(
|
pub(crate) fn create_plan_partial(
|
||||||
call: &ToolCall,
|
call: &ToolCall,
|
||||||
name: &str,
|
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> {
|
pub fn tool_placeholder(name: &str, call_id: &str) -> Result<pb::ToolCall> {
|
||||||
use pb::tool_call::Tool;
|
use pb::tool_call::Tool;
|
||||||
let tool = match normalized(name).as_str() {
|
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()),
|
"delete" => Tool::DeleteToolCall(pb::DeleteToolCall::default()),
|
||||||
"glob" => Tool::GlobToolCall(pb::GlobToolCall::default()),
|
"glob" => Tool::GlobToolCall(pb::GlobToolCall::default()),
|
||||||
"grep" => Tool::GrepToolCall(pb::GrepToolCall::default()),
|
"grep" => Tool::GrepToolCall(pb::GrepToolCall::default()),
|
||||||
@@ -527,3 +570,23 @@ fn now_ms() -> u64 {
|
|||||||
.unwrap_or_default()
|
.unwrap_or_default()
|
||||||
.as_millis() as u64
|
.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)
|
.map(|v| v as i32)
|
||||||
};
|
};
|
||||||
let message = match normalize(&call.name).as_str() {
|
let message = match normalize(&call.name).as_str() {
|
||||||
"shell" => {
|
"shell" | "bash" => {
|
||||||
let command = string("command")?;
|
let command = string("command")?;
|
||||||
let (simple_commands, parsing_result) = shell_command_metadata(&command);
|
let (simple_commands, parsing_result) = shell_command_metadata(&command);
|
||||||
Message::ShellStreamArgs(pb::ShellArgs {
|
Message::ShellStreamArgs(pb::ShellArgs {
|
||||||
@@ -520,3 +520,38 @@ fn prost_value(value: &Value) -> prost_types::Value {
|
|||||||
};
|
};
|
||||||
ProstValue { kind: Some(kind) }
|
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::{
|
cursor::{
|
||||||
protocol::{events, proto::agent::v1 as pb},
|
protocol::{events, proto::agent::v1 as pb},
|
||||||
tools::{
|
tools::{
|
||||||
edit,
|
compat, edit,
|
||||||
runtime::{CursorToolRuntime, ExecStage, PendingExec},
|
runtime::{CursorToolRuntime, ExecStage, PendingExec},
|
||||||
tool_call_result::{self as result, ToolCompletion},
|
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 {
|
let call = match pending.exec_call(message.id).await {
|
||||||
Some(call) => call,
|
Some(call) => call,
|
||||||
None if pending.completed_call(message.id).await.is_some() => {
|
None if pending.completed_call(message.id).await.is_some() => {
|
||||||
return Err(Error::Protocol(format!(
|
tracing::warn!(id = message.id, "ignoring duplicate terminal tool response");
|
||||||
"duplicate terminal ExecClientMessage id: {}",
|
return Ok(ClientExecEvent::Pending);
|
||||||
message.id
|
|
||||||
)))
|
|
||||||
}
|
}
|
||||||
None => {
|
None => {
|
||||||
return Err(Error::Protocol(format!(
|
tracing::warn!(
|
||||||
"unknown ExecClientMessage id: {}",
|
id = message.id,
|
||||||
message.id
|
"ignoring response for unknown tool execution"
|
||||||
)))
|
);
|
||||||
|
return Ok(ClientExecEvent::Pending);
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
let Some(wire_result) = &message.message else {
|
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);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let error = "Cursor Exec stream closed before returning a terminal result";
|
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
|
let command = entry
|
||||||
.call
|
.call
|
||||||
.arguments
|
.arguments
|
||||||
@@ -173,9 +173,22 @@ pub async fn stream_closed(id: u32, pending: &CursorToolRuntime) -> Result<Optio
|
|||||||
}
|
}
|
||||||
let rendered = match &entry.stage {
|
let rendered = match &entry.stage {
|
||||||
ExecStage::DynamicMcp(definition) => {
|
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(
|
Ok(Some(ToolCompletion::from_rendered(
|
||||||
&entry.call,
|
&entry.call,
|
||||||
@@ -211,10 +224,10 @@ async fn advance_edit(
|
|||||||
pb::exec_client_message::Message::ReadResult(result)
|
pb::exec_client_message::Message::ReadResult(result)
|
||||||
| pb::exec_client_message::Message::RedactedReadResult(result) => result,
|
| pb::exec_client_message::Message::RedactedReadResult(result) => result,
|
||||||
_ => {
|
_ => {
|
||||||
return Err(Error::Protocol(format!(
|
let message = format!("expected ReadResult for edit tool {}", entry.call.name);
|
||||||
"expected ReadResult for edit tool {}",
|
return Ok(ClientExecEvent::Completed(Box::new(
|
||||||
entry.call.name
|
compat::failure_with_message(&entry.call, message),
|
||||||
)))
|
)));
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
let write = match edit::after_read(&entry.call, read) {
|
let write = match edit::after_read(&entry.call, read) {
|
||||||
@@ -259,9 +272,14 @@ fn completed(
|
|||||||
pending: PendingExec,
|
pending: PendingExec,
|
||||||
result: pb::exec_client_message::Message,
|
result: pb::exec_client_message::Message,
|
||||||
) -> Result<ClientExecEvent> {
|
) -> Result<ClientExecEvent> {
|
||||||
Ok(ClientExecEvent::Completed(Box::new(result::from_exec(
|
let call = pending.call.clone();
|
||||||
pending, &result,
|
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(
|
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 {
|
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
|
let arguments = call
|
||||||
.arguments
|
.arguments
|
||||||
.as_object()
|
.as_object()
|
||||||
|
|||||||
@@ -183,6 +183,9 @@ fn edit_notebook(call: &ToolCall, before: &str) -> std::result::Result<String, S
|
|||||||
.unwrap_or_default();
|
.unwrap_or_default();
|
||||||
let old =
|
let old =
|
||||||
normalize_newlines(&string(call, "old_string").map_err(|error| error.to_string())?);
|
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 occurrences = source.match_indices(&old).count();
|
||||||
let edited = match occurrences {
|
let edited = match occurrences {
|
||||||
0 => return Err("old_string was not found in the notebook cell".into()),
|
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)
|
.flat_map(char::to_lowercase)
|
||||||
.collect()
|
.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 message_index = first_tool_index + position;
|
||||||
let publish_started = !state.started.contains(&call.call_id);
|
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) {
|
let edit_path = if dynamic_mcp.contains_key(&call.name) {
|
||||||
None
|
None
|
||||||
} else {
|
} 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 {
|
if let Some(path) = edit_path {
|
||||||
let next = self.edit_schedule.lock().await.start_or_defer(
|
let next = self.edit_schedule.lock().await.start_or_defer(
|
||||||
@@ -121,22 +131,28 @@ impl ToolDispatcher {
|
|||||||
let Some(next) = next else {
|
let Some(next) = next else {
|
||||||
continue;
|
continue;
|
||||||
};
|
};
|
||||||
dispatched.push(
|
let started = self
|
||||||
self.start(
|
.start(
|
||||||
&next.call,
|
&next.call,
|
||||||
next.message_index,
|
next.message_index,
|
||||||
next.publish_started,
|
next.publish_started,
|
||||||
dynamic_mcp,
|
dynamic_mcp,
|
||||||
&next.context,
|
&next.context,
|
||||||
)
|
)
|
||||||
.await?,
|
.await;
|
||||||
);
|
dispatched.push(match started {
|
||||||
|
Ok(started) => started,
|
||||||
|
Err(error) => recover_validation_failure(&next.call, error)?,
|
||||||
|
});
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
dispatched.push(
|
let started = self
|
||||||
self.start(call, message_index, publish_started, dynamic_mcp, context)
|
.start(call, message_index, publish_started, dynamic_mcp, context)
|
||||||
.await?,
|
.await;
|
||||||
);
|
dispatched.push(match started {
|
||||||
|
Ok(started) => started,
|
||||||
|
Err(error) => recover_validation_failure(call, error)?,
|
||||||
|
});
|
||||||
}
|
}
|
||||||
Ok(dispatched)
|
Ok(dispatched)
|
||||||
}
|
}
|
||||||
@@ -146,15 +162,19 @@ impl ToolDispatcher {
|
|||||||
let Some(next) = next else {
|
let Some(next) = next else {
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
self.start(
|
match self
|
||||||
&next.call,
|
.start(
|
||||||
next.message_index,
|
&next.call,
|
||||||
next.publish_started,
|
next.message_index,
|
||||||
&BTreeMap::new(),
|
next.publish_started,
|
||||||
&next.context,
|
&BTreeMap::new(),
|
||||||
)
|
&next.context,
|
||||||
.await
|
)
|
||||||
.map(Some)
|
.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> {
|
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 {
|
let pending = match self.runtime.take_interaction(response.id).await {
|
||||||
Some(pending) => pending,
|
Some(pending) => pending,
|
||||||
None if self.runtime.completed_call(response.id).await.is_some() => {
|
None if self.runtime.completed_call(response.id).await.is_some() => {
|
||||||
return Err(Error::Protocol(format!(
|
tracing::warn!(
|
||||||
"duplicate terminal InteractionResponse id: {}",
|
id = response.id,
|
||||||
response.id
|
"ignoring duplicate terminal interaction response"
|
||||||
)));
|
);
|
||||||
|
return Ok(ClientToolEvent::Pending);
|
||||||
}
|
}
|
||||||
None => {
|
None => {
|
||||||
return Err(Error::Protocol(format!(
|
tracing::warn!(
|
||||||
"unknown InteractionResponse id: {}",
|
id = response.id,
|
||||||
response.id
|
"ignoring response for unknown interaction"
|
||||||
)));
|
);
|
||||||
|
return Ok(ClientToolEvent::Pending);
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
Ok(
|
let call = pending.call.clone();
|
||||||
match tool_call_dispatch::resume_interaction(
|
let continuation = match tool_call_dispatch::resume_interaction(
|
||||||
&self.results,
|
&self.results,
|
||||||
&self.search,
|
&self.search,
|
||||||
&self.fetch,
|
&self.fetch,
|
||||||
pending,
|
pending,
|
||||||
response,
|
response,
|
||||||
)
|
|
||||||
.await?
|
|
||||||
{
|
|
||||||
tool_call_dispatch::InteractionContinuation::Completed(completion) => {
|
|
||||||
ClientToolEvent::Completed(completion)
|
|
||||||
}
|
|
||||||
tool_call_dispatch::InteractionContinuation::Pending => ClientToolEvent::Pending,
|
|
||||||
},
|
|
||||||
)
|
)
|
||||||
|
.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),
|
DynamicMcp(pb::McpToolDefinition),
|
||||||
Edit(EditProjection),
|
Edit(EditProjection),
|
||||||
CreatePlan(CreatePlanProjection),
|
CreatePlan(CreatePlanProjection),
|
||||||
|
Task(TaskProjection),
|
||||||
}
|
}
|
||||||
|
|
||||||
struct EditProjection {
|
struct EditProjection {
|
||||||
@@ -38,6 +39,17 @@ struct CreatePlanProjection {
|
|||||||
overview: String,
|
overview: String,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[derive(Default)]
|
||||||
|
struct TaskProjection {
|
||||||
|
fields: JsonStringFields,
|
||||||
|
description: String,
|
||||||
|
prompt: String,
|
||||||
|
subagent_type: String,
|
||||||
|
model: String,
|
||||||
|
resume: String,
|
||||||
|
environment: String,
|
||||||
|
}
|
||||||
|
|
||||||
impl ToolCallStream {
|
impl ToolCallStream {
|
||||||
pub fn new(name: &str, dynamic_mcp: Option<&pb::McpToolDefinition>) -> Self {
|
pub fn new(name: &str, dynamic_mcp: Option<&pb::McpToolDefinition>) -> Self {
|
||||||
let presentation = match dynamic_mcp {
|
let presentation = match dynamic_mcp {
|
||||||
@@ -49,6 +61,7 @@ impl ToolCallStream {
|
|||||||
Presentation::Edit(EditProjection::new("target_notebook", "new_string"))
|
Presentation::Edit(EditProjection::new("target_notebook", "new_string"))
|
||||||
}
|
}
|
||||||
"createplan" => Presentation::CreatePlan(CreatePlanProjection::default()),
|
"createplan" => Presentation::CreatePlan(CreatePlanProjection::default()),
|
||||||
|
"task" => Presentation::Task(TaskProjection::default()),
|
||||||
_ => Presentation::Plain,
|
_ => Presentation::Plain,
|
||||||
},
|
},
|
||||||
};
|
};
|
||||||
@@ -73,10 +86,48 @@ impl ToolCallStream {
|
|||||||
Ok(messages)
|
Ok(messages)
|
||||||
}
|
}
|
||||||
Presentation::CreatePlan(plan) => plan.project(call, raw_delta),
|
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 {
|
impl CreatePlanProjection {
|
||||||
fn project(&mut self, call: &ToolCall, raw_delta: &str) -> Result<Vec<pb::AgentServerMessage>> {
|
fn project(&mut self, call: &ToolCall, raw_delta: &str) -> Result<Vec<pb::AgentServerMessage>> {
|
||||||
let mut completed_field = false;
|
let mut completed_field = false;
|
||||||
@@ -184,3 +235,54 @@ fn normalized(value: &str) -> String {
|
|||||||
.flat_map(char::to_lowercase)
|
.flat_map(char::to_lowercase)
|
||||||
.collect()
|
.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(),
|
name: "WebFetch".into(),
|
||||||
arguments_text: r#"{"url":"https://example.com"}"#.into(),
|
arguments_text: r#"{"url":"https://example.com"}"#.into(),
|
||||||
arguments: json!({"url": "https://example.com"}),
|
arguments: json!({"url": "https://example.com"}),
|
||||||
|
argument_error: None,
|
||||||
},
|
},
|
||||||
started_at_ms: 1,
|
started_at_ms: 1,
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -15,7 +15,7 @@ use crate::{
|
|||||||
Error, Result,
|
Error, Result,
|
||||||
};
|
};
|
||||||
|
|
||||||
use super::OutputHub;
|
use super::{OutputHub, TransportAdmission, TransportLifecycle};
|
||||||
|
|
||||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||||
pub struct TransportParent {
|
pub struct TransportParent {
|
||||||
@@ -30,7 +30,8 @@ pub struct TransportHandle {
|
|||||||
output: Arc<OutputHub>,
|
output: Arc<OutputHub>,
|
||||||
conversation_id: Arc<OnceLock<String>>,
|
conversation_id: Arc<OnceLock<String>>,
|
||||||
parent: Arc<OnceLock<TransportParent>>,
|
parent: Arc<OnceLock<TransportParent>>,
|
||||||
trace: Option<CursorTraceRecorder>,
|
trace: CursorTraceRecorder,
|
||||||
|
lifecycle: TransportLifecycle,
|
||||||
disconnect: CancellationToken,
|
disconnect: CancellationToken,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -39,7 +40,7 @@ impl TransportHandle {
|
|||||||
request_id: String,
|
request_id: String,
|
||||||
commands: mpsc::Sender<TransportCommand>,
|
commands: mpsc::Sender<TransportCommand>,
|
||||||
output: Arc<OutputHub>,
|
output: Arc<OutputHub>,
|
||||||
trace: Option<CursorTraceRecorder>,
|
trace: CursorTraceRecorder,
|
||||||
) -> Self {
|
) -> Self {
|
||||||
Self {
|
Self {
|
||||||
request_id,
|
request_id,
|
||||||
@@ -48,6 +49,7 @@ impl TransportHandle {
|
|||||||
conversation_id: Arc::new(OnceLock::new()),
|
conversation_id: Arc::new(OnceLock::new()),
|
||||||
parent: Arc::new(OnceLock::new()),
|
parent: Arc::new(OnceLock::new()),
|
||||||
trace,
|
trace,
|
||||||
|
lifecycle: TransportLifecycle::new(),
|
||||||
disconnect: CancellationToken::new(),
|
disconnect: CancellationToken::new(),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -127,12 +129,46 @@ impl TransportHandle {
|
|||||||
self.output.close()
|
self.output.close()
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(crate) async fn wait_closed(&self) {
|
pub(crate) fn trace(&self) -> Option<&CursorTraceRecorder> {
|
||||||
self.output.wait_closed().await;
|
Some(&self.trace)
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(crate) fn trace(&self) -> Option<&CursorTraceRecorder> {
|
pub(crate) fn accepting_appends(&self) -> bool {
|
||||||
self.trace.as_ref()
|
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 {
|
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 handle;
|
||||||
mod inbox;
|
mod inbox;
|
||||||
|
mod lifecycle;
|
||||||
mod output;
|
mod output;
|
||||||
mod registry;
|
mod registry;
|
||||||
|
|
||||||
pub use handle::*;
|
pub use handle::*;
|
||||||
pub use inbox::*;
|
pub use inbox::*;
|
||||||
|
pub(crate) use lifecycle::*;
|
||||||
pub use output::*;
|
pub use output::*;
|
||||||
pub use registry::*;
|
pub use registry::*;
|
||||||
|
|||||||
@@ -1,12 +1,11 @@
|
|||||||
//! Buffers, replays, broadcasts, and atomically closes downstream output.
|
//! Buffers, replays, broadcasts, and atomically closes downstream output.
|
||||||
|
|
||||||
use bytes::Bytes;
|
use bytes::Bytes;
|
||||||
use tokio::sync::{mpsc, Notify};
|
use tokio::sync::mpsc;
|
||||||
|
|
||||||
#[derive(Default)]
|
#[derive(Default)]
|
||||||
pub struct OutputHub {
|
pub struct OutputHub {
|
||||||
state: parking_lot::Mutex<OutputState>,
|
state: parking_lot::Mutex<OutputState>,
|
||||||
closed: Notify,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Default)]
|
#[derive(Default)]
|
||||||
@@ -49,17 +48,6 @@ impl OutputHub {
|
|||||||
state.closed = true;
|
state.closed = true;
|
||||||
state.subscribers.clear();
|
state.subscribers.clear();
|
||||||
drop(state);
|
drop(state);
|
||||||
self.closed.notify_waiters();
|
|
||||||
true
|
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.
|
//! 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 tokio::sync::{mpsc, Mutex, Notify};
|
||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
cursor::{
|
cursor::{
|
||||||
conversation::ConversationRegistry, prompting::PromptCompiler,
|
conversation::ConversationRegistry, prompting::PromptCompiler,
|
||||||
services::observability::CursorTraceRecorder,
|
services::observability::CursorTraceService,
|
||||||
},
|
},
|
||||||
plugin::PluginRegistry,
|
plugin::PluginRegistry,
|
||||||
provider::Provider,
|
provider::Provider,
|
||||||
@@ -24,15 +30,23 @@ pub struct TransportRegistry {
|
|||||||
}
|
}
|
||||||
|
|
||||||
struct RegistryInner {
|
struct RegistryInner {
|
||||||
local: Mutex<HashMap<String, TransportHandle>>,
|
local: Mutex<HashMap<String, LocalTransport>>,
|
||||||
|
next_local_generation: AtomicU64,
|
||||||
upstream: Mutex<HashMap<String, u64>>,
|
upstream: Mutex<HashMap<String, u64>>,
|
||||||
route_changed: Notify,
|
route_changed: Notify,
|
||||||
store: Store,
|
store: Store,
|
||||||
|
traces: CursorTraceService,
|
||||||
web_cache: WebCache,
|
web_cache: WebCache,
|
||||||
plugins: Option<PluginRegistry>,
|
plugins: Option<PluginRegistry>,
|
||||||
conversations: ConversationRegistry,
|
conversations: ConversationRegistry,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[derive(Clone)]
|
||||||
|
struct LocalTransport {
|
||||||
|
generation: u64,
|
||||||
|
handle: TransportHandle,
|
||||||
|
}
|
||||||
|
|
||||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||||
pub enum TransportRoute {
|
pub enum TransportRoute {
|
||||||
Local,
|
Local,
|
||||||
@@ -99,8 +113,10 @@ impl TransportRegistry {
|
|||||||
Self {
|
Self {
|
||||||
inner: Arc::new(RegistryInner {
|
inner: Arc::new(RegistryInner {
|
||||||
local: Mutex::new(HashMap::new()),
|
local: Mutex::new(HashMap::new()),
|
||||||
|
next_local_generation: AtomicU64::new(1),
|
||||||
upstream: Mutex::new(HashMap::new()),
|
upstream: Mutex::new(HashMap::new()),
|
||||||
route_changed: Notify::new(),
|
route_changed: Notify::new(),
|
||||||
|
traces: CursorTraceService::new(store.clone()),
|
||||||
conversations: ConversationRegistry::new(
|
conversations: ConversationRegistry::new(
|
||||||
store.clone(),
|
store.clone(),
|
||||||
provider,
|
provider,
|
||||||
@@ -119,6 +135,13 @@ impl TransportRegistry {
|
|||||||
&self.inner.store
|
&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 {
|
pub fn web_cache(&self) -> &WebCache {
|
||||||
&self.inner.web_cache
|
&self.inner.web_cache
|
||||||
}
|
}
|
||||||
@@ -132,18 +155,37 @@ impl TransportRegistry {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub async fn get_or_create(&self, request_id: &str) -> Result<TransportHandle> {
|
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() {
|
self.get_or_create_for_append(request_id, false).await
|
||||||
return Ok(handle);
|
}
|
||||||
|
|
||||||
|
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 (commands, receiver) = mpsc::channel(128);
|
||||||
let output = Arc::new(OutputHub::default());
|
let output = Arc::new(OutputHub::default());
|
||||||
let trace = CursorTraceRecorder::resume(self.inner.store.clone(), request_id).await;
|
let trace = self.inner.traces.recorder(request_id);
|
||||||
let handle = TransportHandle::new(request_id.into(), commands, output.clone(), trace);
|
trace.resume();
|
||||||
let mut local = self.inner.local.lock().await;
|
let handle = TransportHandle::new(request_id.into(), commands, output, trace);
|
||||||
if let Some(existing) = local.get(request_id).cloned() {
|
let generation = self
|
||||||
return Ok(existing);
|
.inner
|
||||||
}
|
.next_local_generation
|
||||||
local.insert(request_id.into(), handle.clone());
|
.fetch_add(1, Ordering::Relaxed);
|
||||||
|
local.insert(
|
||||||
|
request_id.into(),
|
||||||
|
LocalTransport {
|
||||||
|
generation,
|
||||||
|
handle: handle.clone(),
|
||||||
|
},
|
||||||
|
);
|
||||||
drop(local);
|
drop(local);
|
||||||
self.inner.route_changed.notify_waiters();
|
self.inner.route_changed.notify_waiters();
|
||||||
self.inner
|
self.inner
|
||||||
@@ -152,17 +194,29 @@ impl TransportRegistry {
|
|||||||
|
|
||||||
let registry = Arc::downgrade(&self.inner);
|
let registry = Arc::downgrade(&self.inner);
|
||||||
let request_id = request_id.to_string();
|
let request_id = request_id.to_string();
|
||||||
|
let lifecycle = handle.clone();
|
||||||
tokio::spawn(async move {
|
tokio::spawn(async move {
|
||||||
output.wait_closed().await;
|
lifecycle.wait_transport_closed().await;
|
||||||
if let Some(registry) = registry.upgrade() {
|
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)
|
Ok(handle)
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn local(&self, request_id: &str) -> Option<TransportHandle> {
|
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) {
|
pub async fn mark_upstream(&self, request_id: &str) {
|
||||||
@@ -206,10 +260,13 @@ impl TransportRegistry {
|
|||||||
self.inner.conversations.shutdown().await;
|
self.inner.conversations.shutdown().await;
|
||||||
let handles = std::mem::take(&mut *self.inner.local.lock().await);
|
let handles = std::mem::take(&mut *self.inner.local.lock().await);
|
||||||
self.inner.upstream.lock().await.clear();
|
self.inner.upstream.lock().await.clear();
|
||||||
for handle in handles.into_values() {
|
for transport in handles.into_values() {
|
||||||
handle.disconnect().await;
|
transport.handle.disconnect().await;
|
||||||
let _ =
|
let _ = tokio::time::timeout(
|
||||||
tokio::time::timeout(std::time::Duration::from_secs(2), handle.wait_closed()).await;
|
std::time::Duration::from_secs(2),
|
||||||
|
transport.handle.wait_transport_closed(),
|
||||||
|
)
|
||||||
|
.await;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -6,11 +6,10 @@ mod usage {
|
|||||||
|
|
||||||
use serde::{Deserialize, Serialize};
|
use serde::{Deserialize, Serialize};
|
||||||
|
|
||||||
use super::ProviderType;
|
|
||||||
|
|
||||||
#[derive(Clone, Copy, Debug, Default, Serialize, Deserialize, PartialEq, Eq)]
|
#[derive(Clone, Copy, Debug, Default, Serialize, Deserialize, PartialEq, Eq)]
|
||||||
pub struct Usage {
|
pub struct Usage {
|
||||||
pub input_tokens: Option<u64>,
|
pub input_tokens: Option<u64>,
|
||||||
|
pub context_input_tokens: Option<u64>,
|
||||||
pub output_tokens: Option<u64>,
|
pub output_tokens: Option<u64>,
|
||||||
pub total_tokens: Option<u64>,
|
pub total_tokens: Option<u64>,
|
||||||
pub cache_read_tokens: Option<u64>,
|
pub cache_read_tokens: Option<u64>,
|
||||||
@@ -18,24 +17,10 @@ mod usage {
|
|||||||
pub reasoning_tokens: Option<u64>,
|
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 {
|
impl AddAssign for Usage {
|
||||||
fn add_assign(&mut self, rhs: Self) {
|
fn add_assign(&mut self, rhs: Self) {
|
||||||
self.input_tokens = sum(self.input_tokens, rhs.input_tokens);
|
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.output_tokens = sum(self.output_tokens, rhs.output_tokens);
|
||||||
self.total_tokens = sum(self.total_tokens, rhs.total_tokens);
|
self.total_tokens = sum(self.total_tokens, rhs.total_tokens);
|
||||||
self.cache_read_tokens = sum(self.cache_read_tokens, rhs.cache_read_tokens);
|
self.cache_read_tokens = sum(self.cache_read_tokens, rhs.cache_read_tokens);
|
||||||
@@ -53,7 +38,7 @@ pub use usage::*;
|
|||||||
mod llm_call {
|
mod llm_call {
|
||||||
use serde::Serialize;
|
use serde::Serialize;
|
||||||
|
|
||||||
use super::{ProviderType, Usage};
|
use super::ProviderType;
|
||||||
|
|
||||||
#[derive(Clone, Debug)]
|
#[derive(Clone, Debug)]
|
||||||
pub struct NewLlmCall {
|
pub struct NewLlmCall {
|
||||||
@@ -75,14 +60,6 @@ mod llm_call {
|
|||||||
pub detailed: bool,
|
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)]
|
#[derive(Clone, Debug, Serialize)]
|
||||||
pub struct LlmCallSummary {
|
pub struct LlmCallSummary {
|
||||||
pub call_id: String,
|
pub call_id: String,
|
||||||
|
|||||||
@@ -6,8 +6,8 @@ use serde::{Deserialize, Serialize};
|
|||||||
use crate::{Error, Result};
|
use crate::{Error, Result};
|
||||||
|
|
||||||
use super::{
|
use super::{
|
||||||
CanonicalMessage, ContentPart, MessageContent, ProviderReplayState, Role, ToolCallContent,
|
normalize_tool_name, CanonicalMessage, ContentPart, MessageContent, ProviderReplayState, Role,
|
||||||
ToolResultContent,
|
ToolCallContent, ToolResultContent,
|
||||||
};
|
};
|
||||||
|
|
||||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
|
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
|
||||||
@@ -91,7 +91,7 @@ fn project_tool_round(
|
|||||||
"tool round repeats provider replay state".into(),
|
"tool round repeats provider replay state".into(),
|
||||||
));
|
));
|
||||||
}
|
}
|
||||||
calls.extend(part_calls.iter().cloned());
|
calls.extend(part_calls.iter().map(normalized_tool_call));
|
||||||
cursor += 1;
|
cursor += 1;
|
||||||
|
|
||||||
while cursor < messages.len() {
|
while cursor < messages.len() {
|
||||||
@@ -139,7 +139,7 @@ fn project_tool_round(
|
|||||||
.map(|(message_id, result)| ProjectedMessage {
|
.map(|(message_id, result)| ProjectedMessage {
|
||||||
message_id,
|
message_id,
|
||||||
role: Role::Tool,
|
role: Role::Tool,
|
||||||
content: ProjectedContent::ToolResult(result),
|
content: ProjectedContent::ToolResult(normalized_tool_result(&result)),
|
||||||
}),
|
}),
|
||||||
);
|
);
|
||||||
Ok(Some((output, cursor)))
|
Ok(Some((output, cursor)))
|
||||||
@@ -158,9 +158,11 @@ fn project_message(message: &CanonicalMessage) -> ProjectedMessage {
|
|||||||
text: text.clone(),
|
text: text.clone(),
|
||||||
thinking: thinking.clone(),
|
thinking: thinking.clone(),
|
||||||
replay_state: replay_state.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 {
|
ProjectedMessage {
|
||||||
message_id: message.message_id.clone(),
|
message_id: message.message_id.clone(),
|
||||||
@@ -168,3 +170,50 @@ fn project_message(message: &CanonicalMessage) -> ProjectedMessage {
|
|||||||
content,
|
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> {
|
pub(crate) fn parse_token_count(value: &str) -> Option<u64> {
|
||||||
let value = value.trim().to_ascii_lowercase();
|
let value = value.trim().to_ascii_lowercase();
|
||||||
let (number, multiplier) = match value.chars().last()? {
|
let (number, multiplier) = match value.chars().last()? {
|
||||||
@@ -18,3 +113,136 @@ pub(crate) fn format_token_count(tokens: u64) -> String {
|
|||||||
tokens.to_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;
|
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)]
|
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
|
||||||
pub struct ToolDefinition {
|
pub struct ToolDefinition {
|
||||||
pub name: String,
|
pub name: String,
|
||||||
@@ -19,6 +37,7 @@ pub struct ToolCall {
|
|||||||
pub name: String,
|
pub name: String,
|
||||||
pub arguments_text: String,
|
pub arguments_text: String,
|
||||||
pub arguments: Value,
|
pub arguments: Value,
|
||||||
|
pub argument_error: Option<String>,
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
|
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
|
||||||
|
|||||||
+87
-2
@@ -1,8 +1,93 @@
|
|||||||
//! Provides shared network client and transport configuration.
|
//! Owns reusable outbound HTTP clients configured from persisted proxy settings.
|
||||||
//! Outbound HTTP clients configured from persisted application proxy settings.
|
|
||||||
|
use std::{sync::Arc, time::Duration};
|
||||||
|
|
||||||
|
use tokio::sync::RwLock;
|
||||||
|
|
||||||
use crate::{store::Store, Result};
|
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> {
|
pub async fn client_builder(store: &Store) -> Result<reqwest::ClientBuilder> {
|
||||||
let settings = store.proxy_settings_secret().await?;
|
let settings = store.proxy_settings_secret().await?;
|
||||||
// Use the platform TLS stack for compatibility with provider gateways that
|
// Use the platform TLS stack for compatibility with provider gateways that
|
||||||
|
|||||||
@@ -107,9 +107,7 @@ pub struct PluginModelDescriptor {
|
|||||||
pub description: Option<String>,
|
pub description: Option<String>,
|
||||||
pub icon: String,
|
pub icon: String,
|
||||||
pub provider_type: String,
|
pub provider_type: String,
|
||||||
pub context_window_tokens: Option<u64>,
|
|
||||||
pub max_output_tokens: Option<u64>,
|
pub max_output_tokens: Option<u64>,
|
||||||
pub thinking: bool,
|
|
||||||
pub images: bool,
|
pub images: bool,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -209,9 +207,7 @@ impl PluginModelDescriptor {
|
|||||||
description: model.description.clone(),
|
description: model.description.clone(),
|
||||||
icon: icon.to_owned(),
|
icon: icon.to_owned(),
|
||||||
provider_type: provider.provider_type.clone(),
|
provider_type: provider.provider_type.clone(),
|
||||||
context_window_tokens: model.context_window_tokens,
|
|
||||||
max_output_tokens: model.max_output_tokens,
|
max_output_tokens: model.max_output_tokens,
|
||||||
thinking: model.thinking,
|
|
||||||
images: model.images,
|
images: model.images,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -20,7 +20,6 @@ pub use descriptor::{
|
|||||||
};
|
};
|
||||||
pub use registry::{ImportResponse, OAuthBeginResponse, OAuthPollResponse, PluginRegistry};
|
pub use registry::{ImportResponse, OAuthBeginResponse, OAuthPollResponse, PluginRegistry};
|
||||||
pub use runtime::{PluginRuntime, PluginRuntimePhase, PluginRuntimeState, PluginRuntimeStatus};
|
pub use runtime::{PluginRuntime, PluginRuntimePhase, PluginRuntimeState, PluginRuntimeStatus};
|
||||||
pub(crate) use wire::llm_request as plugin_llm_request;
|
|
||||||
|
|
||||||
/// Windows 下阻止 Deno 子进程弹出控制台窗口(CREATE_NO_WINDOW)。
|
/// Windows 下阻止 Deno 子进程弹出控制台窗口(CREATE_NO_WINDOW)。
|
||||||
#[cfg(windows)]
|
#[cfg(windows)]
|
||||||
|
|||||||
@@ -20,8 +20,11 @@ use super::{
|
|||||||
worker::{PluginWorker, WorkerStreamItem},
|
worker::{PluginWorker, WorkerStreamItem},
|
||||||
};
|
};
|
||||||
use crate::{
|
use crate::{
|
||||||
model::ModelInvocation, provider::ModelEvent, provider::ProviderStream, store::Store, Error,
|
model::ModelInvocation,
|
||||||
Result,
|
provider::ProviderStream,
|
||||||
|
provider::{CallRecorder, ModelEvent},
|
||||||
|
store::Store,
|
||||||
|
Error, Result,
|
||||||
};
|
};
|
||||||
|
|
||||||
const OAUTH_SLOW_DOWN_STEP_MS: i64 = 5_000;
|
const OAUTH_SLOW_DOWN_STEP_MS: i64 = 5_000;
|
||||||
@@ -203,6 +206,7 @@ impl PluginRegistry {
|
|||||||
&self,
|
&self,
|
||||||
invocation: ModelInvocation,
|
invocation: ModelInvocation,
|
||||||
cancellation: CancellationToken,
|
cancellation: CancellationToken,
|
||||||
|
recorder: CallRecorder,
|
||||||
) -> ProviderStream {
|
) -> ProviderStream {
|
||||||
let registry = self.clone();
|
let registry = self.clone();
|
||||||
Box::pin(try_stream! {
|
Box::pin(try_stream! {
|
||||||
@@ -232,7 +236,7 @@ impl PluginRegistry {
|
|||||||
"request": request,
|
"request": request,
|
||||||
});
|
});
|
||||||
let worker = registry.worker(&entry, &executable).await;
|
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() };
|
yield ModelEvent::Start { model_call_id: invocation.call_id.clone() };
|
||||||
while let Some(item) = items.recv().await {
|
while let Some(item) = items.recv().await {
|
||||||
match item {
|
match item {
|
||||||
|
|||||||
@@ -2,7 +2,6 @@ import type { JsonValue, PluginContext } from "./plugin.ts";
|
|||||||
import type { ResourceSnapshot } from "./resource.ts";
|
import type { ResourceSnapshot } from "./resource.ts";
|
||||||
|
|
||||||
export type ModelCapabilities = {
|
export type ModelCapabilities = {
|
||||||
thinking?: boolean;
|
|
||||||
images?: boolean;
|
images?: boolean;
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -10,7 +9,6 @@ export type ModelDefinition = {
|
|||||||
id: string;
|
id: string;
|
||||||
displayName: string;
|
displayName: string;
|
||||||
description?: string;
|
description?: string;
|
||||||
contextWindowTokens?: number;
|
|
||||||
maxOutputTokens?: number;
|
maxOutputTokens?: number;
|
||||||
capabilities?: ModelCapabilities;
|
capabilities?: ModelCapabilities;
|
||||||
/** 之后的调用原样传回;永远不会展示给用户。 */
|
/** 之后的调用原样传回;永远不会展示给用户。 */
|
||||||
|
|||||||
@@ -53,7 +53,17 @@ function replayItems(value: JsonValue): JsonValue[] {
|
|||||||
if (!Array.isArray(items)) {
|
if (!Array.isArray(items)) {
|
||||||
throw new Error("OpenAI Responses replay state is missing 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> {
|
export function buildResponsesBody(call: OpenAiResponsesCall): Record<string, JsonValue> {
|
||||||
|
|||||||
@@ -135,12 +135,8 @@ pub struct StoredModel {
|
|||||||
#[serde(default)]
|
#[serde(default)]
|
||||||
pub description: Option<String>,
|
pub description: Option<String>,
|
||||||
#[serde(default)]
|
#[serde(default)]
|
||||||
pub context_window_tokens: Option<u64>,
|
|
||||||
#[serde(default)]
|
|
||||||
pub max_output_tokens: Option<u64>,
|
pub max_output_tokens: Option<u64>,
|
||||||
#[serde(default)]
|
#[serde(default)]
|
||||||
pub thinking: bool,
|
|
||||||
#[serde(default)]
|
|
||||||
pub images: bool,
|
pub images: bool,
|
||||||
#[serde(default)]
|
#[serde(default)]
|
||||||
pub private_data: serde_json::Value,
|
pub private_data: serde_json::Value,
|
||||||
@@ -179,13 +175,9 @@ impl StoredModel {
|
|||||||
.get("description")
|
.get("description")
|
||||||
.and_then(serde_json::Value::as_str)
|
.and_then(serde_json::Value::as_str)
|
||||||
.map(str::to_owned),
|
.map(str::to_owned),
|
||||||
context_window_tokens: object
|
|
||||||
.get("contextWindowTokens")
|
|
||||||
.and_then(serde_json::Value::as_u64),
|
|
||||||
max_output_tokens: object
|
max_output_tokens: object
|
||||||
.get("maxOutputTokens")
|
.get("maxOutputTokens")
|
||||||
.and_then(serde_json::Value::as_u64),
|
.and_then(serde_json::Value::as_u64),
|
||||||
thinking: capability("thinking"),
|
|
||||||
images: capability("images"),
|
images: capability("images"),
|
||||||
private_data: object
|
private_data: object
|
||||||
.get("privateData")
|
.get("privateData")
|
||||||
@@ -200,9 +192,8 @@ impl StoredModel {
|
|||||||
"id": self.id,
|
"id": self.id,
|
||||||
"displayName": self.display_name,
|
"displayName": self.display_name,
|
||||||
"description": self.description,
|
"description": self.description,
|
||||||
"contextWindowTokens": self.context_window_tokens,
|
|
||||||
"maxOutputTokens": self.max_output_tokens,
|
"maxOutputTokens": self.max_output_tokens,
|
||||||
"capabilities": { "thinking": self.thinking, "images": self.images },
|
"capabilities": { "images": self.images },
|
||||||
"privateData": self.private_data,
|
"privateData": self.private_data,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
@@ -450,7 +441,7 @@ mod tests {
|
|||||||
let model = StoredModel::from_definition(&serde_json::json!({
|
let model = StoredModel::from_definition(&serde_json::json!({
|
||||||
"id": "gpt-test",
|
"id": "gpt-test",
|
||||||
"displayName": "GPT Test",
|
"displayName": "GPT Test",
|
||||||
"capabilities": {"thinking": true},
|
"capabilities": {"images": true},
|
||||||
"privateData": {"reasoningEfforts": ["low"]},
|
"privateData": {"reasoningEfforts": ["low"]},
|
||||||
}))
|
}))
|
||||||
.unwrap();
|
.unwrap();
|
||||||
@@ -460,7 +451,7 @@ mod tests {
|
|||||||
.unwrap();
|
.unwrap();
|
||||||
let models = store.models("dev.example", "codex").await.unwrap();
|
let models = store.models("dev.example", "codex").await.unwrap();
|
||||||
assert_eq!(models.len(), 1);
|
assert_eq!(models.len(), 1);
|
||||||
assert!(models[0].thinking);
|
assert!(models[0].images);
|
||||||
assert_eq!(models[0].private_data["reasoningEfforts"][0], "low");
|
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")
|
.get("usage")
|
||||||
.ok_or_else(|| Error::Protocol("plugin usage event requires usage".into()))?;
|
.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 tokens = |name: &str| usage.get(name).and_then(serde_json::Value::as_u64);
|
||||||
|
let input_tokens = tokens("inputTokens");
|
||||||
ModelEvent::Usage(Usage {
|
ModelEvent::Usage(Usage {
|
||||||
input_tokens: tokens("inputTokens"),
|
input_tokens,
|
||||||
|
context_input_tokens: input_tokens,
|
||||||
output_tokens: tokens("outputTokens"),
|
output_tokens: tokens("outputTokens"),
|
||||||
total_tokens: tokens("totalTokens"),
|
total_tokens: tokens("totalTokens"),
|
||||||
cache_read_tokens: tokens("cacheReadTokens"),
|
cache_read_tokens: tokens("cacheReadTokens"),
|
||||||
@@ -285,6 +287,7 @@ mod tests {
|
|||||||
usage,
|
usage,
|
||||||
ModelEvent::Usage(Usage {
|
ModelEvent::Usage(Usage {
|
||||||
input_tokens: Some(10),
|
input_tokens: Some(10),
|
||||||
|
context_input_tokens: Some(10),
|
||||||
output_tokens: Some(2),
|
output_tokens: Some(2),
|
||||||
total_tokens: None,
|
total_tokens: None,
|
||||||
cache_read_tokens: Some(4),
|
cache_read_tokens: Some(4),
|
||||||
|
|||||||
+269
-24
@@ -3,7 +3,10 @@ use std::{
|
|||||||
collections::{HashMap, HashSet},
|
collections::{HashMap, HashSet},
|
||||||
path::PathBuf,
|
path::PathBuf,
|
||||||
process::Stdio,
|
process::Stdio,
|
||||||
sync::Arc,
|
sync::{
|
||||||
|
atomic::{AtomicBool, Ordering},
|
||||||
|
Arc,
|
||||||
|
},
|
||||||
time::Duration,
|
time::Duration,
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -19,7 +22,7 @@ use super::{
|
|||||||
definition::{file_url, PluginDefinitionLoader},
|
definition::{file_url, PluginDefinitionLoader},
|
||||||
protocol::{HostMessage, WorkerMessage},
|
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 INVOCATION_TIMEOUT: Duration = Duration::from_secs(10 * 60);
|
||||||
const MAX_NETWORK_RESPONSE_BYTES: u64 = 16 * 1024 * 1024;
|
const MAX_NETWORK_RESPONSE_BYTES: u64 = 16 * 1024 * 1024;
|
||||||
@@ -56,12 +59,29 @@ struct WorkerProcess {
|
|||||||
stdin: Arc<Mutex<ChildStdin>>,
|
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)]
|
#[derive(Clone)]
|
||||||
struct HostContext {
|
struct HostContext {
|
||||||
plugin_id: String,
|
plugin_id: String,
|
||||||
network_hosts: Arc<HashSet<String>>,
|
network_hosts: Arc<HashSet<String>>,
|
||||||
store: Store,
|
store: Store,
|
||||||
cancellations: Arc<Mutex<HashMap<String, CancellationToken>>>,
|
invocations: Arc<Mutex<HashMap<String, Arc<InvocationState>>>>,
|
||||||
streams: Arc<Mutex<HashMap<String, StreamLines>>>,
|
streams: Arc<Mutex<HashMap<String, StreamLines>>>,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -87,7 +107,7 @@ impl PluginWorker {
|
|||||||
.collect(),
|
.collect(),
|
||||||
),
|
),
|
||||||
store,
|
store,
|
||||||
cancellations: Arc::new(Mutex::new(HashMap::new())),
|
invocations: Arc::new(Mutex::new(HashMap::new())),
|
||||||
streams: Arc::new(Mutex::new(HashMap::new())),
|
streams: Arc::new(Mutex::new(HashMap::new())),
|
||||||
},
|
},
|
||||||
plugin_id,
|
plugin_id,
|
||||||
@@ -108,7 +128,9 @@ impl PluginWorker {
|
|||||||
params: serde_json::Value,
|
params: serde_json::Value,
|
||||||
cancellation: CancellationToken,
|
cancellation: CancellationToken,
|
||||||
) -> Result<serde_json::Value> {
|
) -> 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 {
|
let result = tokio::time::timeout(INVOCATION_TIMEOUT, async {
|
||||||
while let Some(item) = items.recv().await {
|
while let Some(item) = items.recv().await {
|
||||||
if let WorkerStreamItem::Result(result) = item {
|
if let WorkerStreamItem::Result(result) = item {
|
||||||
@@ -137,15 +159,18 @@ impl PluginWorker {
|
|||||||
method: &str,
|
method: &str,
|
||||||
params: serde_json::Value,
|
params: serde_json::Value,
|
||||||
cancellation: CancellationToken,
|
cancellation: CancellationToken,
|
||||||
|
recorder: Option<CallRecorder>,
|
||||||
) -> Result<mpsc::UnboundedReceiver<WorkerStreamItem>> {
|
) -> Result<mpsc::UnboundedReceiver<WorkerStreamItem>> {
|
||||||
let id = uuid::Uuid::new_v4().to_string();
|
let id = uuid::Uuid::new_v4().to_string();
|
||||||
let request_cancellation = CancellationToken::new();
|
let request_cancellation = CancellationToken::new();
|
||||||
self.inner
|
self.inner.host.invocations.lock().await.insert(
|
||||||
.host
|
id.clone(),
|
||||||
.cancellations
|
Arc::new(InvocationState {
|
||||||
.lock()
|
cancellation: request_cancellation.clone(),
|
||||||
.await
|
recorder,
|
||||||
.insert(id.clone(), request_cancellation.clone());
|
recorder_claimed: AtomicBool::new(false),
|
||||||
|
}),
|
||||||
|
);
|
||||||
let (sender, receiver) = mpsc::unbounded_channel();
|
let (sender, receiver) = mpsc::unbounded_channel();
|
||||||
self.inner
|
self.inner
|
||||||
.pending
|
.pending
|
||||||
@@ -182,10 +207,10 @@ impl PluginWorker {
|
|||||||
}
|
}
|
||||||
let _ = sender.send(WorkerStreamItem::Result(Err(Error::Cancelled)));
|
let _ = sender.send(WorkerStreamItem::Result(Err(Error::Cancelled)));
|
||||||
inner.pending.lock().await.remove(&request_id);
|
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() => {
|
_ = 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) {
|
async fn cleanup(&self, id: &str) {
|
||||||
self.inner.pending.lock().await.remove(id);
|
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>>> {
|
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 {
|
impl HostContext {
|
||||||
async fn call(
|
async fn call(
|
||||||
&self,
|
&self,
|
||||||
@@ -421,7 +468,11 @@ impl HostContext {
|
|||||||
&self,
|
&self,
|
||||||
request_id: &str,
|
request_id: &str,
|
||||||
params: &serde_json::Value,
|
params: &serde_json::Value,
|
||||||
) -> Result<(reqwest::RequestBuilder, CancellationToken)> {
|
) -> Result<(
|
||||||
|
reqwest::RequestBuilder,
|
||||||
|
CancellationToken,
|
||||||
|
Option<CallRecorder>,
|
||||||
|
)> {
|
||||||
let raw_url = required_string(params, "url")?;
|
let raw_url = required_string(params, "url")?;
|
||||||
let url = url::Url::parse(raw_url)
|
let url = url::Url::parse(raw_url)
|
||||||
.map_err(|error| Error::Config(format!("invalid plugin network URL: {error}")))?;
|
.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) {
|
if let Some(body) = params.get("body").and_then(serde_json::Value::as_str) {
|
||||||
request = request.body(body.to_owned());
|
request = request.body(body.to_owned());
|
||||||
}
|
}
|
||||||
let cancellation = self
|
let invocation = self.invocations.lock().await.get(request_id).cloned();
|
||||||
.cancellations
|
let cancellation = invocation
|
||||||
.lock()
|
.as_ref()
|
||||||
.await
|
.map(|state| state.cancellation.clone())
|
||||||
.get(request_id)
|
|
||||||
.cloned()
|
|
||||||
.unwrap_or_default();
|
.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(
|
async fn fetch(
|
||||||
@@ -478,13 +532,16 @@ impl HostContext {
|
|||||||
request_id: &str,
|
request_id: &str,
|
||||||
params: serde_json::Value,
|
params: serde_json::Value,
|
||||||
) -> Result<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 request = request.timeout(Duration::from_secs(60));
|
||||||
let response = tokio::select! {
|
let response = tokio::select! {
|
||||||
_ = cancellation.cancelled() => return Err(Error::Cancelled),
|
_ = cancellation.cancelled() => return Err(Error::Cancelled),
|
||||||
response = request.send() => response?,
|
response = request.send() => response?,
|
||||||
};
|
};
|
||||||
let status = response.status().as_u16();
|
let status = response.status().as_u16();
|
||||||
|
if let Some(recorder) = &recorder {
|
||||||
|
recorder.response_headers(status).await?;
|
||||||
|
}
|
||||||
if response
|
if response
|
||||||
.content_length()
|
.content_length()
|
||||||
.is_some_and(|size| size > MAX_NETWORK_RESPONSE_BYTES)
|
.is_some_and(|size| size > MAX_NETWORK_RESPONSE_BYTES)
|
||||||
@@ -503,6 +560,9 @@ impl HostContext {
|
|||||||
"plugin network response is larger than allowed".into(),
|
"plugin network response is larger than allowed".into(),
|
||||||
));
|
));
|
||||||
}
|
}
|
||||||
|
if let Some(recorder) = &recorder {
|
||||||
|
recorder.response_chunk(&body).await?;
|
||||||
|
}
|
||||||
Ok(
|
Ok(
|
||||||
serde_json::json!({ "status": status, "headers": headers, "body": String::from_utf8_lossy(&body) }),
|
serde_json::json!({ "status": status, "headers": headers, "body": String::from_utf8_lossy(&body) }),
|
||||||
)
|
)
|
||||||
@@ -514,12 +574,15 @@ impl HostContext {
|
|||||||
request_id: &str,
|
request_id: &str,
|
||||||
params: serde_json::Value,
|
params: serde_json::Value,
|
||||||
) -> Result<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! {
|
let response = tokio::select! {
|
||||||
_ = cancellation.cancelled() => return Err(Error::Cancelled),
|
_ = cancellation.cancelled() => return Err(Error::Cancelled),
|
||||||
response = request.send() => response?,
|
response = request.send() => response?,
|
||||||
};
|
};
|
||||||
let status = response.status().as_u16();
|
let status = response.status().as_u16();
|
||||||
|
if let Some(recorder) = &recorder {
|
||||||
|
recorder.response_headers(status).await?;
|
||||||
|
}
|
||||||
let headers = header_map(&response);
|
let headers = header_map(&response);
|
||||||
let (sender, receiver) = mpsc::channel::<Result<String>>(256);
|
let (sender, receiver) = mpsc::channel::<Result<String>>(256);
|
||||||
tokio::spawn(async move {
|
tokio::spawn(async move {
|
||||||
@@ -552,6 +615,12 @@ impl HostContext {
|
|||||||
.await;
|
.await;
|
||||||
return;
|
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);
|
buffered.extend_from_slice(&chunk);
|
||||||
while let Some(position) = buffered.iter().position(|byte| *byte == b'\n') {
|
while let Some(position) = buffered.iter().position(|byte| *byte == b'\n') {
|
||||||
let mut line = buffered.drain(..=position).collect::<Vec<u8>>();
|
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)
|
.and_then(serde_json::Value::as_str)
|
||||||
.ok_or_else(|| Error::Protocol(format!("plugin host call requires string '{key}'")))
|
.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::{
|
use super::{
|
||||||
|
attempt::{send_once, Attempt},
|
||||||
map_sse_error, merge_extra_params, provider_event_error,
|
map_sse_error, merge_extra_params, provider_event_error,
|
||||||
recorder::recorded_headers,
|
recorder::recorded_headers,
|
||||||
retry::{send_with_retry, Attempt, RetryPolicy},
|
|
||||||
CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream,
|
CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream,
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -90,17 +90,14 @@ impl Provider for AnthropicProvider {
|
|||||||
if let Some(recorder) = &recorder {
|
if let Some(recorder) = &recorder {
|
||||||
recorder.request(request_headers.clone(), &body).await?;
|
recorder.request(request_headers.clone(), &body).await?;
|
||||||
}
|
}
|
||||||
let attempt = send_with_retry(
|
let attempt = send_once(
|
||||||
"Anthropic",
|
"Anthropic",
|
||||||
|| client.post(&config.request_url)
|
|| client.post(&config.request_url)
|
||||||
.header("x-api-key", &config.api_key).header("anthropic-version", "2023-06-01")
|
.header("x-api-key", &config.api_key).header("anthropic-version", "2023-06-01")
|
||||||
.headers(config.custom_headers.clone())
|
.headers(config.custom_headers.clone())
|
||||||
.json(&body),
|
.json(&body),
|
||||||
RetryPolicy { retries: config.retry_count, ..RetryPolicy::default() },
|
|
||||||
&cancellation,
|
&cancellation,
|
||||||
recorder.as_ref(),
|
recorder.as_ref(),
|
||||||
request_headers,
|
|
||||||
&body,
|
|
||||||
).await?;
|
).await?;
|
||||||
let Attempt::Response(response) = attempt else { return };
|
let Attempt::Response(response) = attempt else { return };
|
||||||
yield ModelEvent::Start { model_call_id: call_id };
|
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) {
|
fn merge_usage(total: &mut Usage, update: Usage) {
|
||||||
merge_usage_field(&mut total.input_tokens, update.input_tokens);
|
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.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_read_tokens, update.cache_read_tokens);
|
||||||
merge_usage_field(&mut total.cache_write_tokens, update.cache_write_tokens);
|
merge_usage_field(&mut total.cache_write_tokens, update.cache_write_tokens);
|
||||||
merge_usage_field(&mut total.reasoning_tokens, update.reasoning_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>) {
|
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 {
|
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 {
|
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),
|
output_tokens: value.get("output_tokens").and_then(Value::as_u64),
|
||||||
total_tokens: value.get("total_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_read_tokens,
|
||||||
cache_write_tokens: value
|
cache_write_tokens,
|
||||||
.get("cache_creation_input_tokens")
|
|
||||||
.and_then(Value::as_u64),
|
|
||||||
reasoning_tokens: None,
|
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.
|
//! Defines the provider interface and exports provider implementations.
|
||||||
mod anthropic;
|
mod anthropic;
|
||||||
|
mod attempt;
|
||||||
mod event;
|
mod event;
|
||||||
mod normalize;
|
mod normalize;
|
||||||
mod openai_chat;
|
mod openai_chat;
|
||||||
mod openai_responses;
|
mod openai_responses;
|
||||||
mod recorder;
|
mod recorder;
|
||||||
mod retry;
|
|
||||||
mod router;
|
mod router;
|
||||||
|
|
||||||
use std::pin::Pin;
|
use std::pin::Pin;
|
||||||
|
|||||||
@@ -17,10 +17,10 @@ use crate::{
|
|||||||
};
|
};
|
||||||
|
|
||||||
use super::{
|
use super::{
|
||||||
apply_body_allowlist, apply_openai_prompt_cache_key, map_sse_error, merge_extra_params,
|
apply_body_allowlist, apply_openai_prompt_cache_key,
|
||||||
provider_event_error,
|
attempt::{send_once, Attempt},
|
||||||
|
map_sse_error, merge_extra_params, provider_event_error,
|
||||||
recorder::recorded_headers,
|
recorder::recorded_headers,
|
||||||
retry::{send_with_retry, Attempt, RetryPolicy},
|
|
||||||
CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream,
|
CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream,
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -92,15 +92,12 @@ impl Provider for OpenAiChatProvider {
|
|||||||
if let Some(recorder) = &recorder {
|
if let Some(recorder) = &recorder {
|
||||||
recorder.request(request_headers.clone(), &body).await?;
|
recorder.request(request_headers.clone(), &body).await?;
|
||||||
}
|
}
|
||||||
let attempt = send_with_retry(
|
let attempt = send_once(
|
||||||
"OpenAI Chat",
|
"OpenAI Chat",
|
||||||
|| client.post(&config.request_url)
|
|| client.post(&config.request_url)
|
||||||
.bearer_auth(&config.api_key).headers(config.custom_headers.clone()).json(&body),
|
.bearer_auth(&config.api_key).headers(config.custom_headers.clone()).json(&body),
|
||||||
RetryPolicy { retries: config.retry_count, ..RetryPolicy::default() },
|
|
||||||
&cancellation,
|
&cancellation,
|
||||||
recorder.as_ref(),
|
recorder.as_ref(),
|
||||||
request_headers,
|
|
||||||
&body,
|
|
||||||
).await?;
|
).await?;
|
||||||
let Attempt::Response(response) = attempt else { return };
|
let Attempt::Response(response) = attempt else { return };
|
||||||
yield ModelEvent::Start { model_call_id: call_id };
|
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 {
|
pub(crate) fn openai_usage(value: &Value) -> Usage {
|
||||||
|
let input_tokens = value.get("prompt_tokens").and_then(Value::as_u64);
|
||||||
Usage {
|
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),
|
output_tokens: value.get("completion_tokens").and_then(Value::as_u64),
|
||||||
total_tokens: value.get("total_tokens").and_then(Value::as_u64),
|
total_tokens: value.get("total_tokens").and_then(Value::as_u64),
|
||||||
cache_read_tokens: value
|
cache_read_tokens: value
|
||||||
|
|||||||
@@ -14,10 +14,10 @@ use crate::{
|
|||||||
};
|
};
|
||||||
|
|
||||||
use super::{
|
use super::{
|
||||||
apply_body_allowlist, apply_openai_prompt_cache_key, map_sse_error, merge_extra_params,
|
apply_body_allowlist, apply_openai_prompt_cache_key,
|
||||||
provider_event_error,
|
attempt::{send_once, Attempt},
|
||||||
|
map_sse_error, merge_extra_params, provider_event_error,
|
||||||
recorder::recorded_headers,
|
recorder::recorded_headers,
|
||||||
retry::{send_with_retry, Attempt, RetryPolicy},
|
|
||||||
CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream,
|
CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream,
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -89,15 +89,12 @@ impl Provider for OpenAiResponsesProvider {
|
|||||||
if let Some(recorder) = &recorder {
|
if let Some(recorder) = &recorder {
|
||||||
recorder.request(request_headers.clone(), &body).await?;
|
recorder.request(request_headers.clone(), &body).await?;
|
||||||
}
|
}
|
||||||
let attempt = send_with_retry(
|
let attempt = send_once(
|
||||||
"OpenAI Responses",
|
"OpenAI Responses",
|
||||||
|| client.post(&config.request_url)
|
|| client.post(&config.request_url)
|
||||||
.bearer_auth(&config.api_key).headers(config.custom_headers.clone()).json(&body),
|
.bearer_auth(&config.api_key).headers(config.custom_headers.clone()).json(&body),
|
||||||
RetryPolicy { retries: config.retry_count, ..RetryPolicy::default() },
|
|
||||||
&cancellation,
|
&cancellation,
|
||||||
recorder.as_ref(),
|
recorder.as_ref(),
|
||||||
request_headers,
|
|
||||||
&body,
|
|
||||||
).await?;
|
).await?;
|
||||||
let Attempt::Response(response) = attempt else { return };
|
let Attempt::Response(response) = attempt else { return };
|
||||||
yield ModelEvent::Start { model_call_id: call_id };
|
yield ModelEvent::Start { model_call_id: call_id };
|
||||||
@@ -423,7 +420,12 @@ fn responses_input(messages: &[ProjectedMessage]) -> Result<Vec<Value>> {
|
|||||||
.ok_or_else(|| {
|
.ok_or_else(|| {
|
||||||
Error::Protocol("OpenAI Responses replay state is missing items".into())
|
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);
|
push_responses_text(&mut input, &message.role, text);
|
||||||
for call in calls {
|
for call in calls {
|
||||||
@@ -440,6 +442,23 @@ fn responses_input(messages: &[ProjectedMessage]) -> Result<Vec<Value>> {
|
|||||||
Ok(input)
|
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<()> {
|
fn push_responses_parts(input: &mut Vec<Value>, role: &Role, parts: &[ContentPart]) -> Result<()> {
|
||||||
let text_type = if *role == Role::Assistant {
|
let text_type = if *role == Role::Assistant {
|
||||||
"output_text"
|
"output_text"
|
||||||
@@ -505,8 +524,10 @@ fn required_u64(value: &Value, name: &str) -> Result<u64> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
fn responses_usage(value: &Value) -> Usage {
|
fn responses_usage(value: &Value) -> Usage {
|
||||||
|
let input_tokens = value.get("input_tokens").and_then(Value::as_u64);
|
||||||
Usage {
|
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),
|
output_tokens: value.get("output_tokens").and_then(Value::as_u64),
|
||||||
total_tokens: value.get("total_tokens").and_then(Value::as_u64),
|
total_tokens: value.get("total_tokens").and_then(Value::as_u64),
|
||||||
cache_read_tokens: value
|
cache_read_tokens: value
|
||||||
@@ -518,3 +539,47 @@ fn responses_usage(value: &Value) -> Usage {
|
|||||||
.and_then(Value::as_u64),
|
.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.
|
//! Records provider requests, responses, usage, and timing.
|
||||||
use std::{
|
use std::{
|
||||||
sync::{
|
sync::{
|
||||||
atomic::{AtomicBool, AtomicI64, AtomicU32, AtomicU64, Ordering},
|
atomic::{AtomicBool, AtomicI64, AtomicU64, Ordering},
|
||||||
Arc,
|
Arc,
|
||||||
},
|
},
|
||||||
time::Instant,
|
time::Instant,
|
||||||
@@ -71,7 +71,6 @@ struct Inner {
|
|||||||
base_call: NewLlmCall,
|
base_call: NewLlmCall,
|
||||||
detailed: bool,
|
detailed: bool,
|
||||||
attempt: Mutex<AttemptState>,
|
attempt: Mutex<AttemptState>,
|
||||||
next_attempt: AtomicU32,
|
|
||||||
next_generation: AtomicU64,
|
next_generation: AtomicU64,
|
||||||
finished: AtomicBool,
|
finished: AtomicBool,
|
||||||
}
|
}
|
||||||
@@ -120,7 +119,6 @@ impl CallRecorder {
|
|||||||
base_call: call.clone(),
|
base_call: call.clone(),
|
||||||
detailed: call.detailed,
|
detailed: call.detailed,
|
||||||
attempt: Mutex::new(AttemptState::new(call.call_id.clone())),
|
attempt: Mutex::new(AttemptState::new(call.call_id.clone())),
|
||||||
next_attempt: AtomicU32::new(0),
|
|
||||||
next_generation: AtomicU64::new(0),
|
next_generation: AtomicU64::new(0),
|
||||||
finished: AtomicBool::new(false),
|
finished: AtomicBool::new(false),
|
||||||
}),
|
}),
|
||||||
@@ -284,34 +282,6 @@ impl CallRecorder {
|
|||||||
self.finish("cancelled", None, None, None).await
|
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(
|
async fn finish(
|
||||||
&self,
|
&self,
|
||||||
status: &str,
|
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,
|
OpenAiChatProvider, OpenAiResponsesProvider, Provider, ProviderStream,
|
||||||
};
|
};
|
||||||
|
|
||||||
const BUILTIN_PROVIDER_RETRIES: u32 = 5;
|
|
||||||
|
|
||||||
pub struct ProviderRouter {
|
pub struct ProviderRouter {
|
||||||
store: Store,
|
store: Store,
|
||||||
plugins: PluginRegistry,
|
plugins: PluginRegistry,
|
||||||
|
clients: crate::network::NetworkClients,
|
||||||
request_timeout: Duration,
|
request_timeout: Duration,
|
||||||
stream_idle_timeout: Duration,
|
stream_idle_timeout: Duration,
|
||||||
}
|
}
|
||||||
@@ -31,12 +30,14 @@ impl ProviderRouter {
|
|||||||
pub fn new(
|
pub fn new(
|
||||||
store: Store,
|
store: Store,
|
||||||
plugins: PluginRegistry,
|
plugins: PluginRegistry,
|
||||||
|
clients: crate::network::NetworkClients,
|
||||||
request_timeout: Duration,
|
request_timeout: Duration,
|
||||||
stream_idle_timeout: Duration,
|
stream_idle_timeout: Duration,
|
||||||
) -> Self {
|
) -> Self {
|
||||||
Self {
|
Self {
|
||||||
store,
|
store,
|
||||||
plugins,
|
plugins,
|
||||||
|
clients,
|
||||||
request_timeout,
|
request_timeout,
|
||||||
stream_idle_timeout,
|
stream_idle_timeout,
|
||||||
}
|
}
|
||||||
@@ -51,6 +52,7 @@ impl Provider for ProviderRouter {
|
|||||||
) -> ProviderStream {
|
) -> ProviderStream {
|
||||||
let store = self.store.clone();
|
let store = self.store.clone();
|
||||||
let plugins = self.plugins.clone();
|
let plugins = self.plugins.clone();
|
||||||
|
let clients = self.clients.clone();
|
||||||
let request_timeout = self.request_timeout;
|
let request_timeout = self.request_timeout;
|
||||||
let stream_idle_timeout = self.stream_idle_timeout;
|
let stream_idle_timeout = self.stream_idle_timeout;
|
||||||
Box::pin(try_stream! {
|
Box::pin(try_stream! {
|
||||||
@@ -64,17 +66,14 @@ impl Provider for ProviderRouter {
|
|||||||
let plan = plugins.plan_model(&selected).await?;
|
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 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();
|
let guard = recorder.cancel_on_drop();
|
||||||
recorder.request(serde_json::json!({}), &crate::plugin::plugin_llm_request(&invocation)?).await?;
|
|
||||||
let mut routed = invocation.clone();
|
let mut routed = invocation.clone();
|
||||||
routed.request.model.display_name = Some(plan.model.display_name.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 {
|
if let Some(tokens) = plan.model.max_output_tokens {
|
||||||
routed.request.model.max_output_tokens.get_or_insert(tokens);
|
routed.request.model.max_output_tokens.get_or_insert(tokens);
|
||||||
}
|
}
|
||||||
let provider: Arc<dyn Provider> = Arc::new(NormalizedProvider::new(Arc::new(PluginModelProvider {
|
let provider: Arc<dyn Provider> = Arc::new(NormalizedProvider::new(Arc::new(PluginModelProvider {
|
||||||
registry: plugins.clone(),
|
registry: plugins.clone(),
|
||||||
|
recorder: recorder.clone(),
|
||||||
})));
|
})));
|
||||||
(recorder, guard, provider.stream(routed, cancellation.clone()))
|
(recorder, guard, provider.stream(routed, cancellation.clone()))
|
||||||
} else {
|
} 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() },
|
custom_headers: if model.custom_headers_enabled { custom_headers(&model.custom_headers)? } else { reqwest::header::HeaderMap::new() },
|
||||||
max_output_tokens: model.max_output_tokens(),
|
max_output_tokens: model.max_output_tokens(),
|
||||||
request_timeout,
|
request_timeout,
|
||||||
retry_count: BUILTIN_PROVIDER_RETRIES,
|
|
||||||
allowed_body_fields: None,
|
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)?;
|
let provider = build_observed(&config, recorder.clone(), client)?;
|
||||||
(recorder, guard, provider.stream(routed, cancellation.clone()))
|
(recorder, guard, provider.stream(routed, cancellation.clone()))
|
||||||
};
|
};
|
||||||
@@ -233,6 +231,7 @@ async fn finish_stream(recorder: &CallRecorder, cancellation: &CancellationToken
|
|||||||
/// 插件模型的 Provider 实现;对路由与规范化层完全等同于内置 Provider。
|
/// 插件模型的 Provider 实现;对路由与规范化层完全等同于内置 Provider。
|
||||||
struct PluginModelProvider {
|
struct PluginModelProvider {
|
||||||
registry: PluginRegistry,
|
registry: PluginRegistry,
|
||||||
|
recorder: CallRecorder,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl Provider for PluginModelProvider {
|
impl Provider for PluginModelProvider {
|
||||||
@@ -241,7 +240,8 @@ impl Provider for PluginModelProvider {
|
|||||||
invocation: ModelInvocation,
|
invocation: ModelInvocation,
|
||||||
cancellation: CancellationToken,
|
cancellation: CancellationToken,
|
||||||
) -> ProviderStream {
|
) -> 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 std::collections::HashSet;
|
||||||
|
|
||||||
use crate::model::{
|
use crate::{
|
||||||
CanonicalMessage, LlmCallUsageAnchor, PreparedRun, ProjectedMessage, RunAction,
|
model::{
|
||||||
|
estimate_context_tokens, estimate_projected_messages_tokens, CanonicalMessage, PreparedRun,
|
||||||
|
ProjectedMessage,
|
||||||
|
},
|
||||||
|
store::ContextUsageAnchor,
|
||||||
};
|
};
|
||||||
|
|
||||||
const FALLBACK_CHARS: usize = 12_000;
|
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 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.";
|
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) fn input_budget(prepared: &PreparedRun) -> Option<u64> {
|
||||||
pub(super) struct ContextUsageAnchor {
|
prepared
|
||||||
input_tokens: u64,
|
.model
|
||||||
message_count: usize,
|
.context_window_tokens
|
||||||
tool_count: usize,
|
.map(|window| window.saturating_sub(RESERVE_TOKENS))
|
||||||
}
|
}
|
||||||
|
|
||||||
impl ContextUsageAnchor {
|
pub(super) fn estimated_tokens(
|
||||||
pub(super) fn from_llm_call(anchor: LlmCallUsageAnchor) -> Option<Self> {
|
prepared: &PreparedRun,
|
||||||
Some(Self {
|
projected_messages: &[ProjectedMessage],
|
||||||
input_tokens: anchor.usage.context_input_tokens(anchor.request_type)?,
|
anchor: Option<ContextUsageAnchor>,
|
||||||
message_count: anchor.message_count,
|
) -> u64 {
|
||||||
tool_count: anchor.tool_count,
|
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(
|
pub(super) fn should_compact(
|
||||||
prepared: &PreparedRun,
|
prepared: &PreparedRun,
|
||||||
messages: &[CanonicalMessage],
|
|
||||||
projected_messages: &[ProjectedMessage],
|
projected_messages: &[ProjectedMessage],
|
||||||
anchor: Option<ContextUsageAnchor>,
|
anchor: Option<ContextUsageAnchor>,
|
||||||
) -> bool {
|
) -> bool {
|
||||||
if prepared.action != RunAction::Start {
|
compaction_estimate(prepared, projected_messages, anchor).is_some()
|
||||||
return false;
|
}
|
||||||
}
|
|
||||||
let Some(context_window) = prepared.model.context_window_tokens else {
|
pub(super) fn validate_compacted(
|
||||||
return false;
|
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() {
|
if estimated <= budget {
|
||||||
return false;
|
return Ok(estimated);
|
||||||
}
|
}
|
||||||
let estimated_input = anchor
|
Err(format!(
|
||||||
.filter(|anchor| {
|
"context overflow after compaction: estimated input {estimated} tokens exceeds budget {budget} tokens"
|
||||||
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
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(super) fn partition(
|
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)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
use super::*;
|
use super::*;
|
||||||
use crate::model::{
|
use crate::model::{
|
||||||
project_messages, CheckpointId, ConversationId, ModelSpec, Origin, PromptSpec, Role, RunId,
|
project_messages, CheckpointId, ConversationId, ModelSpec, Origin, PromptSpec, Role,
|
||||||
RunKind,
|
RunAction, RunId, RunKind,
|
||||||
};
|
};
|
||||||
|
|
||||||
#[test]
|
fn prepared(context_window_tokens: u64) -> PreparedRun {
|
||||||
fn automatic_compaction_starts_only_after_the_context_window_is_exceeded() {
|
|
||||||
let mut model = ModelSpec::new("model");
|
let mut model = ModelSpec::new("model");
|
||||||
model.context_window_tokens = Some(200_000);
|
model.context_window_tokens = Some(context_window_tokens);
|
||||||
let prepared = PreparedRun {
|
PreparedRun {
|
||||||
run_id: RunId::new("run"),
|
run_id: RunId::new("run"),
|
||||||
cursor_request_id: None,
|
cursor_request_id: None,
|
||||||
conversation_id: ConversationId::new("conversation"),
|
conversation_id: ConversationId::new("conversation"),
|
||||||
@@ -134,40 +135,133 @@ mod tests {
|
|||||||
initial_messages: Vec::new(),
|
initial_messages: Vec::new(),
|
||||||
action: RunAction::Start,
|
action: RunAction::Start,
|
||||||
base_checkpoint_id: CheckpointId(1),
|
base_checkpoint_id: CheckpointId(1),
|
||||||
};
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn automatic_compaction_uses_fixed_reserve_for_every_action() {
|
||||||
let messages = vec![CanonicalMessage::text(
|
let messages = vec![CanonicalMessage::text(
|
||||||
"user",
|
"user",
|
||||||
Role::User,
|
Role::User,
|
||||||
Origin::Runtime,
|
Origin::Runtime,
|
||||||
"hello",
|
"x".repeat(40_000),
|
||||||
)];
|
)];
|
||||||
let projected = project_messages(&messages).unwrap();
|
let projected = project_messages(&messages).unwrap();
|
||||||
let tail_tokens = estimate_serialized_tokens(&serde_json::to_string(&projected).unwrap());
|
let estimated = estimate_context_tokens(&prepared(1).prompt, &projected);
|
||||||
let anchor = |estimated_input| {
|
let mut prepared = prepared(estimated + RESERVE_TOKENS);
|
||||||
Some(ContextUsageAnchor {
|
|
||||||
input_tokens: estimated_input - tail_tokens,
|
|
||||||
message_count: 0,
|
|
||||||
tool_count: 0,
|
|
||||||
})
|
|
||||||
};
|
|
||||||
|
|
||||||
|
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(
|
assert!(!should_compact(
|
||||||
&prepared,
|
&prepared(200_000),
|
||||||
&messages,
|
|
||||||
&projected,
|
&projected,
|
||||||
anchor(199_999)
|
Some(anchor)
|
||||||
));
|
|
||||||
assert!(!should_compact(
|
|
||||||
&prepared,
|
|
||||||
&messages,
|
|
||||||
&projected,
|
|
||||||
anchor(200_000)
|
|
||||||
));
|
|
||||||
assert!(should_compact(
|
|
||||||
&prepared,
|
|
||||||
&messages,
|
|
||||||
&projected,
|
|
||||||
anchor(200_001)
|
|
||||||
));
|
));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[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,
|
ToolRoundId, Usage,
|
||||||
},
|
},
|
||||||
provider::Provider,
|
provider::Provider,
|
||||||
store::{RunStatus, Store},
|
store::{ContextUsageAnchor, RunStatus, Store},
|
||||||
};
|
};
|
||||||
|
|
||||||
use super::{
|
use super::{
|
||||||
consume_model_cycle, CommitBarrier, CommitCause, MessagesCommitted, ModelCycleFailure,
|
consume_model_cycle,
|
||||||
RunCommand, RunEvent, RunFailure, RunOutcome, RunPort,
|
model_retry::{should_retry, MODEL_RETRY_DELAY},
|
||||||
|
CommitBarrier, CommitCause, MessagesCommitted, RunCommand, RunEvent, RunFailure, RunOutcome,
|
||||||
|
RunPort,
|
||||||
};
|
};
|
||||||
|
|
||||||
pub struct RunEngine {
|
pub struct RunEngine {
|
||||||
@@ -90,6 +92,14 @@ impl RunEngine {
|
|||||||
cancellation: &CancellationToken,
|
cancellation: &CancellationToken,
|
||||||
) -> (RunOutcome, Option<Usage>) {
|
) -> (RunOutcome, Option<Usage>) {
|
||||||
let mut usage = None;
|
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!(
|
tracing::info!(
|
||||||
checkpoint_id = checkpoint.0,
|
checkpoint_id = checkpoint.0,
|
||||||
"Run claimed conversation ownership"
|
"Run claimed conversation ownership"
|
||||||
@@ -161,7 +171,6 @@ impl RunEngine {
|
|||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
||||||
let mut auto_compacted = prepared.action == RunAction::Compact;
|
|
||||||
'model: loop {
|
'model: loop {
|
||||||
if cancellation.is_cancelled() {
|
if cancellation.is_cancelled() {
|
||||||
return (RunOutcome::Cancelled, usage);
|
return (RunOutcome::Cancelled, usage);
|
||||||
@@ -170,37 +179,32 @@ impl RunEngine {
|
|||||||
Ok(messages) => messages,
|
Ok(messages) => messages,
|
||||||
Err(error) => return (RunOutcome::Failed(error.into()), usage),
|
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) {
|
let history = match crate::model::project_messages(&messages) {
|
||||||
Ok(history) => history,
|
Ok(history) => history,
|
||||||
Err(error) => return (RunOutcome::Failed(error.into()), usage),
|
Err(error) => return (RunOutcome::Failed(error.into()), usage),
|
||||||
};
|
};
|
||||||
if !auto_compacted
|
let compaction_estimate = (prepared.action != RunAction::Compact)
|
||||||
&& super::compaction::should_compact(prepared, &messages, &history, context_anchor)
|
.then(|| {
|
||||||
{
|
super::compaction::compaction_estimate(prepared, &history, context_usage_anchor)
|
||||||
auto_compacted = true;
|
})
|
||||||
|
.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
|
match self
|
||||||
.auto_compact(prepared, checkpoint, &messages, client, cancellation)
|
.auto_compact(prepared, checkpoint, &messages, client, cancellation)
|
||||||
.await
|
.await
|
||||||
{
|
{
|
||||||
Ok((next_checkpoint, compaction_usage)) => {
|
Ok((next_checkpoint, compaction_usage)) => {
|
||||||
checkpoint = next_checkpoint;
|
checkpoint = next_checkpoint;
|
||||||
|
context_usage_anchor = None;
|
||||||
if let Some(compaction_usage) = compaction_usage {
|
if let Some(compaction_usage) = compaction_usage {
|
||||||
accumulate_usage(&mut usage, compaction_usage);
|
accumulate_usage(&mut usage, compaction_usage);
|
||||||
}
|
}
|
||||||
@@ -227,122 +231,236 @@ impl RunEngine {
|
|||||||
model: prepared.model.clone(),
|
model: prepared.model.clone(),
|
||||||
history,
|
history,
|
||||||
};
|
};
|
||||||
let invocation = crate::model::ModelInvocation {
|
let mut retries = 0_u32;
|
||||||
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 pending_insertions = Vec::new();
|
let mut pending_insertions = Vec::new();
|
||||||
let cycle = loop {
|
let cycle = 'attempt: loop {
|
||||||
tokio::select! {
|
let call_id = if retries == 0 {
|
||||||
biased;
|
format!("{}:{provider_call_index}", prepared.run_id)
|
||||||
command = client.commands.recv() => {
|
} else {
|
||||||
let interruption = match command {
|
format!("{}:{provider_call_index}:retry-{retries}", prepared.run_id)
|
||||||
Some(RunCommand::InsertMessages(insertion)) => {
|
};
|
||||||
pending_insertions.push(insertion);
|
let invocation = crate::model::ModelInvocation {
|
||||||
continue;
|
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,
|
if emit(client, RunEvent::CycleInterrupted).await.is_err() {
|
||||||
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);
|
return (client_failure(), usage);
|
||||||
}
|
}
|
||||||
};
|
checkpoint = match super::messages::append_batches(
|
||||||
cycle_cancellation.cancel();
|
&self.store,
|
||||||
let interrupted = cycle.await;
|
prepared,
|
||||||
match interrupted {
|
client,
|
||||||
Ok(cycle) => {
|
cancellation,
|
||||||
if let Some(cycle_usage) = cycle.usage {
|
checkpoint,
|
||||||
accumulate_usage(&mut usage, cycle_usage);
|
std::mem::take(&mut pending_insertions),
|
||||||
}
|
)
|
||||||
}
|
.await
|
||||||
Err(failure) => {
|
{
|
||||||
if let Some(cycle_usage) = failure.usage {
|
Ok((checkpoint, _)) => checkpoint,
|
||||||
accumulate_usage(&mut usage, cycle_usage);
|
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);
|
return (client_failure(), usage);
|
||||||
}
|
}
|
||||||
checkpoint = match super::messages::append_batches(
|
|
||||||
&self.store,
|
let delay = tokio::time::sleep(MODEL_RETRY_DELAY);
|
||||||
prepared,
|
tokio::pin!(delay);
|
||||||
client,
|
loop {
|
||||||
cancellation,
|
tokio::select! {
|
||||||
checkpoint,
|
biased;
|
||||||
std::mem::take(&mut pending_insertions),
|
command = client.commands.recv() => {
|
||||||
)
|
let interruption = match command {
|
||||||
.await
|
Some(RunCommand::InsertMessages(insertion)) => {
|
||||||
{
|
pending_insertions.push(insertion);
|
||||||
Ok((checkpoint, _)) => checkpoint,
|
continue;
|
||||||
Err(outcome) => return (outcome, usage),
|
}
|
||||||
};
|
Some(RunCommand::BreakMessages(messages)) => messages,
|
||||||
checkpoint = match super::messages::append_batches(
|
Some(RunCommand::Cancel) => {
|
||||||
&self.store,
|
let _ = emit(client, RunEvent::CycleInterrupted).await;
|
||||||
prepared,
|
return (RunOutcome::Cancelled, usage);
|
||||||
client,
|
}
|
||||||
cancellation,
|
Some(RunCommand::ToolResult(_)) => {
|
||||||
checkpoint,
|
return (
|
||||||
vec![interruption],
|
RunOutcome::Failed(RunFailure::Protocol(
|
||||||
)
|
"received a tool result while waiting to retry the model".into(),
|
||||||
.await
|
)),
|
||||||
{
|
usage,
|
||||||
Ok((checkpoint, _)) => checkpoint,
|
);
|
||||||
Err(outcome) => return (outcome, usage),
|
}
|
||||||
};
|
None => return (client_failure(), usage),
|
||||||
continue 'model;
|
};
|
||||||
},
|
if emit(client, RunEvent::CycleInterrupted).await.is_err() {
|
||||||
result = &mut cycle => break result,
|
return (client_failure(), usage);
|
||||||
}
|
}
|
||||||
};
|
checkpoint = match super::messages::append_batches(
|
||||||
let cycle = match cycle {
|
&self.store,
|
||||||
Ok(cycle) => cycle,
|
prepared,
|
||||||
Err(ModelCycleFailure {
|
client,
|
||||||
failure,
|
cancellation,
|
||||||
usage: cycle_usage,
|
checkpoint,
|
||||||
..
|
std::mem::take(&mut pending_insertions),
|
||||||
}) => {
|
)
|
||||||
if let Some(cycle_usage) = cycle_usage {
|
.await
|
||||||
accumulate_usage(&mut usage, cycle_usage);
|
{
|
||||||
|
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 {
|
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);
|
accumulate_usage(&mut usage, cycle_usage);
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -583,7 +701,15 @@ impl RunEngine {
|
|||||||
let (compactable, retained_request_context) =
|
let (compactable, retained_request_context) =
|
||||||
super::compaction::partition(messages, ¤t_ids);
|
super::compaction::partition(messages, ¤t_ids);
|
||||||
if compactable.is_empty() {
|
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)
|
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 {
|
let summary_message = CanonicalMessage {
|
||||||
message_id: format!("runtime:{event_id}"),
|
message_id: format!("runtime:{event_id}"),
|
||||||
role: Role::User,
|
role: Role::User,
|
||||||
@@ -704,6 +830,10 @@ impl RunEngine {
|
|||||||
let mut replacement = retained_request_context.into_iter().collect::<Vec<_>>();
|
let mut replacement = retained_request_context.into_iter().collect::<Vec<_>>();
|
||||||
replacement.push(summary_message);
|
replacement.push(summary_message);
|
||||||
replacement.extend(prepared.initial_messages.iter().cloned());
|
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
|
let mut checkpoint = self
|
||||||
.store
|
.store
|
||||||
.replace_checkpoint(
|
.replace_checkpoint(
|
||||||
@@ -730,6 +860,9 @@ impl RunEngine {
|
|||||||
emit(client, RunEvent::AutoCompactionCompleted)
|
emit(client, RunEvent::AutoCompactionCompleted)
|
||||||
.await
|
.await
|
||||||
.map_err(|_| client_failure())?;
|
.map_err(|_| client_failure())?;
|
||||||
|
emit(client, RunEvent::UsageSnapshot(context_usage_snapshot(0)))
|
||||||
|
.await
|
||||||
|
.map_err(|_| client_failure())?;
|
||||||
checkpoint = super::messages::append_batches(
|
checkpoint = super::messages::append_batches(
|
||||||
&self.store,
|
&self.store,
|
||||||
prepared,
|
prepared,
|
||||||
@@ -790,6 +923,31 @@ async fn hydrate_tool_images(
|
|||||||
Ok(())
|
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) {
|
fn accumulate_usage(total: &mut Option<Usage>, usage: Usage) {
|
||||||
match total {
|
match total {
|
||||||
Some(total) => *total += usage,
|
Some(total) => *total += usage,
|
||||||
|
|||||||
@@ -104,6 +104,10 @@ pub enum RunEvent {
|
|||||||
AutoCompactionStarted,
|
AutoCompactionStarted,
|
||||||
AutoCompactionCompleted,
|
AutoCompactionCompleted,
|
||||||
CycleInterrupted,
|
CycleInterrupted,
|
||||||
|
ModelAttemptFailed {
|
||||||
|
attempt: u32,
|
||||||
|
message: String,
|
||||||
|
},
|
||||||
TextStart,
|
TextStart,
|
||||||
TextDelta(String),
|
TextDelta(String),
|
||||||
TextEnd,
|
TextEnd,
|
||||||
@@ -125,6 +129,7 @@ pub enum RunEvent {
|
|||||||
ToolCallEnd {
|
ToolCallEnd {
|
||||||
index: usize,
|
index: usize,
|
||||||
},
|
},
|
||||||
|
UsageSnapshot(Usage),
|
||||||
Usage(Usage),
|
Usage(Usage),
|
||||||
ExecuteToolRound {
|
ExecuteToolRound {
|
||||||
round_id: ToolRoundId,
|
round_id: ToolRoundId,
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ mod event;
|
|||||||
mod handle;
|
mod handle;
|
||||||
mod messages;
|
mod messages;
|
||||||
mod model_cycle;
|
mod model_cycle;
|
||||||
|
mod model_retry;
|
||||||
mod port;
|
mod port;
|
||||||
mod tool_round;
|
mod tool_round;
|
||||||
|
|
||||||
|
|||||||
+130
-16
@@ -6,7 +6,7 @@ use tokio::sync::mpsc;
|
|||||||
use tokio_util::sync::CancellationToken;
|
use tokio_util::sync::CancellationToken;
|
||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
model::{ProviderReplayState, ToolCall, Usage},
|
model::{normalize_tool_name, ProviderReplayState, ToolCall, Usage},
|
||||||
provider::{FinishReason, ModelEvent, ProviderStream},
|
provider::{FinishReason, ModelEvent, ProviderStream},
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -29,6 +29,7 @@ pub struct ModelCycleFailure {
|
|||||||
pub partial_text: String,
|
pub partial_text: String,
|
||||||
pub partial_reasoning: String,
|
pub partial_reasoning: String,
|
||||||
pub usage: Option<Usage>,
|
pub usage: Option<Usage>,
|
||||||
|
pub retryable: bool,
|
||||||
}
|
}
|
||||||
|
|
||||||
struct OpenTool {
|
struct OpenTool {
|
||||||
@@ -164,6 +165,7 @@ pub async fn consume_model_cycle(
|
|||||||
call_id,
|
call_id,
|
||||||
name,
|
name,
|
||||||
} => {
|
} => {
|
||||||
|
let name = normalize_tool_name(&name);
|
||||||
let Some(model_call_id) = model_call_id.as_ref() else {
|
let Some(model_call_id) = model_call_id.as_ref() else {
|
||||||
return Err(failure(
|
return Err(failure(
|
||||||
RunFailure::Protocol("provider emitted content before Start".into()),
|
RunFailure::Protocol("provider emitted content before Start".into()),
|
||||||
@@ -186,6 +188,7 @@ pub async fn consume_model_cycle(
|
|||||||
name: name.clone(),
|
name: name.clone(),
|
||||||
arguments_text: String::new(),
|
arguments_text: String::new(),
|
||||||
arguments: serde_json::Value::Null,
|
arguments: serde_json::Value::Null,
|
||||||
|
argument_error: None,
|
||||||
},
|
},
|
||||||
ended: false,
|
ended: false,
|
||||||
});
|
});
|
||||||
@@ -218,13 +221,26 @@ pub async fn consume_model_cycle(
|
|||||||
serde_json::from_str(&tool.call.arguments_text)
|
serde_json::from_str(&tool.call.arguments_text)
|
||||||
};
|
};
|
||||||
match arguments {
|
match arguments {
|
||||||
Ok(arguments) => {
|
Ok(arguments) if arguments.is_object() => {
|
||||||
tool.call.arguments = arguments;
|
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"),
|
Some(_) => Err("provider emitted duplicate ToolCallEnd"),
|
||||||
None => Err("provider ended an unknown tool index"),
|
None => Err("provider ended an unknown tool index"),
|
||||||
@@ -239,6 +255,13 @@ pub async fn consume_model_cycle(
|
|||||||
ModelEvent::Usage(value) => {
|
ModelEvent::Usage(value) => {
|
||||||
if usage.replace(value).is_some() {
|
if usage.replace(value).is_some() {
|
||||||
Err("provider emitted duplicate Usage")
|
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 {
|
} else {
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
@@ -280,7 +303,7 @@ pub async fn consume_model_cycle(
|
|||||||
.map(|tool| tool.call)
|
.map(|tool| tool.call)
|
||||||
.collect::<Vec<_>>();
|
.collect::<Vec<_>>();
|
||||||
if finish_reason == FinishReason::Length {
|
if finish_reason == FinishReason::Length {
|
||||||
return Err(failure(
|
return Err(terminal_failure(
|
||||||
RunFailure::Provider("model stopped before completing the response".into()),
|
RunFailure::Provider("model stopped before completing the response".into()),
|
||||||
text,
|
text,
|
||||||
reasoning,
|
reasoning,
|
||||||
@@ -296,16 +319,6 @@ pub async fn consume_model_cycle(
|
|||||||
usage,
|
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(|| {
|
let model_call_id = model_call_id.ok_or_else(|| {
|
||||||
failure(
|
failure(
|
||||||
RunFailure::Protocol("provider completed without Start".into()),
|
RunFailure::Protocol("provider completed without Start".into()),
|
||||||
@@ -360,10 +373,111 @@ fn failure(
|
|||||||
partial_reasoning: String,
|
partial_reasoning: String,
|
||||||
usage: Option<Usage>,
|
usage: Option<Usage>,
|
||||||
) -> ModelCycleFailure {
|
) -> ModelCycleFailure {
|
||||||
|
let retryable = matches!(failure, RunFailure::Protocol(_) | RunFailure::Provider(_));
|
||||||
ModelCycleFailure {
|
ModelCycleFailure {
|
||||||
failure,
|
failure,
|
||||||
partial_text,
|
partial_text,
|
||||||
partial_reasoning,
|
partial_reasoning,
|
||||||
usage,
|
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?)
|
.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(
|
pub async fn append_cursor_trace_artifact(
|
||||||
&self,
|
&self,
|
||||||
request_id: &str,
|
request_id: &str,
|
||||||
@@ -144,23 +178,6 @@ impl Store {
|
|||||||
Ok(())
|
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<()> {
|
pub async fn start_cursor_trace_response(&self, request_id: &str, status: u16) -> Result<()> {
|
||||||
let now = now_ms();
|
let now = now_ms();
|
||||||
let _write = self.writes.lock().await;
|
let _write = self.writes.lock().await;
|
||||||
|
|||||||
+100
-41
@@ -1,18 +1,19 @@
|
|||||||
//! Persists provider call payloads, timing, and usage.
|
//! Persists provider call payloads, timing, and usage.
|
||||||
use std::str::FromStr;
|
|
||||||
|
|
||||||
use sqlx::Row;
|
use sqlx::Row;
|
||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
model::{
|
model::{LlmCallRequest, LlmCallResponseChunk, LlmCallSummary, NewLlmCall, Usage},
|
||||||
ConversationId, LlmCallRequest, LlmCallResponseChunk, LlmCallSummary, LlmCallUsageAnchor,
|
|
||||||
NewLlmCall, ProviderType, Usage,
|
|
||||||
},
|
|
||||||
Result,
|
Result,
|
||||||
};
|
};
|
||||||
|
|
||||||
use super::{now_ms, Store};
|
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)]
|
#[derive(Clone, Debug)]
|
||||||
pub(crate) struct BufferedLlmChunk {
|
pub(crate) struct BufferedLlmChunk {
|
||||||
pub(crate) seq: i64,
|
pub(crate) seq: i64,
|
||||||
@@ -271,6 +272,37 @@ impl Store {
|
|||||||
Ok(())
|
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>> {
|
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 ?")
|
let rows = sqlx::query("SELECT * FROM llm_calls ORDER BY created_at_ms DESC LIMIT ?")
|
||||||
.bind(limit.clamp(1, 500))
|
.bind(limit.clamp(1, 500))
|
||||||
@@ -288,41 +320,6 @@ impl Store {
|
|||||||
.transpose()
|
.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>> {
|
pub async fn llm_call_request(&self, call_id: &str) -> Result<Option<LlmCallRequest>> {
|
||||||
let row = sqlx::query(
|
let row = sqlx::query(
|
||||||
"SELECT headers_json, body_json, byte_count FROM llm_call_requests WHERE call_id = ?",
|
"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)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
use super::*;
|
use super::*;
|
||||||
|
use crate::model::ProviderType;
|
||||||
|
|
||||||
/// 插件模型不在 model_configs 中,调用记录必须照常落库并可按其稳定 ID 筛选。
|
/// 插件模型不在 model_configs 中,调用记录必须照常落库并可按其稳定 ID 筛选。
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
@@ -460,4 +458,65 @@ mod tests {
|
|||||||
assert_eq!(overview.metrics.llm_calls, 1);
|
assert_eq!(overview.metrics.llm_calls, 1);
|
||||||
assert_eq!(overview.metrics.successful_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
|
.await
|
||||||
.unwrap();
|
.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!(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!(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