mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-04 02:52:55 +08:00
Compare commits
36
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 | ||
|
|
b807608bf3 | ||
|
|
e7a1cca4c6 | ||
|
|
76baa3b0e7 | ||
|
|
fc79adbb43 | ||
|
|
e673a034df | ||
|
|
6ac666da4f | ||
|
|
9e2d22418b | ||
|
|
97ee138de8 |
@@ -0,0 +1,71 @@
|
||||
import { readFile, writeFile } from "node:fs/promises";
|
||||
import { basename, resolve } from "node:path";
|
||||
import { pathToFileURL } from "node:url";
|
||||
|
||||
function readOptions(args) {
|
||||
const options = new Map();
|
||||
for (let index = 0; index < args.length; index += 2) {
|
||||
const name = args[index];
|
||||
const value = args[index + 1];
|
||||
if (!name?.startsWith("--") || value === undefined) {
|
||||
throw new Error(`invalid argument near ${name ?? "end of command"}`);
|
||||
}
|
||||
options.set(name.slice(2), value);
|
||||
}
|
||||
return options;
|
||||
}
|
||||
|
||||
function required(options, name) {
|
||||
const value = options.get(name)?.trim();
|
||||
if (!value) throw new Error(`--${name} is required`);
|
||||
return value;
|
||||
}
|
||||
|
||||
export function generatePortableUpdate({ version, repository, assetName, signature }) {
|
||||
const normalizedVersion = version.replace(/^v/, "");
|
||||
if (!/^\d+\.\d+\.\d+(?:-[0-9A-Za-z.-]+)?$/.test(normalizedVersion)) {
|
||||
throw new Error(`invalid semantic version: ${normalizedVersion}`);
|
||||
}
|
||||
if (!/^[^/\s]+\/[^/\s]+$/.test(repository)) {
|
||||
throw new Error(`invalid GitHub repository: ${repository}`);
|
||||
}
|
||||
if (!assetName || basename(assetName) !== assetName) {
|
||||
throw new Error("asset name must be a file name");
|
||||
}
|
||||
if (!signature.trim()) throw new Error("signature is required");
|
||||
|
||||
return {
|
||||
version: normalizedVersion,
|
||||
notes: `Cursor BYOK v${normalizedVersion}`,
|
||||
pub_date: new Date().toISOString(),
|
||||
platforms: {
|
||||
"windows-x86_64": {
|
||||
signature: signature.trim(),
|
||||
url: `https://github.com/${repository}/releases/download/v${normalizedVersion}/${encodeURIComponent(assetName)}`,
|
||||
},
|
||||
},
|
||||
};
|
||||
}
|
||||
|
||||
async function main() {
|
||||
const options = readOptions(process.argv.slice(2));
|
||||
const version = required(options, "version");
|
||||
const repository = required(options, "repository");
|
||||
const asset = required(options, "asset");
|
||||
const signaturePath = resolve(required(options, "signature"));
|
||||
const output = resolve(required(options, "output"));
|
||||
const manifest = generatePortableUpdate({
|
||||
version,
|
||||
repository,
|
||||
assetName: basename(asset),
|
||||
signature: await readFile(signaturePath, "utf8"),
|
||||
});
|
||||
await writeFile(output, `${JSON.stringify(manifest, null, 2)}\n`);
|
||||
}
|
||||
|
||||
if (process.argv[1] && import.meta.url === pathToFileURL(resolve(process.argv[1])).href) {
|
||||
main().catch((error) => {
|
||||
console.error(error instanceof Error ? error.message : String(error));
|
||||
process.exitCode = 1;
|
||||
});
|
||||
}
|
||||
@@ -0,0 +1,41 @@
|
||||
import assert from "node:assert/strict";
|
||||
import test from "node:test";
|
||||
import { generatePortableUpdate } from "./generate-portable-update.mjs";
|
||||
|
||||
test("generates a signed Windows portable updater manifest", () => {
|
||||
const manifest = generatePortableUpdate({
|
||||
version: "v1.2.3-beta.1",
|
||||
repository: "owner/repository",
|
||||
assetName: "cursor-byok-1.2.3-beta.1-windows-amd64.zip",
|
||||
signature: "signed-payload\n",
|
||||
});
|
||||
|
||||
assert.equal(manifest.version, "1.2.3-beta.1");
|
||||
assert.deepEqual(Object.keys(manifest.platforms), ["windows-x86_64"]);
|
||||
assert.equal(manifest.platforms["windows-x86_64"].signature, "signed-payload");
|
||||
assert.equal(
|
||||
manifest.platforms["windows-x86_64"].url,
|
||||
"https://github.com/owner/repository/releases/download/v1.2.3-beta.1/cursor-byok-1.2.3-beta.1-windows-amd64.zip",
|
||||
);
|
||||
});
|
||||
|
||||
test("rejects invalid inputs", () => {
|
||||
assert.throws(() => generatePortableUpdate({
|
||||
version: "latest",
|
||||
repository: "owner/repository",
|
||||
assetName: "update.zip",
|
||||
signature: "signature",
|
||||
}), /semantic version/);
|
||||
assert.throws(() => generatePortableUpdate({
|
||||
version: "1.2.3",
|
||||
repository: "owner/repository",
|
||||
assetName: "../update.zip",
|
||||
signature: "signature",
|
||||
}), /file name/);
|
||||
assert.throws(() => generatePortableUpdate({
|
||||
version: "1.2.3",
|
||||
repository: "owner/repository",
|
||||
assetName: "update.zip",
|
||||
signature: " ",
|
||||
}), /signature/);
|
||||
});
|
||||
@@ -158,14 +158,27 @@ jobs:
|
||||
mkdir -p legacy-update
|
||||
tar -czf "legacy-update/cursor-byok-${VERSION}-linux-amd64.tar.gz" -C target/release cursor-byok-desktop
|
||||
|
||||
- name: Package legacy Windows updater asset
|
||||
- name: Package and sign legacy Windows updater asset
|
||||
if: matrix.platform == 'windows-x86_64'
|
||||
shell: pwsh
|
||||
env:
|
||||
VERSION: ${{ needs.prepare.outputs.version }}
|
||||
TAURI_SIGNING_PRIVATE_KEY: ${{ secrets.TAURI_SIGNING_PRIVATE_KEY }}
|
||||
TAURI_SIGNING_PRIVATE_KEY_PASSWORD: ${{ secrets.TAURI_SIGNING_PRIVATE_KEY_PASSWORD }}
|
||||
run: |
|
||||
New-Item -ItemType Directory -Force legacy-update | Out-Null
|
||||
Compress-Archive -LiteralPath target/release/cursor-byok-desktop.exe -DestinationPath "legacy-update/cursor-byok-$env:VERSION-windows-amd64.zip"
|
||||
$asset = "legacy-update/cursor-byok-$env:VERSION-windows-amd64.zip"
|
||||
Compress-Archive -LiteralPath target/release/cursor-byok-desktop.exe -DestinationPath $asset
|
||||
Push-Location apps/desktop
|
||||
npm exec tauri signer sign -- "../../$asset"
|
||||
Pop-Location
|
||||
$entries = @(tar -tf $asset)
|
||||
if ($entries.Count -ne 1 -or [System.IO.Path]::GetFileName($entries[0]) -ne 'cursor-byok-desktop.exe') {
|
||||
throw "Windows updater archive must contain only cursor-byok-desktop.exe"
|
||||
}
|
||||
if (!(Test-Path "$asset.sig")) {
|
||||
throw "Windows updater archive signature was not generated"
|
||||
}
|
||||
|
||||
- name: Package legacy macOS updater asset
|
||||
if: contains(matrix.platform, 'macos')
|
||||
@@ -213,6 +226,20 @@ jobs:
|
||||
--output legacy-update/update.json \
|
||||
--notes "Cursor BYOK v${VERSION}"
|
||||
|
||||
- name: Generate signed Windows portable update manifest
|
||||
env:
|
||||
VERSION: ${{ needs.prepare.outputs.version }}
|
||||
run: |
|
||||
asset="cursor-byok-${VERSION}-windows-amd64.zip"
|
||||
test -f "legacy-update/${asset}"
|
||||
test -f "legacy-update/${asset}.sig"
|
||||
node .github/scripts/generate-portable-update.mjs \
|
||||
--version "${VERSION}" \
|
||||
--repository "${GITHUB_REPOSITORY}" \
|
||||
--asset "${asset}" \
|
||||
--signature "legacy-update/${asset}.sig" \
|
||||
--output legacy-update/portable-latest.json
|
||||
|
||||
- name: Normalize Tauri updater download URLs
|
||||
env:
|
||||
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
|
||||
Generated
+5
-1
@@ -1172,10 +1172,11 @@ checksum = "52560adf09603e58c9a7ee1fe1dcb95a16927b17c127f0ac02d6e768a0e25bc1"
|
||||
|
||||
[[package]]
|
||||
name = "cursor-byok-desktop"
|
||||
version = "0.1.5-beta.1"
|
||||
version = "0.1.6"
|
||||
dependencies = [
|
||||
"axum",
|
||||
"cursor-server",
|
||||
"libc",
|
||||
"rfd",
|
||||
"serde",
|
||||
"serde_json",
|
||||
@@ -1187,12 +1188,15 @@ dependencies = [
|
||||
"tauri-plugin-process",
|
||||
"tauri-plugin-single-instance",
|
||||
"tauri-plugin-updater",
|
||||
"tempfile",
|
||||
"tokio",
|
||||
"tokio-util",
|
||||
"tracing",
|
||||
"tracing-appender",
|
||||
"tracing-subscriber",
|
||||
"url",
|
||||
"windows-sys 0.61.2",
|
||||
"zip",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
||||
Generated
+2
-2
@@ -1,12 +1,12 @@
|
||||
{
|
||||
"name": "cursor-byok-desktop",
|
||||
"version": "0.1.5-beta.1",
|
||||
"version": "0.1.6",
|
||||
"lockfileVersion": 3,
|
||||
"requires": true,
|
||||
"packages": {
|
||||
"": {
|
||||
"name": "cursor-byok-desktop",
|
||||
"version": "0.1.5-beta.1",
|
||||
"version": "0.1.6",
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"@floating-ui/dom": "^1.8.0",
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "cursor-byok-desktop",
|
||||
"version": "0.1.5-beta.1",
|
||||
"version": "0.1.6",
|
||||
"description": "Cursor BYOK desktop management application",
|
||||
"type": "module",
|
||||
"scripts": {
|
||||
@@ -8,7 +8,7 @@
|
||||
"dev": "vite",
|
||||
"typecheck": "tsc --noEmit",
|
||||
"typecheck:node": "tsc --noEmit -p tsconfig.node.json",
|
||||
"i18n:scan": "STATIC_I18N_SCAN=true vite build",
|
||||
"i18n:scan": "cross-env STATIC_I18N_SCAN=true vite build",
|
||||
"build": "vite build",
|
||||
"build:demo": "npm run typecheck && npm run typecheck:node && vite build --config vite.demo.config.ts",
|
||||
"check": "npm run typecheck && npm run typecheck:node && npm run build",
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
[package]
|
||||
name = "cursor-byok-desktop"
|
||||
version = "0.1.5-beta.1"
|
||||
version = "0.1.6"
|
||||
edition = "2021"
|
||||
publish = false
|
||||
|
||||
@@ -14,6 +14,7 @@ tauri-build = { version = "2", features = [] }
|
||||
[dependencies]
|
||||
axum = "0.8"
|
||||
cursor-server = { path = "../../../server" }
|
||||
libc = "0.2"
|
||||
rfd = "0.15"
|
||||
serde = { version = "1", features = ["derive"] }
|
||||
serde_json = "1"
|
||||
@@ -24,9 +25,14 @@ tauri-plugin-opener = "2"
|
||||
tauri-plugin-autostart = "2"
|
||||
tauri-plugin-process = "2"
|
||||
tauri-plugin-updater = "2"
|
||||
tempfile = "3"
|
||||
tokio = { version = "1", features = ["time"] }
|
||||
tokio-util = "0.7"
|
||||
tracing = "0.1"
|
||||
tracing-appender = "0.2"
|
||||
tracing-subscriber = { version = "0.3", features = ["env-filter"] }
|
||||
url = "2"
|
||||
zip = { version = "4", default-features = false, features = ["deflate"] }
|
||||
|
||||
[target.'cfg(windows)'.dependencies]
|
||||
windows-sys = { version = "0.61", features = ["Win32_Foundation", "Win32_System_Threading"] }
|
||||
|
||||
@@ -1,5 +1,9 @@
|
||||
fn main() {
|
||||
let manifest = tauri_build::AppManifest::new().commands(&["open_terminal_with_command"]);
|
||||
let manifest = tauri_build::AppManifest::new().commands(&[
|
||||
"open_terminal_with_command",
|
||||
"check_portable_update",
|
||||
"install_portable_update",
|
||||
]);
|
||||
tauri_build::try_build(tauri_build::Attributes::new().app_manifest(manifest))
|
||||
.expect("failed to build Tauri application")
|
||||
}
|
||||
|
||||
@@ -18,6 +18,8 @@
|
||||
"core:window:allow-close",
|
||||
"core:app:allow-set-dock-visibility",
|
||||
"allow-open-terminal-with-command",
|
||||
"allow-check-portable-update",
|
||||
"allow-install-portable-update",
|
||||
"clipboard-manager:allow-write-text",
|
||||
"autostart:default",
|
||||
"process:allow-restart",
|
||||
|
||||
@@ -0,0 +1,11 @@
|
||||
# Automatically generated - DO NOT EDIT!
|
||||
|
||||
[[permission]]
|
||||
identifier = "allow-check-portable-update"
|
||||
description = "Enables the check_portable_update command without any pre-configured scope."
|
||||
commands.allow = ["check_portable_update"]
|
||||
|
||||
[[permission]]
|
||||
identifier = "deny-check-portable-update"
|
||||
description = "Denies the check_portable_update command without any pre-configured scope."
|
||||
commands.deny = ["check_portable_update"]
|
||||
@@ -0,0 +1,11 @@
|
||||
# Automatically generated - DO NOT EDIT!
|
||||
|
||||
[[permission]]
|
||||
identifier = "allow-install-portable-update"
|
||||
description = "Enables the install_portable_update command without any pre-configured scope."
|
||||
commands.allow = ["install_portable_update"]
|
||||
|
||||
[[permission]]
|
||||
identifier = "deny-install-portable-update"
|
||||
description = "Denies the install_portable_update command without any pre-configured scope."
|
||||
commands.deny = ["install_portable_update"]
|
||||
@@ -149,6 +149,23 @@ pub fn run() -> ExitCode {
|
||||
return ExitCode::FAILURE;
|
||||
}
|
||||
};
|
||||
#[cfg(unix)]
|
||||
{
|
||||
let open_file_limit = match crate::resource_limits::raise_open_file_limit() {
|
||||
Ok(limit) => limit,
|
||||
Err(error) => {
|
||||
diagnostics.report_fatal(&error);
|
||||
return ExitCode::FAILURE;
|
||||
}
|
||||
};
|
||||
tracing::info!(
|
||||
requested = crate::resource_limits::REQUESTED_OPEN_FILE_LIMIT,
|
||||
previous = open_file_limit.previous,
|
||||
effective = open_file_limit.effective,
|
||||
hard = open_file_limit.hard,
|
||||
"open file limit configured"
|
||||
);
|
||||
}
|
||||
tracing::info!(
|
||||
version = env!("CARGO_PKG_VERSION"),
|
||||
os = std::env::consts::OS,
|
||||
@@ -160,7 +177,11 @@ pub fn run() -> ExitCode {
|
||||
let started_by_autostart = std::env::args_os().any(|arg| arg == AUTOSTART_ARG);
|
||||
|
||||
let app = tauri::Builder::default()
|
||||
.invoke_handler(tauri::generate_handler![open_terminal_with_command])
|
||||
.invoke_handler(tauri::generate_handler![
|
||||
open_terminal_with_command,
|
||||
crate::update::check_portable_update,
|
||||
crate::update::install_portable_update,
|
||||
])
|
||||
.plugin(tauri_plugin_single_instance::init(|app, args, _| {
|
||||
if !args.iter().any(|arg| arg == AUTOSTART_ARG) {
|
||||
tray::show_main_window(app);
|
||||
@@ -228,6 +249,7 @@ pub fn run() -> ExitCode {
|
||||
window.set_focus()?;
|
||||
}
|
||||
tray::create(app)?;
|
||||
crate::update::signal_ready_if_requested()?;
|
||||
Ok(())
|
||||
})
|
||||
.build(tauri::generate_context!());
|
||||
|
||||
@@ -1,7 +1,14 @@
|
||||
mod desktop;
|
||||
#[cfg(not(dev))]
|
||||
mod frontend;
|
||||
mod resource_limits;
|
||||
mod startup;
|
||||
mod tray;
|
||||
mod update;
|
||||
|
||||
pub use desktop::run;
|
||||
pub fn run() -> std::process::ExitCode {
|
||||
if let Some(exit_code) = update::run_replacement_if_requested() {
|
||||
return exit_code;
|
||||
}
|
||||
desktop::run()
|
||||
}
|
||||
|
||||
@@ -0,0 +1,56 @@
|
||||
//! Configures process resource limits before the desktop runtime starts.
|
||||
|
||||
#[cfg(unix)]
|
||||
use std::io;
|
||||
|
||||
#[cfg(unix)]
|
||||
pub(crate) const REQUESTED_OPEN_FILE_LIMIT: u64 = 65_536;
|
||||
|
||||
#[cfg(unix)]
|
||||
pub(crate) struct OpenFileLimit {
|
||||
pub(crate) previous: u64,
|
||||
pub(crate) effective: u64,
|
||||
pub(crate) hard: u64,
|
||||
}
|
||||
|
||||
#[cfg(unix)]
|
||||
pub(crate) fn raise_open_file_limit() -> io::Result<OpenFileLimit> {
|
||||
let mut limits = libc::rlimit {
|
||||
rlim_cur: 0,
|
||||
rlim_max: 0,
|
||||
};
|
||||
// SAFETY: `limits` points to writable memory for one `rlimit` value.
|
||||
if unsafe { libc::getrlimit(libc::RLIMIT_NOFILE, &mut limits) } != 0 {
|
||||
return Err(io::Error::last_os_error());
|
||||
}
|
||||
|
||||
let previous = limits.rlim_cur;
|
||||
let target = limits
|
||||
.rlim_max
|
||||
.min(REQUESTED_OPEN_FILE_LIMIT as libc::rlim_t);
|
||||
if previous < target {
|
||||
let requested = libc::rlimit {
|
||||
rlim_cur: target,
|
||||
rlim_max: limits.rlim_max,
|
||||
};
|
||||
// SAFETY: `requested` is a valid `rlimit` value and does not raise the hard limit.
|
||||
if unsafe { libc::setrlimit(libc::RLIMIT_NOFILE, &requested) } != 0 {
|
||||
return Err(io::Error::last_os_error());
|
||||
}
|
||||
}
|
||||
|
||||
let mut effective = libc::rlimit {
|
||||
rlim_cur: 0,
|
||||
rlim_max: 0,
|
||||
};
|
||||
// SAFETY: `effective` points to writable memory for one `rlimit` value.
|
||||
if unsafe { libc::getrlimit(libc::RLIMIT_NOFILE, &mut effective) } != 0 {
|
||||
return Err(io::Error::last_os_error());
|
||||
}
|
||||
|
||||
Ok(OpenFileLimit {
|
||||
previous: previous as u64,
|
||||
effective: effective.rlim_cur as u64,
|
||||
hard: effective.rlim_max as u64,
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,269 @@
|
||||
use std::{
|
||||
fs::{self, OpenOptions},
|
||||
io::{Cursor, Read, Write},
|
||||
path::{Path, PathBuf},
|
||||
process::{Command, ExitCode},
|
||||
};
|
||||
|
||||
use serde::Serialize;
|
||||
use tauri::AppHandle;
|
||||
|
||||
#[cfg(target_os = "windows")]
|
||||
use tauri_plugin_updater::UpdaterExt;
|
||||
|
||||
#[cfg(target_os = "windows")]
|
||||
mod replacement;
|
||||
|
||||
const PORTABLE_UPDATE_ENDPOINT: &str =
|
||||
"https://github.com/leookun/cursor-byok/releases/latest/download/portable-latest.json";
|
||||
const WINDOWS_PAYLOAD_NAME: &str = "cursor-byok-desktop.exe";
|
||||
|
||||
#[derive(Serialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub(crate) struct PortableUpdateInfo {
|
||||
version: String,
|
||||
}
|
||||
|
||||
pub fn run_replacement_if_requested() -> Option<ExitCode> {
|
||||
#[cfg(target_os = "windows")]
|
||||
{
|
||||
match replacement::request_from_args() {
|
||||
Ok(Some(request)) => return Some(replacement::run(request)),
|
||||
Ok(None) => {}
|
||||
Err(error) => {
|
||||
eprintln!("invalid portable update replacement request: {error}");
|
||||
return Some(ExitCode::FAILURE);
|
||||
}
|
||||
}
|
||||
}
|
||||
None
|
||||
}
|
||||
|
||||
pub(crate) fn signal_ready_if_requested() -> std::io::Result<()> {
|
||||
#[cfg(target_os = "windows")]
|
||||
if let Some(path) = replacement::ready_marker_from_args() {
|
||||
if let Some(parent) = path.parent() {
|
||||
fs::create_dir_all(parent)?;
|
||||
}
|
||||
fs::write(path, b"ready")?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tauri::command]
|
||||
pub(crate) async fn check_portable_update(
|
||||
app: AppHandle,
|
||||
) -> Result<Option<PortableUpdateInfo>, String> {
|
||||
#[cfg(target_os = "windows")]
|
||||
{
|
||||
let update = portable_update(&app).await?;
|
||||
return Ok(update.map(|update| PortableUpdateInfo {
|
||||
version: update.version,
|
||||
}));
|
||||
}
|
||||
#[cfg(not(target_os = "windows"))]
|
||||
{
|
||||
let _ = app;
|
||||
Err("portable updates are only supported on Windows".into())
|
||||
}
|
||||
}
|
||||
|
||||
#[tauri::command]
|
||||
pub(crate) async fn install_portable_update(
|
||||
app: AppHandle,
|
||||
expected_version: String,
|
||||
) -> Result<(), String> {
|
||||
#[cfg(target_os = "windows")]
|
||||
{
|
||||
let update = portable_update(&app)
|
||||
.await?
|
||||
.ok_or_else(|| "the selected update is no longer available".to_string())?;
|
||||
if update.version != expected_version {
|
||||
return Err(format!(
|
||||
"available update changed from {expected_version} to {}",
|
||||
update.version
|
||||
));
|
||||
}
|
||||
|
||||
let target = std::env::current_exe()
|
||||
.map_err(|error| format!("failed to locate the running executable: {error}"))?;
|
||||
ensure_target_writable(&target)
|
||||
.map_err(|error| format!("the application directory is not writable: {error}"))?;
|
||||
|
||||
let bytes = update
|
||||
.download(|_, _| {}, || {})
|
||||
.await
|
||||
.map_err(|error| format!("failed to download or verify the update: {error}"))?;
|
||||
let payload = extract_windows_payload(&bytes)
|
||||
.map_err(|error| format!("invalid Windows update archive: {error}"))?;
|
||||
let staged = stage_payload(&target, &payload)
|
||||
.map_err(|error| format!("failed to stage the update: {error}"))?;
|
||||
|
||||
let handshake = staged.with_extension("started");
|
||||
let _ = fs::remove_file(&handshake);
|
||||
let mut replacement = Command::new(&staged)
|
||||
.arg("--apply-portable-update")
|
||||
.arg("--update-target")
|
||||
.arg(&target)
|
||||
.arg("--update-wait-pid")
|
||||
.arg(std::process::id().to_string())
|
||||
.arg("--update-handshake")
|
||||
.arg(&handshake)
|
||||
.spawn()
|
||||
.map_err(|error| format!("failed to start the update replacement process: {error}"))?;
|
||||
wait_for_replacement_start(&mut replacement, &handshake).await?;
|
||||
|
||||
app.exit(0);
|
||||
Ok(())
|
||||
}
|
||||
#[cfg(not(target_os = "windows"))]
|
||||
{
|
||||
let _ = (app, expected_version);
|
||||
Err("portable updates are only supported on Windows".into())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(target_os = "windows")]
|
||||
async fn portable_update(app: &AppHandle) -> Result<Option<tauri_plugin_updater::Update>, String> {
|
||||
let endpoint = PORTABLE_UPDATE_ENDPOINT
|
||||
.parse()
|
||||
.map_err(|error| format!("invalid portable update endpoint: {error}"))?;
|
||||
let updater = app
|
||||
.updater_builder()
|
||||
.endpoints(vec![endpoint])
|
||||
.map_err(|error| format!("failed to configure the updater: {error}"))?
|
||||
.build()
|
||||
.map_err(|error| format!("failed to initialize the updater: {error}"))?;
|
||||
updater
|
||||
.check()
|
||||
.await
|
||||
.map_err(|error| format!("failed to check for updates: {error}"))
|
||||
}
|
||||
|
||||
#[cfg(target_os = "windows")]
|
||||
async fn wait_for_replacement_start(
|
||||
child: &mut std::process::Child,
|
||||
handshake: &Path,
|
||||
) -> Result<(), String> {
|
||||
let wait = async {
|
||||
loop {
|
||||
if handshake.is_file() {
|
||||
return Ok(());
|
||||
}
|
||||
if let Some(status) = child
|
||||
.try_wait()
|
||||
.map_err(|error| format!("failed to inspect replacement process: {error}"))?
|
||||
{
|
||||
return Err(format!(
|
||||
"update replacement process exited before it was ready: {status}"
|
||||
));
|
||||
}
|
||||
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
|
||||
}
|
||||
};
|
||||
let result = match tokio::time::timeout(std::time::Duration::from_secs(5), wait).await {
|
||||
Ok(result) => result,
|
||||
Err(_) => {
|
||||
let _ = child.kill();
|
||||
let _ = child.wait();
|
||||
return Err("update replacement process did not become ready".into());
|
||||
}
|
||||
};
|
||||
if result.is_err() {
|
||||
let _ = child.kill();
|
||||
let _ = child.wait();
|
||||
}
|
||||
result
|
||||
}
|
||||
|
||||
fn ensure_target_writable(target: &Path) -> std::io::Result<()> {
|
||||
let parent = target
|
||||
.parent()
|
||||
.ok_or_else(|| std::io::Error::other("application executable has no parent directory"))?;
|
||||
let probe = parent.join(format!(
|
||||
".cursor-byok-update-write-test-{}",
|
||||
std::process::id()
|
||||
));
|
||||
let mut file = OpenOptions::new()
|
||||
.write(true)
|
||||
.create_new(true)
|
||||
.open(&probe)?;
|
||||
file.write_all(b"test")?;
|
||||
drop(file);
|
||||
fs::remove_file(probe)
|
||||
}
|
||||
|
||||
fn extract_windows_payload(bytes: &[u8]) -> Result<Vec<u8>, String> {
|
||||
let mut archive = zip::ZipArchive::new(Cursor::new(bytes))
|
||||
.map_err(|error| format!("failed to open ZIP: {error}"))?;
|
||||
if archive.len() != 1 {
|
||||
return Err("archive must contain exactly one file".into());
|
||||
}
|
||||
let mut entry = archive
|
||||
.by_index(0)
|
||||
.map_err(|error| format!("failed to read ZIP entry: {error}"))?;
|
||||
let name = Path::new(entry.name())
|
||||
.file_name()
|
||||
.and_then(|name| name.to_str())
|
||||
.ok_or_else(|| "archive entry has an invalid file name".to_string())?;
|
||||
if name != WINDOWS_PAYLOAD_NAME || entry.is_dir() {
|
||||
return Err(format!(
|
||||
"expected {WINDOWS_PAYLOAD_NAME}, found {}",
|
||||
entry.name()
|
||||
));
|
||||
}
|
||||
let mut payload = Vec::with_capacity(entry.size() as usize);
|
||||
entry
|
||||
.read_to_end(&mut payload)
|
||||
.map_err(|error| format!("failed to extract executable: {error}"))?;
|
||||
if payload.len() < 2 || &payload[..2] != b"MZ" {
|
||||
return Err("payload is not a Windows executable".into());
|
||||
}
|
||||
Ok(payload)
|
||||
}
|
||||
|
||||
fn stage_payload(target: &Path, payload: &[u8]) -> std::io::Result<PathBuf> {
|
||||
let directory = tempfile::Builder::new()
|
||||
.prefix("cursor-byok-portable-update-")
|
||||
.tempdir()?;
|
||||
let name = target
|
||||
.file_name()
|
||||
.ok_or_else(|| std::io::Error::other("application executable has no file name"))?;
|
||||
let path = directory.path().join(name);
|
||||
fs::write(&path, payload)?;
|
||||
let _ = directory.keep();
|
||||
Ok(path)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::io::Write;
|
||||
|
||||
fn update_zip(name: &str, payload: &[u8]) -> Vec<u8> {
|
||||
let mut bytes = Cursor::new(Vec::new());
|
||||
{
|
||||
let mut archive = zip::ZipWriter::new(&mut bytes);
|
||||
archive
|
||||
.start_file(name, zip::write::SimpleFileOptions::default())
|
||||
.unwrap();
|
||||
archive.write_all(payload).unwrap();
|
||||
archive.finish().unwrap();
|
||||
}
|
||||
bytes.into_inner()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn extracts_the_single_expected_windows_executable() {
|
||||
let bytes = update_zip(WINDOWS_PAYLOAD_NAME, b"MZpayload");
|
||||
assert_eq!(extract_windows_payload(&bytes).unwrap(), b"MZpayload");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_unexpected_or_non_executable_payloads() {
|
||||
let wrong_name = update_zip("other.exe", b"MZpayload");
|
||||
assert!(extract_windows_payload(&wrong_name).is_err());
|
||||
let wrong_content = update_zip(WINDOWS_PAYLOAD_NAME, b"not an executable");
|
||||
assert!(extract_windows_payload(&wrong_content).is_err());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,296 @@
|
||||
use std::{
|
||||
ffi::{OsStr, OsString},
|
||||
fs, io,
|
||||
path::{Path, PathBuf},
|
||||
process::{Child, Command, ExitCode},
|
||||
thread,
|
||||
time::{Duration, Instant},
|
||||
};
|
||||
|
||||
const APPLY_ARG: &str = "--apply-portable-update";
|
||||
const TARGET_ARG: &str = "--update-target";
|
||||
const PID_ARG: &str = "--update-wait-pid";
|
||||
const HANDSHAKE_ARG: &str = "--update-handshake";
|
||||
pub(super) const READY_ARG: &str = "--portable-update-ready";
|
||||
const PROCESS_WAIT_TIMEOUT: Duration = Duration::from_secs(30);
|
||||
const READY_WAIT_TIMEOUT: Duration = Duration::from_secs(30);
|
||||
|
||||
pub(super) struct ReplacementRequest {
|
||||
target: PathBuf,
|
||||
pid: u32,
|
||||
handshake: PathBuf,
|
||||
}
|
||||
|
||||
pub(super) fn request_from_args() -> Result<Option<ReplacementRequest>, String> {
|
||||
let args = std::env::args_os().collect::<Vec<_>>();
|
||||
if !args.iter().any(|arg| arg == APPLY_ARG) {
|
||||
return Ok(None);
|
||||
}
|
||||
let target = PathBuf::from(
|
||||
argument_value(&args, TARGET_ARG).ok_or_else(|| format!("{TARGET_ARG} is required"))?,
|
||||
);
|
||||
let pid = argument_value(&args, PID_ARG)
|
||||
.ok_or_else(|| format!("{PID_ARG} is required"))?
|
||||
.to_string_lossy()
|
||||
.parse::<u32>()
|
||||
.map_err(|error| format!("invalid {PID_ARG}: {error}"))?;
|
||||
let handshake = PathBuf::from(
|
||||
argument_value(&args, HANDSHAKE_ARG)
|
||||
.ok_or_else(|| format!("{HANDSHAKE_ARG} is required"))?,
|
||||
);
|
||||
validate_target(&target).map_err(|error| error.to_string())?;
|
||||
Ok(Some(ReplacementRequest {
|
||||
target,
|
||||
pid,
|
||||
handshake,
|
||||
}))
|
||||
}
|
||||
|
||||
pub(super) fn ready_marker_from_args() -> Option<PathBuf> {
|
||||
let args = std::env::args_os().collect::<Vec<_>>();
|
||||
argument_value(&args, READY_ARG).map(PathBuf::from)
|
||||
}
|
||||
|
||||
fn argument_value<'a>(args: &'a [OsString], name: &str) -> Option<&'a OsStr> {
|
||||
args.iter()
|
||||
.position(|arg| arg == name)
|
||||
.and_then(|index| args.get(index + 1))
|
||||
.map(OsString::as_os_str)
|
||||
}
|
||||
|
||||
fn validate_target(target: &Path) -> io::Result<()> {
|
||||
if !target.is_absolute() || !target.is_file() {
|
||||
return Err(io::Error::new(
|
||||
io::ErrorKind::InvalidInput,
|
||||
"update target must be an existing absolute file",
|
||||
));
|
||||
}
|
||||
let source_name = std::env::current_exe()?
|
||||
.file_name()
|
||||
.map(OsStr::to_os_string)
|
||||
.ok_or_else(|| io::Error::new(io::ErrorKind::InvalidInput, "updater has no file name"))?;
|
||||
if target.file_name() != Some(source_name.as_os_str()) {
|
||||
return Err(io::Error::new(
|
||||
io::ErrorKind::InvalidInput,
|
||||
"update target file name does not match the updater",
|
||||
));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub(super) fn run(request: ReplacementRequest) -> ExitCode {
|
||||
match run_inner(request) {
|
||||
Ok(()) => ExitCode::SUCCESS,
|
||||
Err(error) => {
|
||||
eprintln!("portable update replacement failed: {error}");
|
||||
ExitCode::FAILURE
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn run_inner(request: ReplacementRequest) -> io::Result<()> {
|
||||
wait_for_process(request.pid, PROCESS_WAIT_TIMEOUT, &request.handshake)?;
|
||||
remove_file_if_exists(&request.handshake)?;
|
||||
let source = std::env::current_exe()?;
|
||||
let backup = backup_path(&request.target);
|
||||
let ready = source.with_extension("ready");
|
||||
remove_file_if_exists(&ready)?;
|
||||
|
||||
if let Err(error) = install_staged(&source, &request.target, &backup) {
|
||||
relaunch(&request.target);
|
||||
return Err(error);
|
||||
}
|
||||
let mut child = match Command::new(&request.target)
|
||||
.arg(READY_ARG)
|
||||
.arg(&ready)
|
||||
.spawn()
|
||||
{
|
||||
Ok(child) => child,
|
||||
Err(error) => {
|
||||
restore_backup(&request.target, &backup)?;
|
||||
relaunch(&request.target);
|
||||
return Err(error);
|
||||
}
|
||||
};
|
||||
|
||||
match wait_until_ready(&mut child, &ready, READY_WAIT_TIMEOUT) {
|
||||
Ok(()) => {
|
||||
let _ = fs::remove_file(&backup);
|
||||
let _ = fs::remove_file(&ready);
|
||||
Ok(())
|
||||
}
|
||||
Err(error) => {
|
||||
let _ = child.kill();
|
||||
let _ = child.wait();
|
||||
restore_backup(&request.target, &backup)?;
|
||||
relaunch(&request.target);
|
||||
Err(error)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn relaunch(target: &Path) {
|
||||
if target.is_file() {
|
||||
let _ = Command::new(target).spawn();
|
||||
}
|
||||
}
|
||||
|
||||
fn remove_file_if_exists(path: &Path) -> io::Result<()> {
|
||||
match fs::remove_file(path) {
|
||||
Ok(()) => Ok(()),
|
||||
Err(error) if error.kind() == io::ErrorKind::NotFound => Ok(()),
|
||||
Err(error) => Err(error),
|
||||
}
|
||||
}
|
||||
|
||||
fn backup_path(target: &Path) -> PathBuf {
|
||||
path_with_suffix(target, ".old")
|
||||
}
|
||||
|
||||
fn pending_path(target: &Path) -> PathBuf {
|
||||
path_with_suffix(target, ".new")
|
||||
}
|
||||
|
||||
fn path_with_suffix(target: &Path, suffix: &str) -> PathBuf {
|
||||
let mut name = target.as_os_str().to_os_string();
|
||||
name.push(suffix);
|
||||
PathBuf::from(name)
|
||||
}
|
||||
|
||||
fn install_staged(source: &Path, target: &Path, backup: &Path) -> io::Result<()> {
|
||||
let pending = pending_path(target);
|
||||
remove_file_if_exists(&pending)?;
|
||||
fs::copy(source, &pending)?;
|
||||
|
||||
let result = activate_pending(&pending, target, backup);
|
||||
if result.is_err() {
|
||||
let _ = fs::remove_file(&pending);
|
||||
}
|
||||
result
|
||||
}
|
||||
|
||||
fn activate_pending(pending: &Path, target: &Path, backup: &Path) -> io::Result<()> {
|
||||
remove_file_if_exists(backup)?;
|
||||
fs::rename(target, backup)?;
|
||||
if let Err(error) = fs::rename(pending, target) {
|
||||
if let Err(restore_error) = restore_backup(target, backup) {
|
||||
return Err(io::Error::other(format!(
|
||||
"failed to install update ({error}) and restore the original executable ({restore_error})"
|
||||
)));
|
||||
}
|
||||
return Err(error);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn restore_backup(target: &Path, backup: &Path) -> io::Result<()> {
|
||||
remove_file_if_exists(target)?;
|
||||
fs::rename(backup, target)
|
||||
}
|
||||
|
||||
fn wait_until_ready(child: &mut Child, marker: &Path, timeout: Duration) -> io::Result<()> {
|
||||
let deadline = Instant::now() + timeout;
|
||||
loop {
|
||||
if marker.is_file() {
|
||||
return Ok(());
|
||||
}
|
||||
if let Some(status) = child.try_wait()? {
|
||||
return Err(io::Error::other(format!(
|
||||
"updated application exited before startup completed: {status}"
|
||||
)));
|
||||
}
|
||||
if Instant::now() >= deadline {
|
||||
return Err(io::Error::new(
|
||||
io::ErrorKind::TimedOut,
|
||||
"updated application did not report a successful startup",
|
||||
));
|
||||
}
|
||||
thread::sleep(Duration::from_millis(100));
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(windows)]
|
||||
fn wait_for_process(pid: u32, timeout: Duration, handshake: &Path) -> io::Result<()> {
|
||||
use windows_sys::Win32::{
|
||||
Foundation::{
|
||||
CloseHandle, GetLastError, ERROR_INVALID_PARAMETER, WAIT_OBJECT_0, WAIT_TIMEOUT,
|
||||
},
|
||||
System::Threading::{OpenProcess, WaitForSingleObject},
|
||||
};
|
||||
|
||||
const SYNCHRONIZE_ACCESS: u32 = 0x0010_0000;
|
||||
let handle = unsafe { OpenProcess(SYNCHRONIZE_ACCESS, 0, pid) };
|
||||
if handle.is_null() {
|
||||
let error = unsafe { GetLastError() };
|
||||
return if error == ERROR_INVALID_PARAMETER {
|
||||
fs::write(handshake, b"started")
|
||||
} else {
|
||||
Err(io::Error::from_raw_os_error(error as i32))
|
||||
};
|
||||
}
|
||||
if let Err(error) = fs::write(handshake, b"started") {
|
||||
unsafe { CloseHandle(handle) };
|
||||
return Err(error);
|
||||
}
|
||||
let milliseconds = timeout.as_millis().min(u32::MAX as u128) as u32;
|
||||
let result = unsafe { WaitForSingleObject(handle, milliseconds) };
|
||||
unsafe { CloseHandle(handle) };
|
||||
match result {
|
||||
WAIT_OBJECT_0 => Ok(()),
|
||||
WAIT_TIMEOUT => Err(io::Error::new(
|
||||
io::ErrorKind::TimedOut,
|
||||
"running application did not exit before the update timeout",
|
||||
)),
|
||||
_ => Err(io::Error::last_os_error()),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn staged_file_can_be_restored() {
|
||||
let directory = tempfile::tempdir().unwrap();
|
||||
let source = directory.path().join("source.exe");
|
||||
let target = directory.path().join("target.exe");
|
||||
let backup = backup_path(&target);
|
||||
fs::write(&source, b"new").unwrap();
|
||||
fs::write(&target, b"old").unwrap();
|
||||
|
||||
install_staged(&source, &target, &backup).unwrap();
|
||||
assert_eq!(fs::read(&target).unwrap(), b"new");
|
||||
assert_eq!(fs::read(&backup).unwrap(), b"old");
|
||||
|
||||
restore_backup(&target, &backup).unwrap();
|
||||
assert_eq!(fs::read(&target).unwrap(), b"old");
|
||||
assert!(!backup.exists());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn missing_pending_file_restores_original_after_backup() {
|
||||
let directory = tempfile::tempdir().unwrap();
|
||||
let target = directory.path().join("target.exe");
|
||||
let pending = pending_path(&target);
|
||||
let backup = backup_path(&target);
|
||||
fs::write(&target, b"old").unwrap();
|
||||
|
||||
assert!(activate_pending(&pending, &target, &backup).is_err());
|
||||
assert_eq!(fs::read(&target).unwrap(), b"old");
|
||||
assert!(!backup.exists());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn missing_source_preserves_original_file() {
|
||||
let directory = tempfile::tempdir().unwrap();
|
||||
let source = directory.path().join("missing.exe");
|
||||
let target = directory.path().join("target.exe");
|
||||
let backup = backup_path(&target);
|
||||
fs::write(&target, b"old").unwrap();
|
||||
|
||||
assert!(install_staged(&source, &target, &backup).is_err());
|
||||
assert_eq!(fs::read(&target).unwrap(), b"old");
|
||||
assert!(!backup.exists());
|
||||
assert!(!pending_path(&target).exists());
|
||||
}
|
||||
}
|
||||
@@ -1,7 +1,7 @@
|
||||
{
|
||||
"$schema": "https://schema.tauri.app/config/2",
|
||||
"productName": "Cursor BYOK",
|
||||
"version": "0.1.5-beta.1",
|
||||
"version": "0.1.6",
|
||||
"identifier": "dev.cursorbyok.desktop",
|
||||
"build": {
|
||||
"beforeDevCommand": "npm run dev",
|
||||
@@ -12,7 +12,7 @@
|
||||
"app": {
|
||||
"windows": [],
|
||||
"security": {
|
||||
"csp": "default-src 'self'; img-src 'self' asset: http://asset.localhost data:; style-src 'self' 'unsafe-inline'; connect-src 'self' http://127.0.0.1:*",
|
||||
"csp": "default-src 'self'; img-src 'self' asset: http://asset.localhost data: https:; style-src 'self' 'unsafe-inline'; connect-src 'self' http://127.0.0.1:*",
|
||||
"dangerousDisableAssetCspModification": [
|
||||
"style-src"
|
||||
]
|
||||
|
||||
@@ -185,6 +185,7 @@ function createModel({ hash, order, name, type, url, modelId, endpoint = "/v1/re
|
||||
model_hash: hash,
|
||||
sort_order: order,
|
||||
display_name: name,
|
||||
group_name: null,
|
||||
type,
|
||||
base_url: url,
|
||||
use_full_url: false,
|
||||
|
||||
@@ -32,6 +32,7 @@ type CursorModelCardsProps = {
|
||||
onTestPluginModel: (model: PluginModelDescriptor) => void;
|
||||
onPluginSettings: (model: PluginModelDescriptor) => void;
|
||||
onReorder: (modelHashes: string[]) => void;
|
||||
onGroupSettings: (group: CursorModelGroup) => void;
|
||||
};
|
||||
|
||||
type ModelGridProps = Omit<CursorModelCardsProps, "grouping" | "pluginModels" | "onTestPluginModel" | "onPluginSettings"> & {
|
||||
@@ -60,6 +61,7 @@ export function CursorModelCards(props: CursorModelCardsProps) {
|
||||
key={group.key}
|
||||
label={group.label}
|
||||
icon={group.icon}
|
||||
onSettings={props.grouping === "provider" ? () => props.onGroupSettings(group) : undefined}
|
||||
>
|
||||
{group.models.map((model) => <ModelListRow
|
||||
key={model.model_hash}
|
||||
@@ -107,25 +109,37 @@ function pluginGroups(models: PluginModelDescriptor[]) {
|
||||
return groups;
|
||||
}
|
||||
|
||||
function CollapsibleGroup({ label, icon, iconSrc, children }: {
|
||||
function CollapsibleGroup({ label, icon, iconSrc, onSettings, children }: {
|
||||
label: string;
|
||||
icon?: IconifyIcon;
|
||||
iconSrc?: string;
|
||||
onSettings?: () => void;
|
||||
children: ReactNode;
|
||||
}) {
|
||||
const [open, setOpen] = useState(true);
|
||||
return <Card className={styles.groupCard}>
|
||||
<button
|
||||
type="button"
|
||||
className={styles.groupToggle}
|
||||
aria-expanded={open}
|
||||
onClick={() => setOpen((current) => !current)}
|
||||
>
|
||||
{icon && <Icon icon={icon} size="1.1em" />}
|
||||
{iconSrc && <Icon src={iconSrc} size="1.1em" />}
|
||||
<span className={styles.groupLabel}>{label}</span>
|
||||
<Icon icon={open ? chevronDownIcon : chevronRightIcon} size="1em" />
|
||||
</button>
|
||||
<div className={styles.groupHeader}>
|
||||
<button
|
||||
type="button"
|
||||
className={styles.groupToggle}
|
||||
aria-expanded={open}
|
||||
onClick={() => setOpen((current) => !current)}
|
||||
>
|
||||
{icon && <Icon icon={icon} size="1.1em" />}
|
||||
{iconSrc && <Icon src={iconSrc} size="1.1em" />}
|
||||
<span className={styles.groupLabel}>{label}</span>
|
||||
</button>
|
||||
{onSettings && <button type="button" className={styles.groupSettings} onClick={onSettings}>{t("分组设置")}</button>}
|
||||
<button
|
||||
type="button"
|
||||
className={styles.groupChevron}
|
||||
tabIndex={-1}
|
||||
aria-hidden="true"
|
||||
onClick={() => setOpen((current) => !current)}
|
||||
>
|
||||
<Icon icon={open ? chevronDownIcon : chevronRightIcon} size="1em" />
|
||||
</button>
|
||||
</div>
|
||||
{open && <div className={styles.modelList}>{children}</div>}
|
||||
</Card>;
|
||||
}
|
||||
@@ -273,8 +287,9 @@ function ModelGrid({
|
||||
}
|
||||
|
||||
function providerGroup(model: Model) {
|
||||
const label = providerDomain(model.base_url);
|
||||
return { key: label, label, icon: flatColorOrganizationIcon };
|
||||
const key = providerDomain(model.base_url);
|
||||
const label = model.group_name?.trim() || key;
|
||||
return { key, label, icon: flatColorOrganizationIcon };
|
||||
}
|
||||
|
||||
function providerDomain(baseUrl: string) {
|
||||
|
||||
@@ -24,6 +24,7 @@ export const emptyCursorModelDraft = (): CursorModelDraft => ({
|
||||
model: {
|
||||
sort_order: 0,
|
||||
display_name: "",
|
||||
group_name: null,
|
||||
type: "openai",
|
||||
base_url: "",
|
||||
use_full_url: false,
|
||||
|
||||
@@ -73,8 +73,34 @@
|
||||
.groupCard {
|
||||
padding: 4px 12px 8px;
|
||||
}
|
||||
.groupHeader {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 8px;
|
||||
}
|
||||
.groupSettings {
|
||||
padding: 4px 2px;
|
||||
color: var(--vscode-descriptionForeground);
|
||||
background: none;
|
||||
border: none;
|
||||
font-size: type.$font-size-xs;
|
||||
white-space: nowrap;
|
||||
cursor: pointer;
|
||||
|
||||
&:hover { color: var(--vscode-foreground); }
|
||||
}
|
||||
.groupChevron {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
padding: 4px 0;
|
||||
color: var(--vscode-foreground);
|
||||
background: none;
|
||||
border: none;
|
||||
cursor: pointer;
|
||||
}
|
||||
.groupToggle {
|
||||
width: 100%;
|
||||
flex: 1;
|
||||
min-width: 0;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 8px;
|
||||
|
||||
@@ -2,13 +2,14 @@ import { useCallback, useEffect, useRef, useState } from "react";
|
||||
import { useNavigate } from "react-router-dom";
|
||||
import { api, configuredPluginModels, type Model, type ModelInput } from "../../shared/api";
|
||||
import { CursorCaGate, CursorCaProvider, CursorModelGate, CursorModelProvider } from "./CursorGates";
|
||||
import { CursorModelCards, cursorModelGroups, type CursorModelGrouping } from "./CursorModelCards";
|
||||
import { CursorModelCards, cursorModelGroups, type CursorModelGroup, type CursorModelGrouping } from "./CursorModelCards";
|
||||
import { CursorModelEditor, emptyCursorModelDraft, type CursorModelDraft } from "./CursorModelEditor";
|
||||
import { CursorModelTestResult, type CursorModelTestState } from "./CursorModelTestResult";
|
||||
import styles from "./CursorSettings.module.scss";
|
||||
import { PageContent } from "../../shell/layout/PageContent";
|
||||
import { LegacyModelImport } from "./LegacyModelImport";
|
||||
import { ConfirmDialog } from "../../shared/ui/ConfirmDialog";
|
||||
import { FormField, SecretTextInput, TextInput } from "../../shared/ui/FormControls";
|
||||
import controls from "../../shared/ui/Controls.module.scss";
|
||||
import { Icon } from "../../shared/ui/Icon";
|
||||
import { Modal } from "../../shared/ui/Modal";
|
||||
@@ -34,6 +35,11 @@ export function CursorSettingsPage() {
|
||||
const [savingAndTesting, setSavingAndTesting] = useState(false);
|
||||
const [batchTesting, setBatchTesting] = useState(false);
|
||||
const [grouping, setGrouping] = useState<CursorModelGrouping>("flat");
|
||||
const [settingsGroup, setSettingsGroup] = useState<CursorModelGroup | null>(null);
|
||||
const [groupNameDraft, setGroupNameDraft] = useState("");
|
||||
const [groupBaseUrlDraft, setGroupBaseUrlDraft] = useState("");
|
||||
const [groupApiKeyDraft, setGroupApiKeyDraft] = useState("");
|
||||
const [groupSettingsBusy, setGroupSettingsBusy] = useState(false);
|
||||
const activeModelTests = useRef(new Map<string, { testId: string; controller: AbortController; cancelling: boolean }>());
|
||||
const caReady = cursorHarness?.ca === "ready";
|
||||
const pluginModels = configuredPluginModels(plugins);
|
||||
@@ -208,6 +214,39 @@ export function CursorSettingsPage() {
|
||||
}]);
|
||||
if (created) message(t("模型已复制"));
|
||||
};
|
||||
const openGroupSettings = (group: CursorModelGroup) => {
|
||||
setGroupNameDraft(group.models.find((model) => model.group_name?.trim())?.group_name?.trim() ?? "");
|
||||
setGroupBaseUrlDraft(sharedValue(group.models.map((model) => model.base_url)) ?? "");
|
||||
setGroupApiKeyDraft(sharedValue(group.models.map((model) => model.api_key)) ?? "");
|
||||
setSettingsGroup(group);
|
||||
};
|
||||
const saveGroupSettings = async () => {
|
||||
if (!settingsGroup) return;
|
||||
const group_name = groupNameDraft.trim() || null;
|
||||
const base_url = groupBaseUrlDraft.trim();
|
||||
const api_key = groupApiKeyDraft.trim();
|
||||
setGroupSettingsBusy(true);
|
||||
try {
|
||||
for (const model of settingsGroup.models) {
|
||||
const input: ModelInput = {
|
||||
...modelInput(model),
|
||||
group_name,
|
||||
...(base_url ? { base_url } : {}),
|
||||
...(api_key ? { api_key } : {}),
|
||||
};
|
||||
if (input.group_name === (model.group_name ?? null)
|
||||
&& input.base_url === model.base_url
|
||||
&& input.api_key === model.api_key) continue;
|
||||
await api.updateModel(model.model_hash, input);
|
||||
}
|
||||
await appStore.refresh();
|
||||
setSettingsGroup(null);
|
||||
} catch (cause) {
|
||||
message(errorText(cause));
|
||||
} finally {
|
||||
setGroupSettingsBusy(false);
|
||||
}
|
||||
};
|
||||
const reorderModels = useCallback(async (modelHashes: string[]) => {
|
||||
if (!await appStore.reorderCursorModels(modelHashes)) {
|
||||
message(appStore.getSnapshot().error || t("排序失败"));
|
||||
@@ -228,6 +267,7 @@ export function CursorSettingsPage() {
|
||||
onTestPluginModel={(model) => void testModel({ model_hash: model.id, display_name: model.displayName })}
|
||||
onPluginSettings={() => navigate("/plugins")}
|
||||
onReorder={reorderModels}
|
||||
onGroupSettings={openGroupSettings}
|
||||
/>;
|
||||
|
||||
const refreshCa = async () => {
|
||||
@@ -274,6 +314,19 @@ export function CursorSettingsPage() {
|
||||
<ConfirmDialog open={caCommand !== null} title={t("安装本地 CA")} cancelLabel={t("关闭")} confirmLabel={t("打开终端")} onCancel={() => setCaCommand(null)} onConfirm={openCaTerminal}>
|
||||
<div className={styles.editor}><strong>{t("需要授权安装证书")}</strong><span>{t("安装命令已自动复制。点击“打开终端”,将命令粘贴到终端中执行,并按提示输入密码。")}</span><pre className={styles.command}>{caCommand}</pre></div>
|
||||
</ConfirmDialog>
|
||||
<Modal open={settingsGroup !== null} title={t("分组设置")} busy={groupSettingsBusy || cursorBusy} onClose={() => setSettingsGroup(null)} onSubmit={() => void saveGroupSettings()} submitLabel={t("保存")}>
|
||||
{settingsGroup && <div className={styles.editor}>
|
||||
<FormField label={t("分组名称")} hint={t("应用于该分组下的全部模型,并作为 Cursor 模型选择器中的徽章标签;清空则恢复显示服务器域名。")}>
|
||||
<TextInput placeholder={settingsGroup.key} value={groupNameDraft} onChange={(event) => setGroupNameDraft(event.target.value)} />
|
||||
</FormField>
|
||||
<FormField label={t("服务器地址")} hint={t("修改后应用于该分组下的全部模型;留空保持各模型现有配置不变。")}>
|
||||
<TextInput placeholder={t("留空保持不变")} value={groupBaseUrlDraft} onChange={(event) => setGroupBaseUrlDraft(event.target.value)} />
|
||||
</FormField>
|
||||
<FormField label="API Key" hint={t("修改后应用于该分组下的全部模型;留空保持各模型现有配置不变。")}>
|
||||
<SecretTextInput placeholder={t("留空保持不变")} autoComplete="off" value={groupApiKeyDraft} onChange={(event) => setGroupApiKeyDraft(event.target.value)} />
|
||||
</FormField>
|
||||
</div>}
|
||||
</Modal>
|
||||
<ConfirmDialog open={deleting !== null} title={t("删除模型")} cancelLabel={t("取消")} confirmLabel={t("删除")} onCancel={() => setDeleting(null)} onConfirm={() => { if (deleting) void appStore.deleteModel(deleting.model_hash); setDeleting(null); }}><p>{t("确定删除这个模型吗?")}</p></ConfirmDialog>
|
||||
</>;
|
||||
}
|
||||
@@ -283,6 +336,13 @@ function modelInput(model: Model): ModelInput {
|
||||
return input;
|
||||
}
|
||||
|
||||
/** 组内所有模型取值一致时返回该值,否则返回 null(表单留空表示保持不变)。 */
|
||||
function sharedValue(values: string[]): string | null {
|
||||
const [first, ...rest] = values;
|
||||
if (first === undefined) return null;
|
||||
return rest.every((value) => value === first) ? first : null;
|
||||
}
|
||||
|
||||
function draftInput(draft: CursorModelDraft): ModelInput {
|
||||
const model = {
|
||||
...draft.model,
|
||||
|
||||
@@ -95,7 +95,8 @@ export function AppLifecycleSettingsCard() {
|
||||
const nextVersion = await updateStore.check();
|
||||
message(nextVersion ? t("发现新版本 {version}", { version: nextVersion }) : t("当前已是最新版本"));
|
||||
} catch (cause) {
|
||||
message(cause instanceof Error ? cause.message : String(cause));
|
||||
const error = cause instanceof Error ? cause.message : String(cause);
|
||||
message(t("检查更新失败:{error}", { error }));
|
||||
}
|
||||
};
|
||||
|
||||
@@ -103,7 +104,8 @@ export function AppLifecycleSettingsCard() {
|
||||
try {
|
||||
await updateStore.install();
|
||||
} catch (cause) {
|
||||
message(cause instanceof Error ? cause.message : String(cause));
|
||||
const error = cause instanceof Error ? cause.message : String(cause);
|
||||
message(t("安装更新失败:{error}", { error }));
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -79,6 +79,7 @@
|
||||
"2f9daa828907b93f": "Delete",
|
||||
"2fe5a8d0eee9f14c": "Invalid",
|
||||
"303c30f301514250": "Search resources",
|
||||
"3260348163d03b8e": "Leave blank to keep unchanged",
|
||||
"32896fdaaaa4c106": "Account saved and the model catalog is synced.",
|
||||
"346ff60e6c7c5181": "Reading…",
|
||||
"36f33adaf0942634": "Confirm",
|
||||
@@ -220,6 +221,7 @@
|
||||
"91aaf184cfc17ffd": "Overview",
|
||||
"91af6e57e7453fbe": "Add account",
|
||||
"92156a483d4ba248": "Only request, response, and trace attachments are deleted; call summaries, metrics, and configuration are kept.",
|
||||
"92e26b27d5ea8f0e": "Failed to check for updates: {error}",
|
||||
"940a168911ade998": "Items per page",
|
||||
"945fb1c67eca8493": "Installing the plugin runtime",
|
||||
"946b3ffc02f026c0": "Delete this model?",
|
||||
@@ -258,6 +260,7 @@
|
||||
"a748cc074f78de00": "View details",
|
||||
"a7617f42f898b2bf": "Use complete request URL",
|
||||
"a8036485f9227f2c": "Drag to reorder",
|
||||
"a80b53f8848e6d27": "Failed to install update: {error}",
|
||||
"a98585871c5313ff": "Display name",
|
||||
"ab9084a640fbb864": "Deselect all",
|
||||
"abecab6701177721": "Launch at login enabled",
|
||||
@@ -270,6 +273,7 @@
|
||||
"b06325c5660f0c29": "Direct",
|
||||
"b16c3b2ecedd6fe1": "Cursor integration is active. Add a model configuration to use a BYOK model.",
|
||||
"b254ff315d861346": "Try initializing again",
|
||||
"b2617bf9ae663752": "Group settings",
|
||||
"b4411558b932266f": "Provider type",
|
||||
"b4c9e08870d41aa2": "Initialize the plugin runtime first",
|
||||
"b502b1d414664337": "Prompt: {tokens}",
|
||||
@@ -290,6 +294,7 @@
|
||||
"bb7efdcb6af6e805": "Default dark",
|
||||
"bda62ce1d5e4ace9": "Tell us why",
|
||||
"bda74b5674b6a57d": "Initialize plugins",
|
||||
"be961dc60ab610da": "Applies to every model in this group when changed; leave blank to keep each model's current configuration.",
|
||||
"bf57afd709694b55": "Overview time range",
|
||||
"bfc01caf9fe0c841": "Cache hit rate {rate}",
|
||||
"c0b3fbff51ccc40b": "Done",
|
||||
@@ -322,6 +327,7 @@
|
||||
"d60669bb26a22f5d": "Leave blank to use the default",
|
||||
"d6b1f203680f5496": "Leave blank to use adaptive thinking",
|
||||
"d766536c18e8e990": "Plugin runtime {version} is installed and ready to use.",
|
||||
"d7e266bdc8064193": "Group name",
|
||||
"d86fa42c3848c680": "Use system proxy",
|
||||
"d8c47e9776cf1082": "Main menu",
|
||||
"da521d1c1cbd36af": "Authorization is required to install the certificate",
|
||||
@@ -361,6 +367,7 @@
|
||||
"ea26b760e930a7ca": "Call observability",
|
||||
"eb11e2df1d8ae387": "Provider URL",
|
||||
"eb1be07f2ca6e506": "Estimated using Claude Opus 4.7 pricing.",
|
||||
"eb4a3db23661fb52": "Applies to every model in this group and is used as the badge label in Cursor's model picker; clear it to fall back to the server domain.",
|
||||
"eb77492c9f76a7e1": "The install command has been copied. Click “Open terminal”, paste it into the terminal, and enter your password when prompted.",
|
||||
"eba54690937bc532": "Manage accounts",
|
||||
"ed31fbb483ee1b0a": "Actions",
|
||||
|
||||
@@ -79,6 +79,7 @@
|
||||
"2f9daa828907b93f": "删除",
|
||||
"2fe5a8d0eee9f14c": "已失效",
|
||||
"303c30f301514250": "搜索资源",
|
||||
"3260348163d03b8e": "留空保持不变",
|
||||
"32896fdaaaa4c106": "账号已保存,模型目录已同步。",
|
||||
"346ff60e6c7c5181": "读取中…",
|
||||
"36f33adaf0942634": "确认",
|
||||
@@ -220,6 +221,7 @@
|
||||
"91aaf184cfc17ffd": "数据概览",
|
||||
"91af6e57e7453fbe": "添加账号",
|
||||
"92156a483d4ba248": "仅删除请求、响应和追踪附件等详细内容,保留调用汇总、统计指标和配置。",
|
||||
"92e26b27d5ea8f0e": "检查更新失败:{error}",
|
||||
"940a168911ade998": "每页条数",
|
||||
"945fb1c67eca8493": "正在安装插件运行时",
|
||||
"946b3ffc02f026c0": "确定删除这个模型吗?",
|
||||
@@ -258,6 +260,7 @@
|
||||
"a748cc074f78de00": "查看详情",
|
||||
"a7617f42f898b2bf": "使用完整请求地址",
|
||||
"a8036485f9227f2c": "拖动排序",
|
||||
"a80b53f8848e6d27": "安装更新失败:{error}",
|
||||
"a98585871c5313ff": "显示名称",
|
||||
"ab9084a640fbb864": "全不选",
|
||||
"abecab6701177721": "已开启开机启动",
|
||||
@@ -270,6 +273,7 @@
|
||||
"b06325c5660f0c29": "直连",
|
||||
"b16c3b2ecedd6fe1": "Cursor 接管已生效;添加模型配置后即可使用 BYOK 模型。",
|
||||
"b254ff315d861346": "请重试初始化",
|
||||
"b2617bf9ae663752": "分组设置",
|
||||
"b4411558b932266f": "上游类型",
|
||||
"b4c9e08870d41aa2": "需要先初始化插件运行时",
|
||||
"b502b1d414664337": "提示词:{tokens}",
|
||||
@@ -290,6 +294,7 @@
|
||||
"bb7efdcb6af6e805": "默认暗色",
|
||||
"bda62ce1d5e4ace9": "可以告诉我们原因",
|
||||
"bda74b5674b6a57d": "初始化插件",
|
||||
"be961dc60ab610da": "修改后应用于该分组下的全部模型;留空保持各模型现有配置不变。",
|
||||
"bf57afd709694b55": "概览时间范围",
|
||||
"bfc01caf9fe0c841": "缓存命中率 {rate}",
|
||||
"c0b3fbff51ccc40b": "完成",
|
||||
@@ -322,6 +327,7 @@
|
||||
"d60669bb26a22f5d": "留空使用默认值",
|
||||
"d6b1f203680f5496": "留空使用 adaptive thinking",
|
||||
"d766536c18e8e990": "插件运行时 {version} 已安装,可以开始使用插件。",
|
||||
"d7e266bdc8064193": "分组名称",
|
||||
"d86fa42c3848c680": "使用系统代理",
|
||||
"d8c47e9776cf1082": "主菜单",
|
||||
"da521d1c1cbd36af": "需要授权安装证书",
|
||||
@@ -361,6 +367,7 @@
|
||||
"ea26b760e930a7ca": "调用观测",
|
||||
"eb11e2df1d8ae387": "上游地址",
|
||||
"eb1be07f2ca6e506": "按 Claude Opus 4.7 价格估算。",
|
||||
"eb4a3db23661fb52": "应用于该分组下的全部模型,并作为 Cursor 模型选择器中的徽章标签;清空则恢复显示服务器域名。",
|
||||
"eb77492c9f76a7e1": "安装命令已自动复制。点击“打开终端”,将命令粘贴到终端中执行,并按提示输入密码。",
|
||||
"eba54690937bc532": "账号管理",
|
||||
"ed31fbb483ee1b0a": "操作",
|
||||
|
||||
@@ -7,6 +7,7 @@ export interface Model {
|
||||
model_hash: string;
|
||||
sort_order: number;
|
||||
display_name: string;
|
||||
group_name: string | null;
|
||||
type: ModelType;
|
||||
base_url: string;
|
||||
use_full_url: boolean;
|
||||
@@ -33,6 +34,7 @@ export interface Model {
|
||||
export interface ModelInput {
|
||||
sort_order: number;
|
||||
display_name: string;
|
||||
group_name: string | null;
|
||||
type: ModelType;
|
||||
base_url: string;
|
||||
use_full_url: boolean;
|
||||
@@ -233,9 +235,7 @@ export interface PluginModelDescriptor {
|
||||
description: string | null;
|
||||
icon: string;
|
||||
providerType: string;
|
||||
contextWindowTokens: number | null;
|
||||
maxOutputTokens: number | null;
|
||||
thinking: boolean;
|
||||
images: boolean;
|
||||
}
|
||||
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
import { getVersion, setDockVisibility } from "@tauri-apps/api/app";
|
||||
import { isTauri } from "@tauri-apps/api/core";
|
||||
import { invoke, isTauri } from "@tauri-apps/api/core";
|
||||
import { disable, enable, isEnabled } from "@tauri-apps/plugin-autostart";
|
||||
import { relaunch } from "@tauri-apps/plugin-process";
|
||||
import { check, type Update } from "@tauri-apps/plugin-updater";
|
||||
@@ -46,11 +46,39 @@ export async function writeDockIconVisibility(visible: boolean): Promise<void> {
|
||||
}
|
||||
}
|
||||
|
||||
export async function checkForUpdate(): Promise<Update | null> {
|
||||
return check();
|
||||
export type AppUpdate = {
|
||||
version: string;
|
||||
install(): Promise<void>;
|
||||
close(): Promise<void>;
|
||||
};
|
||||
|
||||
type PortableUpdateInfo = {
|
||||
version: string;
|
||||
};
|
||||
|
||||
export async function checkForUpdate(): Promise<AppUpdate | null> {
|
||||
if (desktopPlatform() === "windows") {
|
||||
const update = await invoke<PortableUpdateInfo | null>("check_portable_update");
|
||||
if (!update) return null;
|
||||
return {
|
||||
version: update.version,
|
||||
install: () => invoke("install_portable_update", { expectedVersion: update.version }),
|
||||
close: async () => {},
|
||||
};
|
||||
}
|
||||
|
||||
const update: Update | null = await check();
|
||||
if (!update) return null;
|
||||
return {
|
||||
version: update.version,
|
||||
install: async () => {
|
||||
await update.downloadAndInstall();
|
||||
await relaunch();
|
||||
},
|
||||
close: () => update.close(),
|
||||
};
|
||||
}
|
||||
|
||||
export async function installUpdate(update: Update): Promise<void> {
|
||||
await update.downloadAndInstall();
|
||||
await relaunch();
|
||||
export async function installUpdate(update: AppUpdate): Promise<void> {
|
||||
await update.install();
|
||||
}
|
||||
|
||||
@@ -1,9 +1,9 @@
|
||||
import { useSyncExternalStore } from "react";
|
||||
import type { Update } from "@tauri-apps/plugin-updater";
|
||||
import {
|
||||
checkForUpdate,
|
||||
hasNativeAppLifecycle,
|
||||
installUpdate,
|
||||
type AppUpdate,
|
||||
} from "../native/appLifecycle";
|
||||
|
||||
export type UpdateSnapshot = {
|
||||
@@ -17,7 +17,7 @@ let snapshot: UpdateSnapshot = {
|
||||
checking: false,
|
||||
installing: false,
|
||||
};
|
||||
let availableUpdate: Update | null = null;
|
||||
let availableUpdate: AppUpdate | null = null;
|
||||
let pendingCheck: Promise<string | null> | null = null;
|
||||
const listeners = new Set<() => void>();
|
||||
|
||||
@@ -26,7 +26,7 @@ function update(patch: Partial<UpdateSnapshot>) {
|
||||
listeners.forEach((listener) => listener());
|
||||
}
|
||||
|
||||
async function replaceAvailableUpdate(next: Update | null) {
|
||||
async function replaceAvailableUpdate(next: AppUpdate | null) {
|
||||
const previous = availableUpdate;
|
||||
availableUpdate = next;
|
||||
update({ availableVersion: next?.version ?? null });
|
||||
|
||||
@@ -1,9 +1,4 @@
|
||||
以下是重构后完整目标版本
|
||||
实现时,先创建所有目录和文件固化,每个文件头部都写好注释再实现
|
||||
旧服务已被备份为server_backup,/Users/leokun/Documents/cursor-byok/server 目录已创建
|
||||
行数均为目标估算,使用 `≈` 标记;不包含测试、生成代码和空行。
|
||||
实现时可做略微调整,测试要求相对于目标文件旁边的独立文件,禁止码内测试
|
||||
本文档目录 /Users/leokun/Documents/cursor-byok/cursor.md
|
||||
|
||||
## 完整目录
|
||||
|
||||
```text
|
||||
@@ -656,54 +651,7 @@ store ─X→ cursor
|
||||
model ─X→ cursor
|
||||
```
|
||||
|
||||
## 当前代码迁移
|
||||
|
||||
```text
|
||||
当前 目标
|
||||
|
||||
cursor/bidi_append.rs → api/cursor/bidi.rs
|
||||
cursor/run_sse.rs → api/cursor/run_sse.rs
|
||||
cursor/handlers.rs → api/cursor/handlers.rs
|
||||
cursor/proxy.rs → api/cursor/proxy.rs
|
||||
|
||||
cursor/sessions.rs → cursor/transport/registry.rs
|
||||
+ cursor/transport/handle.rs
|
||||
+ cursor/transport/output.rs
|
||||
|
||||
cursor/inbox.rs → cursor/transport/inbox.rs
|
||||
|
||||
cursor/actor.rs → cursor/transport/
|
||||
+ cursor/conversation/runtime.rs
|
||||
+ cursor/conversation/delivery.rs
|
||||
|
||||
cursor/session.rs → cursor/conversation/runtime.rs
|
||||
+ cursor/conversation/output.rs
|
||||
+ cursor/checkpoint/
|
||||
+ cursor/tools/
|
||||
|
||||
cursor/request/prepare.rs → cursor/compile/run.rs
|
||||
cursor/request/context.rs → cursor/compile/context.rs
|
||||
cursor/request/background.rs → cursor/compile/insert_messages.rs
|
||||
cursor/request/runtime.rs → cursor/compile/break_messages.rs
|
||||
cursor/request/images.rs → cursor/compile/images.rs
|
||||
cursor/request/model.rs → cursor/compile/model.rs
|
||||
|
||||
cursor/interaction/mod.rs → cursor/protocol/events.rs
|
||||
cursor/interaction/query.rs → cursor/tools/codec/query.rs
|
||||
cursor/interaction/render.rs → cursor/tools/codec/render.rs
|
||||
|
||||
cursor/projection/decode.rs → cursor/checkpoint/messages/decode.rs
|
||||
cursor/projection/encode.rs → cursor/checkpoint/messages/encode.rs
|
||||
cursor/projection/tests.rs → cursor/checkpoint/messages/tests.rs
|
||||
|
||||
cursor/presentation.rs → cursor/checkpoint/steps.rs
|
||||
|
||||
run/runtime.rs RunRegistry → cursor/conversation/registry.rs
|
||||
run/runtime.rs RunActor → run/engine.rs + run/handle.rs
|
||||
run/port.rs → run/command.rs + run/event.rs + run/port.rs
|
||||
|
||||
store/revisions.rs → store/checkpoints.rs
|
||||
```
|
||||
|
||||
|
||||
## 最终核心
|
||||
@@ -719,4 +667,4 @@ Bidi
|
||||
→ Checkpoint
|
||||
→ Transport
|
||||
→ RunSSE
|
||||
```
|
||||
```
|
||||
|
||||
@@ -0,0 +1,4 @@
|
||||
-- Custom provider-group display name shared by models with the same upstream host.
|
||||
-- NULL means no custom name; the UI falls back to the base_url hostname and the
|
||||
-- Cursor model picker badge falls back to the model type label.
|
||||
ALTER TABLE model_configs ADD COLUMN group_name TEXT;
|
||||
@@ -0,0 +1 @@
|
||||
ALTER TABLE tool_round_calls ADD COLUMN argument_error TEXT;
|
||||
@@ -8,6 +8,7 @@ import type { LlmRequest, ModelEvent } from "cursor-byok:provider";
|
||||
import type { ResourceSnapshot } from "cursor-byok:resource";
|
||||
import { codexDeviceOAuth } from "./oauth.ts";
|
||||
import { parseOfficialModels } from "./models.ts";
|
||||
import { buildResponsesBody } from "cursor-byok:protocol/openai-responses";
|
||||
import { codexProvider, isQuotaError } from "./provider.ts";
|
||||
import {
|
||||
accountIdentity,
|
||||
@@ -164,7 +165,10 @@ Deno.test("official model discovery excludes hidden models and puts the default
|
||||
display_name: "GPT First",
|
||||
supported_in_api: true,
|
||||
visibility: "list",
|
||||
supported_reasoning_efforts: ["low", "medium"],
|
||||
supported_reasoning_levels: [
|
||||
{ effort: "low", description: "Fast responses" },
|
||||
{ effort: "medium", description: "Balanced" },
|
||||
],
|
||||
},
|
||||
{ slug: "gpt-second", supported_in_api: true, visibility: "list" },
|
||||
{ slug: "gpt-hidden", supported_in_api: true, visibility: "hidden" },
|
||||
@@ -172,7 +176,7 @@ Deno.test("official model discovery excludes hidden models and puts the default
|
||||
],
|
||||
});
|
||||
assertEquals(models.map((model) => model.id), ["gpt-second", "gpt-first"]);
|
||||
assertEquals(models[1].capabilities, { thinking: true, images: true });
|
||||
assertEquals(models[1].capabilities, { images: true });
|
||||
assertEquals(models[1].privateData, { reasoningEfforts: ["low", "medium"] });
|
||||
});
|
||||
|
||||
@@ -302,6 +306,43 @@ Deno.test("invoke streams normalized events from the Codex Responses API", async
|
||||
]);
|
||||
});
|
||||
|
||||
Deno.test("reasoning replay projects response items to valid input items", () => {
|
||||
const replayRequest = request();
|
||||
replayRequest.messages = [{
|
||||
role: "assistant",
|
||||
text: "",
|
||||
thinking: "",
|
||||
replayState: {
|
||||
providerKind: "openai_responses",
|
||||
value: {
|
||||
items: [{
|
||||
type: "reasoning",
|
||||
id: "item-1",
|
||||
status: "completed",
|
||||
summary: [{ type: "summary_text", text: "why" }],
|
||||
content: [],
|
||||
encrypted_content: "opaque",
|
||||
output_only: true,
|
||||
}],
|
||||
},
|
||||
},
|
||||
toolCalls: [],
|
||||
}];
|
||||
|
||||
const body = buildResponsesBody({
|
||||
url: "https://example.com/responses",
|
||||
model: "gpt-test",
|
||||
request: replayRequest,
|
||||
});
|
||||
assertEquals(body.input, [{
|
||||
type: "reasoning",
|
||||
id: "item-1",
|
||||
summary: [{ type: "summary_text", text: "why" }],
|
||||
content: [],
|
||||
encrypted_content: "opaque",
|
||||
}]);
|
||||
});
|
||||
|
||||
Deno.test("invoke streams incremental tool calls and replays reasoning items", async () => {
|
||||
const token = jwt({ "https://api.openai.com/auth": { chatgpt_account_id: "acct-1" } });
|
||||
const draft = await credentialDraft({
|
||||
|
||||
@@ -24,7 +24,11 @@ function positiveInteger(value: unknown): number | null {
|
||||
}
|
||||
|
||||
function parseReasoningEfforts(model: Record<string, unknown>): string[] {
|
||||
const source = model.supported_reasoning_efforts ??
|
||||
const source = model.supported_reasoning_levels ??
|
||||
model.supportedReasoningLevels ??
|
||||
model.reasoning_levels ??
|
||||
model.reasoningLevels ??
|
||||
model.supported_reasoning_efforts ??
|
||||
model.supportedReasoningEfforts ??
|
||||
model.reasoning_efforts ??
|
||||
model.reasoningEfforts;
|
||||
@@ -61,10 +65,6 @@ export function parseOfficialModels(body: unknown): ModelDefinition[] {
|
||||
seen.add(id);
|
||||
const efforts = parseReasoningEfforts(model);
|
||||
const description = text(model.description);
|
||||
const contextWindowTokens = positiveInteger(
|
||||
model.context_window_tokens ?? model.contextWindowTokens ?? model.context_window ??
|
||||
model.contextWindow,
|
||||
);
|
||||
const maxOutputTokens = positiveInteger(
|
||||
model.max_output_tokens ?? model.maxOutputTokens ?? model.max_completion_tokens ??
|
||||
model.maxCompletionTokens,
|
||||
@@ -74,9 +74,8 @@ export function parseOfficialModels(body: unknown): ModelDefinition[] {
|
||||
displayName: text(model.display_name ?? model.displayName ?? model.title ?? model.name) ??
|
||||
id,
|
||||
...(description ? { description } : {}),
|
||||
...(contextWindowTokens !== null ? { contextWindowTokens } : {}),
|
||||
...(maxOutputTokens !== null ? { maxOutputTokens } : {}),
|
||||
capabilities: { thinking: efforts.length > 0, images: true },
|
||||
capabilities: { images: true },
|
||||
privateData: { reasoningEfforts: efforts },
|
||||
});
|
||||
}
|
||||
|
||||
@@ -160,9 +160,8 @@ Deno.test("model discovery parses both language-models and standard list shapes"
|
||||
});
|
||||
assertEquals(richModels.map((model) => model.id), ["grok-4", "grok-3-mini"]);
|
||||
assertEquals(richModels[0].displayName, "Grok 4");
|
||||
assertEquals(richModels[0].contextWindowTokens, 256_000);
|
||||
assertEquals(richModels[0].capabilities, { thinking: false, images: true });
|
||||
assertEquals(richModels[1].capabilities, { thinking: false, images: false });
|
||||
assertEquals(richModels[0].capabilities, { images: true });
|
||||
assertEquals(richModels[1].capabilities, { images: false });
|
||||
|
||||
const plainModels = parseGrokModels({ data: [{ id: "grok-4-fast" }] });
|
||||
assertEquals(plainModels.map((model) => model.id), ["grok-4-fast"]);
|
||||
|
||||
@@ -9,12 +9,12 @@ export const FALLBACK_MODELS: ModelDefinition[] = [
|
||||
{
|
||||
id: "grok-4.6",
|
||||
displayName: "Grok 4.6",
|
||||
capabilities: { thinking: false, images: true },
|
||||
capabilities: { images: true },
|
||||
},
|
||||
{
|
||||
id: "grok-4.5",
|
||||
displayName: "Grok 4.5",
|
||||
capabilities: { thinking: false, images: true },
|
||||
capabilities: { images: true },
|
||||
},
|
||||
];
|
||||
|
||||
@@ -28,15 +28,6 @@ function text(value: unknown): string | null {
|
||||
return typeof value === "string" && value.trim() ? value.trim() : null;
|
||||
}
|
||||
|
||||
function positiveInteger(value: unknown): number | null {
|
||||
const parsed = typeof value === "number"
|
||||
? value
|
||||
: typeof value === "string"
|
||||
? Number(value)
|
||||
: NaN;
|
||||
return Number.isFinite(parsed) && parsed > 0 ? Math.floor(parsed) : null;
|
||||
}
|
||||
|
||||
function modalities(value: unknown): string[] {
|
||||
return Array.isArray(value)
|
||||
? value.flatMap((item) => (typeof item === "string" ? [item.toLowerCase()] : []))
|
||||
@@ -66,15 +57,10 @@ export function parseGrokModels(body: unknown): ModelDefinition[] {
|
||||
if (!id || seen.has(id)) continue;
|
||||
seen.add(id);
|
||||
const inputs = modalities(model?.input_modalities ?? model?.inputModalities);
|
||||
const contextWindowTokens = positiveInteger(
|
||||
model?.context_window ?? model?.contextWindow ?? model?.max_prompt_length,
|
||||
);
|
||||
models.push({
|
||||
id,
|
||||
displayName: displayName(id),
|
||||
...(contextWindowTokens !== null ? { contextWindowTokens } : {}),
|
||||
capabilities: {
|
||||
thinking: false,
|
||||
images: inputs.length === 0 || inputs.includes("image"),
|
||||
},
|
||||
});
|
||||
|
||||
@@ -184,7 +184,11 @@ pub async fn append(
|
||||
request: DecodedAppend,
|
||||
parent: Option<TransportParent>,
|
||||
) -> Result<ai::BidiAppendResponse> {
|
||||
let handle = registry.get_or_create(&request.request_id).await?;
|
||||
let replace_closing = request.model_id().is_some();
|
||||
let handle = registry
|
||||
.get_or_create_for_append(&request.request_id, replace_closing)
|
||||
.await?;
|
||||
let _admission = handle.admit()?;
|
||||
if let Some(conversation_id) = request.conversation_id() {
|
||||
handle.set_conversation_id(conversation_id)?;
|
||||
}
|
||||
|
||||
@@ -19,18 +19,26 @@ use crate::{
|
||||
connect,
|
||||
proto::{agent::v1 as agent, aiserver::v1 as ai},
|
||||
},
|
||||
services::{account, analytics, model_catalog, observability::CursorTraceRecorder, tab},
|
||||
services::{account, analytics, knowledge, model_catalog, tab},
|
||||
transport::{TransportParent, TransportRegistry},
|
||||
},
|
||||
Result,
|
||||
};
|
||||
|
||||
pub fn router(registry: TransportRegistry) -> Result<Router> {
|
||||
let proxy = CursorProxy::cursor(registry.store().clone())?;
|
||||
Ok(router_with_proxy(registry, proxy))
|
||||
pub fn router(
|
||||
registry: TransportRegistry,
|
||||
clients: crate::network::NetworkClients,
|
||||
) -> Result<Router> {
|
||||
let proxy = CursorProxy::cursor(clients);
|
||||
let knowledge = knowledge::KnowledgeService::managed()?;
|
||||
Ok(router_with_proxy(registry, proxy, knowledge))
|
||||
}
|
||||
|
||||
fn router_with_proxy(registry: TransportRegistry, proxy: CursorProxy) -> Router {
|
||||
fn router_with_proxy(
|
||||
registry: TransportRegistry,
|
||||
proxy: CursorProxy,
|
||||
knowledge_service: knowledge::KnowledgeService,
|
||||
) -> Router {
|
||||
let web_cache = registry.web_cache().router();
|
||||
Router::new()
|
||||
.route("/__byok-api__/healthz", get(health))
|
||||
@@ -69,6 +77,22 @@ fn router_with_proxy(registry: TransportRegistry, proxy: CursorProxy) -> Router
|
||||
"/aiserver.v1.DashboardService/GetUsageLimitStatusAndActiveGrants",
|
||||
post(account::usage_limit_status),
|
||||
)
|
||||
.route(
|
||||
"/aiserver.v1.AiService/KnowledgeBaseAdd",
|
||||
post(knowledge::add),
|
||||
)
|
||||
.route(
|
||||
"/aiserver.v1.AiService/KnowledgeBaseList",
|
||||
post(knowledge::list),
|
||||
)
|
||||
.route(
|
||||
"/aiserver.v1.AiService/KnowledgeBaseUpdate",
|
||||
post(knowledge::update),
|
||||
)
|
||||
.route(
|
||||
"/aiserver.v1.AiService/KnowledgeBaseRemove",
|
||||
post(knowledge::remove),
|
||||
)
|
||||
.route(
|
||||
analytics::BOOTSTRAP_STATSIG_PATH,
|
||||
post(analytics::bootstrap_statsig),
|
||||
@@ -80,6 +104,7 @@ fn router_with_proxy(registry: TransportRegistry, proxy: CursorProxy) -> Router
|
||||
.fallback(proxy::forward)
|
||||
.method_not_allowed_fallback(proxy::forward)
|
||||
.layer(Extension(proxy))
|
||||
.layer(Extension(knowledge_service))
|
||||
.with_state(registry)
|
||||
.merge(web_cache)
|
||||
}
|
||||
@@ -96,16 +121,13 @@ async fn run_sse_handler(
|
||||
let (parts, body) = buffered(request).await?;
|
||||
let request: agent::BidiRequestId = connect::decode_unary(&body)?;
|
||||
let route = registry.wait_route(&request.request_id).await;
|
||||
let trace = CursorTraceRecorder::resume(registry.store().clone(), &request.request_id).await;
|
||||
if let Some(trace) = &trace {
|
||||
trace
|
||||
.request(
|
||||
"run_sse_request",
|
||||
&body,
|
||||
serde_json::json!({"request_id": request.request_id}),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
let trace = registry.trace(&request.request_id);
|
||||
trace.resume();
|
||||
trace.request(
|
||||
"run_sse_request",
|
||||
body.clone(),
|
||||
serde_json::json!({"request_id": request.request_id}),
|
||||
);
|
||||
match route {
|
||||
crate::cursor::transport::TransportRoute::Local => {
|
||||
run_sse::stream(®istry, &request.request_id).await
|
||||
@@ -116,7 +138,14 @@ async fn run_sse_handler(
|
||||
Request::from_parts(parts, Body::from(body)),
|
||||
)
|
||||
.await?;
|
||||
Ok(run_sse::upstream(registry, request.request_id, generation, response, trace).await)
|
||||
Ok(run_sse::upstream(
|
||||
registry,
|
||||
request.request_id,
|
||||
generation,
|
||||
response,
|
||||
Some(trace),
|
||||
)
|
||||
.await)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -132,6 +161,7 @@ async fn bidi_handler(
|
||||
let first_model = decoded.model_id().map(str::to_owned);
|
||||
let conversation_id = decoded.conversation_id().map(str::to_owned);
|
||||
let trace_metadata = decoded.trace_metadata();
|
||||
let trace = registry.trace(&decoded.request_id);
|
||||
let local = if let Some(model_id) = decoded.model_id() {
|
||||
// 插件模型 ID 只在本地有意义,永远不转发到 Cursor 官方上游。
|
||||
if model_id.starts_with(crate::plugin::ADAPTER_ID_PREFIX)
|
||||
@@ -156,14 +186,18 @@ async fn bidi_handler(
|
||||
} else if registry.upstream(&decoded.request_id).await {
|
||||
false
|
||||
} else {
|
||||
trace.resume();
|
||||
trace.request(
|
||||
"bidi_request",
|
||||
body.clone(),
|
||||
trace_outcome(trace_metadata, false, "missing_transport", None),
|
||||
);
|
||||
return Err(crate::Error::Protocol(
|
||||
"first BidiAppend message must select a model".into(),
|
||||
));
|
||||
};
|
||||
let trace = if first_model.is_some() {
|
||||
CursorTraceRecorder::begin(
|
||||
registry.store().clone(),
|
||||
&decoded.request_id,
|
||||
if first_model.is_some() {
|
||||
trace.begin(
|
||||
conversation_id.as_deref(),
|
||||
if local {
|
||||
"local_byok"
|
||||
@@ -171,26 +205,61 @@ async fn bidi_handler(
|
||||
"cursor_official"
|
||||
},
|
||||
first_model.as_deref(),
|
||||
)
|
||||
.await
|
||||
);
|
||||
} else {
|
||||
CursorTraceRecorder::resume(registry.store().clone(), &decoded.request_id).await
|
||||
};
|
||||
if let Some(trace) = &trace {
|
||||
trace.request("bidi_request", &body, trace_metadata).await;
|
||||
trace.resume();
|
||||
}
|
||||
if !local {
|
||||
if first_model.is_some() {
|
||||
registry.mark_upstream(&decoded.request_id).await;
|
||||
}
|
||||
trace.request(
|
||||
"bidi_request",
|
||||
body.clone(),
|
||||
trace_outcome(trace_metadata, true, "upstream", None),
|
||||
);
|
||||
return proxy::forward(
|
||||
Extension(proxy),
|
||||
Request::from_parts(parts, Body::from(body)),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
let parent = parent_headers(&parts.headers)?;
|
||||
bidi::append(®istry, decoded, parent).await?;
|
||||
let parent = match parent_headers(&parts.headers) {
|
||||
Ok(parent) => parent,
|
||||
Err(error) => {
|
||||
trace.request(
|
||||
"bidi_request",
|
||||
body,
|
||||
trace_outcome(
|
||||
trace_metadata,
|
||||
false,
|
||||
"invalid_parent",
|
||||
Some(error.to_string()),
|
||||
),
|
||||
);
|
||||
return Err(error);
|
||||
}
|
||||
};
|
||||
match bidi::append(®istry, decoded, parent).await {
|
||||
Ok(_) => trace.request(
|
||||
"bidi_request",
|
||||
body,
|
||||
trace_outcome(trace_metadata, true, "local", None),
|
||||
),
|
||||
Err(error) => {
|
||||
trace.request(
|
||||
"bidi_request",
|
||||
body,
|
||||
trace_outcome(
|
||||
trace_metadata,
|
||||
false,
|
||||
"command_rejected",
|
||||
Some(error.to_string()),
|
||||
),
|
||||
);
|
||||
return Err(error);
|
||||
}
|
||||
}
|
||||
let mut response = Response::new(axum::body::Body::empty());
|
||||
*response.status_mut() = StatusCode::OK;
|
||||
response.headers_mut().insert(
|
||||
@@ -200,6 +269,22 @@ async fn bidi_handler(
|
||||
Ok(response)
|
||||
}
|
||||
|
||||
fn trace_outcome(
|
||||
mut metadata: serde_json::Value,
|
||||
accepted: bool,
|
||||
route_outcome: &str,
|
||||
error: Option<String>,
|
||||
) -> serde_json::Value {
|
||||
if let Some(metadata) = metadata.as_object_mut() {
|
||||
metadata.insert("accepted".into(), accepted.into());
|
||||
metadata.insert("route_outcome".into(), route_outcome.into());
|
||||
if let Some(error) = error {
|
||||
metadata.insert("error".into(), error.into());
|
||||
}
|
||||
}
|
||||
metadata
|
||||
}
|
||||
|
||||
async fn buffered(request: Request<Body>) -> Result<(axum::http::request::Parts, Bytes)> {
|
||||
let (parts, body) = request.into_parts();
|
||||
let body = to_bytes(body, usize::MAX)
|
||||
|
||||
@@ -14,8 +14,7 @@ pub const UPSTREAM_URL_HEADER: &str = "x-server-upstream-url";
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct CursorProxy {
|
||||
client: Option<reqwest::Client>,
|
||||
store: Option<crate::store::Store>,
|
||||
clients: crate::network::NetworkClients,
|
||||
upstream: String,
|
||||
}
|
||||
|
||||
@@ -47,23 +46,15 @@ impl BufferedResponse {
|
||||
}
|
||||
|
||||
impl CursorProxy {
|
||||
pub fn cursor(store: crate::store::Store) -> Result<Self> {
|
||||
Ok(Self {
|
||||
client: None,
|
||||
store: Some(store),
|
||||
pub fn cursor(clients: crate::network::NetworkClients) -> Self {
|
||||
Self {
|
||||
clients,
|
||||
upstream: CURSOR_UPSTREAM.into(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
async fn client(&self) -> Result<reqwest::Client> {
|
||||
match (&self.client, &self.store) {
|
||||
(Some(client), _) => Ok(client.clone()),
|
||||
(_, Some(store)) => Ok(crate::network::client_builder(store)
|
||||
.await?
|
||||
.redirect(reqwest::redirect::Policy::none())
|
||||
.build()?),
|
||||
_ => unreachable!("Cursor proxy always has a client or store"),
|
||||
}
|
||||
self.clients.cursor_client().await
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -22,7 +22,7 @@ pub async fn stream(registry: &TransportRegistry, request_id: &str) -> Result<Re
|
||||
let receiver = handle.subscribe();
|
||||
let trace = handle.trace().cloned();
|
||||
if let Some(trace) = &trace {
|
||||
trace.response_started(StatusCode::OK.as_u16()).await;
|
||||
trace.response_started(StatusCode::OK.as_u16());
|
||||
}
|
||||
let body_stream = local_body_stream(receiver, handle, trace);
|
||||
let mut response = Response::new(Body::from_stream(body_stream));
|
||||
@@ -133,7 +133,7 @@ pub async fn upstream(
|
||||
) -> Response<Body> {
|
||||
let (parts, body) = response.into_parts();
|
||||
if let Some(trace) = &trace {
|
||||
trace.response_started(parts.status.as_u16()).await;
|
||||
trace.response_started(parts.status.as_u16());
|
||||
}
|
||||
let stream = async_stream::stream! {
|
||||
let _guard = UpstreamRunGuard {
|
||||
@@ -180,15 +180,15 @@ impl TraceStreamSink {
|
||||
while let Some(event) = receiver.recv().await {
|
||||
match event {
|
||||
TraceStreamEvent::Chunk(chunk) => {
|
||||
trace.response_chunk(source, &chunk).await;
|
||||
trace.response_chunk(source, chunk);
|
||||
}
|
||||
TraceStreamEvent::Finish(error) => {
|
||||
trace.finish(error.as_deref()).await;
|
||||
trace.finish(error.as_deref());
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
trace.finish(None).await;
|
||||
trace.finish(None);
|
||||
});
|
||||
Self {
|
||||
sender: Some(sender),
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
//! Builds the top-level server router.
|
||||
|
||||
use crate::{cursor::transport::TransportRegistry, Result};
|
||||
use crate::{cursor::transport::TransportRegistry, network::NetworkClients, Result};
|
||||
|
||||
pub fn router(registry: TransportRegistry) -> Result<axum::Router> {
|
||||
super::cursor::router(registry)
|
||||
pub fn router(registry: TransportRegistry, clients: NetworkClients) -> Result<axum::Router> {
|
||||
super::cursor::router(registry, clients)
|
||||
}
|
||||
|
||||
+12
-3
@@ -44,9 +44,11 @@ impl App {
|
||||
plugin_runtime.clone(),
|
||||
config.app_version.clone(),
|
||||
)?;
|
||||
let clients = crate::network::NetworkClients::new(store.clone());
|
||||
let provider = std::sync::Arc::new(ProviderRouter::new(
|
||||
store.clone(),
|
||||
plugins.clone(),
|
||||
clients.clone(),
|
||||
config.provider_request_timeout,
|
||||
config.provider_stream_idle_timeout,
|
||||
));
|
||||
@@ -56,11 +58,18 @@ impl App {
|
||||
compiler,
|
||||
WebCache::managed()?,
|
||||
plugins.clone(),
|
||||
crate::config::managed_data_dir()?.join("rules"),
|
||||
);
|
||||
let control =
|
||||
control::ControlService::new(store.clone(), provider, plugin_runtime, plugins)?;
|
||||
let control = control::ControlService::new(
|
||||
store.clone(),
|
||||
provider,
|
||||
plugin_runtime,
|
||||
plugins,
|
||||
clients.clone(),
|
||||
config.app_version.clone(),
|
||||
)?;
|
||||
let harness = control.cursor_harness().clone();
|
||||
let mut router = api::router(registry.clone())?;
|
||||
let mut router = api::router(registry.clone(), clients)?;
|
||||
router = match &config.console {
|
||||
Some(ConsoleSource::Directory(directory)) => {
|
||||
router.merge(control::web_router(control.clone(), directory))
|
||||
|
||||
@@ -46,7 +46,6 @@ pub struct ProviderConfig {
|
||||
pub custom_headers: reqwest::header::HeaderMap,
|
||||
pub max_output_tokens: Option<u64>,
|
||||
pub request_timeout: Duration,
|
||||
pub retry_count: u32,
|
||||
pub allowed_body_fields: Option<std::collections::HashSet<String>>,
|
||||
}
|
||||
|
||||
|
||||
@@ -17,6 +17,7 @@ use super::ControlService;
|
||||
// 开源项目广告为作者唯一收入来源,PR时请勿去除此广告
|
||||
|
||||
pub(super) const ADS_ENDPOINT: &str = "https://tab.leokun.cn/api/v1/ads?placement=menu";
|
||||
// pub(super) const ADS_ENDPOINT: &str = "http://localhost:8080/api/v1/ads?placement=menu";
|
||||
pub(super) const DEVICE_ID_HEADER: &str = "X-Cursor-Assistant-Device-ID";
|
||||
pub(super) const OS_HEADER: &str = "X-Cursor-Assistant-OS";
|
||||
pub(super) const APP_VERSION_HEADER: &str = "X-Cursor-Assistant-Version";
|
||||
|
||||
@@ -41,6 +41,8 @@ pub struct ControlService {
|
||||
provider: Arc<dyn Provider>,
|
||||
plugin_runtime: PluginRuntime,
|
||||
plugins: PluginRegistry,
|
||||
clients: crate::network::NetworkClients,
|
||||
app_version: String,
|
||||
model_tests: Arc<Mutex<BTreeMap<String, CancellationToken>>>,
|
||||
}
|
||||
|
||||
@@ -151,6 +153,8 @@ impl ControlService {
|
||||
provider: Arc<dyn Provider>,
|
||||
plugin_runtime: PluginRuntime,
|
||||
plugins: PluginRegistry,
|
||||
clients: crate::network::NetworkClients,
|
||||
app_version: String,
|
||||
) -> Result<Self> {
|
||||
Ok(Self {
|
||||
cursor_harness: CursorHarness::new(store.clone())?,
|
||||
@@ -158,6 +162,8 @@ impl ControlService {
|
||||
provider,
|
||||
plugin_runtime,
|
||||
plugins,
|
||||
clients,
|
||||
app_version,
|
||||
model_tests: Arc::new(Mutex::new(BTreeMap::new())),
|
||||
})
|
||||
}
|
||||
@@ -256,13 +262,13 @@ impl ControlService {
|
||||
disabled_ad_ids: Option<&str>,
|
||||
language: &str,
|
||||
) -> Result<AdRuntime> {
|
||||
let client = crate::network::client(&self.store).await?;
|
||||
let client = self.clients.default_client().await?;
|
||||
let installation_id = self.store.installation_id().await?;
|
||||
let mut request = client
|
||||
.get(ADS_ENDPOINT)
|
||||
.header(DEVICE_ID_HEADER, installation_id)
|
||||
.header(OS_HEADER, std::env::consts::OS)
|
||||
.header(APP_VERSION_HEADER, env!("CARGO_PKG_VERSION"))
|
||||
.header(APP_VERSION_HEADER, &self.app_version)
|
||||
.header(LANGUAGE_HEADER, language)
|
||||
.timeout(std::time::Duration::from_secs(60));
|
||||
if let Some(disabled_ad_ids) = disabled_ad_ids.filter(|value| !value.is_empty()) {
|
||||
@@ -281,7 +287,7 @@ impl ControlService {
|
||||
}
|
||||
|
||||
pub(super) async fn dismiss_ad(&self, ad_id: &str, input: &AdDismissalInput) -> Result<()> {
|
||||
let client = crate::network::client(&self.store).await?;
|
||||
let client = self.clients.default_client().await?;
|
||||
let installation_id = self.store.installation_id().await?;
|
||||
let mut endpoint = Url::parse(ADS_ENDPOINT).map_err(|error| {
|
||||
Error::Config(format!("advertisement endpoint is invalid: {error}"))
|
||||
@@ -296,7 +302,7 @@ impl ControlService {
|
||||
.post(endpoint)
|
||||
.header(DEVICE_ID_HEADER, installation_id)
|
||||
.header(OS_HEADER, std::env::consts::OS)
|
||||
.header(APP_VERSION_HEADER, env!("CARGO_PKG_VERSION"))
|
||||
.header(APP_VERSION_HEADER, &self.app_version)
|
||||
.json(input)
|
||||
.timeout(std::time::Duration::from_secs(5))
|
||||
.send()
|
||||
@@ -392,7 +398,6 @@ impl ControlService {
|
||||
if model_hash.starts_with(crate::plugin::ADAPTER_ID_PREFIX) {
|
||||
let descriptor = self.plugins.model_descriptor(model_hash).await?;
|
||||
model.display_name = Some(descriptor.display_name);
|
||||
model.context_window_tokens = descriptor.context_window_tokens;
|
||||
model.max_output_tokens = Some(descriptor.max_output_tokens.unwrap_or(65_536));
|
||||
} else {
|
||||
let configured = self
|
||||
@@ -511,7 +516,7 @@ impl ControlService {
|
||||
}
|
||||
|
||||
pub async fn discover_models(&self, input: &ModelDiscoveryInput) -> Result<DiscoveredModels> {
|
||||
let client = crate::network::client(&self.store).await?;
|
||||
let client = self.clients.default_client().await?;
|
||||
let base_url = crate::model::normalize_request_url(&input.base_url)?;
|
||||
discover_models_from_endpoint(
|
||||
&client,
|
||||
@@ -678,7 +683,9 @@ impl ControlService {
|
||||
}
|
||||
|
||||
pub async fn set_proxy_settings(&self, settings: ProxySettingsInput) -> Result<ProxySettings> {
|
||||
self.store.set_proxy_settings(settings).await
|
||||
let settings = self.store.set_proxy_settings(settings).await?;
|
||||
self.clients.invalidate().await;
|
||||
Ok(settings)
|
||||
}
|
||||
|
||||
pub async fn tab_settings(&self) -> Result<TabSettings> {
|
||||
|
||||
@@ -278,19 +278,17 @@ impl CheckpointBuilder {
|
||||
),
|
||||
});
|
||||
if let Some(trace) = handle.trace() {
|
||||
trace
|
||||
.artifact(
|
||||
"checkpoint",
|
||||
"byok_server",
|
||||
&checkpoint.encode_to_vec(),
|
||||
serde_json::json!({
|
||||
"root_message_count": checkpoint.root_prompt_messages_json.len(),
|
||||
"turn_count": checkpoint.turns.len(),
|
||||
"pending_tool_call_count": checkpoint.pending_tool_calls.len(),
|
||||
"emit_status": if result.is_ok() { "sent" } else { "error" },
|
||||
}),
|
||||
)
|
||||
.await;
|
||||
trace.artifact(
|
||||
"checkpoint",
|
||||
"byok_server",
|
||||
&checkpoint.encode_to_vec(),
|
||||
serde_json::json!({
|
||||
"root_message_count": checkpoint.root_prompt_messages_json.len(),
|
||||
"turn_count": checkpoint.turns.len(),
|
||||
"pending_tool_call_count": checkpoint.pending_tool_calls.len(),
|
||||
"emit_status": if result.is_ok() { "sent" } else { "error" },
|
||||
}),
|
||||
);
|
||||
}
|
||||
result
|
||||
}
|
||||
|
||||
@@ -102,6 +102,12 @@ pub fn decode_pending(value: &str) -> Result<RecoveredToolRound> {
|
||||
.into_iter()
|
||||
.enumerate()
|
||||
.map(|(index, call)| {
|
||||
let argument_error = wire
|
||||
.pointer("/providerOptions/cursor/pendingToolExecutionContracts")
|
||||
.and_then(|contracts| contracts.get(&call.call_id))
|
||||
.and_then(|contract| contract.get("argumentError"))
|
||||
.and_then(Value::as_str)
|
||||
.map(str::to_string);
|
||||
Ok(ToolCall {
|
||||
index,
|
||||
call_id: call.call_id,
|
||||
@@ -109,6 +115,7 @@ pub fn decode_pending(value: &str) -> Result<RecoveredToolRound> {
|
||||
name: call.name,
|
||||
arguments_text: serde_json::to_string(&call.arguments)?,
|
||||
arguments: call.arguments,
|
||||
argument_error,
|
||||
})
|
||||
})
|
||||
.collect::<Result<Vec<_>>>()?;
|
||||
|
||||
@@ -71,6 +71,7 @@ pub fn staged_tool_round(
|
||||
allowed_tools,
|
||||
dynamic_tools,
|
||||
started_at_ms,
|
||||
tool_calls: Some(calls),
|
||||
}),
|
||||
)?)?)
|
||||
}
|
||||
@@ -96,6 +97,7 @@ pub fn staged_final(
|
||||
allowed_tools,
|
||||
dynamic_tools,
|
||||
started_at_ms,
|
||||
tool_calls: None,
|
||||
}),
|
||||
)?)?)
|
||||
}
|
||||
@@ -105,6 +107,7 @@ pub(super) struct PendingContext<'a> {
|
||||
allowed_tools: &'a [String],
|
||||
dynamic_tools: &'a HashSet<String>,
|
||||
started_at_ms: u64,
|
||||
tool_calls: Option<&'a [ToolCall]>,
|
||||
}
|
||||
|
||||
pub(super) fn wire_message(
|
||||
@@ -132,16 +135,23 @@ pub(super) fn wire_message(
|
||||
calls
|
||||
.iter()
|
||||
.map(|call| {
|
||||
(
|
||||
call.call_id.clone(),
|
||||
json!({
|
||||
"toolCallId": call.call_id,
|
||||
"outerToolName": call.name,
|
||||
"toolIdentifier": tool_identifier(&call.name, pending.dynamic_tools),
|
||||
"isDynamic": pending.dynamic_tools.contains(&call.name),
|
||||
"allowedToolNames": pending.allowed_tools,
|
||||
}),
|
||||
)
|
||||
let mut contract = json!({
|
||||
"toolCallId": call.call_id,
|
||||
"outerToolName": call.name,
|
||||
"toolIdentifier": tool_identifier(&call.name, pending.dynamic_tools),
|
||||
"isDynamic": pending.dynamic_tools.contains(&call.name),
|
||||
"allowedToolNames": pending.allowed_tools,
|
||||
});
|
||||
if let Some(error) = pending
|
||||
.tool_calls
|
||||
.and_then(|calls| {
|
||||
calls.iter().find(|candidate| candidate.call_id == call.call_id)
|
||||
})
|
||||
.and_then(|call| call.argument_error.as_deref())
|
||||
{
|
||||
contract["argumentError"] = Value::String(error.into());
|
||||
}
|
||||
(call.call_id.clone(), contract)
|
||||
})
|
||||
.collect(),
|
||||
),
|
||||
|
||||
@@ -30,6 +30,7 @@ fn pending_tool_round_is_one_complete_assistant_message_and_round_trips() {
|
||||
name: "Read".into(),
|
||||
arguments_text: r#"{"path":"/a"}"#.into(),
|
||||
arguments: json!({"path":"/a"}),
|
||||
argument_error: Some("Read arguments are not valid JSON".into()),
|
||||
},
|
||||
ToolCall {
|
||||
index: 1,
|
||||
@@ -38,6 +39,7 @@ fn pending_tool_round_is_one_complete_assistant_message_and_round_trips() {
|
||||
name: "Grep".into(),
|
||||
arguments_text: r#"{"pattern":"x"}"#.into(),
|
||||
arguments: json!({"pattern":"x"}),
|
||||
argument_error: None,
|
||||
},
|
||||
];
|
||||
let pending = staged_tool_round(
|
||||
@@ -55,6 +57,10 @@ fn pending_tool_round_is_one_complete_assistant_message_and_round_trips() {
|
||||
wire["providerOptions"]["cursor"]["pendingToolExecutionContracts"]["a"]["toolIdentifier"],
|
||||
"READ"
|
||||
);
|
||||
assert_eq!(
|
||||
wire["providerOptions"]["cursor"]["pendingToolExecutionContracts"]["a"]["argumentError"],
|
||||
"Read arguments are not valid JSON"
|
||||
);
|
||||
assert_eq!(wire["role"], "assistant");
|
||||
assert_eq!(
|
||||
wire["providerOptions"]["cursor"]["pendingToolExecutionContracts"]
|
||||
@@ -77,6 +83,10 @@ fn pending_tool_round_is_one_complete_assistant_message_and_round_trips() {
|
||||
assert_eq!(recovered.assistant.replay_state, Some(replay_state));
|
||||
assert_eq!(recovered.calls.len(), 2);
|
||||
assert_eq!(recovered.calls[0].call_id, "a");
|
||||
assert_eq!(
|
||||
recovered.calls[0].argument_error.as_deref(),
|
||||
Some("Read arguments are not valid JSON")
|
||||
);
|
||||
assert_eq!(recovered.calls[1].call_id, "b");
|
||||
}
|
||||
|
||||
|
||||
@@ -75,6 +75,11 @@ impl StepBuffer {
|
||||
});
|
||||
}
|
||||
|
||||
pub fn finish_model_attempt(&mut self) {
|
||||
self.finish_text();
|
||||
self.finish_thinking(Duration::ZERO);
|
||||
}
|
||||
|
||||
pub fn discard_model_output(&mut self) {
|
||||
self.text.clear();
|
||||
self.thinking.clear();
|
||||
@@ -101,6 +106,28 @@ impl StepBuffer {
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn failed_attempt_output_is_retained_for_the_next_checkpoint() {
|
||||
let mut buffer = StepBuffer::default();
|
||||
buffer.text_delta("partial answer");
|
||||
buffer.thinking_delta("partial reasoning");
|
||||
|
||||
buffer.finish_model_attempt();
|
||||
|
||||
let steps = buffer.take().steps;
|
||||
assert_eq!(steps.len(), 2);
|
||||
assert!(matches!(
|
||||
&steps[0].message,
|
||||
Some(pb::conversation_step::Message::AssistantMessage(message))
|
||||
if message.text == "partial answer"
|
||||
));
|
||||
assert!(matches!(
|
||||
&steps[1].message,
|
||||
Some(pb::conversation_step::Message::ThinkingMessage(message))
|
||||
if message.text == "partial reasoning"
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn interrupted_model_output_is_not_persisted_as_checkpoint_steps() {
|
||||
let mut buffer = StepBuffer::default();
|
||||
|
||||
@@ -3,7 +3,7 @@ use prost::Message;
|
||||
|
||||
use crate::{
|
||||
cursor::{checkpoint::PendingSteps, protocol::proto::agent::v1 as pb},
|
||||
model::CanonicalMessage,
|
||||
model::{estimate_context_tokens, project_messages, CanonicalMessage, PromptSpec},
|
||||
store::{BlobEdge, BlobId},
|
||||
Error, Result,
|
||||
};
|
||||
@@ -85,6 +85,12 @@ impl CheckpointBuilder {
|
||||
.push(archive_id.as_bytes().to_vec());
|
||||
}
|
||||
self.base.self_summary_count = self.base.self_summary_count.saturating_add(1);
|
||||
let projected = project_messages(messages)?;
|
||||
let prompt = PromptSpec {
|
||||
instructions: self.instructions.clone(),
|
||||
tools: self.tool_definitions.clone(),
|
||||
};
|
||||
self.record_context_tokens(Some(estimate_context_tokens(&prompt, &projected)));
|
||||
if let Some(details) = self.base.token_details.as_mut() {
|
||||
details.breakdown = Some(crate::cursor::services::usage::breakdown(
|
||||
details.used_tokens,
|
||||
|
||||
@@ -12,7 +12,7 @@ use crate::{
|
||||
protocol::proto::agent::v1 as pb, services::context_sync::RequestContextSynchronizer,
|
||||
tools::runtime::McpRoute,
|
||||
},
|
||||
model::ToolDefinition,
|
||||
model::{normalize_tool_name, ToolDefinition},
|
||||
store::BlobId,
|
||||
Error, Result,
|
||||
};
|
||||
@@ -146,6 +146,37 @@ async fn decode_part<T: Message + Default>(
|
||||
.map_err(|error| Error::Protocol(format!("invalid {name} context Blob: {error}")))
|
||||
}
|
||||
|
||||
/// 把本地 md 规则目录(rules 服务的存储)合并进请求上下文,
|
||||
/// 使 BYOK 运行在 IDE 未携带这些规则时也能消费它们。
|
||||
/// 与 IDE 已发规则按内容去重;读取失败只告警,不影响运行。
|
||||
pub fn merge_local_rules(context: &mut pb::RequestContext, rules_dir: &Path) {
|
||||
let records = match crate::cursor::services::knowledge::RuleStore::open(rules_dir.into())
|
||||
.and_then(|store| store.list())
|
||||
{
|
||||
Ok(records) => records,
|
||||
Err(error) => {
|
||||
tracing::warn!(%error, "cannot read local rules; continuing without them");
|
||||
return;
|
||||
}
|
||||
};
|
||||
let existing = context
|
||||
.rules
|
||||
.iter()
|
||||
.chain(context.non_file_rules.iter())
|
||||
.map(|rule| rule.content.trim().to_owned())
|
||||
.chain(context.cloud_rule.iter().map(|rule| rule.trim().to_owned()))
|
||||
.collect::<HashSet<_>>();
|
||||
for record in records {
|
||||
if record.knowledge.trim().is_empty() || existing.contains(record.knowledge.trim()) {
|
||||
continue;
|
||||
}
|
||||
context.non_file_rules.push(pb::CursorRule {
|
||||
content: record.knowledge,
|
||||
..Default::default()
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
pub fn request_context(request: &pb::AgentRunRequest) -> Option<&pb::RequestContext> {
|
||||
let action = request.action.as_ref()?;
|
||||
action
|
||||
@@ -462,7 +493,7 @@ pub fn dynamic_mcp(
|
||||
})?),
|
||||
};
|
||||
let parameters = normalize_mcp_parameters(&wire.name, parameters)?;
|
||||
let name = model_tool_name(&wire.name);
|
||||
let name = normalize_tool_name(&wire.name);
|
||||
let definition = ToolDefinition {
|
||||
name: name.clone(),
|
||||
description: wire.description.clone(),
|
||||
@@ -522,18 +553,6 @@ fn invalid_mcp_parameters(tool_name: &str) -> Error {
|
||||
))
|
||||
}
|
||||
|
||||
fn model_tool_name(name: &str) -> String {
|
||||
name.chars()
|
||||
.map(|character| {
|
||||
if character.is_ascii_alphanumeric() || matches!(character, '_' | '-') {
|
||||
character
|
||||
} else {
|
||||
'_'
|
||||
}
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn prost_value(value: &prost_types::Value) -> Value {
|
||||
use prost_types::value::Kind;
|
||||
match value.kind.as_ref() {
|
||||
@@ -563,3 +582,48 @@ fn xml(value: &str) -> String {
|
||||
.replace('<', "<")
|
||||
.replace('>', ">")
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn rule(content: &str) -> pb::CursorRule {
|
||||
pb::CursorRule {
|
||||
content: content.into(),
|
||||
..Default::default()
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn merge_local_rules_appends_and_dedupes_by_content() {
|
||||
let directory = tempfile::tempdir().unwrap();
|
||||
std::fs::write(directory.path().join("a.md"), "shared rule").unwrap();
|
||||
std::fs::write(directory.path().join("b.md"), "local only rule").unwrap();
|
||||
std::fs::write(directory.path().join("c.md"), " \n").unwrap();
|
||||
|
||||
let mut context = pb::RequestContext {
|
||||
non_file_rules: vec![rule(" shared rule ")],
|
||||
..Default::default()
|
||||
};
|
||||
merge_local_rules(&mut context, directory.path());
|
||||
|
||||
let contents = context
|
||||
.non_file_rules
|
||||
.iter()
|
||||
.map(|rule| rule.content.as_str())
|
||||
.collect::<Vec<_>>();
|
||||
assert_eq!(
|
||||
contents,
|
||||
[" shared rule ", "local only rule"],
|
||||
"IDE-sent duplicate is kept once and blank local rules are skipped"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn merge_local_rules_survives_a_missing_directory() {
|
||||
let directory = tempfile::tempdir().unwrap();
|
||||
let mut context = pb::RequestContext::default();
|
||||
merge_local_rules(&mut context, &directory.path().join("nested/rules"));
|
||||
assert!(context.non_file_rules.is_empty());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -51,6 +51,7 @@ pub(crate) struct PrepareDependencies<'a> {
|
||||
pub checkpoint: &'a CheckpointBuilder,
|
||||
pub blob_sync: &'a BlobSynchronizer,
|
||||
pub context_sync: &'a RequestContextSynchronizer,
|
||||
pub local_rules_dir: Option<&'a std::path::Path>,
|
||||
}
|
||||
|
||||
pub(crate) async fn prepare(
|
||||
@@ -64,6 +65,7 @@ pub(crate) async fn prepare(
|
||||
checkpoint,
|
||||
blob_sync,
|
||||
context_sync,
|
||||
local_rules_dir,
|
||||
} = dependencies;
|
||||
checkpoint
|
||||
.import_prefetched(&request.pre_fetched_blobs)
|
||||
@@ -115,11 +117,13 @@ pub(crate) async fn prepare(
|
||||
"selected_source": "root_prompt_messages_json",
|
||||
});
|
||||
let encoded = serde_json::to_vec(&summary)?;
|
||||
trace
|
||||
.artifact("history_projection", "byok_server", &encoded, summary)
|
||||
.await;
|
||||
trace.artifact("history_projection", "byok_server", &encoded, summary);
|
||||
}
|
||||
let request_context = context::hydrate(request, context_sync).await?;
|
||||
let mut request_context = context::hydrate(request, context_sync).await?;
|
||||
if let Some(rules_dir) = local_rules_dir {
|
||||
context::merge_local_rules(&mut request_context, rules_dir);
|
||||
}
|
||||
let request_context = request_context;
|
||||
let ActionProjection {
|
||||
mode: mode_number,
|
||||
mut turn_user,
|
||||
|
||||
@@ -1,6 +1,19 @@
|
||||
//! Defines commands accepted by a Conversation runtime.
|
||||
|
||||
use crate::cursor::protocol::proto::agent::v1 as pb;
|
||||
use crate::{cursor::protocol::proto::agent::v1 as pb, Error};
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum RunFinish {
|
||||
TurnCompleted,
|
||||
Transport(TransportFinish),
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum TransportFinish {
|
||||
Success,
|
||||
Failed(Error),
|
||||
Cancelled,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum TransportCommand {
|
||||
@@ -8,6 +21,9 @@ pub enum TransportCommand {
|
||||
seqno: i64,
|
||||
message: Box<pb::AgentClientMessage>,
|
||||
},
|
||||
RunFinished {
|
||||
generation: u64,
|
||||
finish: RunFinish,
|
||||
},
|
||||
Disconnect,
|
||||
Close,
|
||||
}
|
||||
|
||||
@@ -21,7 +21,7 @@ use crate::{
|
||||
protocol::proto::agent::v1 as pb,
|
||||
services::blob_sync::BlobSynchronizer,
|
||||
tools::{
|
||||
codec,
|
||||
codec, compat,
|
||||
runtime::CursorToolRuntime,
|
||||
stream::ToolCallStream,
|
||||
tool_call_result::{ToolCompletion, ToolResultReceiver},
|
||||
@@ -34,7 +34,7 @@ use crate::{
|
||||
Error, Result,
|
||||
};
|
||||
|
||||
use super::{CompiledMessages, ConversationRegistry, MessageDelivery};
|
||||
use super::{CompiledMessages, ConversationRegistry, MessageDelivery, RunFinish, TransportFinish};
|
||||
use crate::cursor::transport::TransportHandle;
|
||||
|
||||
pub struct ConversationOutput {
|
||||
@@ -110,7 +110,7 @@ impl ConversationOutput {
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn run(mut self) -> Result<()> {
|
||||
pub async fn run(mut self) -> Result<RunFinish> {
|
||||
let result = self.run_inner().await;
|
||||
if let Err(error) = &result {
|
||||
if !self.superseded.is_cancelled() {
|
||||
@@ -143,7 +143,7 @@ impl ConversationOutput {
|
||||
result
|
||||
}
|
||||
|
||||
async fn run_inner(&mut self) -> Result<()> {
|
||||
async fn run_inner(&mut self) -> Result<RunFinish> {
|
||||
if self.context.compacting {
|
||||
self.handle.emit(&events::summary_started())?;
|
||||
}
|
||||
@@ -175,7 +175,7 @@ impl ConversationOutput {
|
||||
if self.superseded.is_cancelled() {
|
||||
worker.abort();
|
||||
self.abort_execs().await;
|
||||
return Ok(());
|
||||
return Ok(RunFinish::Transport(TransportFinish::Cancelled));
|
||||
}
|
||||
let input = if let Ok(action) = self.runtime_actions.try_recv() {
|
||||
Input::RuntimeAction(Some(Box::new(action)))
|
||||
@@ -187,7 +187,7 @@ impl ConversationOutput {
|
||||
_ = self.superseded.cancelled() => {
|
||||
worker.abort();
|
||||
self.abort_execs().await;
|
||||
return Ok(());
|
||||
return Ok(RunFinish::Transport(TransportFinish::Cancelled));
|
||||
}
|
||||
action = self.runtime_actions.recv() => Input::RuntimeAction(action.map(Box::new)),
|
||||
event = self.core.events.recv() => Input::Event(event),
|
||||
@@ -264,6 +264,32 @@ impl ConversationOutput {
|
||||
streams.clear();
|
||||
presentation.discard_model_output();
|
||||
}
|
||||
RunEvent::ModelAttemptFailed { attempt, message } => {
|
||||
tracing::warn!(
|
||||
run_id = %self.run.run_id(),
|
||||
attempt,
|
||||
%message,
|
||||
"retrying model call from current checkpoint"
|
||||
);
|
||||
presentation.finish_model_attempt();
|
||||
for call in calls.values_mut() {
|
||||
if call.arguments.is_null() {
|
||||
call.arguments = serde_json::from_str(&call.arguments_text)
|
||||
.unwrap_or_else(|_| serde_json::json!({}));
|
||||
}
|
||||
let completion = compat::failure_with_message(
|
||||
call,
|
||||
format!("Model attempt failed before tool completion: {message}"),
|
||||
);
|
||||
self.handle
|
||||
.emit(&codec::tool_completed(call, &completion))?;
|
||||
presentation.tool_completed(&completion);
|
||||
}
|
||||
response_text.clear();
|
||||
response_thinking.clear();
|
||||
calls.clear();
|
||||
streams.clear();
|
||||
}
|
||||
RunEvent::TextStart => {}
|
||||
RunEvent::TextEnd => {
|
||||
if !self.context.compacting {
|
||||
@@ -312,6 +338,7 @@ impl ConversationOutput {
|
||||
name: name.clone(),
|
||||
arguments_text: String::new(),
|
||||
arguments: serde_json::Value::Null,
|
||||
argument_error: None,
|
||||
};
|
||||
self.emit_model_event(
|
||||
crate::provider::ModelEvent::ToolCallStart {
|
||||
@@ -335,15 +362,50 @@ impl ConversationOutput {
|
||||
let stream = streams.get_mut(&index).ok_or_else(|| {
|
||||
Error::Protocol(format!("missing Cursor tool stream: {index}"))
|
||||
})?;
|
||||
for message in stream.arguments_delta(call, &delta)? {
|
||||
self.handle.emit(&message)?;
|
||||
match stream.arguments_delta(call, &delta) {
|
||||
Ok(messages) => {
|
||||
for message in messages {
|
||||
self.handle.emit(&message)?;
|
||||
}
|
||||
}
|
||||
Err(Error::Protocol(message)) => {
|
||||
tracing::warn!(
|
||||
call_id = %call.call_id,
|
||||
%message,
|
||||
"ignoring invalid streaming tool arguments until completion"
|
||||
);
|
||||
}
|
||||
Err(Error::Json(error)) => {
|
||||
tracing::warn!(
|
||||
call_id = %call.call_id,
|
||||
%error,
|
||||
"ignoring invalid streaming tool arguments until completion"
|
||||
);
|
||||
}
|
||||
Err(error) => return Err(error),
|
||||
}
|
||||
}
|
||||
RunEvent::ToolCallEnd { index } => {
|
||||
let call = calls.get_mut(&index).ok_or_else(|| {
|
||||
Error::Protocol(format!("unknown completed tool index: {index}"))
|
||||
})?;
|
||||
call.arguments = serde_json::from_str(&call.arguments_text)?;
|
||||
// A tool call with no arguments streams no argument text.
|
||||
// Treat empty text as an empty object, matching the model
|
||||
// cycle, instead of failing the run on `from_str("")`.
|
||||
call.arguments = if call.arguments_text.trim().is_empty() {
|
||||
serde_json::json!({})
|
||||
} else {
|
||||
serde_json::from_str(&call.arguments_text)
|
||||
.unwrap_or_else(|_| serde_json::json!({}))
|
||||
};
|
||||
}
|
||||
RunEvent::UsageSnapshot(usage) => {
|
||||
if !self.context.compacting {
|
||||
if let Some(output_tokens) = usage.output_tokens {
|
||||
self.handle.emit(&events::token_delta(output_tokens))?;
|
||||
}
|
||||
context_tokens = usage.context_input_tokens;
|
||||
}
|
||||
}
|
||||
RunEvent::Usage(usage) => {
|
||||
if !self.context.compacting {
|
||||
@@ -352,10 +414,7 @@ impl ConversationOutput {
|
||||
}
|
||||
}
|
||||
if !self.context.compacting {
|
||||
context_tokens = usage
|
||||
.input_tokens
|
||||
.zip(usage.output_tokens)
|
||||
.and_then(|(input, output)| input.checked_add(output));
|
||||
context_tokens = usage.context_input_tokens;
|
||||
}
|
||||
match &mut turn_usage {
|
||||
Some(total) => *total += usage,
|
||||
@@ -424,21 +483,28 @@ impl ConversationOutput {
|
||||
streams.clear();
|
||||
}
|
||||
if let CommitCause::RuntimeEvent { event_id } = &state.cause {
|
||||
if let Some(injection_id) = event_id.strip_prefix("inject-context:") {
|
||||
if let Some(pending) = self.pending_injections.remove(injection_id)
|
||||
{
|
||||
let delivered_at_ms = crate::cursor::tools::runtime::now_ms()
|
||||
.min(i64::MAX as u64)
|
||||
as i64;
|
||||
self.handle.emit(&events::context_injection_delivered(
|
||||
injection_id.to_owned(),
|
||||
pending.delivery_batch_id.clone(),
|
||||
delivered_at_ms,
|
||||
))?;
|
||||
if let Some(user_message) = pending.user_message {
|
||||
self.handle
|
||||
.emit(&events::user_message_appended(user_message))?;
|
||||
}
|
||||
// Injections key `pending_injections` by their raw
|
||||
// injection id and commit under `inject-context:{id}`,
|
||||
// while runtime user messages key it by (and commit
|
||||
// under) the full `user-message:{id}` event id. Strip
|
||||
// the injection prefix when present and otherwise use
|
||||
// the event id verbatim so both are cleared and emit
|
||||
// their delivered/appended events.
|
||||
let injection_id = event_id
|
||||
.strip_prefix("inject-context:")
|
||||
.unwrap_or(event_id.as_str());
|
||||
if let Some(pending) = self.pending_injections.remove(injection_id) {
|
||||
let delivered_at_ms = crate::cursor::tools::runtime::now_ms()
|
||||
.min(i64::MAX as u64)
|
||||
as i64;
|
||||
self.handle.emit(&events::context_injection_delivered(
|
||||
injection_id.to_owned(),
|
||||
pending.delivery_batch_id.clone(),
|
||||
delivered_at_ms,
|
||||
))?;
|
||||
if let Some(user_message) = pending.user_message {
|
||||
self.handle
|
||||
.emit(&events::user_message_appended(user_message))?;
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -509,6 +575,7 @@ impl ConversationOutput {
|
||||
.map_err(|_| Error::Protocol("checkpoint worker stopped".into()))?
|
||||
{
|
||||
Ok(checkpoint) => {
|
||||
context_tokens = checkpoint_context_tokens(&checkpoint);
|
||||
compaction_checkpoint = Some(checkpoint);
|
||||
state.barrier.complete(Ok(()));
|
||||
}
|
||||
@@ -652,7 +719,7 @@ impl ConversationOutput {
|
||||
if self.superseded.is_cancelled() {
|
||||
worker.abort();
|
||||
self.abort_execs().await;
|
||||
return Ok(());
|
||||
return Ok(RunFinish::Transport(TransportFinish::Cancelled));
|
||||
}
|
||||
return match outcome {
|
||||
RunOutcome::Completed => {
|
||||
@@ -668,8 +735,7 @@ impl ConversationOutput {
|
||||
for _ in 0..3 {
|
||||
self.checkpoint.publish(&self.handle, &checkpoint).await?;
|
||||
}
|
||||
finish_success(&self.handle);
|
||||
return Ok(());
|
||||
return Ok(RunFinish::TurnCompleted);
|
||||
}
|
||||
let checkpoints = final_checkpoint.take().ok_or_else(|| {
|
||||
Error::Protocol("Completed without final state".into())
|
||||
@@ -685,18 +751,19 @@ impl ConversationOutput {
|
||||
ttft_breakdown: None,
|
||||
message: Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(checkpoints.settled)),
|
||||
})?;
|
||||
finish_success(&self.handle);
|
||||
Ok(())
|
||||
Ok(RunFinish::TurnCompleted)
|
||||
}
|
||||
RunOutcome::Cancelled => {
|
||||
worker.abort();
|
||||
self.abort_execs().await;
|
||||
finish_cancelled(&self.handle)
|
||||
Ok(RunFinish::Transport(TransportFinish::Cancelled))
|
||||
}
|
||||
RunOutcome::Failed(failure) => {
|
||||
worker.abort();
|
||||
self.abort_execs().await;
|
||||
finish_failed(&self.handle, &cursor_error(failure))
|
||||
Ok(RunFinish::Transport(TransportFinish::Failed(cursor_error(
|
||||
failure,
|
||||
))))
|
||||
}
|
||||
};
|
||||
}
|
||||
@@ -1015,10 +1082,34 @@ pub(crate) fn finish_cancelled(handle: &TransportHandle) -> Result<()> {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn checkpoint_context_tokens(checkpoint: &pb::ConversationStateStructure) -> Option<u64> {
|
||||
checkpoint
|
||||
.token_details
|
||||
.as_ref()
|
||||
.map(|details| u64::from(details.used_tokens))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::accept_tool_completion;
|
||||
use crate::{run::CommandResult, Error};
|
||||
use super::{accept_tool_completion, checkpoint_context_tokens};
|
||||
use crate::{cursor::protocol::proto::agent::v1 as pb, run::CommandResult, Error};
|
||||
|
||||
#[test]
|
||||
fn compacted_checkpoint_replaces_the_in_memory_context_usage() {
|
||||
let compacted = pb::ConversationStateStructure {
|
||||
token_details: Some(pb::ConversationTokenDetails {
|
||||
used_tokens: 20_000,
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
assert_eq!(checkpoint_context_tokens(&compacted), Some(20_000));
|
||||
assert_eq!(
|
||||
checkpoint_context_tokens(&pb::ConversationStateStructure::default()),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn closing_and_ended_runs_ignore_known_tool_completions() {
|
||||
|
||||
@@ -26,6 +26,8 @@ pub(crate) struct ConversationDependencies {
|
||||
pub provider: Arc<dyn Provider>,
|
||||
pub compiler: PromptCompiler,
|
||||
pub web_cache: WebCache,
|
||||
/// 本地 rules 服务的 md 存储目录;编译请求上下文时合并其中的规则。
|
||||
pub local_rules_dir: Option<std::path::PathBuf>,
|
||||
}
|
||||
|
||||
struct RegistryInner {
|
||||
@@ -47,6 +49,7 @@ impl ConversationRegistry {
|
||||
provider: Arc<dyn Provider>,
|
||||
compiler: PromptCompiler,
|
||||
web_cache: WebCache,
|
||||
local_rules_dir: Option<std::path::PathBuf>,
|
||||
) -> Self {
|
||||
Self {
|
||||
inner: Arc::new(RegistryInner {
|
||||
@@ -58,6 +61,7 @@ impl ConversationRegistry {
|
||||
provider,
|
||||
compiler,
|
||||
web_cache,
|
||||
local_rules_dir,
|
||||
},
|
||||
}),
|
||||
}
|
||||
|
||||
@@ -11,25 +11,29 @@ use crate::{
|
||||
protocol::proto::agent::v1 as pb,
|
||||
services::{blob_sync::BlobSynchronizer, context_sync::RequestContextSynchronizer},
|
||||
tools::{
|
||||
codec,
|
||||
codec, compat,
|
||||
runtime::CursorToolRuntime,
|
||||
tool_call_result::{tool_result_channel, ToolResultReceiver, ToolResultSender},
|
||||
ClientToolEvent, ToolDispatcher,
|
||||
},
|
||||
transport::{OrderedInbox, TransportHandle},
|
||||
},
|
||||
run::{CommandResult, RunEngine, RunHandle},
|
||||
run::{CommandResult, RunEngine, RunHandle, RunPhase},
|
||||
};
|
||||
|
||||
use super::{
|
||||
CompiledMessages, ConversationDependencies, ConversationOutput, ConversationOutputDependencies,
|
||||
ConversationRegistry, MessageDelivery, TransportCommand,
|
||||
ConversationRegistry, MessageDelivery, RunFinish, TransportCommand, TransportFinish,
|
||||
};
|
||||
|
||||
pub struct ConversationRuntime;
|
||||
|
||||
const CONTINUATION_IDLE_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(2);
|
||||
|
||||
#[derive(Clone)]
|
||||
struct RunGeneration {
|
||||
id: u64,
|
||||
request: pb::AgentRunRequest,
|
||||
superseded: CancellationToken,
|
||||
finished: CancellationToken,
|
||||
run: Arc<parking_lot::Mutex<Option<RunHandle>>>,
|
||||
@@ -41,6 +45,14 @@ struct RunGeneration {
|
||||
|
||||
struct FinishGeneration(CancellationToken);
|
||||
|
||||
struct TransportActorGuard(TransportHandle);
|
||||
|
||||
impl Drop for TransportActorGuard {
|
||||
fn drop(&mut self) {
|
||||
self.0.close_transport();
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for FinishGeneration {
|
||||
fn drop(&mut self) {
|
||||
self.0.cancel();
|
||||
@@ -54,6 +66,7 @@ impl ConversationRuntime {
|
||||
mut receiver: mpsc::Receiver<TransportCommand>,
|
||||
) {
|
||||
tokio::spawn(async move {
|
||||
let _actor_guard = TransportActorGuard(handle.clone());
|
||||
let dependencies = registry.dependencies().clone();
|
||||
let blob_sync = BlobSynchronizer::new(
|
||||
handle.request_id().into(),
|
||||
@@ -65,19 +78,70 @@ impl ConversationRuntime {
|
||||
let context_sync =
|
||||
RequestContextSynchronizer::new(handle.clone(), dependencies.store.clone());
|
||||
let mut current = None::<RunGeneration>;
|
||||
let mut next_generation = 1_u64;
|
||||
let mut pending_finish = None::<(u64, TransportFinish)>;
|
||||
let mut draining = false;
|
||||
let mut waiting_for_action = false;
|
||||
loop {
|
||||
let command = match receiver.recv().await {
|
||||
Some(command) => command,
|
||||
None => {
|
||||
handle.mark_disconnected();
|
||||
if let Some(generation) = current.as_ref() {
|
||||
generation.superseded.cancel();
|
||||
if let Some(run) = generation.run.lock().clone() {
|
||||
run.cancel();
|
||||
let command = if draining {
|
||||
if !handle.admissions_drained() {
|
||||
tokio::select! {
|
||||
command = receiver.recv() => match command {
|
||||
Some(command) => command,
|
||||
None => {
|
||||
finish_pending(&handle, ¤t, pending_finish.take());
|
||||
break;
|
||||
}
|
||||
},
|
||||
_ = handle.wait_admissions_drained() => continue,
|
||||
}
|
||||
} else {
|
||||
handle.mark_draining();
|
||||
match receiver.try_recv() {
|
||||
Ok(command) => command,
|
||||
Err(mpsc::error::TryRecvError::Empty)
|
||||
| Err(mpsc::error::TryRecvError::Disconnected) => {
|
||||
finish_pending(&handle, ¤t, pending_finish.take());
|
||||
break;
|
||||
}
|
||||
}
|
||||
super::finish_cancelled(&handle).ok();
|
||||
break;
|
||||
}
|
||||
} else if waiting_for_action {
|
||||
tokio::select! {
|
||||
command = receiver.recv() => match command {
|
||||
Some(command) => command,
|
||||
None => {
|
||||
handle.mark_disconnected();
|
||||
super::finish_success(&handle);
|
||||
break;
|
||||
}
|
||||
},
|
||||
_ = tokio::time::sleep(CONTINUATION_IDLE_TIMEOUT) => {
|
||||
let Some(generation) = current.as_ref() else {
|
||||
super::finish_success(&handle);
|
||||
break;
|
||||
};
|
||||
handle.begin_close();
|
||||
pending_finish = Some((generation.id, TransportFinish::Success));
|
||||
draining = true;
|
||||
waiting_for_action = false;
|
||||
continue;
|
||||
}
|
||||
}
|
||||
} else {
|
||||
match receiver.recv().await {
|
||||
Some(command) => command,
|
||||
None => {
|
||||
handle.mark_disconnected();
|
||||
if let Some(generation) = current.as_ref() {
|
||||
generation.superseded.cancel();
|
||||
if let Some(run) = generation.run.lock().clone() {
|
||||
run.cancel();
|
||||
}
|
||||
}
|
||||
super::finish_cancelled(&handle).ok();
|
||||
break;
|
||||
}
|
||||
}
|
||||
};
|
||||
match command {
|
||||
@@ -92,11 +156,41 @@ impl ConversationRuntime {
|
||||
let _ = handle.emit(&codec::abort(id));
|
||||
}
|
||||
}
|
||||
super::finish_cancelled(&handle).ok();
|
||||
let turn_completed = waiting_for_action
|
||||
|| current.as_ref().is_some_and(|generation| {
|
||||
generation
|
||||
.run
|
||||
.lock()
|
||||
.as_ref()
|
||||
.is_none_or(|run| run.phase() != RunPhase::Running)
|
||||
});
|
||||
if turn_completed {
|
||||
super::finish_success(&handle);
|
||||
} else {
|
||||
super::finish_cancelled(&handle).ok();
|
||||
}
|
||||
break;
|
||||
}
|
||||
TransportCommand::Close => {
|
||||
break;
|
||||
TransportCommand::RunFinished { generation, finish } => {
|
||||
if !current
|
||||
.as_ref()
|
||||
.is_some_and(|current| current.id == generation)
|
||||
{
|
||||
continue;
|
||||
}
|
||||
match finish {
|
||||
RunFinish::TurnCompleted => {
|
||||
pending_finish = None;
|
||||
draining = false;
|
||||
waiting_for_action = true;
|
||||
}
|
||||
RunFinish::Transport(finish) => {
|
||||
waiting_for_action = false;
|
||||
handle.begin_close();
|
||||
pending_finish = Some((generation, finish));
|
||||
draining = true;
|
||||
}
|
||||
}
|
||||
}
|
||||
TransportCommand::Append { seqno, message } => {
|
||||
for (_seqno, message) in inbox.push(seqno, *message) {
|
||||
@@ -105,6 +199,12 @@ impl ConversationRuntime {
|
||||
Some(pb::agent_client_message::Message::RunRequest(
|
||||
request,
|
||||
)) => {
|
||||
waiting_for_action = false;
|
||||
if draining {
|
||||
handle.reopen();
|
||||
draining = false;
|
||||
pending_finish = None;
|
||||
}
|
||||
if let Some(conversation_id) =
|
||||
request.conversation_id.as_deref()
|
||||
{
|
||||
@@ -117,60 +217,21 @@ impl ConversationRuntime {
|
||||
"invalid Cursor conversation id"
|
||||
);
|
||||
let _ = super::finish_failed(&handle, &error);
|
||||
let _ =
|
||||
handle.command(TransportCommand::Close).await;
|
||||
return;
|
||||
}
|
||||
}
|
||||
let previous_finished =
|
||||
if let Some(previous) = current.take() {
|
||||
previous.superseded.cancel();
|
||||
if let Some(run) = previous.run.lock().clone() {
|
||||
run.cancel();
|
||||
}
|
||||
for id in previous
|
||||
.tool_runtime
|
||||
.interrupt_for_run_replacement()
|
||||
.await
|
||||
{
|
||||
let _ = handle.emit(&codec::abort(id));
|
||||
}
|
||||
Some(previous.finished.clone())
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let (results, result_receiver) = tool_result_channel();
|
||||
let (runtime_actions, runtime_action_receiver) =
|
||||
mpsc::unbounded_channel::<compile::RuntimeAction>();
|
||||
let tool_runtime = tool_runtime_factory.next_run();
|
||||
let tools = ToolDispatcher::with_results(
|
||||
tool_runtime.clone(),
|
||||
results.clone(),
|
||||
dependencies.store.clone(),
|
||||
dependencies.web_cache.clone(),
|
||||
);
|
||||
let generation = RunGeneration {
|
||||
superseded: CancellationToken::new(),
|
||||
finished: CancellationToken::new(),
|
||||
run: Arc::new(parking_lot::Mutex::new(None)),
|
||||
results,
|
||||
runtime_actions,
|
||||
tool_runtime,
|
||||
tools,
|
||||
};
|
||||
current = Some(generation.clone());
|
||||
spawn_run_request(
|
||||
registry.clone(),
|
||||
handle.clone(),
|
||||
start_generation(
|
||||
®istry,
|
||||
&handle,
|
||||
&dependencies,
|
||||
&blob_sync,
|
||||
&context_sync,
|
||||
&tool_runtime_factory,
|
||||
&mut current,
|
||||
&mut next_generation,
|
||||
request,
|
||||
dependencies.clone(),
|
||||
blob_sync.clone(),
|
||||
context_sync.clone(),
|
||||
generation,
|
||||
previous_finished,
|
||||
result_receiver,
|
||||
runtime_action_receiver,
|
||||
);
|
||||
)
|
||||
.await;
|
||||
}
|
||||
Some(pb::agent_client_message::Message::ExecClientMessage(
|
||||
message,
|
||||
@@ -262,17 +323,18 @@ impl ConversationRuntime {
|
||||
.take_exec(throw.id)
|
||||
.await
|
||||
{
|
||||
Some(pending) => generation.results.send_error(
|
||||
crate::Error::Protocol(format!(
|
||||
"Exec {} failed: {}",
|
||||
pending.call.call_id, throw.error
|
||||
)),
|
||||
Some(pending) => generation.results.send(
|
||||
compat::failure_with_message(
|
||||
&pending.call,
|
||||
format!(
|
||||
"Exec {} failed: {}",
|
||||
pending.call.call_id, throw.error
|
||||
),
|
||||
),
|
||||
),
|
||||
None => generation.results.send_error(
|
||||
crate::Error::Protocol(format!(
|
||||
"unknown ExecClientThrow id: {}",
|
||||
throw.id
|
||||
)),
|
||||
None => tracing::warn!(
|
||||
id = throw.id,
|
||||
"ignoring failure for unknown tool execution"
|
||||
),
|
||||
}
|
||||
}
|
||||
@@ -330,27 +392,48 @@ impl ConversationRuntime {
|
||||
// return an explicit Protocol Error rather than falling through silently.
|
||||
Some(
|
||||
pb::agent_client_message::Message::ConversationAction(
|
||||
action,
|
||||
conversation_action,
|
||||
),
|
||||
) => match action.action {
|
||||
) => match conversation_action.action.clone() {
|
||||
Some(
|
||||
pb::conversation_action::Action::UserMessageAction(
|
||||
action,
|
||||
),
|
||||
) => {
|
||||
let Some(generation) = current.as_ref() else {
|
||||
let delivered_to_active_run =
|
||||
current.as_ref().is_some_and(|generation| {
|
||||
generation.run.lock().as_ref().is_some_and(
|
||||
|run| run.phase() == RunPhase::Running,
|
||||
) && generation
|
||||
.runtime_actions
|
||||
.send(compile::RuntimeAction::UserMessage(
|
||||
action.clone(),
|
||||
))
|
||||
.is_ok()
|
||||
});
|
||||
if delivered_to_active_run {
|
||||
continue;
|
||||
}
|
||||
let Some(previous) = current.as_ref() else {
|
||||
continue;
|
||||
};
|
||||
if generation
|
||||
.runtime_actions
|
||||
.send(compile::RuntimeAction::UserMessage(action))
|
||||
.is_err()
|
||||
{
|
||||
generation.results.send_error(crate::Error::Protocol(
|
||||
"UserMessageAction arrived without an active Run"
|
||||
.into(),
|
||||
));
|
||||
}
|
||||
let mut request = previous.request.clone();
|
||||
request.action = Some(conversation_action);
|
||||
request.conversation_state = None;
|
||||
request.pre_fetched_blobs.clear();
|
||||
waiting_for_action = false;
|
||||
start_generation(
|
||||
®istry,
|
||||
&handle,
|
||||
&dependencies,
|
||||
&blob_sync,
|
||||
&context_sync,
|
||||
&tool_runtime_factory,
|
||||
&mut current,
|
||||
&mut next_generation,
|
||||
request,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
Some(pb::conversation_action::Action::CancelAction(_)) => {
|
||||
if let Some(generation) = current.as_ref() {
|
||||
@@ -427,6 +510,95 @@ impl ConversationRuntime {
|
||||
}
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
async fn start_generation(
|
||||
registry: &ConversationRegistry,
|
||||
handle: &TransportHandle,
|
||||
dependencies: &ConversationDependencies,
|
||||
blob_sync: &BlobSynchronizer,
|
||||
context_sync: &RequestContextSynchronizer,
|
||||
tool_runtime_factory: &CursorToolRuntime,
|
||||
current: &mut Option<RunGeneration>,
|
||||
next_generation: &mut u64,
|
||||
request: pb::AgentRunRequest,
|
||||
) {
|
||||
let previous_finished = if let Some(previous) = current.take() {
|
||||
previous.superseded.cancel();
|
||||
if let Some(run) = previous.run.lock().clone() {
|
||||
run.cancel();
|
||||
}
|
||||
for id in previous.tool_runtime.interrupt_for_run_replacement().await {
|
||||
let _ = handle.emit(&codec::abort(id));
|
||||
}
|
||||
Some(previous.finished.clone())
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let (results, result_receiver) = tool_result_channel();
|
||||
let (runtime_actions, runtime_action_receiver) =
|
||||
mpsc::unbounded_channel::<compile::RuntimeAction>();
|
||||
let tool_runtime = tool_runtime_factory.next_run();
|
||||
let tools = ToolDispatcher::with_results(
|
||||
tool_runtime.clone(),
|
||||
results.clone(),
|
||||
dependencies.store.clone(),
|
||||
dependencies.web_cache.clone(),
|
||||
);
|
||||
let generation = RunGeneration {
|
||||
id: *next_generation,
|
||||
request: request.clone(),
|
||||
superseded: CancellationToken::new(),
|
||||
finished: CancellationToken::new(),
|
||||
run: Arc::new(parking_lot::Mutex::new(None)),
|
||||
results,
|
||||
runtime_actions,
|
||||
tool_runtime,
|
||||
tools,
|
||||
};
|
||||
*next_generation = next_generation.saturating_add(1);
|
||||
*current = Some(generation.clone());
|
||||
spawn_run_request(
|
||||
registry.clone(),
|
||||
handle.clone(),
|
||||
request,
|
||||
dependencies.clone(),
|
||||
blob_sync.clone(),
|
||||
context_sync.clone(),
|
||||
generation,
|
||||
previous_finished,
|
||||
result_receiver,
|
||||
runtime_action_receiver,
|
||||
);
|
||||
}
|
||||
|
||||
fn finish_pending(
|
||||
handle: &TransportHandle,
|
||||
current: &Option<RunGeneration>,
|
||||
pending: Option<(u64, TransportFinish)>,
|
||||
) {
|
||||
let Some((generation, finish)) = pending else {
|
||||
return;
|
||||
};
|
||||
if current
|
||||
.as_ref()
|
||||
.is_some_and(|current| current.id == generation)
|
||||
{
|
||||
finish_transport(handle, finish);
|
||||
}
|
||||
}
|
||||
|
||||
fn finish_transport(handle: &TransportHandle, finish: TransportFinish) {
|
||||
match finish {
|
||||
TransportFinish::Success => super::finish_success(handle),
|
||||
TransportFinish::Failed(error) => {
|
||||
let _ = super::finish_failed(handle, &error);
|
||||
}
|
||||
TransportFinish::Cancelled => {
|
||||
let _ = super::finish_cancelled(handle);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
fn spawn_run_request(
|
||||
registry: ConversationRegistry,
|
||||
@@ -471,6 +643,7 @@ fn spawn_run_request(
|
||||
checkpoint: &checkpoint,
|
||||
blob_sync: &blob_sync,
|
||||
context_sync: &context_sync,
|
||||
local_rules_dir: dependencies.local_rules_dir.as_deref(),
|
||||
},
|
||||
) => prepared,
|
||||
};
|
||||
@@ -485,8 +658,12 @@ fn spawn_run_request(
|
||||
%error,
|
||||
"failed to prepare Cursor Run"
|
||||
);
|
||||
let _ = super::finish_failed(&handle, &error);
|
||||
let _ = handle.command(TransportCommand::Close).await;
|
||||
let _ = handle
|
||||
.command(TransportCommand::RunFinished {
|
||||
generation: generation.id,
|
||||
finish: RunFinish::Transport(TransportFinish::Failed(error)),
|
||||
})
|
||||
.await;
|
||||
return;
|
||||
}
|
||||
};
|
||||
@@ -518,8 +695,12 @@ fn spawn_run_request(
|
||||
{
|
||||
CommandResult::Applied | CommandResult::Duplicate => {
|
||||
if !generation.superseded.is_cancelled() {
|
||||
super::finish_success(&handle);
|
||||
let _ = handle.command(TransportCommand::Close).await;
|
||||
let _ = handle
|
||||
.command(TransportCommand::RunFinished {
|
||||
generation: generation.id,
|
||||
finish: RunFinish::Transport(TransportFinish::Success),
|
||||
})
|
||||
.await;
|
||||
}
|
||||
return;
|
||||
}
|
||||
@@ -543,8 +724,12 @@ fn spawn_run_request(
|
||||
}
|
||||
CommandResult::StaleTarget => {
|
||||
if !generation.superseded.is_cancelled() {
|
||||
super::finish_success(&handle);
|
||||
let _ = handle.command(TransportCommand::Close).await;
|
||||
let _ = handle
|
||||
.command(TransportCommand::RunFinished {
|
||||
generation: generation.id,
|
||||
finish: RunFinish::Transport(TransportFinish::Success),
|
||||
})
|
||||
.await;
|
||||
}
|
||||
return;
|
||||
}
|
||||
@@ -603,16 +788,21 @@ fn spawn_run_request(
|
||||
tool_runtime: generation.tool_runtime.clone(),
|
||||
},
|
||||
);
|
||||
if let Err(error) = output.run().await {
|
||||
if !generation.superseded.is_cancelled() {
|
||||
tracing::error!(
|
||||
request_id = handle.request_id(),
|
||||
%error,
|
||||
"Cursor session failed"
|
||||
);
|
||||
let _ = super::finish_failed(&handle, &error);
|
||||
let finish = match output.run().await {
|
||||
Ok(finish) => finish,
|
||||
Err(error) => {
|
||||
if generation.superseded.is_cancelled() {
|
||||
RunFinish::Transport(TransportFinish::Cancelled)
|
||||
} else {
|
||||
tracing::error!(
|
||||
request_id = handle.request_id(),
|
||||
%error,
|
||||
"Cursor session failed"
|
||||
);
|
||||
RunFinish::Transport(TransportFinish::Failed(error))
|
||||
}
|
||||
}
|
||||
}
|
||||
};
|
||||
let _ = core_run.await;
|
||||
registry.release(&conversation_id, &run_id).await;
|
||||
if generation
|
||||
@@ -624,7 +814,12 @@ fn spawn_run_request(
|
||||
*generation.run.lock() = None;
|
||||
}
|
||||
if !generation.superseded.is_cancelled() {
|
||||
let _ = handle.command(TransportCommand::Close).await;
|
||||
let _ = handle
|
||||
.command(TransportCommand::RunFinished {
|
||||
generation: generation.id,
|
||||
finish,
|
||||
})
|
||||
.await;
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
@@ -20,6 +20,9 @@ use crate::{
|
||||
|
||||
type BlobSetSender = oneshot::Sender<Result<()>>;
|
||||
|
||||
const SET_TIMEOUT: Duration = Duration::from_secs(30 * 60);
|
||||
const GET_TIMEOUT: Duration = Duration::from_secs(10 * 60);
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct BlobSynchronizer {
|
||||
inner: Arc<Inner>,
|
||||
@@ -73,22 +76,20 @@ impl BlobSynchronizer {
|
||||
let id = self.inner.store.put_blob(data, edges).await?;
|
||||
let result = self.ensure_set(&id, data).await;
|
||||
if let Some(trace) = self.inner.handle.trace() {
|
||||
trace
|
||||
.linked_blob(
|
||||
"blob_set",
|
||||
"byok_server",
|
||||
&id,
|
||||
serde_json::json!({
|
||||
"byte_count": data.len(),
|
||||
"status": if result.is_ok() { "acknowledged" } else { "error" },
|
||||
"error": result.as_ref().err().map(ToString::to_string),
|
||||
"edges": edges.iter().map(|edge| serde_json::json!({
|
||||
"child_blob_id": edge.child.to_base64(),
|
||||
"field_name": edge.field_name,
|
||||
})).collect::<Vec<_>>(),
|
||||
}),
|
||||
)
|
||||
.await;
|
||||
trace.linked_blob(
|
||||
"blob_set",
|
||||
"byok_server",
|
||||
&id,
|
||||
serde_json::json!({
|
||||
"byte_count": data.len(),
|
||||
"status": if result.is_ok() { "acknowledged" } else { "error" },
|
||||
"error": result.as_ref().err().map(ToString::to_string),
|
||||
"edges": edges.iter().map(|edge| serde_json::json!({
|
||||
"child_blob_id": edge.child.to_base64(),
|
||||
"field_name": edge.field_name,
|
||||
})).collect::<Vec<_>>(),
|
||||
}),
|
||||
);
|
||||
}
|
||||
result?;
|
||||
Ok(id)
|
||||
@@ -130,7 +131,7 @@ impl BlobSynchronizer {
|
||||
let result = tokio::select! {
|
||||
result = receiver => result.map_err(|_| Error::Protocol("KV SET response channel closed".into()))?,
|
||||
_ = cancellation.cancelled() => Err(Error::Cancelled),
|
||||
_ = tokio::time::sleep(Duration::from_secs(60)) => Err(Error::Protocol(format!("KV SET timed out: {}", blob_id.to_base64()))),
|
||||
_ = tokio::time::sleep(SET_TIMEOUT) => Err(Error::Protocol(format!("KV SET timed out: {}", blob_id.to_base64()))),
|
||||
};
|
||||
if result.is_err() {
|
||||
self.inner.set_requests.lock().await.remove(&id);
|
||||
@@ -141,18 +142,16 @@ impl BlobSynchronizer {
|
||||
pub async fn get(&self, blob_id: &BlobId) -> Result<Option<Vec<u8>>> {
|
||||
if let Some(data) = self.inner.store.get_blob(blob_id).await? {
|
||||
if let Some(trace) = self.inner.handle.trace() {
|
||||
trace
|
||||
.linked_blob(
|
||||
"blob_get",
|
||||
"byok_server",
|
||||
blob_id,
|
||||
serde_json::json!({
|
||||
"byte_count": data.len(),
|
||||
"source": "local_store",
|
||||
"status": "found",
|
||||
}),
|
||||
)
|
||||
.await;
|
||||
trace.linked_blob(
|
||||
"blob_get",
|
||||
"byok_server",
|
||||
blob_id,
|
||||
serde_json::json!({
|
||||
"byte_count": data.len(),
|
||||
"source": "local_store",
|
||||
"status": "found",
|
||||
}),
|
||||
);
|
||||
}
|
||||
return Ok(Some(data));
|
||||
}
|
||||
@@ -183,7 +182,7 @@ impl BlobSynchronizer {
|
||||
let result = tokio::select! {
|
||||
result = receiver => result.map_err(|_| Error::Protocol("KV GET response channel closed".into()))?,
|
||||
_ = cancellation.cancelled() => Err(Error::Cancelled),
|
||||
_ = tokio::time::sleep(Duration::from_secs(60)) => Err(Error::Protocol(format!("KV GET timed out: {}", blob_id.to_base64()))),
|
||||
_ = tokio::time::sleep(GET_TIMEOUT) => Err(Error::Protocol(format!("KV GET timed out: {}", blob_id.to_base64()))),
|
||||
};
|
||||
if result.is_err() {
|
||||
self.inner.get_requests.lock().await.remove(&id);
|
||||
@@ -191,45 +190,39 @@ impl BlobSynchronizer {
|
||||
if let Some(trace) = self.inner.handle.trace() {
|
||||
match &result {
|
||||
Ok(Some(data)) => {
|
||||
trace
|
||||
.linked_blob(
|
||||
"blob_get",
|
||||
"cursor_client",
|
||||
blob_id,
|
||||
serde_json::json!({
|
||||
"byte_count": data.len(),
|
||||
"source": "cursor_client",
|
||||
"status": "found",
|
||||
}),
|
||||
)
|
||||
.await;
|
||||
trace.linked_blob(
|
||||
"blob_get",
|
||||
"cursor_client",
|
||||
blob_id,
|
||||
serde_json::json!({
|
||||
"byte_count": data.len(),
|
||||
"source": "cursor_client",
|
||||
"status": "found",
|
||||
}),
|
||||
);
|
||||
}
|
||||
Ok(None) => {
|
||||
trace
|
||||
.artifact(
|
||||
"blob_get",
|
||||
"cursor_client",
|
||||
&[],
|
||||
serde_json::json!({
|
||||
"blob_id": blob_id.to_base64(),
|
||||
"status": "missing",
|
||||
}),
|
||||
)
|
||||
.await;
|
||||
trace.artifact(
|
||||
"blob_get",
|
||||
"cursor_client",
|
||||
&[],
|
||||
serde_json::json!({
|
||||
"blob_id": blob_id.to_base64(),
|
||||
"status": "missing",
|
||||
}),
|
||||
);
|
||||
}
|
||||
Err(error) => {
|
||||
trace
|
||||
.artifact(
|
||||
"blob_get",
|
||||
"cursor_client",
|
||||
&[],
|
||||
serde_json::json!({
|
||||
"blob_id": blob_id.to_base64(),
|
||||
"status": "error",
|
||||
"error": error.to_string(),
|
||||
}),
|
||||
)
|
||||
.await;
|
||||
trace.artifact(
|
||||
"blob_get",
|
||||
"cursor_client",
|
||||
&[],
|
||||
serde_json::json!({
|
||||
"blob_id": blob_id.to_base64(),
|
||||
"status": "error",
|
||||
"error": error.to_string(),
|
||||
}),
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -318,3 +311,18 @@ impl BlobSynchronizer {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn set_timeout_allows_slow_cursor_acknowledgements() {
|
||||
assert_eq!(SET_TIMEOUT, Duration::from_secs(30 * 60));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn get_timeout_allows_slow_cursor_responses() {
|
||||
assert_eq!(GET_TIMEOUT, Duration::from_secs(10 * 60));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,356 @@
|
||||
//! Serves Cursor user rules: upstream-first with an offline markdown cache.
|
||||
//!
|
||||
//! 每个请求先回放离线日志再尝试上游;上游成功时把结果写穿到本地镜像,
|
||||
//! 上游不可达时降级为本地 md 存储并记录日志等待回放。
|
||||
mod store;
|
||||
mod sync;
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use axum::{
|
||||
body::{to_bytes, Body, Bytes},
|
||||
extract::Extension,
|
||||
http::{header, Request, Response},
|
||||
};
|
||||
use prost::Message;
|
||||
|
||||
use crate::{api::cursor::proxy, config, cursor::protocol::connect, Result};
|
||||
|
||||
pub(crate) use store::{RuleRecord, RuleStore};
|
||||
|
||||
#[derive(Clone, PartialEq, Message)]
|
||||
pub(crate) struct KnowledgeBaseAddRequest {
|
||||
#[prost(string, tag = "1")]
|
||||
knowledge: String,
|
||||
#[prost(string, tag = "2")]
|
||||
title: String,
|
||||
#[prost(string, tag = "3")]
|
||||
git_origin: String,
|
||||
#[prost(string, optional, tag = "4")]
|
||||
composer_id: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, Message)]
|
||||
pub(crate) struct KnowledgeBaseAddResponse {
|
||||
#[prost(bool, tag = "1")]
|
||||
success: bool,
|
||||
#[prost(string, tag = "2")]
|
||||
id: String,
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, Message)]
|
||||
pub(crate) struct KnowledgeBaseListRequest {
|
||||
#[prost(int32, optional, tag = "1")]
|
||||
limit: Option<i32>,
|
||||
#[prost(string, optional, tag = "2")]
|
||||
git_origin: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, Message)]
|
||||
pub(crate) struct KnowledgeBaseListResponse {
|
||||
#[prost(bool, tag = "1")]
|
||||
success: bool,
|
||||
#[prost(message, repeated, tag = "2")]
|
||||
all_results: Vec<KnowledgeBaseListItem>,
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, Message)]
|
||||
pub(crate) struct KnowledgeBaseListItem {
|
||||
#[prost(string, tag = "1")]
|
||||
id: String,
|
||||
#[prost(string, tag = "2")]
|
||||
knowledge: String,
|
||||
#[prost(string, tag = "3")]
|
||||
title: String,
|
||||
#[prost(string, tag = "4")]
|
||||
created_at: String,
|
||||
#[prost(bool, tag = "5")]
|
||||
is_generated: bool,
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, Message)]
|
||||
pub(crate) struct KnowledgeBaseUpdateRequest {
|
||||
#[prost(string, tag = "1")]
|
||||
id: String,
|
||||
#[prost(string, tag = "2")]
|
||||
knowledge: String,
|
||||
#[prost(string, tag = "3")]
|
||||
title: String,
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, Message)]
|
||||
pub(crate) struct KnowledgeBaseUpdateResponse {
|
||||
#[prost(bool, tag = "1")]
|
||||
success: bool,
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, Message)]
|
||||
pub(crate) struct KnowledgeBaseRemoveRequest {
|
||||
#[prost(string, tag = "1")]
|
||||
id: String,
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, Message)]
|
||||
pub(crate) struct KnowledgeBaseRemoveResponse {
|
||||
#[prost(bool, tag = "1")]
|
||||
success: bool,
|
||||
}
|
||||
|
||||
/// 规则存储与并发锁;经 axum Extension 注入四个 handler。
|
||||
#[derive(Clone)]
|
||||
pub struct KnowledgeService {
|
||||
inner: Arc<Inner>,
|
||||
}
|
||||
|
||||
struct Inner {
|
||||
store: RuleStore,
|
||||
lock: tokio::sync::Mutex<()>,
|
||||
}
|
||||
|
||||
impl KnowledgeService {
|
||||
pub fn managed() -> Result<Self> {
|
||||
Self::with_root(config::managed_data_dir()?.join("rules"))
|
||||
}
|
||||
|
||||
/// 指定存储根目录构造;managed() 与集成测试共用。
|
||||
pub fn with_root(root: std::path::PathBuf) -> Result<Self> {
|
||||
Ok(Self {
|
||||
inner: Arc::new(Inner {
|
||||
store: RuleStore::open(root)?,
|
||||
lock: tokio::sync::Mutex::new(()),
|
||||
}),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn add(
|
||||
Extension(upstream): Extension<proxy::CursorProxy>,
|
||||
Extension(service): Extension<KnowledgeService>,
|
||||
request: Request<Body>,
|
||||
) -> Result<Response<Body>> {
|
||||
let (parts, body) = buffered(request).await?;
|
||||
let message: KnowledgeBaseAddRequest = connect::decode_unary(&body)?;
|
||||
let _guard = service.inner.lock.lock().await;
|
||||
let store = &service.inner.store;
|
||||
|
||||
if sync::replay(&upstream, &parts.headers, store).await? {
|
||||
match proxy::forward_buffered(&upstream, Request::from_parts(parts, Body::from(body))).await
|
||||
{
|
||||
Ok(response) if response.status.is_success() => {
|
||||
if let Ok(reply) = connect::decode_unary::<KnowledgeBaseAddResponse>(&response.body)
|
||||
{
|
||||
if reply.success && !reply.id.is_empty() {
|
||||
store.upsert(&RuleRecord {
|
||||
id: reply.id,
|
||||
knowledge: message.knowledge,
|
||||
title: message.title,
|
||||
created_at: now(),
|
||||
is_generated: false,
|
||||
git_origin: message.git_origin,
|
||||
})?;
|
||||
}
|
||||
}
|
||||
return Ok(response.into_response());
|
||||
}
|
||||
Ok(response) => {
|
||||
tracing::warn!(status = %response.status, "rules upstream rejected add; storing locally");
|
||||
}
|
||||
Err(error) => {
|
||||
tracing::warn!(%error, "rules upstream unavailable for add; storing locally");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let id = format!("{}{}", store::LOCAL_ID_PREFIX, uuid::Uuid::new_v4());
|
||||
store.upsert(&RuleRecord {
|
||||
id: id.clone(),
|
||||
knowledge: message.knowledge,
|
||||
title: message.title,
|
||||
created_at: now(),
|
||||
is_generated: false,
|
||||
git_origin: message.git_origin,
|
||||
})?;
|
||||
store.record_add(&id)?;
|
||||
proto(KnowledgeBaseAddResponse { success: true, id })
|
||||
}
|
||||
|
||||
pub async fn list(
|
||||
Extension(upstream): Extension<proxy::CursorProxy>,
|
||||
Extension(service): Extension<KnowledgeService>,
|
||||
request: Request<Body>,
|
||||
) -> Result<Response<Body>> {
|
||||
let (parts, body) = buffered(request).await?;
|
||||
let message: KnowledgeBaseListRequest = connect::decode_unary(&body)?;
|
||||
let _guard = service.inner.lock.lock().await;
|
||||
let store = &service.inner.store;
|
||||
let git_origin = message.git_origin.unwrap_or_default();
|
||||
|
||||
if sync::replay(&upstream, &parts.headers, store).await? {
|
||||
match proxy::forward_buffered(&upstream, Request::from_parts(parts, Body::from(body))).await
|
||||
{
|
||||
Ok(response) if response.status.is_success() => {
|
||||
if let Ok(reply) =
|
||||
connect::decode_unary::<KnowledgeBaseListResponse>(&response.body)
|
||||
{
|
||||
// 带 git_origin 过滤的列表只是子集,整体覆盖会误删其他规则。
|
||||
if reply.success && git_origin.is_empty() {
|
||||
sync::mirror(store, reply.all_results)?;
|
||||
}
|
||||
}
|
||||
return Ok(response.into_response());
|
||||
}
|
||||
Ok(response) => {
|
||||
tracing::warn!(status = %response.status, "rules upstream rejected list; serving local cache");
|
||||
}
|
||||
Err(error) => {
|
||||
tracing::warn!(%error, "rules upstream unavailable for list; serving local cache");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let mut records = store.list()?;
|
||||
if !git_origin.is_empty() {
|
||||
records.retain(|record| record.git_origin == git_origin);
|
||||
}
|
||||
if let Some(limit) = message.limit {
|
||||
if limit >= 0 {
|
||||
records.truncate(limit as usize);
|
||||
}
|
||||
}
|
||||
proto(KnowledgeBaseListResponse {
|
||||
success: true,
|
||||
all_results: records
|
||||
.into_iter()
|
||||
.map(|record| KnowledgeBaseListItem {
|
||||
id: record.id,
|
||||
knowledge: record.knowledge,
|
||||
title: record.title,
|
||||
created_at: record.created_at,
|
||||
is_generated: record.is_generated,
|
||||
})
|
||||
.collect(),
|
||||
})
|
||||
}
|
||||
|
||||
pub async fn update(
|
||||
Extension(upstream): Extension<proxy::CursorProxy>,
|
||||
Extension(service): Extension<KnowledgeService>,
|
||||
request: Request<Body>,
|
||||
) -> Result<Response<Body>> {
|
||||
let (parts, body) = buffered(request).await?;
|
||||
let message: KnowledgeBaseUpdateRequest = connect::decode_unary(&body)?;
|
||||
let _guard = service.inner.lock.lock().await;
|
||||
let store = &service.inner.store;
|
||||
|
||||
if sync::replay(&upstream, &parts.headers, store).await? {
|
||||
match proxy::forward_buffered(&upstream, Request::from_parts(parts, Body::from(body))).await
|
||||
{
|
||||
Ok(response) if response.status.is_success() => {
|
||||
if let Ok(reply) =
|
||||
connect::decode_unary::<KnowledgeBaseUpdateResponse>(&response.body)
|
||||
{
|
||||
if reply.success {
|
||||
let existing = store.get(&message.id)?;
|
||||
store.upsert(&RuleRecord {
|
||||
id: message.id,
|
||||
knowledge: message.knowledge,
|
||||
title: message.title,
|
||||
created_at: existing
|
||||
.as_ref()
|
||||
.map_or_else(now, |record| record.created_at.clone()),
|
||||
is_generated: existing
|
||||
.as_ref()
|
||||
.is_some_and(|record| record.is_generated),
|
||||
git_origin: existing
|
||||
.map(|record| record.git_origin)
|
||||
.unwrap_or_default(),
|
||||
})?;
|
||||
}
|
||||
}
|
||||
return Ok(response.into_response());
|
||||
}
|
||||
Ok(response) => {
|
||||
tracing::warn!(status = %response.status, "rules upstream rejected update; storing locally");
|
||||
}
|
||||
Err(error) => {
|
||||
tracing::warn!(%error, "rules upstream unavailable for update; storing locally");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let Some(mut record) = store.get(&message.id)? else {
|
||||
return proto(KnowledgeBaseUpdateResponse { success: false });
|
||||
};
|
||||
record.knowledge = message.knowledge;
|
||||
record.title = message.title;
|
||||
store.upsert(&record)?;
|
||||
store.record_update(&message.id)?;
|
||||
proto(KnowledgeBaseUpdateResponse { success: true })
|
||||
}
|
||||
|
||||
pub async fn remove(
|
||||
Extension(upstream): Extension<proxy::CursorProxy>,
|
||||
Extension(service): Extension<KnowledgeService>,
|
||||
request: Request<Body>,
|
||||
) -> Result<Response<Body>> {
|
||||
let (parts, body) = buffered(request).await?;
|
||||
let message: KnowledgeBaseRemoveRequest = connect::decode_unary(&body)?;
|
||||
let _guard = service.inner.lock.lock().await;
|
||||
let store = &service.inner.store;
|
||||
|
||||
if sync::replay(&upstream, &parts.headers, store).await? {
|
||||
match proxy::forward_buffered(&upstream, Request::from_parts(parts, Body::from(body))).await
|
||||
{
|
||||
Ok(response) if response.status.is_success() => {
|
||||
if let Ok(reply) =
|
||||
connect::decode_unary::<KnowledgeBaseRemoveResponse>(&response.body)
|
||||
{
|
||||
if reply.success {
|
||||
store.remove(&message.id)?;
|
||||
}
|
||||
}
|
||||
return Ok(response.into_response());
|
||||
}
|
||||
Ok(response) => {
|
||||
tracing::warn!(status = %response.status, "rules upstream rejected remove; removing locally");
|
||||
}
|
||||
Err(error) => {
|
||||
tracing::warn!(%error, "rules upstream unavailable for remove; removing locally");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
store.remove(&message.id)?;
|
||||
store.record_remove(&message.id)?;
|
||||
proto(KnowledgeBaseRemoveResponse { success: true })
|
||||
}
|
||||
|
||||
async fn buffered(request: Request<Body>) -> Result<(axum::http::request::Parts, Bytes)> {
|
||||
let (parts, body) = request.into_parts();
|
||||
let body = to_bytes(body, usize::MAX)
|
||||
.await
|
||||
.map_err(|error| crate::Error::Protocol(format!("cannot read request body: {error}")))?;
|
||||
Ok((parts, body))
|
||||
}
|
||||
|
||||
fn now() -> String {
|
||||
chrono::Utc::now().to_rfc3339_opts(chrono::SecondsFormat::Millis, true)
|
||||
}
|
||||
|
||||
fn proto(message: impl Message) -> Result<Response<Body>> {
|
||||
let body = message.encode_to_vec();
|
||||
let length = body.len();
|
||||
let mut response = Response::new(Body::from(body));
|
||||
response.headers_mut().insert(
|
||||
header::CONTENT_TYPE,
|
||||
axum::http::HeaderValue::from_static("application/proto"),
|
||||
);
|
||||
response.headers_mut().insert(
|
||||
header::CONTENT_LENGTH,
|
||||
length
|
||||
.to_string()
|
||||
.parse()
|
||||
.expect("body length is always a valid header value"),
|
||||
);
|
||||
Ok(response)
|
||||
}
|
||||
@@ -0,0 +1,487 @@
|
||||
//! Persists rules as markdown files with a JSON metadata sidecar.
|
||||
use std::{
|
||||
collections::BTreeMap,
|
||||
path::{Path, PathBuf},
|
||||
};
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use crate::{Error, Result};
|
||||
|
||||
const META_FILE: &str = "meta.json";
|
||||
const RULE_EXTENSION: &str = "md";
|
||||
pub const LOCAL_ID_PREFIX: &str = "local-";
|
||||
|
||||
/// 一条规则的完整视图:knowledge 来自 md 文件,其余字段来自 meta.json。
|
||||
#[derive(Clone, Debug, PartialEq)]
|
||||
pub struct RuleRecord {
|
||||
pub id: String,
|
||||
pub knowledge: String,
|
||||
pub title: String,
|
||||
pub created_at: String,
|
||||
pub is_generated: bool,
|
||||
pub git_origin: String,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum JournalOp {
|
||||
Add,
|
||||
Update,
|
||||
Remove,
|
||||
}
|
||||
|
||||
/// 离线期间未同步到上游的一次变更。
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
pub struct JournalEntry {
|
||||
pub op: JournalOp,
|
||||
pub id: String,
|
||||
}
|
||||
|
||||
#[derive(Default, Serialize, Deserialize)]
|
||||
struct Meta {
|
||||
#[serde(default)]
|
||||
rules: BTreeMap<String, RuleMeta>,
|
||||
#[serde(default)]
|
||||
journal: Vec<JournalEntry>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Default, Serialize, Deserialize)]
|
||||
struct RuleMeta {
|
||||
#[serde(default)]
|
||||
title: String,
|
||||
#[serde(default)]
|
||||
created_at: String,
|
||||
#[serde(default)]
|
||||
is_generated: bool,
|
||||
#[serde(default)]
|
||||
git_origin: String,
|
||||
}
|
||||
|
||||
/// md 文件为核心的规则存储;调用方需自行串行化并发访问。
|
||||
pub struct RuleStore {
|
||||
root: PathBuf,
|
||||
}
|
||||
|
||||
impl RuleStore {
|
||||
pub fn open(root: PathBuf) -> Result<Self> {
|
||||
std::fs::create_dir_all(&root)?;
|
||||
Ok(Self { root })
|
||||
}
|
||||
|
||||
pub fn list(&self) -> Result<Vec<RuleRecord>> {
|
||||
let meta = self.read_meta();
|
||||
let mut records = Vec::new();
|
||||
for entry in std::fs::read_dir(&self.root)? {
|
||||
let path = entry?.path();
|
||||
if path.extension().and_then(|value| value.to_str()) != Some(RULE_EXTENSION) {
|
||||
continue;
|
||||
}
|
||||
let Some(id) = path.file_stem().and_then(|value| value.to_str()) else {
|
||||
continue;
|
||||
};
|
||||
if validate_id(id).is_err() {
|
||||
continue;
|
||||
}
|
||||
let knowledge = std::fs::read_to_string(&path)?;
|
||||
records.push(assemble(id, knowledge, meta.rules.get(id), &path));
|
||||
}
|
||||
records.sort_by(|left, right| {
|
||||
timestamp(&right.created_at)
|
||||
.cmp(×tamp(&left.created_at))
|
||||
.then_with(|| left.id.cmp(&right.id))
|
||||
});
|
||||
Ok(records)
|
||||
}
|
||||
|
||||
pub fn get(&self, id: &str) -> Result<Option<RuleRecord>> {
|
||||
validate_id(id)?;
|
||||
let path = self.rule_path(id);
|
||||
let knowledge = match std::fs::read_to_string(&path) {
|
||||
Ok(knowledge) => knowledge,
|
||||
Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(None),
|
||||
Err(error) => return Err(error.into()),
|
||||
};
|
||||
let meta = self.read_meta();
|
||||
Ok(Some(assemble(id, knowledge, meta.rules.get(id), &path)))
|
||||
}
|
||||
|
||||
pub fn upsert(&self, record: &RuleRecord) -> Result<()> {
|
||||
validate_id(&record.id)?;
|
||||
write_atomic(&self.rule_path(&record.id), record.knowledge.as_bytes())?;
|
||||
let mut meta = self.read_meta();
|
||||
meta.rules.insert(record.id.clone(), rule_meta(record));
|
||||
self.write_meta(&meta)
|
||||
}
|
||||
|
||||
pub fn remove(&self, id: &str) -> Result<()> {
|
||||
validate_id(id)?;
|
||||
remove_file_if_exists(&self.rule_path(id))?;
|
||||
let mut meta = self.read_meta();
|
||||
if meta.rules.remove(id).is_some() {
|
||||
self.write_meta(&meta)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 离线新增的规则在上游落地后,把本地临时 id 换成上游分配的真实 id。
|
||||
pub fn promote(&self, old_id: &str, new_id: &str) -> Result<()> {
|
||||
validate_id(old_id)?;
|
||||
validate_id(new_id)?;
|
||||
let source = self.rule_path(old_id);
|
||||
let target = self.rule_path(new_id);
|
||||
#[cfg(windows)]
|
||||
remove_file_if_exists(&target)?;
|
||||
std::fs::rename(&source, &target)?;
|
||||
let mut meta = self.read_meta();
|
||||
if let Some(rule) = meta.rules.remove(old_id) {
|
||||
meta.rules.insert(new_id.into(), rule);
|
||||
}
|
||||
for entry in &mut meta.journal {
|
||||
if entry.id == old_id {
|
||||
entry.id = new_id.into();
|
||||
}
|
||||
}
|
||||
self.write_meta(&meta)
|
||||
}
|
||||
|
||||
/// 用上游的完整列表覆盖本地镜像;仅应在日志为空(已全部回放)时调用。
|
||||
pub fn replace_all(&self, records: &[RuleRecord]) -> Result<()> {
|
||||
let mut meta = self.read_meta();
|
||||
meta.rules.clear();
|
||||
for record in records {
|
||||
validate_id(&record.id)?;
|
||||
write_atomic(&self.rule_path(&record.id), record.knowledge.as_bytes())?;
|
||||
meta.rules.insert(record.id.clone(), rule_meta(record));
|
||||
}
|
||||
for entry in std::fs::read_dir(&self.root)? {
|
||||
let path = entry?.path();
|
||||
if path.extension().and_then(|value| value.to_str()) != Some(RULE_EXTENSION) {
|
||||
continue;
|
||||
}
|
||||
let keep = path
|
||||
.file_stem()
|
||||
.and_then(|value| value.to_str())
|
||||
.is_some_and(|id| meta.rules.contains_key(id));
|
||||
if !keep {
|
||||
remove_file_if_exists(&path)?;
|
||||
}
|
||||
}
|
||||
self.write_meta(&meta)
|
||||
}
|
||||
|
||||
pub fn journal_front(&self) -> Result<Option<JournalEntry>> {
|
||||
Ok(self.read_meta().journal.first().cloned())
|
||||
}
|
||||
|
||||
pub fn pop_journal(&self) -> Result<()> {
|
||||
let mut meta = self.read_meta();
|
||||
if !meta.journal.is_empty() {
|
||||
meta.journal.remove(0);
|
||||
self.write_meta(&meta)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn record_add(&self, id: &str) -> Result<()> {
|
||||
let mut meta = self.read_meta();
|
||||
meta.journal.push(JournalEntry {
|
||||
op: JournalOp::Add,
|
||||
id: id.into(),
|
||||
});
|
||||
self.write_meta(&meta)
|
||||
}
|
||||
|
||||
pub fn record_update(&self, id: &str) -> Result<()> {
|
||||
let mut meta = self.read_meta();
|
||||
if journal_contains(&meta.journal, id, JournalOp::Add) {
|
||||
// 回放 add 时会读取最新内容,无需单独的 update 日志。
|
||||
return Ok(());
|
||||
}
|
||||
let op = if id.starts_with(LOCAL_ID_PREFIX) {
|
||||
// 本地临时 id 没有对应的 add 日志(如镜像覆盖后的残留),按新增回放。
|
||||
JournalOp::Add
|
||||
} else {
|
||||
JournalOp::Update
|
||||
};
|
||||
if !journal_contains(&meta.journal, id, op) {
|
||||
meta.journal.push(JournalEntry { op, id: id.into() });
|
||||
self.write_meta(&meta)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn record_remove(&self, id: &str) -> Result<()> {
|
||||
let mut meta = self.read_meta();
|
||||
let never_synced = journal_contains(&meta.journal, id, JournalOp::Add);
|
||||
meta.journal.retain(|entry| entry.id != id);
|
||||
if !never_synced && !id.starts_with(LOCAL_ID_PREFIX) {
|
||||
meta.journal.push(JournalEntry {
|
||||
op: JournalOp::Remove,
|
||||
id: id.into(),
|
||||
});
|
||||
}
|
||||
self.write_meta(&meta)
|
||||
}
|
||||
|
||||
fn rule_path(&self, id: &str) -> PathBuf {
|
||||
self.root.join(format!("{id}.{RULE_EXTENSION}"))
|
||||
}
|
||||
|
||||
fn meta_path(&self) -> PathBuf {
|
||||
self.root.join(META_FILE)
|
||||
}
|
||||
|
||||
fn read_meta(&self) -> Meta {
|
||||
match std::fs::read(self.meta_path()) {
|
||||
Ok(bytes) => serde_json::from_slice(&bytes).unwrap_or_else(|error| {
|
||||
tracing::warn!(%error, "rules meta.json is corrupt; starting from empty metadata");
|
||||
Meta::default()
|
||||
}),
|
||||
Err(_) => Meta::default(),
|
||||
}
|
||||
}
|
||||
|
||||
fn write_meta(&self, meta: &Meta) -> Result<()> {
|
||||
write_atomic(&self.meta_path(), &serde_json::to_vec_pretty(meta)?)
|
||||
}
|
||||
}
|
||||
|
||||
fn assemble(id: &str, knowledge: String, meta: Option<&RuleMeta>, path: &Path) -> RuleRecord {
|
||||
match meta {
|
||||
Some(meta) => RuleRecord {
|
||||
id: id.into(),
|
||||
knowledge,
|
||||
title: meta.title.clone(),
|
||||
created_at: meta.created_at.clone(),
|
||||
is_generated: meta.is_generated,
|
||||
git_origin: meta.git_origin.clone(),
|
||||
},
|
||||
// 用户手放的 md 文件没有元数据,用文件名当标题、修改时间当创建时间。
|
||||
None => RuleRecord {
|
||||
id: id.into(),
|
||||
knowledge,
|
||||
title: id.into(),
|
||||
created_at: file_modified_at(path),
|
||||
is_generated: false,
|
||||
git_origin: String::new(),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
fn rule_meta(record: &RuleRecord) -> RuleMeta {
|
||||
RuleMeta {
|
||||
title: record.title.clone(),
|
||||
created_at: record.created_at.clone(),
|
||||
is_generated: record.is_generated,
|
||||
git_origin: record.git_origin.clone(),
|
||||
}
|
||||
}
|
||||
|
||||
fn journal_contains(journal: &[JournalEntry], id: &str, op: JournalOp) -> bool {
|
||||
journal.iter().any(|entry| entry.id == id && entry.op == op)
|
||||
}
|
||||
|
||||
fn timestamp(created_at: &str) -> i64 {
|
||||
chrono::DateTime::parse_from_rfc3339(created_at)
|
||||
.map(|time| time.timestamp_millis())
|
||||
.unwrap_or(0)
|
||||
}
|
||||
|
||||
fn file_modified_at(path: &Path) -> String {
|
||||
let modified = std::fs::metadata(path)
|
||||
.and_then(|meta| meta.modified())
|
||||
.unwrap_or_else(|_| std::time::SystemTime::now());
|
||||
chrono::DateTime::<chrono::Utc>::from(modified)
|
||||
.to_rfc3339_opts(chrono::SecondsFormat::Millis, true)
|
||||
}
|
||||
|
||||
fn validate_id(id: &str) -> Result<()> {
|
||||
if id.is_empty()
|
||||
|| id.len() > 128
|
||||
|| !id
|
||||
.bytes()
|
||||
.all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_'))
|
||||
{
|
||||
return Err(Error::Protocol(format!("invalid rule id: {id:?}")));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn remove_file_if_exists(path: &Path) -> Result<()> {
|
||||
match std::fs::remove_file(path) {
|
||||
Ok(()) => Ok(()),
|
||||
Err(error) if error.kind() == std::io::ErrorKind::NotFound => Ok(()),
|
||||
Err(error) => Err(error.into()),
|
||||
}
|
||||
}
|
||||
|
||||
fn write_atomic(path: &Path, bytes: &[u8]) -> Result<()> {
|
||||
use std::io::Write;
|
||||
let directory = path.parent().expect("rule path has a parent");
|
||||
let temporary = directory.join(format!(".{}.tmp", uuid::Uuid::new_v4()));
|
||||
let mut file = std::fs::File::create(&temporary)?;
|
||||
file.write_all(bytes)?;
|
||||
file.sync_all()?;
|
||||
drop(file);
|
||||
#[cfg(windows)]
|
||||
remove_file_if_exists(path)?;
|
||||
std::fs::rename(&temporary, path).inspect_err(|_| {
|
||||
let _ = std::fs::remove_file(&temporary);
|
||||
})?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn record(id: &str, knowledge: &str, created_at: &str) -> RuleRecord {
|
||||
RuleRecord {
|
||||
id: id.into(),
|
||||
knowledge: knowledge.into(),
|
||||
title: format!("title-{id}"),
|
||||
created_at: created_at.into(),
|
||||
is_generated: false,
|
||||
git_origin: String::new(),
|
||||
}
|
||||
}
|
||||
|
||||
fn journal(store: &RuleStore) -> Vec<JournalEntry> {
|
||||
store.read_meta().journal
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn upserts_lists_and_removes_rules() {
|
||||
let root = tempfile::tempdir().unwrap();
|
||||
let store = RuleStore::open(root.path().join("rules")).unwrap();
|
||||
store
|
||||
.upsert(&record("100", "older", "2026-01-01T00:00:00.000Z"))
|
||||
.unwrap();
|
||||
store
|
||||
.upsert(&record("200", "newer", "2026-02-01T00:00:00.000Z"))
|
||||
.unwrap();
|
||||
|
||||
let listed = store.list().unwrap();
|
||||
assert_eq!(
|
||||
listed
|
||||
.iter()
|
||||
.map(|rule| rule.id.as_str())
|
||||
.collect::<Vec<_>>(),
|
||||
["200", "100"],
|
||||
"list is sorted by created_at descending"
|
||||
);
|
||||
assert_eq!(listed[0].knowledge, "newer");
|
||||
assert_eq!(listed[0].title, "title-200");
|
||||
|
||||
store.remove("200").unwrap();
|
||||
assert!(store.get("200").unwrap().is_none());
|
||||
assert_eq!(store.list().unwrap().len(), 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_path_traversal_ids() {
|
||||
let root = tempfile::tempdir().unwrap();
|
||||
let store = RuleStore::open(root.path().join("rules")).unwrap();
|
||||
assert!(store.get("../escape").is_err());
|
||||
assert!(store.get("a/b").is_err());
|
||||
assert!(store.get("").is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn compacts_offline_journal() {
|
||||
let root = tempfile::tempdir().unwrap();
|
||||
let store = RuleStore::open(root.path().join("rules")).unwrap();
|
||||
|
||||
// 离线新增后再更新:回放 add 即可携带最新内容,不产生 update 日志。
|
||||
store
|
||||
.upsert(&record("local-a", "v1", "2026-01-01T00:00:00.000Z"))
|
||||
.unwrap();
|
||||
store.record_add("local-a").unwrap();
|
||||
store.record_update("local-a").unwrap();
|
||||
assert_eq!(
|
||||
journal(&store),
|
||||
vec![JournalEntry {
|
||||
op: JournalOp::Add,
|
||||
id: "local-a".into()
|
||||
}]
|
||||
);
|
||||
|
||||
// 离线新增后又删除:上游从未见过它,日志清空。
|
||||
store.record_remove("local-a").unwrap();
|
||||
assert!(journal(&store).is_empty());
|
||||
|
||||
// 更新上游已有规则:多次更新合并为一条;删除后 update 日志被顶替。
|
||||
store.record_update("42").unwrap();
|
||||
store.record_update("42").unwrap();
|
||||
assert_eq!(
|
||||
journal(&store),
|
||||
vec![JournalEntry {
|
||||
op: JournalOp::Update,
|
||||
id: "42".into()
|
||||
}]
|
||||
);
|
||||
store.record_remove("42").unwrap();
|
||||
assert_eq!(
|
||||
journal(&store),
|
||||
vec![JournalEntry {
|
||||
op: JournalOp::Remove,
|
||||
id: "42".into()
|
||||
}]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn promote_renames_rule_and_journal_ids() {
|
||||
let root = tempfile::tempdir().unwrap();
|
||||
let store = RuleStore::open(root.path().join("rules")).unwrap();
|
||||
store
|
||||
.upsert(&record("local-a", "content", "2026-01-01T00:00:00.000Z"))
|
||||
.unwrap();
|
||||
store.record_add("local-a").unwrap();
|
||||
|
||||
store.promote("local-a", "17353272").unwrap();
|
||||
|
||||
assert!(store.get("local-a").unwrap().is_none());
|
||||
let promoted = store.get("17353272").unwrap().unwrap();
|
||||
assert_eq!(promoted.knowledge, "content");
|
||||
assert_eq!(promoted.title, "title-local-a");
|
||||
assert_eq!(journal(&store)[0].id, "17353272");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn replace_all_mirrors_upstream_state() {
|
||||
let root = tempfile::tempdir().unwrap();
|
||||
let store = RuleStore::open(root.path().join("rules")).unwrap();
|
||||
store
|
||||
.upsert(&record("stale", "gone soon", "2026-01-01T00:00:00.000Z"))
|
||||
.unwrap();
|
||||
|
||||
store
|
||||
.replace_all(&[record(
|
||||
"17353272",
|
||||
"from upstream",
|
||||
"2026-02-01T00:00:00.000Z",
|
||||
)])
|
||||
.unwrap();
|
||||
|
||||
let listed = store.list().unwrap();
|
||||
assert_eq!(listed.len(), 1);
|
||||
assert_eq!(listed[0].id, "17353272");
|
||||
assert_eq!(listed[0].knowledge, "from upstream");
|
||||
assert!(store.get("stale").unwrap().is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn lists_hand_written_markdown_without_metadata() {
|
||||
let root = tempfile::tempdir().unwrap();
|
||||
let store = RuleStore::open(root.path().join("rules")).unwrap();
|
||||
std::fs::write(root.path().join("rules/manual_rule.md"), "hand written").unwrap();
|
||||
|
||||
let listed = store.list().unwrap();
|
||||
assert_eq!(listed.len(), 1);
|
||||
assert_eq!(listed[0].id, "manual_rule");
|
||||
assert_eq!(listed[0].title, "manual_rule");
|
||||
assert_eq!(listed[0].knowledge, "hand written");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,184 @@
|
||||
//! Replays the offline journal to upstream and mirrors upstream list state.
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http::{header, HeaderMap, HeaderValue, Method, Request},
|
||||
};
|
||||
use prost::Message;
|
||||
|
||||
use crate::{api::cursor::proxy, cursor::protocol::connect, Result};
|
||||
|
||||
use super::{
|
||||
store::{JournalOp, RuleRecord, RuleStore},
|
||||
KnowledgeBaseAddRequest, KnowledgeBaseAddResponse, KnowledgeBaseListItem,
|
||||
KnowledgeBaseRemoveRequest, KnowledgeBaseRemoveResponse, KnowledgeBaseUpdateRequest,
|
||||
KnowledgeBaseUpdateResponse,
|
||||
};
|
||||
|
||||
const ADD_PATH: &str = "/aiserver.v1.AiService/KnowledgeBaseAdd";
|
||||
const UPDATE_PATH: &str = "/aiserver.v1.AiService/KnowledgeBaseUpdate";
|
||||
const REMOVE_PATH: &str = "/aiserver.v1.AiService/KnowledgeBaseRemove";
|
||||
|
||||
/// 逐条把离线日志推送到上游。返回 true 表示日志已清空(上游可用),
|
||||
/// false 表示上游不可达,剩余日志保留、调用方应降级到本地。
|
||||
pub async fn replay(
|
||||
upstream: &proxy::CursorProxy,
|
||||
headers: &HeaderMap,
|
||||
store: &RuleStore,
|
||||
) -> Result<bool> {
|
||||
while let Some(entry) = store.journal_front()? {
|
||||
let advanced = match entry.op {
|
||||
JournalOp::Add => replay_add(upstream, headers, store, &entry.id).await?,
|
||||
JournalOp::Update => replay_update(upstream, headers, store, &entry.id).await?,
|
||||
JournalOp::Remove => replay_remove(upstream, headers, store, &entry.id).await?,
|
||||
};
|
||||
if !advanced {
|
||||
return Ok(false);
|
||||
}
|
||||
}
|
||||
Ok(true)
|
||||
}
|
||||
|
||||
/// 用上游返回的完整列表覆盖本地镜像。仅应在日志已清空时调用。
|
||||
pub fn mirror(store: &RuleStore, items: Vec<KnowledgeBaseListItem>) -> Result<()> {
|
||||
let records = items
|
||||
.into_iter()
|
||||
.map(|item| RuleRecord {
|
||||
id: item.id,
|
||||
knowledge: item.knowledge,
|
||||
title: item.title,
|
||||
created_at: item.created_at,
|
||||
is_generated: item.is_generated,
|
||||
git_origin: String::new(),
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
store.replace_all(&records)
|
||||
}
|
||||
|
||||
async fn replay_add(
|
||||
upstream: &proxy::CursorProxy,
|
||||
headers: &HeaderMap,
|
||||
store: &RuleStore,
|
||||
id: &str,
|
||||
) -> Result<bool> {
|
||||
let Some(record) = store.get(id)? else {
|
||||
// 规则文件已不在(被手动删除等),日志作废。
|
||||
store.pop_journal()?;
|
||||
return Ok(true);
|
||||
};
|
||||
let message = KnowledgeBaseAddRequest {
|
||||
knowledge: record.knowledge,
|
||||
title: record.title,
|
||||
git_origin: record.git_origin,
|
||||
composer_id: None,
|
||||
};
|
||||
let Some(body) = send(upstream, headers, ADD_PATH, &message).await else {
|
||||
return Ok(false);
|
||||
};
|
||||
let Ok(reply) = connect::decode_unary::<KnowledgeBaseAddResponse>(&body) else {
|
||||
return Ok(false);
|
||||
};
|
||||
if !reply.success || reply.id.is_empty() {
|
||||
tracing::warn!(
|
||||
id,
|
||||
"rules upstream declined replayed add; dropping journal entry"
|
||||
);
|
||||
store.pop_journal()?;
|
||||
return Ok(true);
|
||||
}
|
||||
store.promote(id, &reply.id)?;
|
||||
store.pop_journal()?;
|
||||
tracing::info!(
|
||||
local_id = id,
|
||||
upstream_id = reply.id,
|
||||
"replayed offline rule add to upstream"
|
||||
);
|
||||
Ok(true)
|
||||
}
|
||||
|
||||
async fn replay_update(
|
||||
upstream: &proxy::CursorProxy,
|
||||
headers: &HeaderMap,
|
||||
store: &RuleStore,
|
||||
id: &str,
|
||||
) -> Result<bool> {
|
||||
let Some(record) = store.get(id)? else {
|
||||
store.pop_journal()?;
|
||||
return Ok(true);
|
||||
};
|
||||
let message = KnowledgeBaseUpdateRequest {
|
||||
id: id.into(),
|
||||
knowledge: record.knowledge,
|
||||
title: record.title,
|
||||
};
|
||||
let Some(body) = send(upstream, headers, UPDATE_PATH, &message).await else {
|
||||
return Ok(false);
|
||||
};
|
||||
let Ok(reply) = connect::decode_unary::<KnowledgeBaseUpdateResponse>(&body) else {
|
||||
return Ok(false);
|
||||
};
|
||||
if !reply.success {
|
||||
tracing::warn!(
|
||||
id,
|
||||
"rules upstream declined replayed update; dropping journal entry"
|
||||
);
|
||||
}
|
||||
store.pop_journal()?;
|
||||
Ok(true)
|
||||
}
|
||||
|
||||
async fn replay_remove(
|
||||
upstream: &proxy::CursorProxy,
|
||||
headers: &HeaderMap,
|
||||
store: &RuleStore,
|
||||
id: &str,
|
||||
) -> Result<bool> {
|
||||
let message = KnowledgeBaseRemoveRequest { id: id.into() };
|
||||
let Some(body) = send(upstream, headers, REMOVE_PATH, &message).await else {
|
||||
return Ok(false);
|
||||
};
|
||||
let Ok(reply) = connect::decode_unary::<KnowledgeBaseRemoveResponse>(&body) else {
|
||||
return Ok(false);
|
||||
};
|
||||
if !reply.success {
|
||||
tracing::warn!(
|
||||
id,
|
||||
"rules upstream declined replayed remove; dropping journal entry"
|
||||
);
|
||||
}
|
||||
store.pop_journal()?;
|
||||
Ok(true)
|
||||
}
|
||||
|
||||
/// 以当前请求的头为模板向上游发起一次 unary RPC。
|
||||
/// 成功(2xx)返回响应体;不可达或被拒绝返回 None,由调用方保留日志。
|
||||
async fn send(
|
||||
upstream: &proxy::CursorProxy,
|
||||
template: &HeaderMap,
|
||||
path: &str,
|
||||
message: &impl Message,
|
||||
) -> Option<Bytes> {
|
||||
let mut headers = template.clone();
|
||||
// 模板里的上游 URL 头指向原始 RPC 路径,必须移除才能命中回放路径。
|
||||
headers.remove(proxy::UPSTREAM_URL_HEADER);
|
||||
headers.remove(header::CONTENT_LENGTH);
|
||||
headers.insert(
|
||||
header::CONTENT_TYPE,
|
||||
HeaderValue::from_static("application/proto"),
|
||||
);
|
||||
let mut request = Request::new(Body::from(message.encode_to_vec()));
|
||||
*request.method_mut() = Method::POST;
|
||||
*request.uri_mut() = path.parse().expect("replay path is a valid URI");
|
||||
*request.headers_mut() = headers;
|
||||
|
||||
match proxy::forward_buffered(upstream, request).await {
|
||||
Ok(response) if response.status.is_success() => Some(response.body),
|
||||
Ok(response) => {
|
||||
tracing::warn!(path, status = %response.status, "rules journal replay rejected by upstream");
|
||||
None
|
||||
}
|
||||
Err(error) => {
|
||||
tracing::warn!(path, %error, "rules journal replay cannot reach upstream");
|
||||
None
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -4,6 +4,7 @@ pub mod account;
|
||||
pub mod analytics;
|
||||
pub mod blob_sync;
|
||||
pub mod context_sync;
|
||||
pub mod knowledge;
|
||||
pub mod model_catalog;
|
||||
pub mod observability;
|
||||
pub mod tab;
|
||||
|
||||
@@ -10,7 +10,7 @@ use prost::Message;
|
||||
use crate::{
|
||||
api::cursor::proxy::{self, CursorProxy},
|
||||
cursor::{protocol::proto::agent::v1 as agent, transport::TransportRegistry},
|
||||
model::{format_token_count, parse_token_count, ModelConfig, ModelType},
|
||||
model::{format_token_count, parse_token_count, ModelConfig},
|
||||
plugin::PluginModelDescriptor,
|
||||
Error, Result,
|
||||
};
|
||||
@@ -378,16 +378,25 @@ fn available_model(model: &ModelConfig) -> AvailableModel {
|
||||
display_name: "Cursor".into(),
|
||||
}),
|
||||
model_picker_badges: vec![ModelPickerBadge {
|
||||
label: match model.model_type {
|
||||
ModelType::OpenAi => "OpenAI".into(),
|
||||
ModelType::Anthropic => "Anthropic".into(),
|
||||
},
|
||||
label: model
|
||||
.group_name
|
||||
.clone()
|
||||
.unwrap_or_else(|| provider_host(&model.base_url)),
|
||||
variant: 1,
|
||||
dismiss_on_selection: false,
|
||||
}],
|
||||
}
|
||||
}
|
||||
|
||||
/// 徽章回退标签:base_url 的主机名。入库时已校验为带主机的 HTTP(S) URL,
|
||||
/// 解析失败仅是理论分支,此时原样返回 base_url。
|
||||
fn provider_host(base_url: &str) -> String {
|
||||
reqwest::Url::parse(base_url.trim())
|
||||
.ok()
|
||||
.and_then(|url| url.host_str().map(str::to_lowercase))
|
||||
.unwrap_or_else(|| base_url.trim().into())
|
||||
}
|
||||
|
||||
fn model_parameters(
|
||||
contexts: &[(String, String)],
|
||||
thinking: bool,
|
||||
@@ -571,14 +580,9 @@ fn available_plugin_model(model: &PluginModelDescriptor) -> AvailableModel {
|
||||
let tooltip = TooltipData {
|
||||
markdown_content: model.description.clone(),
|
||||
};
|
||||
let contexts = context_options(model.context_window_tokens);
|
||||
let variants = model_variants(
|
||||
&model.id,
|
||||
&model.display_name,
|
||||
&tooltip,
|
||||
&contexts,
|
||||
model.thinking,
|
||||
);
|
||||
// Effort 与上下文档位由宿主统一提供,与内置模型一致;插件不再声明这两项。
|
||||
let contexts = context_options(None);
|
||||
let variants = model_variants(&model.id, &model.display_name, &tooltip, &contexts, true);
|
||||
let legacy_slugs = variants
|
||||
.iter()
|
||||
.filter_map(|variant| variant.legacy_slug.clone())
|
||||
@@ -589,7 +593,7 @@ fn available_plugin_model(model: &PluginModelDescriptor) -> AvailableModel {
|
||||
supports_agent: Some(true),
|
||||
degradation_status: Some(0),
|
||||
tooltip_data: Some(tooltip.clone()),
|
||||
supports_thinking: Some(model.thinking),
|
||||
supports_thinking: Some(true),
|
||||
supports_images: Some(model.images),
|
||||
supports_max_mode: Some(false),
|
||||
client_display_name: Some(model.display_name.clone()),
|
||||
@@ -601,7 +605,7 @@ fn available_plugin_model(model: &PluginModelDescriptor) -> AvailableModel {
|
||||
inputbox_short_model_name: Some(model.display_name.clone()),
|
||||
supports_sandboxing: Some(true),
|
||||
supports_cmd_k: Some(false),
|
||||
parameter_definitions: model_parameters(&contexts, model.thinking),
|
||||
parameter_definitions: model_parameters(&contexts, true),
|
||||
variants,
|
||||
legacy_slugs,
|
||||
named_model_section_index: Some(1),
|
||||
@@ -611,7 +615,7 @@ fn available_plugin_model(model: &PluginModelDescriptor) -> AvailableModel {
|
||||
display_name: model.provider_type.clone(),
|
||||
}),
|
||||
model_picker_badges: vec![ModelPickerBadge {
|
||||
label: model.provider_type.clone(),
|
||||
label: model.plugin_name.clone(),
|
||||
variant: 1,
|
||||
dismiss_on_selection: false,
|
||||
}],
|
||||
@@ -624,7 +628,7 @@ fn usable_plugin_model(model: &PluginModelDescriptor) -> agent::ModelDetails {
|
||||
display_model_id: model.id.clone(),
|
||||
display_name: model.display_name.clone(),
|
||||
display_name_short: model.display_name.clone(),
|
||||
thinking_details: model.thinking.then(agent::ThinkingDetails::default),
|
||||
thinking_details: Some(agent::ThinkingDetails::default()),
|
||||
..Default::default()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,225 +0,0 @@
|
||||
//! Records Cursor request traces and artifacts.
|
||||
use std::{
|
||||
sync::{
|
||||
atomic::{AtomicBool, Ordering},
|
||||
Arc,
|
||||
},
|
||||
time::{Duration, Instant},
|
||||
};
|
||||
|
||||
use tokio::sync::Mutex;
|
||||
|
||||
use crate::store::{BlobId, BufferedCursorTraceChunk, Store};
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct CursorTraceRecorder {
|
||||
store: Store,
|
||||
request_id: String,
|
||||
chunks: Arc<Mutex<TraceChunkBuffer>>,
|
||||
finished: Arc<AtomicBool>,
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct TraceChunkBuffer {
|
||||
chunks: Vec<BufferedCursorTraceChunk>,
|
||||
bytes: usize,
|
||||
first_chunk_at: Option<Instant>,
|
||||
generation: u64,
|
||||
}
|
||||
|
||||
const MAX_BUFFERED_CHUNKS: usize = 32;
|
||||
const MAX_BUFFERED_BYTES: usize = 256 * 1024;
|
||||
const MAX_BUFFER_AGE: Duration = Duration::from_millis(50);
|
||||
|
||||
impl CursorTraceRecorder {
|
||||
pub async fn begin(
|
||||
store: Store,
|
||||
request_id: &str,
|
||||
conversation_id: Option<&str>,
|
||||
route: &str,
|
||||
model_id: Option<&str>,
|
||||
) -> Option<Self> {
|
||||
match store
|
||||
.start_cursor_trace_if_detailed(request_id, conversation_id, route, model_id)
|
||||
.await
|
||||
{
|
||||
Ok(true) => Some(Self {
|
||||
store,
|
||||
request_id: request_id.into(),
|
||||
chunks: Arc::new(Mutex::new(TraceChunkBuffer::default())),
|
||||
finished: Arc::new(AtomicBool::new(false)),
|
||||
}),
|
||||
Ok(false) => None,
|
||||
Err(error) => {
|
||||
tracing::warn!(request_id, %error, "failed to start Cursor trace");
|
||||
None
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn resume(store: Store, request_id: &str) -> Option<Self> {
|
||||
match store.cursor_trace_exists(request_id).await {
|
||||
Ok(true) => Some(Self {
|
||||
store,
|
||||
request_id: request_id.into(),
|
||||
chunks: Arc::new(Mutex::new(TraceChunkBuffer::default())),
|
||||
finished: Arc::new(AtomicBool::new(false)),
|
||||
}),
|
||||
Ok(false) => None,
|
||||
Err(error) => {
|
||||
tracing::warn!(request_id, %error, "failed to resume Cursor trace");
|
||||
None
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn request_id(&self) -> &str {
|
||||
&self.request_id
|
||||
}
|
||||
|
||||
pub async fn request(&self, artifact_type: &str, data: &[u8], metadata: serde_json::Value) {
|
||||
if let Err(error) = self
|
||||
.store
|
||||
.append_cursor_trace_artifact(
|
||||
&self.request_id,
|
||||
artifact_type,
|
||||
"cursor_client",
|
||||
data,
|
||||
&metadata,
|
||||
)
|
||||
.await
|
||||
{
|
||||
tracing::warn!(request_id = self.request_id, %error, "failed to record Cursor request artifact");
|
||||
return;
|
||||
}
|
||||
if let Err(error) = self
|
||||
.store
|
||||
.add_cursor_trace_request_bytes(&self.request_id, data.len())
|
||||
.await
|
||||
{
|
||||
tracing::warn!(request_id = self.request_id, %error, "failed to update Cursor request trace size");
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn artifact(
|
||||
&self,
|
||||
artifact_type: &str,
|
||||
source: &str,
|
||||
data: &[u8],
|
||||
metadata: serde_json::Value,
|
||||
) {
|
||||
if let Err(error) = self
|
||||
.store
|
||||
.append_cursor_trace_artifact(&self.request_id, artifact_type, source, data, &metadata)
|
||||
.await
|
||||
{
|
||||
tracing::warn!(request_id = self.request_id, %error, artifact_type, "failed to record Cursor trace artifact");
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn linked_blob(
|
||||
&self,
|
||||
artifact_type: &str,
|
||||
source: &str,
|
||||
blob_id: &BlobId,
|
||||
metadata: serde_json::Value,
|
||||
) {
|
||||
if let Err(error) = self
|
||||
.store
|
||||
.link_cursor_trace_artifact(&self.request_id, artifact_type, source, blob_id, &metadata)
|
||||
.await
|
||||
{
|
||||
tracing::warn!(request_id = self.request_id, %error, artifact_type, "failed to link Cursor trace Blob");
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn response_started(&self, status: u16) {
|
||||
if let Err(error) = self
|
||||
.store
|
||||
.start_cursor_trace_response(&self.request_id, status)
|
||||
.await
|
||||
{
|
||||
tracing::warn!(request_id = self.request_id, %error, "failed to start Cursor response trace");
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn response_chunk(&self, source: &str, data: &[u8]) {
|
||||
let mut buffer = self.chunks.lock().await;
|
||||
if self.finished.load(Ordering::Acquire) {
|
||||
return;
|
||||
}
|
||||
let schedule_flush = if buffer.chunks.is_empty() {
|
||||
buffer.generation = buffer.generation.wrapping_add(1);
|
||||
buffer.first_chunk_at = Some(Instant::now());
|
||||
Some(buffer.generation)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
buffer.bytes += data.len();
|
||||
buffer
|
||||
.chunks
|
||||
.push(BufferedCursorTraceChunk::new(source, data));
|
||||
let expired = buffer
|
||||
.first_chunk_at
|
||||
.is_some_and(|started| started.elapsed() >= MAX_BUFFER_AGE);
|
||||
if buffer.chunks.len() >= MAX_BUFFERED_CHUNKS
|
||||
|| buffer.bytes >= MAX_BUFFERED_BYTES
|
||||
|| expired
|
||||
{
|
||||
if let Err(error) = self.flush_locked(&mut buffer).await {
|
||||
tracing::warn!(request_id = self.request_id, %error, "failed to record Cursor response chunk");
|
||||
}
|
||||
}
|
||||
drop(buffer);
|
||||
if let Some(generation) = schedule_flush {
|
||||
let recorder = self.clone();
|
||||
tokio::spawn(async move {
|
||||
tokio::time::sleep(MAX_BUFFER_AGE).await;
|
||||
let mut buffer = recorder.chunks.lock().await;
|
||||
if buffer.generation == generation {
|
||||
if let Err(error) = recorder.flush_locked(&mut buffer).await {
|
||||
tracing::warn!(request_id = recorder.request_id, %error, "failed to flush Cursor response chunks");
|
||||
}
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn finish(&self, error: Option<&str>) {
|
||||
if self.finished.swap(true, Ordering::AcqRel) {
|
||||
return;
|
||||
}
|
||||
let mut buffer = self.chunks.lock().await;
|
||||
if let Err(store_error) = self.flush_locked(&mut buffer).await {
|
||||
tracing::warn!(request_id = self.request_id, %store_error, "failed to flush Cursor response chunks");
|
||||
}
|
||||
drop(buffer);
|
||||
if let Err(store_error) = self
|
||||
.store
|
||||
.finish_cursor_trace(&self.request_id, error)
|
||||
.await
|
||||
{
|
||||
tracing::warn!(request_id = self.request_id, %store_error, "failed to finish Cursor trace");
|
||||
}
|
||||
}
|
||||
|
||||
async fn flush_locked(&self, buffer: &mut TraceChunkBuffer) -> crate::Result<()> {
|
||||
if buffer.chunks.is_empty() {
|
||||
return Ok(());
|
||||
}
|
||||
let chunks = std::mem::take(&mut buffer.chunks);
|
||||
buffer.bytes = 0;
|
||||
buffer.first_chunk_at = None;
|
||||
if let Err(error) = self
|
||||
.store
|
||||
.add_cursor_trace_response_chunks(&self.request_id, &chunks)
|
||||
.await
|
||||
{
|
||||
buffer.bytes = chunks.iter().map(|chunk| chunk.data.len()).sum();
|
||||
buffer.first_chunk_at = Some(Instant::now());
|
||||
buffer.chunks = chunks;
|
||||
return Err(error);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,71 @@
|
||||
use std::sync::{atomic::AtomicU8, Arc};
|
||||
|
||||
use bytes::Bytes;
|
||||
|
||||
use crate::store::BlobId;
|
||||
|
||||
pub(super) const TRACE_UNKNOWN: u8 = 0;
|
||||
pub(super) const TRACE_ACTIVE: u8 = 1;
|
||||
pub(super) const TRACE_DISABLED: u8 = 2;
|
||||
|
||||
pub(super) enum TraceEvent {
|
||||
Begin {
|
||||
request_id: String,
|
||||
activation: Arc<AtomicU8>,
|
||||
conversation_id: Option<String>,
|
||||
route: String,
|
||||
model_id: Option<String>,
|
||||
},
|
||||
Resume {
|
||||
request_id: String,
|
||||
activation: Arc<AtomicU8>,
|
||||
},
|
||||
Request {
|
||||
request_id: String,
|
||||
artifact_type: String,
|
||||
data: Bytes,
|
||||
metadata: serde_json::Value,
|
||||
},
|
||||
Artifact {
|
||||
request_id: String,
|
||||
artifact_type: String,
|
||||
source: String,
|
||||
data: Bytes,
|
||||
metadata: serde_json::Value,
|
||||
},
|
||||
LinkedBlob {
|
||||
request_id: String,
|
||||
artifact_type: String,
|
||||
source: String,
|
||||
blob_id: BlobId,
|
||||
metadata: serde_json::Value,
|
||||
},
|
||||
ResponseStarted {
|
||||
request_id: String,
|
||||
status: u16,
|
||||
},
|
||||
ResponseChunk {
|
||||
request_id: String,
|
||||
source: String,
|
||||
data: Bytes,
|
||||
},
|
||||
Finish {
|
||||
request_id: String,
|
||||
error: Option<String>,
|
||||
},
|
||||
}
|
||||
|
||||
impl TraceEvent {
|
||||
pub(super) fn request_id(&self) -> &str {
|
||||
match self {
|
||||
Self::Begin { request_id, .. }
|
||||
| Self::Resume { request_id, .. }
|
||||
| Self::Request { request_id, .. }
|
||||
| Self::Artifact { request_id, .. }
|
||||
| Self::LinkedBlob { request_id, .. }
|
||||
| Self::ResponseStarted { request_id, .. }
|
||||
| Self::ResponseChunk { request_id, .. }
|
||||
| Self::Finish { request_id, .. } => request_id,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,157 @@
|
||||
//! Records Cursor request traces without blocking request or runtime paths.
|
||||
|
||||
mod event;
|
||||
mod worker;
|
||||
|
||||
use std::sync::{
|
||||
atomic::{AtomicBool, AtomicU8, Ordering},
|
||||
Arc,
|
||||
};
|
||||
|
||||
use bytes::Bytes;
|
||||
use tokio::sync::mpsc;
|
||||
|
||||
use crate::store::{BlobId, Store};
|
||||
|
||||
use event::{TraceEvent, TRACE_DISABLED, TRACE_UNKNOWN};
|
||||
|
||||
const TRACE_QUEUE_CAPACITY: usize = 512;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct CursorTraceService {
|
||||
sender: mpsc::Sender<TraceEvent>,
|
||||
}
|
||||
|
||||
impl CursorTraceService {
|
||||
pub fn new(store: Store) -> Self {
|
||||
let (sender, receiver) = mpsc::channel(TRACE_QUEUE_CAPACITY);
|
||||
tokio::spawn(worker::run(store, receiver));
|
||||
Self { sender }
|
||||
}
|
||||
|
||||
pub fn recorder(&self, request_id: &str) -> CursorTraceRecorder {
|
||||
CursorTraceRecorder {
|
||||
request_id: Arc::from(request_id),
|
||||
sender: self.sender.clone(),
|
||||
finished: Arc::new(AtomicBool::new(false)),
|
||||
activation: Arc::new(AtomicU8::new(TRACE_UNKNOWN)),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct CursorTraceRecorder {
|
||||
request_id: Arc<str>,
|
||||
sender: mpsc::Sender<TraceEvent>,
|
||||
finished: Arc<AtomicBool>,
|
||||
activation: Arc<AtomicU8>,
|
||||
}
|
||||
|
||||
impl CursorTraceRecorder {
|
||||
pub fn request_id(&self) -> &str {
|
||||
&self.request_id
|
||||
}
|
||||
|
||||
pub fn begin(&self, conversation_id: Option<&str>, route: &str, model_id: Option<&str>) {
|
||||
self.send_control(TraceEvent::Begin {
|
||||
request_id: self.request_id.to_string(),
|
||||
activation: self.activation.clone(),
|
||||
conversation_id: conversation_id.map(str::to_owned),
|
||||
route: route.to_owned(),
|
||||
model_id: model_id.map(str::to_owned),
|
||||
});
|
||||
}
|
||||
|
||||
pub fn resume(&self) {
|
||||
self.send_control(TraceEvent::Resume {
|
||||
request_id: self.request_id.to_string(),
|
||||
activation: self.activation.clone(),
|
||||
});
|
||||
}
|
||||
|
||||
pub fn request(&self, artifact_type: &str, data: Bytes, metadata: serde_json::Value) {
|
||||
self.send(TraceEvent::Request {
|
||||
request_id: self.request_id.to_string(),
|
||||
artifact_type: artifact_type.to_owned(),
|
||||
data,
|
||||
metadata,
|
||||
});
|
||||
}
|
||||
|
||||
pub fn artifact(
|
||||
&self,
|
||||
artifact_type: &str,
|
||||
source: &str,
|
||||
data: &[u8],
|
||||
metadata: serde_json::Value,
|
||||
) {
|
||||
self.send(TraceEvent::Artifact {
|
||||
request_id: self.request_id.to_string(),
|
||||
artifact_type: artifact_type.to_owned(),
|
||||
source: source.to_owned(),
|
||||
data: Bytes::copy_from_slice(data),
|
||||
metadata,
|
||||
});
|
||||
}
|
||||
|
||||
pub fn linked_blob(
|
||||
&self,
|
||||
artifact_type: &str,
|
||||
source: &str,
|
||||
blob_id: &BlobId,
|
||||
metadata: serde_json::Value,
|
||||
) {
|
||||
self.send(TraceEvent::LinkedBlob {
|
||||
request_id: self.request_id.to_string(),
|
||||
artifact_type: artifact_type.to_owned(),
|
||||
source: source.to_owned(),
|
||||
blob_id: blob_id.clone(),
|
||||
metadata,
|
||||
});
|
||||
}
|
||||
|
||||
pub fn response_started(&self, status: u16) {
|
||||
self.send(TraceEvent::ResponseStarted {
|
||||
request_id: self.request_id.to_string(),
|
||||
status,
|
||||
});
|
||||
}
|
||||
|
||||
pub fn response_chunk(&self, source: &str, data: Bytes) {
|
||||
if self.finished.load(Ordering::Acquire) {
|
||||
return;
|
||||
}
|
||||
self.send(TraceEvent::ResponseChunk {
|
||||
request_id: self.request_id.to_string(),
|
||||
source: source.to_owned(),
|
||||
data,
|
||||
});
|
||||
}
|
||||
|
||||
pub fn finish(&self, error: Option<&str>) {
|
||||
if self.finished.swap(true, Ordering::AcqRel) {
|
||||
return;
|
||||
}
|
||||
self.send_control(TraceEvent::Finish {
|
||||
request_id: self.request_id.to_string(),
|
||||
error: error.map(str::to_owned),
|
||||
});
|
||||
}
|
||||
|
||||
fn send(&self, event: TraceEvent) {
|
||||
if self.activation.load(Ordering::Acquire) == TRACE_DISABLED {
|
||||
return;
|
||||
}
|
||||
self.send_control(event);
|
||||
}
|
||||
|
||||
fn send_control(&self, event: TraceEvent) {
|
||||
if let Err(error) = self.sender.try_send(event) {
|
||||
tracing::warn!(
|
||||
request_id = %self.request_id,
|
||||
%error,
|
||||
"dropping Cursor trace event"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,349 @@
|
||||
use std::{
|
||||
collections::{BTreeMap, HashMap},
|
||||
sync::atomic::Ordering,
|
||||
time::Duration,
|
||||
};
|
||||
|
||||
use tokio::sync::mpsc;
|
||||
|
||||
use crate::store::{BufferedCursorTraceChunk, Store};
|
||||
|
||||
use super::event::{TraceEvent, TRACE_ACTIVE, TRACE_DISABLED};
|
||||
|
||||
const MAX_BUFFERED_CHUNKS: usize = 32;
|
||||
const MAX_BUFFERED_BYTES: usize = 256 * 1024;
|
||||
const FLUSH_INTERVAL: Duration = Duration::from_millis(50);
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
enum TraceState {
|
||||
Active,
|
||||
Disabled,
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct ResponseBuffer {
|
||||
chunks: Vec<BufferedCursorTraceChunk>,
|
||||
bytes: usize,
|
||||
}
|
||||
|
||||
struct BufferedRequest {
|
||||
artifact_type: String,
|
||||
data: bytes::Bytes,
|
||||
metadata: serde_json::Value,
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct RequestOrder {
|
||||
next: i64,
|
||||
pending: BTreeMap<i64, Vec<BufferedRequest>>,
|
||||
}
|
||||
|
||||
pub(super) async fn run(store: Store, mut receiver: mpsc::Receiver<TraceEvent>) {
|
||||
let mut states = HashMap::<String, TraceState>::new();
|
||||
let mut buffers = HashMap::<String, ResponseBuffer>::new();
|
||||
let mut request_orders = HashMap::<String, RequestOrder>::new();
|
||||
let mut interval = tokio::time::interval(FLUSH_INTERVAL);
|
||||
interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
|
||||
|
||||
loop {
|
||||
tokio::select! {
|
||||
event = receiver.recv() => {
|
||||
let Some(event) = event else {
|
||||
flush_all(&store, &mut buffers).await;
|
||||
return;
|
||||
};
|
||||
process(&store, &mut states, &mut buffers, &mut request_orders, event).await;
|
||||
}
|
||||
_ = interval.tick() => flush_all(&store, &mut buffers).await,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn process(
|
||||
store: &Store,
|
||||
states: &mut HashMap<String, TraceState>,
|
||||
buffers: &mut HashMap<String, ResponseBuffer>,
|
||||
request_orders: &mut HashMap<String, RequestOrder>,
|
||||
event: TraceEvent,
|
||||
) {
|
||||
let request_id = event.request_id().to_owned();
|
||||
let finishes_trace = matches!(&event, TraceEvent::Finish { .. });
|
||||
match event {
|
||||
TraceEvent::Begin {
|
||||
request_id,
|
||||
activation,
|
||||
conversation_id,
|
||||
route,
|
||||
model_id,
|
||||
} => {
|
||||
let state = match store
|
||||
.start_cursor_trace_if_detailed(
|
||||
&request_id,
|
||||
conversation_id.as_deref(),
|
||||
&route,
|
||||
model_id.as_deref(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(true) => TraceState::Active,
|
||||
Ok(false) => TraceState::Disabled,
|
||||
Err(error) => {
|
||||
tracing::warn!(%request_id, %error, "failed to start Cursor trace");
|
||||
TraceState::Disabled
|
||||
}
|
||||
};
|
||||
activation.store(
|
||||
match state {
|
||||
TraceState::Active => TRACE_ACTIVE,
|
||||
TraceState::Disabled => TRACE_DISABLED,
|
||||
},
|
||||
Ordering::Release,
|
||||
);
|
||||
states.insert(request_id, state);
|
||||
return;
|
||||
}
|
||||
TraceEvent::Resume {
|
||||
request_id,
|
||||
activation,
|
||||
} => {
|
||||
let state = ensure_state(store, states, &request_id).await;
|
||||
activation.store(
|
||||
match state {
|
||||
TraceState::Active => TRACE_ACTIVE,
|
||||
TraceState::Disabled => TRACE_DISABLED,
|
||||
},
|
||||
Ordering::Release,
|
||||
);
|
||||
return;
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
|
||||
if !matches!(
|
||||
ensure_state(store, states, &request_id).await,
|
||||
TraceState::Active
|
||||
) {
|
||||
if finishes_trace {
|
||||
states.remove(&request_id);
|
||||
buffers.remove(&request_id);
|
||||
request_orders.remove(&request_id);
|
||||
}
|
||||
return;
|
||||
}
|
||||
|
||||
let result = match event {
|
||||
TraceEvent::Request {
|
||||
artifact_type,
|
||||
data,
|
||||
metadata,
|
||||
..
|
||||
} => {
|
||||
append_request(
|
||||
store,
|
||||
request_orders,
|
||||
&request_id,
|
||||
artifact_type,
|
||||
data,
|
||||
metadata,
|
||||
)
|
||||
.await
|
||||
}
|
||||
TraceEvent::Artifact {
|
||||
artifact_type,
|
||||
source,
|
||||
data,
|
||||
metadata,
|
||||
..
|
||||
} => {
|
||||
store
|
||||
.append_cursor_trace_artifact(
|
||||
&request_id,
|
||||
&artifact_type,
|
||||
&source,
|
||||
&data,
|
||||
&metadata,
|
||||
)
|
||||
.await
|
||||
}
|
||||
TraceEvent::LinkedBlob {
|
||||
artifact_type,
|
||||
source,
|
||||
blob_id,
|
||||
metadata,
|
||||
..
|
||||
} => {
|
||||
store
|
||||
.link_cursor_trace_artifact(
|
||||
&request_id,
|
||||
&artifact_type,
|
||||
&source,
|
||||
&blob_id,
|
||||
&metadata,
|
||||
)
|
||||
.await
|
||||
}
|
||||
TraceEvent::ResponseStarted { status, .. } => {
|
||||
store.start_cursor_trace_response(&request_id, status).await
|
||||
}
|
||||
TraceEvent::ResponseChunk { source, data, .. } => {
|
||||
let buffer = buffers.entry(request_id.clone()).or_default();
|
||||
buffer.bytes += data.len();
|
||||
buffer
|
||||
.chunks
|
||||
.push(BufferedCursorTraceChunk::new(&source, &data));
|
||||
if buffer.chunks.len() >= MAX_BUFFERED_CHUNKS || buffer.bytes >= MAX_BUFFERED_BYTES {
|
||||
flush_one(store, buffers, &request_id).await;
|
||||
}
|
||||
return;
|
||||
}
|
||||
TraceEvent::Finish { error, .. } => {
|
||||
flush_request_order(store, request_orders, &request_id).await;
|
||||
flush_one(store, buffers, &request_id).await;
|
||||
store
|
||||
.finish_cursor_trace(&request_id, error.as_deref())
|
||||
.await
|
||||
}
|
||||
TraceEvent::Begin { .. } | TraceEvent::Resume { .. } => unreachable!(),
|
||||
};
|
||||
if let Err(error) = result {
|
||||
tracing::warn!(%request_id, %error, "failed to record Cursor trace event");
|
||||
}
|
||||
if finishes_trace {
|
||||
states.remove(&request_id);
|
||||
buffers.remove(&request_id);
|
||||
}
|
||||
}
|
||||
|
||||
async fn append_request(
|
||||
store: &Store,
|
||||
request_orders: &mut HashMap<String, RequestOrder>,
|
||||
request_id: &str,
|
||||
artifact_type: String,
|
||||
data: bytes::Bytes,
|
||||
metadata: serde_json::Value,
|
||||
) -> crate::Result<()> {
|
||||
let append_seqno = metadata
|
||||
.get("append_seqno")
|
||||
.and_then(serde_json::Value::as_i64);
|
||||
let ordered = artifact_type == "bidi_request"
|
||||
&& metadata
|
||||
.get("accepted")
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
== Some(true)
|
||||
&& metadata
|
||||
.get("route_outcome")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
== Some("local");
|
||||
let Some(append_seqno) = append_seqno.filter(|_| ordered) else {
|
||||
return store
|
||||
.append_cursor_trace_request(
|
||||
request_id,
|
||||
&artifact_type,
|
||||
"cursor_client",
|
||||
&data,
|
||||
&metadata,
|
||||
)
|
||||
.await;
|
||||
};
|
||||
|
||||
let request = BufferedRequest {
|
||||
artifact_type,
|
||||
data,
|
||||
metadata,
|
||||
};
|
||||
let order = request_orders.entry(request_id.to_owned()).or_default();
|
||||
if append_seqno < order.next {
|
||||
return store
|
||||
.append_cursor_trace_request(
|
||||
request_id,
|
||||
&request.artifact_type,
|
||||
"cursor_client",
|
||||
&request.data,
|
||||
&request.metadata,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
order.pending.entry(append_seqno).or_default().push(request);
|
||||
while let Some(requests) = order.pending.remove(&order.next) {
|
||||
for request in requests {
|
||||
store
|
||||
.append_cursor_trace_request(
|
||||
request_id,
|
||||
&request.artifact_type,
|
||||
"cursor_client",
|
||||
&request.data,
|
||||
&request.metadata,
|
||||
)
|
||||
.await?;
|
||||
}
|
||||
order.next = order.next.saturating_add(1);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn flush_request_order(
|
||||
store: &Store,
|
||||
request_orders: &mut HashMap<String, RequestOrder>,
|
||||
request_id: &str,
|
||||
) {
|
||||
let Some(order) = request_orders.remove(request_id) else {
|
||||
return;
|
||||
};
|
||||
for requests in order.pending.into_values() {
|
||||
for request in requests {
|
||||
if let Err(error) = store
|
||||
.append_cursor_trace_request(
|
||||
request_id,
|
||||
&request.artifact_type,
|
||||
"cursor_client",
|
||||
&request.data,
|
||||
&request.metadata,
|
||||
)
|
||||
.await
|
||||
{
|
||||
tracing::warn!(%request_id, %error, "failed to flush ordered Cursor request trace");
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn ensure_state(
|
||||
store: &Store,
|
||||
states: &mut HashMap<String, TraceState>,
|
||||
request_id: &str,
|
||||
) -> TraceState {
|
||||
if let Some(state) = states.get(request_id).copied() {
|
||||
return state;
|
||||
}
|
||||
let state = match store.cursor_trace_exists(request_id).await {
|
||||
Ok(true) => TraceState::Active,
|
||||
Ok(false) => TraceState::Disabled,
|
||||
Err(error) => {
|
||||
tracing::warn!(%request_id, %error, "failed to resume Cursor trace");
|
||||
TraceState::Disabled
|
||||
}
|
||||
};
|
||||
states.insert(request_id.to_owned(), state);
|
||||
state
|
||||
}
|
||||
|
||||
async fn flush_one(store: &Store, buffers: &mut HashMap<String, ResponseBuffer>, request_id: &str) {
|
||||
let Some(mut buffer) = buffers.remove(request_id) else {
|
||||
return;
|
||||
};
|
||||
if let Err(error) = store
|
||||
.add_cursor_trace_response_chunks(request_id, &buffer.chunks)
|
||||
.await
|
||||
{
|
||||
tracing::warn!(%request_id, %error, "failed to flush Cursor response chunks");
|
||||
buffer.bytes = buffer.chunks.iter().map(|chunk| chunk.data.len()).sum();
|
||||
buffers.insert(request_id.to_owned(), buffer);
|
||||
}
|
||||
}
|
||||
|
||||
async fn flush_all(store: &Store, buffers: &mut HashMap<String, ResponseBuffer>) {
|
||||
let request_ids = buffers.keys().cloned().collect::<Vec<_>>();
|
||||
for request_id in request_ids {
|
||||
flush_one(store, buffers, &request_id).await;
|
||||
}
|
||||
}
|
||||
@@ -12,7 +12,7 @@ use crate::{
|
||||
};
|
||||
|
||||
pub use query::tool_query;
|
||||
pub(crate) use render::{create_plan_partial, edit_content_delta, edit_path_partial};
|
||||
pub(crate) use render::{create_plan_partial, edit_content_delta, edit_path_partial, task_partial};
|
||||
pub use render::{dynamic_mcp_placeholder, render_dynamic_mcp, tool_completed};
|
||||
pub use request::{abort, mcp_request, mcp_state_request, request};
|
||||
pub(crate) use request::{edit_read_request, json_object_to_prost, mcp_meta_request};
|
||||
|
||||
@@ -54,6 +54,49 @@ pub(crate) fn edit_content_delta(call: &ToolCall, content: String) -> pb::AgentS
|
||||
)))
|
||||
}
|
||||
|
||||
pub(crate) fn task_partial(
|
||||
call: &ToolCall,
|
||||
description: &str,
|
||||
prompt: &str,
|
||||
subagent: &str,
|
||||
model: &str,
|
||||
resume: &str,
|
||||
environment: &str,
|
||||
) -> pb::AgentServerMessage {
|
||||
server_interaction(pb::interaction_update::Message::PartialToolCall(
|
||||
pb::PartialToolCallUpdate {
|
||||
call_id: call.call_id.clone(),
|
||||
tool_call: Some(pb::ToolCall {
|
||||
hook_additional_contexts: Vec::new(),
|
||||
tool_call_id: Some(call.call_id.clone()),
|
||||
started_at_ms: None,
|
||||
completed_at_ms: None,
|
||||
tool: Some(pb::tool_call::Tool::TaskToolCall(pb::TaskToolCall {
|
||||
args: Some(pb::TaskArgs {
|
||||
description: description.into(),
|
||||
prompt: prompt.into(),
|
||||
subagent_type: Some(subagent_type(subagent)),
|
||||
model: (!model.is_empty()).then(|| model.into()),
|
||||
resume: (!resume.is_empty()).then(|| resume.into()),
|
||||
agent_id: None,
|
||||
attachments: Vec::new(),
|
||||
mode: 0,
|
||||
responding_to_message_ids: Vec::new(),
|
||||
environment: execution_environment(
|
||||
(!environment.is_empty()).then_some(environment),
|
||||
),
|
||||
machine: None,
|
||||
}),
|
||||
result: None,
|
||||
..Default::default()
|
||||
})),
|
||||
}),
|
||||
args_text_delta: String::new(),
|
||||
model_call_id: call.model_call_id.clone(),
|
||||
},
|
||||
))
|
||||
}
|
||||
|
||||
pub(crate) fn create_plan_partial(
|
||||
call: &ToolCall,
|
||||
name: &str,
|
||||
@@ -167,7 +210,7 @@ pub fn tool_completed(call: &ToolCall, completion: &ToolCompletion) -> pb::Agent
|
||||
pub fn tool_placeholder(name: &str, call_id: &str) -> Result<pb::ToolCall> {
|
||||
use pb::tool_call::Tool;
|
||||
let tool = match normalized(name).as_str() {
|
||||
"shell" => Tool::ShellToolCall(pb::ShellToolCall::default()),
|
||||
"shell" | "bash" => Tool::ShellToolCall(pb::ShellToolCall::default()),
|
||||
"delete" => Tool::DeleteToolCall(pb::DeleteToolCall::default()),
|
||||
"glob" => Tool::GlobToolCall(pb::GlobToolCall::default()),
|
||||
"grep" => Tool::GrepToolCall(pb::GrepToolCall::default()),
|
||||
@@ -527,3 +570,23 @@ fn now_ms() -> u64 {
|
||||
.unwrap_or_default()
|
||||
.as_millis() as u64
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::tool_placeholder;
|
||||
use crate::cursor::protocol::proto::agent::v1 as pb;
|
||||
|
||||
#[test]
|
||||
fn bash_renders_as_a_shell_placeholder() {
|
||||
// The dispatcher treats `bash`/`Bash` as a Shell alias, so the streaming
|
||||
// placeholder must too; otherwise a `Bash` tool call aborts the turn with
|
||||
// `unsupported tool: bash` before it ever runs.
|
||||
for name in ["shell", "Shell", "bash", "Bash"] {
|
||||
let tool = tool_placeholder(name, "call-1").unwrap().tool;
|
||||
assert!(
|
||||
matches!(tool, Some(pb::tool_call::Tool::ShellToolCall(_))),
|
||||
"{name} should render as a Shell tool"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -35,7 +35,7 @@ pub fn request(id: u32, call: &ToolCall, context: &ExecContext) -> Result<pb::Ag
|
||||
.map(|v| v as i32)
|
||||
};
|
||||
let message = match normalize(&call.name).as_str() {
|
||||
"shell" => {
|
||||
"shell" | "bash" => {
|
||||
let command = string("command")?;
|
||||
let (simple_commands, parsing_result) = shell_command_metadata(&command);
|
||||
Message::ShellStreamArgs(pb::ShellArgs {
|
||||
@@ -520,3 +520,38 @@ fn prost_value(value: &Value) -> prost_types::Value {
|
||||
};
|
||||
ProstValue { kind: Some(kind) }
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::request;
|
||||
use crate::cursor::protocol::proto::agent::v1 as pb;
|
||||
use crate::cursor::tools::runtime::ExecContext;
|
||||
use crate::model::ToolCall;
|
||||
|
||||
#[test]
|
||||
fn bash_is_encoded_as_a_shell_exec_request() {
|
||||
// The dispatcher routes `bash`/`Bash` to the shell executor, so the
|
||||
// request codec must encode it as a Shell stream instead of erroring
|
||||
// with `tool bash is not executed through ExecServerMessage`.
|
||||
let call = ToolCall {
|
||||
index: 0,
|
||||
call_id: "call-1".into(),
|
||||
model_call_id: "model-1".into(),
|
||||
name: "Bash".into(),
|
||||
arguments_text: String::new(),
|
||||
arguments: json!({ "command": "ls -la" }),
|
||||
argument_error: None,
|
||||
};
|
||||
let message = request(1, &call, &ExecContext::default()).unwrap();
|
||||
let Some(pb::agent_server_message::Message::ExecServerMessage(exec)) = message.message
|
||||
else {
|
||||
panic!("expected an ExecServerMessage");
|
||||
};
|
||||
let Some(pb::exec_server_message::Message::ShellStreamArgs(args)) = exec.message else {
|
||||
panic!("expected ShellStreamArgs");
|
||||
};
|
||||
assert_eq!(args.command, "ls -la");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -3,7 +3,7 @@ use crate::{
|
||||
cursor::{
|
||||
protocol::{events, proto::agent::v1 as pb},
|
||||
tools::{
|
||||
edit,
|
||||
compat, edit,
|
||||
runtime::{CursorToolRuntime, ExecStage, PendingExec},
|
||||
tool_call_result::{self as result, ToolCompletion},
|
||||
},
|
||||
@@ -34,16 +34,15 @@ pub async fn client_event(
|
||||
let call = match pending.exec_call(message.id).await {
|
||||
Some(call) => call,
|
||||
None if pending.completed_call(message.id).await.is_some() => {
|
||||
return Err(Error::Protocol(format!(
|
||||
"duplicate terminal ExecClientMessage id: {}",
|
||||
message.id
|
||||
)))
|
||||
tracing::warn!(id = message.id, "ignoring duplicate terminal tool response");
|
||||
return Ok(ClientExecEvent::Pending);
|
||||
}
|
||||
None => {
|
||||
return Err(Error::Protocol(format!(
|
||||
"unknown ExecClientMessage id: {}",
|
||||
message.id
|
||||
)))
|
||||
tracing::warn!(
|
||||
id = message.id,
|
||||
"ignoring response for unknown tool execution"
|
||||
);
|
||||
return Ok(ClientExecEvent::Pending);
|
||||
}
|
||||
};
|
||||
let Some(wire_result) = &message.message else {
|
||||
@@ -144,7 +143,8 @@ pub async fn stream_closed(id: u32, pending: &CursorToolRuntime) -> Result<Optio
|
||||
return Ok(None);
|
||||
};
|
||||
let error = "Cursor Exec stream closed before returning a terminal result";
|
||||
if entry.call.name.eq_ignore_ascii_case("Shell") {
|
||||
if entry.call.name.eq_ignore_ascii_case("Shell") || entry.call.name.eq_ignore_ascii_case("Bash")
|
||||
{
|
||||
let command = entry
|
||||
.call
|
||||
.arguments
|
||||
@@ -173,9 +173,22 @@ pub async fn stream_closed(id: u32, pending: &CursorToolRuntime) -> Result<Optio
|
||||
}
|
||||
let rendered = match &entry.stage {
|
||||
ExecStage::DynamicMcp(definition) => {
|
||||
super::render_dynamic_mcp(&entry.call, definition, false)
|
||||
Ok(super::render_dynamic_mcp(&entry.call, definition, false))
|
||||
}
|
||||
_ => super::render_tool_call(&entry.call, false)?,
|
||||
_ => super::render_tool_call(&entry.call, false),
|
||||
};
|
||||
let rendered = match rendered {
|
||||
Ok(rendered) => rendered,
|
||||
Err(Error::Protocol(message)) => {
|
||||
return Ok(Some(compat::failure_with_message(&entry.call, message)));
|
||||
}
|
||||
Err(Error::Json(error)) => {
|
||||
return Ok(Some(compat::failure_with_message(
|
||||
&entry.call,
|
||||
error.to_string(),
|
||||
)));
|
||||
}
|
||||
Err(error) => return Err(error),
|
||||
};
|
||||
Ok(Some(ToolCompletion::from_rendered(
|
||||
&entry.call,
|
||||
@@ -211,10 +224,10 @@ async fn advance_edit(
|
||||
pb::exec_client_message::Message::ReadResult(result)
|
||||
| pb::exec_client_message::Message::RedactedReadResult(result) => result,
|
||||
_ => {
|
||||
return Err(Error::Protocol(format!(
|
||||
"expected ReadResult for edit tool {}",
|
||||
entry.call.name
|
||||
)))
|
||||
let message = format!("expected ReadResult for edit tool {}", entry.call.name);
|
||||
return Ok(ClientExecEvent::Completed(Box::new(
|
||||
compat::failure_with_message(&entry.call, message),
|
||||
)));
|
||||
}
|
||||
};
|
||||
let write = match edit::after_read(&entry.call, read) {
|
||||
@@ -259,9 +272,14 @@ fn completed(
|
||||
pending: PendingExec,
|
||||
result: pb::exec_client_message::Message,
|
||||
) -> Result<ClientExecEvent> {
|
||||
Ok(ClientExecEvent::Completed(Box::new(result::from_exec(
|
||||
pending, &result,
|
||||
)?)))
|
||||
let call = pending.call.clone();
|
||||
let completion = match result::from_exec(pending, &result) {
|
||||
Ok(completion) => completion,
|
||||
Err(Error::Protocol(message)) => compat::failure_with_message(&call, message),
|
||||
Err(Error::Json(error)) => compat::failure_with_message(&call, error.to_string()),
|
||||
Err(error) => return Err(error),
|
||||
};
|
||||
Ok(ClientExecEvent::Completed(Box::new(completion)))
|
||||
}
|
||||
|
||||
fn shell_exit_result(
|
||||
|
||||
@@ -49,7 +49,10 @@ pub(crate) fn render(call: &ToolCall, completed: bool) -> pb::ToolCall {
|
||||
}
|
||||
|
||||
pub(crate) fn failure(call: &ToolCall) -> ToolCompletion {
|
||||
let error = failure_message(&call.name);
|
||||
failure_with_message(call, failure_message(&call.name))
|
||||
}
|
||||
|
||||
pub(crate) fn failure_with_message(call: &ToolCall, error: String) -> ToolCompletion {
|
||||
let arguments = call
|
||||
.arguments
|
||||
.as_object()
|
||||
|
||||
@@ -183,6 +183,9 @@ fn edit_notebook(call: &ToolCall, before: &str) -> std::result::Result<String, S
|
||||
.unwrap_or_default();
|
||||
let old =
|
||||
normalize_newlines(&string(call, "old_string").map_err(|error| error.to_string())?);
|
||||
if old.is_empty() {
|
||||
return Err("old_string must not be empty".into());
|
||||
}
|
||||
let occurrences = source.match_indices(&old).count();
|
||||
let edited = match occurrences {
|
||||
0 => return Err("old_string was not found in the notebook cell".into()),
|
||||
@@ -241,3 +244,50 @@ fn normalized(value: &str) -> String {
|
||||
.flat_map(char::to_lowercase)
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::edit_notebook;
|
||||
use crate::model::ToolCall;
|
||||
|
||||
fn notebook_call(old_string: &str) -> ToolCall {
|
||||
ToolCall {
|
||||
index: 0,
|
||||
call_id: "call".into(),
|
||||
model_call_id: "model".into(),
|
||||
name: "EditNotebook".into(),
|
||||
arguments_text: String::new(),
|
||||
arguments: json!({
|
||||
"target_notebook": "/notebook.ipynb",
|
||||
"cell_idx": 0,
|
||||
"old_string": old_string,
|
||||
"new_string": "replacement",
|
||||
}),
|
||||
argument_error: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn single_cell_notebook() -> String {
|
||||
json!({
|
||||
"cells": [{"cell_type": "code", "source": ["print('hi')\n"]}],
|
||||
})
|
||||
.to_string()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn edit_notebook_rejects_empty_old_string() {
|
||||
// StrReplace rejects an empty old_string; EditNotebook must do the same
|
||||
// instead of prepending new_string (empty cell) or reporting a
|
||||
// misleading "not unique" error (non-empty cell).
|
||||
let error = edit_notebook(¬ebook_call(""), &single_cell_notebook()).unwrap_err();
|
||||
assert_eq!(error, "old_string must not be empty");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn edit_notebook_replaces_a_unique_old_string() {
|
||||
let edited = edit_notebook(¬ebook_call("hi"), &single_cell_notebook()).unwrap();
|
||||
assert!(edited.contains("print('replacement')"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -103,10 +103,20 @@ impl ToolDispatcher {
|
||||
}
|
||||
let message_index = first_tool_index + position;
|
||||
let publish_started = !state.started.contains(&call.call_id);
|
||||
if let Some(error) = &call.argument_error {
|
||||
dispatched.push(validation_failure(call, error.clone()));
|
||||
continue;
|
||||
}
|
||||
let edit_path = if dynamic_mcp.contains_key(&call.name) {
|
||||
None
|
||||
} else {
|
||||
edit::execution_path(call)?
|
||||
match edit::execution_path(call) {
|
||||
Ok(path) => path,
|
||||
Err(error) => {
|
||||
dispatched.push(recover_validation_failure(call, error)?);
|
||||
continue;
|
||||
}
|
||||
}
|
||||
};
|
||||
if let Some(path) = edit_path {
|
||||
let next = self.edit_schedule.lock().await.start_or_defer(
|
||||
@@ -121,22 +131,28 @@ impl ToolDispatcher {
|
||||
let Some(next) = next else {
|
||||
continue;
|
||||
};
|
||||
dispatched.push(
|
||||
self.start(
|
||||
let started = self
|
||||
.start(
|
||||
&next.call,
|
||||
next.message_index,
|
||||
next.publish_started,
|
||||
dynamic_mcp,
|
||||
&next.context,
|
||||
)
|
||||
.await?,
|
||||
);
|
||||
.await;
|
||||
dispatched.push(match started {
|
||||
Ok(started) => started,
|
||||
Err(error) => recover_validation_failure(&next.call, error)?,
|
||||
});
|
||||
continue;
|
||||
}
|
||||
dispatched.push(
|
||||
self.start(call, message_index, publish_started, dynamic_mcp, context)
|
||||
.await?,
|
||||
);
|
||||
let started = self
|
||||
.start(call, message_index, publish_started, dynamic_mcp, context)
|
||||
.await;
|
||||
dispatched.push(match started {
|
||||
Ok(started) => started,
|
||||
Err(error) => recover_validation_failure(call, error)?,
|
||||
});
|
||||
}
|
||||
Ok(dispatched)
|
||||
}
|
||||
@@ -146,15 +162,19 @@ impl ToolDispatcher {
|
||||
let Some(next) = next else {
|
||||
return Ok(None);
|
||||
};
|
||||
self.start(
|
||||
&next.call,
|
||||
next.message_index,
|
||||
next.publish_started,
|
||||
&BTreeMap::new(),
|
||||
&next.context,
|
||||
)
|
||||
.await
|
||||
.map(Some)
|
||||
match self
|
||||
.start(
|
||||
&next.call,
|
||||
next.message_index,
|
||||
next.publish_started,
|
||||
&BTreeMap::new(),
|
||||
&next.context,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(started) => Ok(Some(started)),
|
||||
Err(error) => recover_validation_failure(&next.call, error).map(Some),
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn interrupt_for_message(&self) -> Vec<u32> {
|
||||
@@ -203,34 +223,64 @@ impl ToolDispatcher {
|
||||
let pending = match self.runtime.take_interaction(response.id).await {
|
||||
Some(pending) => pending,
|
||||
None if self.runtime.completed_call(response.id).await.is_some() => {
|
||||
return Err(Error::Protocol(format!(
|
||||
"duplicate terminal InteractionResponse id: {}",
|
||||
response.id
|
||||
)));
|
||||
tracing::warn!(
|
||||
id = response.id,
|
||||
"ignoring duplicate terminal interaction response"
|
||||
);
|
||||
return Ok(ClientToolEvent::Pending);
|
||||
}
|
||||
None => {
|
||||
return Err(Error::Protocol(format!(
|
||||
"unknown InteractionResponse id: {}",
|
||||
response.id
|
||||
)));
|
||||
tracing::warn!(
|
||||
id = response.id,
|
||||
"ignoring response for unknown interaction"
|
||||
);
|
||||
return Ok(ClientToolEvent::Pending);
|
||||
}
|
||||
};
|
||||
Ok(
|
||||
match tool_call_dispatch::resume_interaction(
|
||||
&self.results,
|
||||
&self.search,
|
||||
&self.fetch,
|
||||
pending,
|
||||
response,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
tool_call_dispatch::InteractionContinuation::Completed(completion) => {
|
||||
ClientToolEvent::Completed(completion)
|
||||
}
|
||||
tool_call_dispatch::InteractionContinuation::Pending => ClientToolEvent::Pending,
|
||||
},
|
||||
let call = pending.call.clone();
|
||||
let continuation = match tool_call_dispatch::resume_interaction(
|
||||
&self.results,
|
||||
&self.search,
|
||||
&self.fetch,
|
||||
pending,
|
||||
response,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(continuation) => continuation,
|
||||
Err(Error::Protocol(message)) => {
|
||||
return Ok(ClientToolEvent::Completed(Box::new(
|
||||
compat::failure_with_message(&call, message),
|
||||
)));
|
||||
}
|
||||
Err(Error::Json(error)) => {
|
||||
return Ok(ClientToolEvent::Completed(Box::new(
|
||||
compat::failure_with_message(&call, error.to_string()),
|
||||
)));
|
||||
}
|
||||
Err(error) => return Err(error),
|
||||
};
|
||||
Ok(match continuation {
|
||||
tool_call_dispatch::InteractionContinuation::Completed(completion) => {
|
||||
ClientToolEvent::Completed(completion)
|
||||
}
|
||||
tool_call_dispatch::InteractionContinuation::Pending => ClientToolEvent::Pending,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn validation_failure(call: &ToolCall, message: String) -> DispatchedTool {
|
||||
DispatchedTool {
|
||||
messages: Vec::new(),
|
||||
completion: Some(compat::failure_with_message(call, message)),
|
||||
}
|
||||
}
|
||||
|
||||
fn recover_validation_failure(call: &ToolCall, error: Error) -> Result<DispatchedTool> {
|
||||
match error {
|
||||
Error::Protocol(message) => Ok(validation_failure(call, message)),
|
||||
Error::Json(error) => Ok(validation_failure(call, error.to_string())),
|
||||
error => Err(error),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -20,6 +20,7 @@ enum Presentation {
|
||||
DynamicMcp(pb::McpToolDefinition),
|
||||
Edit(EditProjection),
|
||||
CreatePlan(CreatePlanProjection),
|
||||
Task(TaskProjection),
|
||||
}
|
||||
|
||||
struct EditProjection {
|
||||
@@ -38,6 +39,17 @@ struct CreatePlanProjection {
|
||||
overview: String,
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct TaskProjection {
|
||||
fields: JsonStringFields,
|
||||
description: String,
|
||||
prompt: String,
|
||||
subagent_type: String,
|
||||
model: String,
|
||||
resume: String,
|
||||
environment: String,
|
||||
}
|
||||
|
||||
impl ToolCallStream {
|
||||
pub fn new(name: &str, dynamic_mcp: Option<&pb::McpToolDefinition>) -> Self {
|
||||
let presentation = match dynamic_mcp {
|
||||
@@ -49,6 +61,7 @@ impl ToolCallStream {
|
||||
Presentation::Edit(EditProjection::new("target_notebook", "new_string"))
|
||||
}
|
||||
"createplan" => Presentation::CreatePlan(CreatePlanProjection::default()),
|
||||
"task" => Presentation::Task(TaskProjection::default()),
|
||||
_ => Presentation::Plain,
|
||||
},
|
||||
};
|
||||
@@ -73,10 +86,48 @@ impl ToolCallStream {
|
||||
Ok(messages)
|
||||
}
|
||||
Presentation::CreatePlan(plan) => plan.project(call, raw_delta),
|
||||
Presentation::Task(task) => task.project(call, raw_delta),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl TaskProjection {
|
||||
fn project(&mut self, call: &ToolCall, raw_delta: &str) -> Result<Vec<pb::AgentServerMessage>> {
|
||||
let mut description_completed = false;
|
||||
for event in self.fields.push(raw_delta)? {
|
||||
match event {
|
||||
StringFieldEvent::Delta { name, text } => match name.as_str() {
|
||||
"description" => self.description.push_str(&text),
|
||||
"prompt" => self.prompt.push_str(&text),
|
||||
"subagent_type" => self.subagent_type.push_str(&text),
|
||||
"model" => self.model.push_str(&text),
|
||||
"resume" => self.resume.push_str(&text),
|
||||
"environment" => self.environment.push_str(&text),
|
||||
_ => {}
|
||||
},
|
||||
StringFieldEvent::End { name } if name == "description" => {
|
||||
description_completed = true
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
Ok(description_completed
|
||||
.then(|| {
|
||||
interaction::task_partial(
|
||||
call,
|
||||
&self.description,
|
||||
&self.prompt,
|
||||
&self.subagent_type,
|
||||
&self.model,
|
||||
&self.resume,
|
||||
&self.environment,
|
||||
)
|
||||
})
|
||||
.into_iter()
|
||||
.collect())
|
||||
}
|
||||
}
|
||||
|
||||
impl CreatePlanProjection {
|
||||
fn project(&mut self, call: &ToolCall, raw_delta: &str) -> Result<Vec<pb::AgentServerMessage>> {
|
||||
let mut completed_field = false;
|
||||
@@ -184,3 +235,54 @@ fn normalized(value: &str) -> String {
|
||||
.flat_map(char::to_lowercase)
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn task_call(arguments_text: &str) -> ToolCall {
|
||||
ToolCall {
|
||||
index: 0,
|
||||
call_id: "task-1".into(),
|
||||
model_call_id: "model-1".into(),
|
||||
name: "Task".into(),
|
||||
arguments_text: arguments_text.into(),
|
||||
arguments: serde_json::Value::Null,
|
||||
argument_error: None,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn task_description_projects_a_visible_partial_card_before_execution() {
|
||||
let mut stream = ToolCallStream::new("Task", None);
|
||||
let first = r#"{"description":"Review K10"#;
|
||||
assert!(stream
|
||||
.arguments_delta(&task_call(first), first)
|
||||
.unwrap()
|
||||
.is_empty());
|
||||
|
||||
let closing = "\",";
|
||||
let messages = stream
|
||||
.arguments_delta(&task_call(&format!("{first}{closing}")), closing)
|
||||
.unwrap();
|
||||
assert_eq!(messages.len(), 1);
|
||||
let Some(pb::agent_server_message::Message::InteractionUpdate(update)) =
|
||||
messages[0].message.as_ref()
|
||||
else {
|
||||
panic!("expected interaction update")
|
||||
};
|
||||
let Some(pb::interaction_update::Message::PartialToolCall(partial)) =
|
||||
update.message.as_ref()
|
||||
else {
|
||||
panic!("expected partial tool call")
|
||||
};
|
||||
let tool_call = partial.tool_call.as_ref().expect("expected tool call");
|
||||
assert_eq!(tool_call.started_at_ms, None);
|
||||
let Some(pb::tool_call::Tool::TaskToolCall(task)) = tool_call.tool.as_ref() else {
|
||||
panic!("expected Task tool call")
|
||||
};
|
||||
let args = task.args.as_ref().expect("expected partial Task args");
|
||||
assert_eq!(args.description, "Review K10");
|
||||
assert_eq!(args.prompt, "");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -406,6 +406,7 @@ mod tests {
|
||||
name: "WebFetch".into(),
|
||||
arguments_text: r#"{"url":"https://example.com"}"#.into(),
|
||||
arguments: json!({"url": "https://example.com"}),
|
||||
argument_error: None,
|
||||
},
|
||||
started_at_ms: 1,
|
||||
}
|
||||
|
||||
@@ -15,7 +15,7 @@ use crate::{
|
||||
Error, Result,
|
||||
};
|
||||
|
||||
use super::OutputHub;
|
||||
use super::{OutputHub, TransportAdmission, TransportLifecycle};
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub struct TransportParent {
|
||||
@@ -30,7 +30,8 @@ pub struct TransportHandle {
|
||||
output: Arc<OutputHub>,
|
||||
conversation_id: Arc<OnceLock<String>>,
|
||||
parent: Arc<OnceLock<TransportParent>>,
|
||||
trace: Option<CursorTraceRecorder>,
|
||||
trace: CursorTraceRecorder,
|
||||
lifecycle: TransportLifecycle,
|
||||
disconnect: CancellationToken,
|
||||
}
|
||||
|
||||
@@ -39,7 +40,7 @@ impl TransportHandle {
|
||||
request_id: String,
|
||||
commands: mpsc::Sender<TransportCommand>,
|
||||
output: Arc<OutputHub>,
|
||||
trace: Option<CursorTraceRecorder>,
|
||||
trace: CursorTraceRecorder,
|
||||
) -> Self {
|
||||
Self {
|
||||
request_id,
|
||||
@@ -48,6 +49,7 @@ impl TransportHandle {
|
||||
conversation_id: Arc::new(OnceLock::new()),
|
||||
parent: Arc::new(OnceLock::new()),
|
||||
trace,
|
||||
lifecycle: TransportLifecycle::new(),
|
||||
disconnect: CancellationToken::new(),
|
||||
}
|
||||
}
|
||||
@@ -127,12 +129,46 @@ impl TransportHandle {
|
||||
self.output.close()
|
||||
}
|
||||
|
||||
pub(crate) async fn wait_closed(&self) {
|
||||
self.output.wait_closed().await;
|
||||
pub(crate) fn trace(&self) -> Option<&CursorTraceRecorder> {
|
||||
Some(&self.trace)
|
||||
}
|
||||
|
||||
pub(crate) fn trace(&self) -> Option<&CursorTraceRecorder> {
|
||||
self.trace.as_ref()
|
||||
pub(crate) fn accepting_appends(&self) -> bool {
|
||||
self.lifecycle.is_open()
|
||||
}
|
||||
|
||||
pub(crate) fn admit(&self) -> Result<TransportAdmission> {
|
||||
self.lifecycle
|
||||
.admit()
|
||||
.ok_or_else(|| Error::RunNotFound(self.request_id.clone()))
|
||||
}
|
||||
|
||||
pub(crate) fn begin_close(&self) {
|
||||
self.lifecycle.begin_close();
|
||||
}
|
||||
|
||||
pub(crate) fn admissions_drained(&self) -> bool {
|
||||
self.lifecycle.admissions_drained()
|
||||
}
|
||||
|
||||
pub(crate) async fn wait_admissions_drained(&self) {
|
||||
self.lifecycle.wait_admissions_drained().await;
|
||||
}
|
||||
|
||||
pub(crate) fn mark_draining(&self) {
|
||||
self.lifecycle.mark_draining();
|
||||
}
|
||||
|
||||
pub(crate) fn reopen(&self) {
|
||||
self.lifecycle.reopen();
|
||||
}
|
||||
|
||||
pub(crate) fn close_transport(&self) {
|
||||
self.lifecycle.close();
|
||||
}
|
||||
|
||||
pub(crate) async fn wait_transport_closed(&self) {
|
||||
self.lifecycle.wait_closed().await;
|
||||
}
|
||||
|
||||
pub(crate) fn disconnect_token(&self) -> CancellationToken {
|
||||
|
||||
@@ -0,0 +1,170 @@
|
||||
//! Coordinates append admission with transport shutdown.
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use tokio::sync::Notify;
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub(crate) enum TransportState {
|
||||
Open,
|
||||
Closing,
|
||||
Draining,
|
||||
Closed,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub(crate) struct TransportLifecycle {
|
||||
inner: Arc<LifecycleInner>,
|
||||
}
|
||||
|
||||
struct LifecycleInner {
|
||||
state: parking_lot::Mutex<LifecycleState>,
|
||||
admissions_drained: Notify,
|
||||
closed: Notify,
|
||||
}
|
||||
|
||||
struct LifecycleState {
|
||||
state: TransportState,
|
||||
admissions: usize,
|
||||
}
|
||||
|
||||
pub(crate) struct TransportAdmission {
|
||||
inner: Arc<LifecycleInner>,
|
||||
}
|
||||
|
||||
impl TransportLifecycle {
|
||||
pub(crate) fn new() -> Self {
|
||||
Self {
|
||||
inner: Arc::new(LifecycleInner {
|
||||
state: parking_lot::Mutex::new(LifecycleState {
|
||||
state: TransportState::Open,
|
||||
admissions: 0,
|
||||
}),
|
||||
admissions_drained: Notify::new(),
|
||||
closed: Notify::new(),
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn is_open(&self) -> bool {
|
||||
self.inner.state.lock().state == TransportState::Open
|
||||
}
|
||||
|
||||
pub(crate) fn admit(&self) -> Option<TransportAdmission> {
|
||||
let mut lifecycle = self.inner.state.lock();
|
||||
if lifecycle.state != TransportState::Open {
|
||||
return None;
|
||||
}
|
||||
lifecycle.admissions += 1;
|
||||
Some(TransportAdmission {
|
||||
inner: self.inner.clone(),
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn begin_close(&self) {
|
||||
let mut lifecycle = self.inner.state.lock();
|
||||
if lifecycle.state != TransportState::Open {
|
||||
return;
|
||||
}
|
||||
lifecycle.state = TransportState::Closing;
|
||||
let drained = lifecycle.admissions == 0;
|
||||
drop(lifecycle);
|
||||
if drained {
|
||||
self.inner.admissions_drained.notify_waiters();
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn admissions_drained(&self) -> bool {
|
||||
self.inner.state.lock().admissions == 0
|
||||
}
|
||||
|
||||
pub(crate) async fn wait_admissions_drained(&self) {
|
||||
loop {
|
||||
let notified = self.inner.admissions_drained.notified();
|
||||
if self.inner.state.lock().admissions == 0 {
|
||||
return;
|
||||
}
|
||||
notified.await;
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn mark_draining(&self) {
|
||||
let mut lifecycle = self.inner.state.lock();
|
||||
if lifecycle.state == TransportState::Closing && lifecycle.admissions == 0 {
|
||||
lifecycle.state = TransportState::Draining;
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn reopen(&self) {
|
||||
let mut lifecycle = self.inner.state.lock();
|
||||
if matches!(
|
||||
lifecycle.state,
|
||||
TransportState::Closing | TransportState::Draining
|
||||
) {
|
||||
lifecycle.state = TransportState::Open;
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn close(&self) {
|
||||
let mut lifecycle = self.inner.state.lock();
|
||||
if lifecycle.state == TransportState::Closed {
|
||||
return;
|
||||
}
|
||||
lifecycle.state = TransportState::Closed;
|
||||
drop(lifecycle);
|
||||
self.inner.closed.notify_waiters();
|
||||
}
|
||||
|
||||
pub(crate) async fn wait_closed(&self) {
|
||||
loop {
|
||||
let notified = self.inner.closed.notified();
|
||||
if self.inner.state.lock().state == TransportState::Closed {
|
||||
return;
|
||||
}
|
||||
notified.await;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for TransportAdmission {
|
||||
fn drop(&mut self) {
|
||||
let mut lifecycle = self.inner.state.lock();
|
||||
lifecycle.admissions = lifecycle.admissions.saturating_sub(1);
|
||||
let drained = lifecycle.state == TransportState::Closing && lifecycle.admissions == 0;
|
||||
drop(lifecycle);
|
||||
if drained {
|
||||
self.inner.admissions_drained.notify_waiters();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{TransportLifecycle, TransportState};
|
||||
|
||||
#[tokio::test]
|
||||
async fn closing_waits_for_existing_admissions() {
|
||||
let lifecycle = TransportLifecycle::new();
|
||||
let admission = lifecycle.admit().unwrap();
|
||||
lifecycle.begin_close();
|
||||
assert!(lifecycle.admit().is_none());
|
||||
drop(admission);
|
||||
lifecycle.wait_admissions_drained().await;
|
||||
lifecycle.mark_draining();
|
||||
assert_eq!(lifecycle.inner.state.lock().state, TransportState::Draining);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn an_admitted_continuation_reopens_the_transport() {
|
||||
let lifecycle = TransportLifecycle::new();
|
||||
let admission = lifecycle.admit().unwrap();
|
||||
lifecycle.begin_close();
|
||||
drop(admission);
|
||||
lifecycle.wait_admissions_drained().await;
|
||||
lifecycle.mark_draining();
|
||||
lifecycle.reopen();
|
||||
|
||||
assert_eq!(lifecycle.inner.state.lock().state, TransportState::Open);
|
||||
assert!(lifecycle.admit().is_some());
|
||||
}
|
||||
}
|
||||
@@ -2,10 +2,12 @@
|
||||
|
||||
mod handle;
|
||||
mod inbox;
|
||||
mod lifecycle;
|
||||
mod output;
|
||||
mod registry;
|
||||
|
||||
pub use handle::*;
|
||||
pub use inbox::*;
|
||||
pub(crate) use lifecycle::*;
|
||||
pub use output::*;
|
||||
pub use registry::*;
|
||||
|
||||
@@ -1,12 +1,11 @@
|
||||
//! Buffers, replays, broadcasts, and atomically closes downstream output.
|
||||
|
||||
use bytes::Bytes;
|
||||
use tokio::sync::{mpsc, Notify};
|
||||
use tokio::sync::mpsc;
|
||||
|
||||
#[derive(Default)]
|
||||
pub struct OutputHub {
|
||||
state: parking_lot::Mutex<OutputState>,
|
||||
closed: Notify,
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
@@ -49,17 +48,6 @@ impl OutputHub {
|
||||
state.closed = true;
|
||||
state.subscribers.clear();
|
||||
drop(state);
|
||||
self.closed.notify_waiters();
|
||||
true
|
||||
}
|
||||
|
||||
pub async fn wait_closed(&self) {
|
||||
loop {
|
||||
let notified = self.closed.notified();
|
||||
if self.state.lock().closed {
|
||||
return;
|
||||
}
|
||||
notified.await;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,13 +1,19 @@
|
||||
//! Maps request IDs to active transport handles.
|
||||
|
||||
use std::{collections::HashMap, sync::Arc};
|
||||
use std::{
|
||||
collections::HashMap,
|
||||
sync::{
|
||||
atomic::{AtomicU64, Ordering},
|
||||
Arc,
|
||||
},
|
||||
};
|
||||
|
||||
use tokio::sync::{mpsc, Mutex, Notify};
|
||||
|
||||
use crate::{
|
||||
cursor::{
|
||||
conversation::ConversationRegistry, prompting::PromptCompiler,
|
||||
services::observability::CursorTraceRecorder,
|
||||
services::observability::CursorTraceService,
|
||||
},
|
||||
plugin::PluginRegistry,
|
||||
provider::Provider,
|
||||
@@ -24,15 +30,23 @@ pub struct TransportRegistry {
|
||||
}
|
||||
|
||||
struct RegistryInner {
|
||||
local: Mutex<HashMap<String, TransportHandle>>,
|
||||
local: Mutex<HashMap<String, LocalTransport>>,
|
||||
next_local_generation: AtomicU64,
|
||||
upstream: Mutex<HashMap<String, u64>>,
|
||||
route_changed: Notify,
|
||||
store: Store,
|
||||
traces: CursorTraceService,
|
||||
web_cache: WebCache,
|
||||
plugins: Option<PluginRegistry>,
|
||||
conversations: ConversationRegistry,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct LocalTransport {
|
||||
generation: u64,
|
||||
handle: TransportHandle,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub enum TransportRoute {
|
||||
Local,
|
||||
@@ -50,7 +64,24 @@ impl TransportRegistry {
|
||||
compiler: PromptCompiler,
|
||||
web_cache: WebCache,
|
||||
) -> Self {
|
||||
Self::build(store, provider, compiler, web_cache, None)
|
||||
Self::build(store, provider, compiler, web_cache, None, None)
|
||||
}
|
||||
|
||||
/// 附带本地 rules 目录的构造;编译请求上下文时会合并该目录下的 md 规则。
|
||||
pub fn with_local_rules(
|
||||
store: Store,
|
||||
provider: Arc<dyn Provider>,
|
||||
compiler: PromptCompiler,
|
||||
local_rules_dir: std::path::PathBuf,
|
||||
) -> Self {
|
||||
Self::build(
|
||||
store,
|
||||
provider,
|
||||
compiler,
|
||||
WebCache::default(),
|
||||
None,
|
||||
Some(local_rules_dir),
|
||||
)
|
||||
}
|
||||
|
||||
pub fn with_plugins(
|
||||
@@ -59,8 +90,16 @@ impl TransportRegistry {
|
||||
compiler: PromptCompiler,
|
||||
web_cache: WebCache,
|
||||
plugins: PluginRegistry,
|
||||
local_rules_dir: std::path::PathBuf,
|
||||
) -> Self {
|
||||
Self::build(store, provider, compiler, web_cache, Some(plugins))
|
||||
Self::build(
|
||||
store,
|
||||
provider,
|
||||
compiler,
|
||||
web_cache,
|
||||
Some(plugins),
|
||||
Some(local_rules_dir),
|
||||
)
|
||||
}
|
||||
|
||||
fn build(
|
||||
@@ -69,17 +108,21 @@ impl TransportRegistry {
|
||||
compiler: PromptCompiler,
|
||||
web_cache: WebCache,
|
||||
plugins: Option<PluginRegistry>,
|
||||
local_rules_dir: Option<std::path::PathBuf>,
|
||||
) -> Self {
|
||||
Self {
|
||||
inner: Arc::new(RegistryInner {
|
||||
local: Mutex::new(HashMap::new()),
|
||||
next_local_generation: AtomicU64::new(1),
|
||||
upstream: Mutex::new(HashMap::new()),
|
||||
route_changed: Notify::new(),
|
||||
traces: CursorTraceService::new(store.clone()),
|
||||
conversations: ConversationRegistry::new(
|
||||
store.clone(),
|
||||
provider,
|
||||
compiler,
|
||||
web_cache.clone(),
|
||||
local_rules_dir,
|
||||
),
|
||||
store,
|
||||
web_cache,
|
||||
@@ -92,6 +135,13 @@ impl TransportRegistry {
|
||||
&self.inner.store
|
||||
}
|
||||
|
||||
pub fn trace(
|
||||
&self,
|
||||
request_id: &str,
|
||||
) -> crate::cursor::services::observability::CursorTraceRecorder {
|
||||
self.inner.traces.recorder(request_id)
|
||||
}
|
||||
|
||||
pub fn web_cache(&self) -> &WebCache {
|
||||
&self.inner.web_cache
|
||||
}
|
||||
@@ -105,18 +155,37 @@ impl TransportRegistry {
|
||||
}
|
||||
|
||||
pub async fn get_or_create(&self, request_id: &str) -> Result<TransportHandle> {
|
||||
if let Some(handle) = self.inner.local.lock().await.get(request_id).cloned() {
|
||||
return Ok(handle);
|
||||
self.get_or_create_for_append(request_id, false).await
|
||||
}
|
||||
|
||||
pub(crate) async fn get_or_create_for_append(
|
||||
&self,
|
||||
request_id: &str,
|
||||
replace_closing: bool,
|
||||
) -> Result<TransportHandle> {
|
||||
let mut local = self.inner.local.lock().await;
|
||||
if let Some(transport) = local.get(request_id) {
|
||||
if transport.handle.accepting_appends() || !replace_closing {
|
||||
return Ok(transport.handle.clone());
|
||||
}
|
||||
}
|
||||
local.remove(request_id);
|
||||
let (commands, receiver) = mpsc::channel(128);
|
||||
let output = Arc::new(OutputHub::default());
|
||||
let trace = CursorTraceRecorder::resume(self.inner.store.clone(), request_id).await;
|
||||
let handle = TransportHandle::new(request_id.into(), commands, output.clone(), trace);
|
||||
let mut local = self.inner.local.lock().await;
|
||||
if let Some(existing) = local.get(request_id).cloned() {
|
||||
return Ok(existing);
|
||||
}
|
||||
local.insert(request_id.into(), handle.clone());
|
||||
let trace = self.inner.traces.recorder(request_id);
|
||||
trace.resume();
|
||||
let handle = TransportHandle::new(request_id.into(), commands, output, trace);
|
||||
let generation = self
|
||||
.inner
|
||||
.next_local_generation
|
||||
.fetch_add(1, Ordering::Relaxed);
|
||||
local.insert(
|
||||
request_id.into(),
|
||||
LocalTransport {
|
||||
generation,
|
||||
handle: handle.clone(),
|
||||
},
|
||||
);
|
||||
drop(local);
|
||||
self.inner.route_changed.notify_waiters();
|
||||
self.inner
|
||||
@@ -125,17 +194,29 @@ impl TransportRegistry {
|
||||
|
||||
let registry = Arc::downgrade(&self.inner);
|
||||
let request_id = request_id.to_string();
|
||||
let lifecycle = handle.clone();
|
||||
tokio::spawn(async move {
|
||||
output.wait_closed().await;
|
||||
lifecycle.wait_transport_closed().await;
|
||||
if let Some(registry) = registry.upgrade() {
|
||||
registry.local.lock().await.remove(&request_id);
|
||||
let mut local = registry.local.lock().await;
|
||||
if local
|
||||
.get(&request_id)
|
||||
.is_some_and(|transport| transport.generation == generation)
|
||||
{
|
||||
local.remove(&request_id);
|
||||
}
|
||||
}
|
||||
});
|
||||
Ok(handle)
|
||||
}
|
||||
|
||||
pub async fn local(&self, request_id: &str) -> Option<TransportHandle> {
|
||||
self.inner.local.lock().await.get(request_id).cloned()
|
||||
self.inner
|
||||
.local
|
||||
.lock()
|
||||
.await
|
||||
.get(request_id)
|
||||
.map(|transport| transport.handle.clone())
|
||||
}
|
||||
|
||||
pub async fn mark_upstream(&self, request_id: &str) {
|
||||
@@ -179,10 +260,13 @@ impl TransportRegistry {
|
||||
self.inner.conversations.shutdown().await;
|
||||
let handles = std::mem::take(&mut *self.inner.local.lock().await);
|
||||
self.inner.upstream.lock().await.clear();
|
||||
for handle in handles.into_values() {
|
||||
handle.disconnect().await;
|
||||
let _ =
|
||||
tokio::time::timeout(std::time::Duration::from_secs(2), handle.wait_closed()).await;
|
||||
for transport in handles.into_values() {
|
||||
transport.handle.disconnect().await;
|
||||
let _ = tokio::time::timeout(
|
||||
std::time::Duration::from_secs(2),
|
||||
transport.handle.wait_transport_closed(),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -161,6 +161,10 @@ fn is_local_path(path: &str) -> bool {
|
||||
| "/aiserver.v1.DashboardService/GetUserProfile"
|
||||
| "/aiserver.v1.DashboardService/GetCurrentPeriodUsage"
|
||||
| "/aiserver.v1.DashboardService/GetUsageLimitStatusAndActiveGrants"
|
||||
| "/aiserver.v1.AiService/KnowledgeBaseAdd"
|
||||
| "/aiserver.v1.AiService/KnowledgeBaseList"
|
||||
| "/aiserver.v1.AiService/KnowledgeBaseUpdate"
|
||||
| "/aiserver.v1.AiService/KnowledgeBaseRemove"
|
||||
| "/aiserver.v1.AnalyticsService/BootstrapStatsig"
|
||||
| "/auth/full_stripe_profile"
|
||||
)
|
||||
|
||||
@@ -87,6 +87,9 @@ pub struct ModelConfigInput {
|
||||
#[serde(default)]
|
||||
pub sort_order: i64,
|
||||
pub display_name: String,
|
||||
/// 供应商分组的自定义显示名;同一 base_url 主机下的模型共享。
|
||||
#[serde(default)]
|
||||
pub group_name: Option<String>,
|
||||
#[serde(rename = "type")]
|
||||
pub model_type: ModelType,
|
||||
pub base_url: String,
|
||||
@@ -124,6 +127,7 @@ pub struct ModelConfig {
|
||||
pub model_hash: String,
|
||||
pub sort_order: i64,
|
||||
pub display_name: String,
|
||||
pub group_name: Option<String>,
|
||||
#[serde(rename = "type")]
|
||||
pub model_type: ModelType,
|
||||
pub base_url: String,
|
||||
@@ -204,6 +208,12 @@ impl ModelConfig {
|
||||
|
||||
pub fn normalize_model_input(input: &ModelConfigInput) -> Result<ModelConfigInput> {
|
||||
let display_name = required(&input.display_name, "model display name")?;
|
||||
let group_name = input
|
||||
.group_name
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(String::from);
|
||||
let base_url = normalize_request_url(&input.base_url)?;
|
||||
let api_key = required(&input.api_key, "model API key")?;
|
||||
let tooltip_data = required(&input.tooltip_data, "model tooltip")?;
|
||||
@@ -230,6 +240,7 @@ pub fn normalize_model_input(input: &ModelConfigInput) -> Result<ModelConfigInpu
|
||||
let normalized = ModelConfigInput {
|
||||
sort_order: input.sort_order.max(0),
|
||||
display_name,
|
||||
group_name,
|
||||
model_type: input.model_type,
|
||||
base_url,
|
||||
use_full_url: input.use_full_url,
|
||||
|
||||
@@ -6,11 +6,10 @@ mod usage {
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use super::ProviderType;
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, Serialize, Deserialize, PartialEq, Eq)]
|
||||
pub struct Usage {
|
||||
pub input_tokens: Option<u64>,
|
||||
pub context_input_tokens: Option<u64>,
|
||||
pub output_tokens: Option<u64>,
|
||||
pub total_tokens: Option<u64>,
|
||||
pub cache_read_tokens: Option<u64>,
|
||||
@@ -18,24 +17,10 @@ mod usage {
|
||||
pub reasoning_tokens: Option<u64>,
|
||||
}
|
||||
|
||||
impl Usage {
|
||||
/// Returns the provider-visible input context without counting cached tokens twice.
|
||||
pub(crate) fn context_input_tokens(self, provider: ProviderType) -> Option<u64> {
|
||||
let input = self.input_tokens?;
|
||||
match provider {
|
||||
ProviderType::OpenAiChat | ProviderType::OpenAiResponses | ProviderType::Plugin => {
|
||||
Some(input)
|
||||
}
|
||||
ProviderType::Anthropic => input
|
||||
.checked_add(self.cache_read_tokens.unwrap_or_default())?
|
||||
.checked_add(self.cache_write_tokens.unwrap_or_default()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl AddAssign for Usage {
|
||||
fn add_assign(&mut self, rhs: Self) {
|
||||
self.input_tokens = sum(self.input_tokens, rhs.input_tokens);
|
||||
self.context_input_tokens = sum(self.context_input_tokens, rhs.context_input_tokens);
|
||||
self.output_tokens = sum(self.output_tokens, rhs.output_tokens);
|
||||
self.total_tokens = sum(self.total_tokens, rhs.total_tokens);
|
||||
self.cache_read_tokens = sum(self.cache_read_tokens, rhs.cache_read_tokens);
|
||||
@@ -53,7 +38,7 @@ pub use usage::*;
|
||||
mod llm_call {
|
||||
use serde::Serialize;
|
||||
|
||||
use super::{ProviderType, Usage};
|
||||
use super::ProviderType;
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct NewLlmCall {
|
||||
@@ -75,14 +60,6 @@ mod llm_call {
|
||||
pub detailed: bool,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub(crate) struct LlmCallUsageAnchor {
|
||||
pub request_type: ProviderType,
|
||||
pub usage: Usage,
|
||||
pub message_count: usize,
|
||||
pub tool_count: usize,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize)]
|
||||
pub struct LlmCallSummary {
|
||||
pub call_id: String,
|
||||
|
||||
@@ -6,8 +6,8 @@ use serde::{Deserialize, Serialize};
|
||||
use crate::{Error, Result};
|
||||
|
||||
use super::{
|
||||
CanonicalMessage, ContentPart, MessageContent, ProviderReplayState, Role, ToolCallContent,
|
||||
ToolResultContent,
|
||||
normalize_tool_name, CanonicalMessage, ContentPart, MessageContent, ProviderReplayState, Role,
|
||||
ToolCallContent, ToolResultContent,
|
||||
};
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
|
||||
@@ -91,7 +91,7 @@ fn project_tool_round(
|
||||
"tool round repeats provider replay state".into(),
|
||||
));
|
||||
}
|
||||
calls.extend(part_calls.iter().cloned());
|
||||
calls.extend(part_calls.iter().map(normalized_tool_call));
|
||||
cursor += 1;
|
||||
|
||||
while cursor < messages.len() {
|
||||
@@ -139,7 +139,7 @@ fn project_tool_round(
|
||||
.map(|(message_id, result)| ProjectedMessage {
|
||||
message_id,
|
||||
role: Role::Tool,
|
||||
content: ProjectedContent::ToolResult(result),
|
||||
content: ProjectedContent::ToolResult(normalized_tool_result(&result)),
|
||||
}),
|
||||
);
|
||||
Ok(Some((output, cursor)))
|
||||
@@ -158,9 +158,11 @@ fn project_message(message: &CanonicalMessage) -> ProjectedMessage {
|
||||
text: text.clone(),
|
||||
thinking: thinking.clone(),
|
||||
replay_state: replay_state.clone(),
|
||||
calls: tool_calls.clone(),
|
||||
calls: tool_calls.iter().map(normalized_tool_call).collect(),
|
||||
},
|
||||
MessageContent::ToolResult(result) => ProjectedContent::ToolResult(result.clone()),
|
||||
MessageContent::ToolResult(result) => {
|
||||
ProjectedContent::ToolResult(normalized_tool_result(result))
|
||||
}
|
||||
};
|
||||
ProjectedMessage {
|
||||
message_id: message.message_id.clone(),
|
||||
@@ -168,3 +170,50 @@ fn project_message(message: &CanonicalMessage) -> ProjectedMessage {
|
||||
content,
|
||||
}
|
||||
}
|
||||
|
||||
fn normalized_tool_call(call: &ToolCallContent) -> ToolCallContent {
|
||||
let mut call = call.clone();
|
||||
call.name = normalize_tool_name(&call.name);
|
||||
call
|
||||
}
|
||||
|
||||
fn normalized_tool_result(result: &ToolResultContent) -> ToolResultContent {
|
||||
let mut result = result.clone();
|
||||
result.name = normalize_tool_name(&result.name);
|
||||
result
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::model::Origin;
|
||||
use serde_json::json;
|
||||
|
||||
#[test]
|
||||
fn tool_names_are_normalized_before_provider_dispatch() {
|
||||
let messages = [CanonicalMessage {
|
||||
message_id: "assistant-1".into(),
|
||||
role: Role::Assistant,
|
||||
origin: Origin::Assistant,
|
||||
content: MessageContent::Assistant {
|
||||
text: String::new(),
|
||||
thinking: String::new(),
|
||||
tool_round_id: None,
|
||||
replay_state: None,
|
||||
tool_calls: vec![ToolCallContent {
|
||||
index: 0,
|
||||
call_id: "call-1".into(),
|
||||
name: "multi_tool_use.parallel".into(),
|
||||
arguments: json!({}),
|
||||
}],
|
||||
},
|
||||
runtime_event_id: None,
|
||||
}];
|
||||
|
||||
let projected = project_messages(&messages).unwrap();
|
||||
let ProjectedContent::Assistant { calls, .. } = &projected[0].content else {
|
||||
panic!("expected assistant projection");
|
||||
};
|
||||
assert_eq!(calls[0].name, "multi_tool_use_parallel");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,4 +1,99 @@
|
||||
//! Estimates and records model token usage.
|
||||
//! Estimates provider-visible context size and formats configured token counts.
|
||||
|
||||
use super::{ContentPart, ProjectedContent, ProjectedMessage, PromptSpec};
|
||||
|
||||
const TOKENS_PER_MESSAGE_OVERHEAD: u64 = 8;
|
||||
const TOKENS_PER_TOOL_CALL_OVERHEAD: u64 = 6;
|
||||
const TOKENS_PER_IMAGE: u64 = 1_024;
|
||||
|
||||
pub(crate) fn estimate_context_tokens(prompt: &PromptSpec, messages: &[ProjectedMessage]) -> u64 {
|
||||
let instructions = estimate_text_tokens(&prompt.instructions);
|
||||
let tools = prompt.tools.iter().fold(0_u64, |total, tool| {
|
||||
total.saturating_add(estimate_json_tokens(tool))
|
||||
});
|
||||
instructions
|
||||
.saturating_add(tools)
|
||||
.saturating_add(estimate_projected_messages_tokens(messages))
|
||||
}
|
||||
|
||||
pub(crate) fn estimate_projected_messages_tokens(messages: &[ProjectedMessage]) -> u64 {
|
||||
messages.iter().fold(0_u64, |total, message| {
|
||||
total.saturating_add(estimate_message_tokens(message))
|
||||
})
|
||||
}
|
||||
|
||||
fn estimate_message_tokens(message: &ProjectedMessage) -> u64 {
|
||||
let content = match &message.content {
|
||||
ProjectedContent::Parts(parts) => estimate_parts_tokens(parts),
|
||||
ProjectedContent::Assistant {
|
||||
text,
|
||||
thinking,
|
||||
replay_state: _,
|
||||
calls,
|
||||
} => {
|
||||
let calls = calls.iter().fold(0_u64, |total, call| {
|
||||
total
|
||||
.saturating_add(TOKENS_PER_TOOL_CALL_OVERHEAD)
|
||||
.saturating_add(estimate_text_tokens(&call.call_id))
|
||||
.saturating_add(estimate_text_tokens(&call.name))
|
||||
.saturating_add(estimate_json_tokens(&call.arguments))
|
||||
});
|
||||
estimate_text_tokens(text)
|
||||
.saturating_add(estimate_text_tokens(thinking))
|
||||
.saturating_add(calls)
|
||||
}
|
||||
ProjectedContent::ToolResult(result) => {
|
||||
let content = if result.provider_parts.is_empty() {
|
||||
estimate_text_tokens(&result.content).saturating_add(
|
||||
result
|
||||
.image
|
||||
.as_ref()
|
||||
.map(|image| {
|
||||
TOKENS_PER_IMAGE.saturating_add(estimate_text_tokens(&image.mime_type))
|
||||
})
|
||||
.unwrap_or_default(),
|
||||
)
|
||||
} else {
|
||||
estimate_parts_tokens(&result.provider_parts)
|
||||
};
|
||||
estimate_text_tokens(&result.call_id)
|
||||
.saturating_add(estimate_text_tokens(&result.name))
|
||||
.saturating_add(content)
|
||||
}
|
||||
};
|
||||
TOKENS_PER_MESSAGE_OVERHEAD.saturating_add(content)
|
||||
}
|
||||
|
||||
fn estimate_parts_tokens(parts: &[ContentPart]) -> u64 {
|
||||
parts.iter().fold(0_u64, |total, part| {
|
||||
let tokens = match part {
|
||||
ContentPart::Text { text } => estimate_text_tokens(text),
|
||||
ContentPart::Image { mime_type, .. } => {
|
||||
TOKENS_PER_IMAGE.saturating_add(estimate_text_tokens(mime_type))
|
||||
}
|
||||
};
|
||||
total.saturating_add(tokens)
|
||||
})
|
||||
}
|
||||
|
||||
fn estimate_json_tokens(value: &impl serde::Serialize) -> u64 {
|
||||
serde_json::to_string(value)
|
||||
.map(|value| estimate_text_tokens(&value))
|
||||
.unwrap_or_default()
|
||||
}
|
||||
|
||||
fn estimate_text_tokens(text: &str) -> u64 {
|
||||
let text = text.trim();
|
||||
if text.is_empty() {
|
||||
return 0;
|
||||
}
|
||||
let characters = text.chars().count() as u64;
|
||||
characters
|
||||
.div_ceil(4)
|
||||
.saturating_add(text.bytes().filter(|byte| *byte == b'\n').count() as u64)
|
||||
.max(1)
|
||||
}
|
||||
|
||||
pub(crate) fn parse_token_count(value: &str) -> Option<u64> {
|
||||
let value = value.trim().to_ascii_lowercase();
|
||||
let (number, multiplier) = match value.chars().last()? {
|
||||
@@ -18,3 +113,136 @@ pub(crate) fn format_token_count(tokens: u64) -> String {
|
||||
tokens.to_string()
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::model::{
|
||||
ProviderReplayState, Role, ToolCallContent, ToolDefinition, ToolResultContent,
|
||||
};
|
||||
|
||||
fn prompt() -> PromptSpec {
|
||||
PromptSpec {
|
||||
instructions: "system instructions".into(),
|
||||
tools: vec![ToolDefinition {
|
||||
name: "Read".into(),
|
||||
description: "Read a file".into(),
|
||||
parameters: serde_json::json!({"type": "object", "properties": {"path": {"type": "string"}}}),
|
||||
}],
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn context_estimate_grows_with_provider_visible_text_and_tools() {
|
||||
let short = vec![ProjectedMessage {
|
||||
message_id: "short".into(),
|
||||
role: Role::User,
|
||||
content: ProjectedContent::Parts(vec![ContentPart::Text {
|
||||
text: "hello".into(),
|
||||
}]),
|
||||
}];
|
||||
let long = vec![ProjectedMessage {
|
||||
message_id: "long".into(),
|
||||
role: Role::User,
|
||||
content: ProjectedContent::Parts(vec![ContentPart::Text {
|
||||
text: "x".repeat(40_000),
|
||||
}]),
|
||||
}];
|
||||
let without_tools = PromptSpec {
|
||||
instructions: prompt().instructions,
|
||||
tools: Vec::new(),
|
||||
};
|
||||
|
||||
assert!(
|
||||
estimate_context_tokens(&prompt(), &short)
|
||||
> estimate_context_tokens(&without_tools, &short)
|
||||
);
|
||||
assert!(
|
||||
estimate_context_tokens(&prompt(), &long) > estimate_context_tokens(&prompt(), &short)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn context_estimate_counts_tool_calls_results_and_images() {
|
||||
let assistant = ProjectedMessage {
|
||||
message_id: "assistant".into(),
|
||||
role: Role::Assistant,
|
||||
content: ProjectedContent::Assistant {
|
||||
text: String::new(),
|
||||
thinking: "reasoning".into(),
|
||||
replay_state: None,
|
||||
calls: vec![ToolCallContent {
|
||||
index: 0,
|
||||
call_id: "call-1".into(),
|
||||
name: "Read".into(),
|
||||
arguments: serde_json::json!({"path": "/tmp/file"}),
|
||||
}],
|
||||
},
|
||||
};
|
||||
let text_result = ProjectedMessage {
|
||||
message_id: "result-text".into(),
|
||||
role: Role::Tool,
|
||||
content: ProjectedContent::ToolResult(ToolResultContent {
|
||||
call_id: "call-1".into(),
|
||||
name: "Read".into(),
|
||||
content: "file contents".into(),
|
||||
is_error: false,
|
||||
image: None,
|
||||
provider_parts: Vec::new(),
|
||||
}),
|
||||
};
|
||||
let image_result = ProjectedMessage {
|
||||
message_id: "result-image".into(),
|
||||
role: Role::Tool,
|
||||
content: ProjectedContent::ToolResult(ToolResultContent {
|
||||
call_id: "call-1".into(),
|
||||
name: "Read".into(),
|
||||
content: "file contents".into(),
|
||||
is_error: false,
|
||||
image: None,
|
||||
provider_parts: vec![ContentPart::Image {
|
||||
mime_type: "image/png".into(),
|
||||
data: vec![0; 32],
|
||||
}],
|
||||
}),
|
||||
};
|
||||
|
||||
let base = estimate_context_tokens(&prompt(), &[]);
|
||||
let with_call = estimate_context_tokens(&prompt(), std::slice::from_ref(&assistant));
|
||||
let with_text = estimate_context_tokens(&prompt(), &[assistant.clone(), text_result]);
|
||||
let with_image = estimate_context_tokens(&prompt(), &[assistant, image_result]);
|
||||
assert!(with_call > base);
|
||||
assert!(with_text > with_call);
|
||||
assert!(with_image > with_text);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn assistant_replay_state_does_not_duplicate_thinking_or_count_signature() {
|
||||
let assistant = |replay_state| ProjectedMessage {
|
||||
message_id: "assistant".into(),
|
||||
role: Role::Assistant,
|
||||
content: ProjectedContent::Assistant {
|
||||
text: "answer".into(),
|
||||
thinking: "reasoning".repeat(1_000),
|
||||
replay_state,
|
||||
calls: Vec::new(),
|
||||
},
|
||||
};
|
||||
let without_replay = assistant(None);
|
||||
let with_replay = assistant(Some(ProviderReplayState {
|
||||
provider_kind: "anthropic".into(),
|
||||
value: serde_json::json!({
|
||||
"blocks": [{
|
||||
"type": "thinking",
|
||||
"thinking": "reasoning".repeat(1_000),
|
||||
"signature": "s".repeat(282_100)
|
||||
}]
|
||||
}),
|
||||
}));
|
||||
|
||||
assert_eq!(
|
||||
estimate_projected_messages_tokens(&[without_replay]),
|
||||
estimate_projected_messages_tokens(&[with_replay])
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -4,6 +4,24 @@ use serde_json::Value;
|
||||
|
||||
use super::ProviderReplayState;
|
||||
|
||||
pub fn normalize_tool_name(name: &str) -> String {
|
||||
let normalized = name
|
||||
.chars()
|
||||
.map(|character| {
|
||||
if character.is_ascii_alphanumeric() || matches!(character, '_' | '-') {
|
||||
character
|
||||
} else {
|
||||
'_'
|
||||
}
|
||||
})
|
||||
.collect::<String>();
|
||||
if normalized.is_empty() {
|
||||
"_".into()
|
||||
} else {
|
||||
normalized
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
|
||||
pub struct ToolDefinition {
|
||||
pub name: String,
|
||||
@@ -19,6 +37,7 @@ pub struct ToolCall {
|
||||
pub name: String,
|
||||
pub arguments_text: String,
|
||||
pub arguments: Value,
|
||||
pub argument_error: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
|
||||
|
||||
+87
-2
@@ -1,8 +1,93 @@
|
||||
//! Provides shared network client and transport configuration.
|
||||
//! Outbound HTTP clients configured from persisted application proxy settings.
|
||||
//! Owns reusable outbound HTTP clients configured from persisted proxy settings.
|
||||
|
||||
use std::{sync::Arc, time::Duration};
|
||||
|
||||
use tokio::sync::RwLock;
|
||||
|
||||
use crate::{store::Store, Result};
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct NetworkClients {
|
||||
store: Store,
|
||||
cache: Arc<RwLock<ClientCache>>,
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct ClientCache {
|
||||
default: Option<reqwest::Client>,
|
||||
cursor: Option<reqwest::Client>,
|
||||
provider: Option<(Duration, reqwest::Client)>,
|
||||
}
|
||||
|
||||
impl NetworkClients {
|
||||
pub fn new(store: Store) -> Self {
|
||||
Self {
|
||||
store,
|
||||
cache: Arc::new(RwLock::new(ClientCache::default())),
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn default_client(&self) -> Result<reqwest::Client> {
|
||||
if let Some(client) = self.cache.read().await.default.clone() {
|
||||
return Ok(client);
|
||||
}
|
||||
let mut cache = self.cache.write().await;
|
||||
if let Some(client) = cache.default.clone() {
|
||||
return Ok(client);
|
||||
}
|
||||
let client = client_builder(&self.store).await?.build()?;
|
||||
cache.default = Some(client.clone());
|
||||
Ok(client)
|
||||
}
|
||||
|
||||
pub async fn cursor_client(&self) -> Result<reqwest::Client> {
|
||||
if let Some(client) = self.cache.read().await.cursor.clone() {
|
||||
return Ok(client);
|
||||
}
|
||||
let mut cache = self.cache.write().await;
|
||||
if let Some(client) = cache.cursor.clone() {
|
||||
return Ok(client);
|
||||
}
|
||||
let client = client_builder(&self.store)
|
||||
.await?
|
||||
.redirect(reqwest::redirect::Policy::none())
|
||||
.build()?;
|
||||
cache.cursor = Some(client.clone());
|
||||
Ok(client)
|
||||
}
|
||||
|
||||
pub async fn provider_client(&self, timeout: Duration) -> Result<reqwest::Client> {
|
||||
if let Some((_, client)) = self
|
||||
.cache
|
||||
.read()
|
||||
.await
|
||||
.provider
|
||||
.as_ref()
|
||||
.filter(|(cached_timeout, _)| *cached_timeout == timeout)
|
||||
{
|
||||
return Ok(client.clone());
|
||||
}
|
||||
let mut cache = self.cache.write().await;
|
||||
if let Some((_, client)) = cache
|
||||
.provider
|
||||
.as_ref()
|
||||
.filter(|(cached_timeout, _)| *cached_timeout == timeout)
|
||||
{
|
||||
return Ok(client.clone());
|
||||
}
|
||||
let client = client_builder(&self.store)
|
||||
.await?
|
||||
.timeout(timeout)
|
||||
.build()?;
|
||||
cache.provider = Some((timeout, client.clone()));
|
||||
Ok(client)
|
||||
}
|
||||
|
||||
pub async fn invalidate(&self) {
|
||||
*self.cache.write().await = ClientCache::default();
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn client_builder(store: &Store) -> Result<reqwest::ClientBuilder> {
|
||||
let settings = store.proxy_settings_secret().await?;
|
||||
// Use the platform TLS stack for compatibility with provider gateways that
|
||||
|
||||
@@ -107,9 +107,7 @@ pub struct PluginModelDescriptor {
|
||||
pub description: Option<String>,
|
||||
pub icon: String,
|
||||
pub provider_type: String,
|
||||
pub context_window_tokens: Option<u64>,
|
||||
pub max_output_tokens: Option<u64>,
|
||||
pub thinking: bool,
|
||||
pub images: bool,
|
||||
}
|
||||
|
||||
@@ -209,9 +207,7 @@ impl PluginModelDescriptor {
|
||||
description: model.description.clone(),
|
||||
icon: icon.to_owned(),
|
||||
provider_type: provider.provider_type.clone(),
|
||||
context_window_tokens: model.context_window_tokens,
|
||||
max_output_tokens: model.max_output_tokens,
|
||||
thinking: model.thinking,
|
||||
images: model.images,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -20,7 +20,6 @@ pub use descriptor::{
|
||||
};
|
||||
pub use registry::{ImportResponse, OAuthBeginResponse, OAuthPollResponse, PluginRegistry};
|
||||
pub use runtime::{PluginRuntime, PluginRuntimePhase, PluginRuntimeState, PluginRuntimeStatus};
|
||||
pub(crate) use wire::llm_request as plugin_llm_request;
|
||||
|
||||
/// Windows 下阻止 Deno 子进程弹出控制台窗口(CREATE_NO_WINDOW)。
|
||||
#[cfg(windows)]
|
||||
|
||||
@@ -20,8 +20,11 @@ use super::{
|
||||
worker::{PluginWorker, WorkerStreamItem},
|
||||
};
|
||||
use crate::{
|
||||
model::ModelInvocation, provider::ModelEvent, provider::ProviderStream, store::Store, Error,
|
||||
Result,
|
||||
model::ModelInvocation,
|
||||
provider::ProviderStream,
|
||||
provider::{CallRecorder, ModelEvent},
|
||||
store::Store,
|
||||
Error, Result,
|
||||
};
|
||||
|
||||
const OAUTH_SLOW_DOWN_STEP_MS: i64 = 5_000;
|
||||
@@ -203,6 +206,7 @@ impl PluginRegistry {
|
||||
&self,
|
||||
invocation: ModelInvocation,
|
||||
cancellation: CancellationToken,
|
||||
recorder: CallRecorder,
|
||||
) -> ProviderStream {
|
||||
let registry = self.clone();
|
||||
Box::pin(try_stream! {
|
||||
@@ -232,7 +236,7 @@ impl PluginRegistry {
|
||||
"request": request,
|
||||
});
|
||||
let worker = registry.worker(&entry, &executable).await;
|
||||
let mut items = worker.invoke_streaming("provider.invoke", params, cancellation.clone()).await?;
|
||||
let mut items = worker.invoke_streaming("provider.invoke", params, cancellation.clone(), Some(recorder)).await?;
|
||||
yield ModelEvent::Start { model_call_id: invocation.call_id.clone() };
|
||||
while let Some(item) = items.recv().await {
|
||||
match item {
|
||||
|
||||
@@ -2,7 +2,6 @@ import type { JsonValue, PluginContext } from "./plugin.ts";
|
||||
import type { ResourceSnapshot } from "./resource.ts";
|
||||
|
||||
export type ModelCapabilities = {
|
||||
thinking?: boolean;
|
||||
images?: boolean;
|
||||
};
|
||||
|
||||
@@ -10,7 +9,6 @@ export type ModelDefinition = {
|
||||
id: string;
|
||||
displayName: string;
|
||||
description?: string;
|
||||
contextWindowTokens?: number;
|
||||
maxOutputTokens?: number;
|
||||
capabilities?: ModelCapabilities;
|
||||
/** 之后的调用原样传回;永远不会展示给用户。 */
|
||||
|
||||
@@ -53,7 +53,17 @@ function replayItems(value: JsonValue): JsonValue[] {
|
||||
if (!Array.isArray(items)) {
|
||||
throw new Error("OpenAI Responses replay state is missing items");
|
||||
}
|
||||
return items;
|
||||
return items.map((item) => {
|
||||
const source = record(item);
|
||||
if (source?.type !== "reasoning") {
|
||||
throw new Error("OpenAI Responses replay state contains a non-reasoning item");
|
||||
}
|
||||
const projected: Record<string, JsonValue> = { type: "reasoning" };
|
||||
for (const field of ["id", "summary", "content", "encrypted_content"] as const) {
|
||||
if (field in source) projected[field] = source[field] as JsonValue;
|
||||
}
|
||||
return projected;
|
||||
});
|
||||
}
|
||||
|
||||
export function buildResponsesBody(call: OpenAiResponsesCall): Record<string, JsonValue> {
|
||||
|
||||
@@ -135,12 +135,8 @@ pub struct StoredModel {
|
||||
#[serde(default)]
|
||||
pub description: Option<String>,
|
||||
#[serde(default)]
|
||||
pub context_window_tokens: Option<u64>,
|
||||
#[serde(default)]
|
||||
pub max_output_tokens: Option<u64>,
|
||||
#[serde(default)]
|
||||
pub thinking: bool,
|
||||
#[serde(default)]
|
||||
pub images: bool,
|
||||
#[serde(default)]
|
||||
pub private_data: serde_json::Value,
|
||||
@@ -179,13 +175,9 @@ impl StoredModel {
|
||||
.get("description")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::to_owned),
|
||||
context_window_tokens: object
|
||||
.get("contextWindowTokens")
|
||||
.and_then(serde_json::Value::as_u64),
|
||||
max_output_tokens: object
|
||||
.get("maxOutputTokens")
|
||||
.and_then(serde_json::Value::as_u64),
|
||||
thinking: capability("thinking"),
|
||||
images: capability("images"),
|
||||
private_data: object
|
||||
.get("privateData")
|
||||
@@ -200,9 +192,8 @@ impl StoredModel {
|
||||
"id": self.id,
|
||||
"displayName": self.display_name,
|
||||
"description": self.description,
|
||||
"contextWindowTokens": self.context_window_tokens,
|
||||
"maxOutputTokens": self.max_output_tokens,
|
||||
"capabilities": { "thinking": self.thinking, "images": self.images },
|
||||
"capabilities": { "images": self.images },
|
||||
"privateData": self.private_data,
|
||||
})
|
||||
}
|
||||
@@ -450,7 +441,7 @@ mod tests {
|
||||
let model = StoredModel::from_definition(&serde_json::json!({
|
||||
"id": "gpt-test",
|
||||
"displayName": "GPT Test",
|
||||
"capabilities": {"thinking": true},
|
||||
"capabilities": {"images": true},
|
||||
"privateData": {"reasoningEfforts": ["low"]},
|
||||
}))
|
||||
.unwrap();
|
||||
@@ -460,7 +451,7 @@ mod tests {
|
||||
.unwrap();
|
||||
let models = store.models("dev.example", "codex").await.unwrap();
|
||||
assert_eq!(models.len(), 1);
|
||||
assert!(models[0].thinking);
|
||||
assert!(models[0].images);
|
||||
assert_eq!(models[0].private_data["reasoningEfforts"][0], "low");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -147,8 +147,10 @@ pub fn model_event(value: &serde_json::Value) -> Result<ModelEvent> {
|
||||
.get("usage")
|
||||
.ok_or_else(|| Error::Protocol("plugin usage event requires usage".into()))?;
|
||||
let tokens = |name: &str| usage.get(name).and_then(serde_json::Value::as_u64);
|
||||
let input_tokens = tokens("inputTokens");
|
||||
ModelEvent::Usage(Usage {
|
||||
input_tokens: tokens("inputTokens"),
|
||||
input_tokens,
|
||||
context_input_tokens: input_tokens,
|
||||
output_tokens: tokens("outputTokens"),
|
||||
total_tokens: tokens("totalTokens"),
|
||||
cache_read_tokens: tokens("cacheReadTokens"),
|
||||
@@ -285,6 +287,7 @@ mod tests {
|
||||
usage,
|
||||
ModelEvent::Usage(Usage {
|
||||
input_tokens: Some(10),
|
||||
context_input_tokens: Some(10),
|
||||
output_tokens: Some(2),
|
||||
total_tokens: None,
|
||||
cache_read_tokens: Some(4),
|
||||
|
||||
+269
-24
@@ -3,7 +3,10 @@ use std::{
|
||||
collections::{HashMap, HashSet},
|
||||
path::PathBuf,
|
||||
process::Stdio,
|
||||
sync::Arc,
|
||||
sync::{
|
||||
atomic::{AtomicBool, Ordering},
|
||||
Arc,
|
||||
},
|
||||
time::Duration,
|
||||
};
|
||||
|
||||
@@ -19,7 +22,7 @@ use super::{
|
||||
definition::{file_url, PluginDefinitionLoader},
|
||||
protocol::{HostMessage, WorkerMessage},
|
||||
};
|
||||
use crate::{store::Store, Error, Result};
|
||||
use crate::{provider::CallRecorder, store::Store, Error, Result};
|
||||
|
||||
const INVOCATION_TIMEOUT: Duration = Duration::from_secs(10 * 60);
|
||||
const MAX_NETWORK_RESPONSE_BYTES: u64 = 16 * 1024 * 1024;
|
||||
@@ -56,12 +59,29 @@ struct WorkerProcess {
|
||||
stdin: Arc<Mutex<ChildStdin>>,
|
||||
}
|
||||
|
||||
struct InvocationState {
|
||||
cancellation: CancellationToken,
|
||||
recorder: Option<CallRecorder>,
|
||||
recorder_claimed: AtomicBool,
|
||||
}
|
||||
|
||||
impl InvocationState {
|
||||
fn claim_recorder(&self) -> Option<CallRecorder> {
|
||||
self.recorder.as_ref().and_then(|recorder| {
|
||||
self.recorder_claimed
|
||||
.compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire)
|
||||
.ok()
|
||||
.map(|_| recorder.clone())
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct HostContext {
|
||||
plugin_id: String,
|
||||
network_hosts: Arc<HashSet<String>>,
|
||||
store: Store,
|
||||
cancellations: Arc<Mutex<HashMap<String, CancellationToken>>>,
|
||||
invocations: Arc<Mutex<HashMap<String, Arc<InvocationState>>>>,
|
||||
streams: Arc<Mutex<HashMap<String, StreamLines>>>,
|
||||
}
|
||||
|
||||
@@ -87,7 +107,7 @@ impl PluginWorker {
|
||||
.collect(),
|
||||
),
|
||||
store,
|
||||
cancellations: Arc::new(Mutex::new(HashMap::new())),
|
||||
invocations: Arc::new(Mutex::new(HashMap::new())),
|
||||
streams: Arc::new(Mutex::new(HashMap::new())),
|
||||
},
|
||||
plugin_id,
|
||||
@@ -108,7 +128,9 @@ impl PluginWorker {
|
||||
params: serde_json::Value,
|
||||
cancellation: CancellationToken,
|
||||
) -> Result<serde_json::Value> {
|
||||
let mut items = self.invoke_streaming(method, params, cancellation).await?;
|
||||
let mut items = self
|
||||
.invoke_streaming(method, params, cancellation, None)
|
||||
.await?;
|
||||
let result = tokio::time::timeout(INVOCATION_TIMEOUT, async {
|
||||
while let Some(item) = items.recv().await {
|
||||
if let WorkerStreamItem::Result(result) = item {
|
||||
@@ -137,15 +159,18 @@ impl PluginWorker {
|
||||
method: &str,
|
||||
params: serde_json::Value,
|
||||
cancellation: CancellationToken,
|
||||
recorder: Option<CallRecorder>,
|
||||
) -> Result<mpsc::UnboundedReceiver<WorkerStreamItem>> {
|
||||
let id = uuid::Uuid::new_v4().to_string();
|
||||
let request_cancellation = CancellationToken::new();
|
||||
self.inner
|
||||
.host
|
||||
.cancellations
|
||||
.lock()
|
||||
.await
|
||||
.insert(id.clone(), request_cancellation.clone());
|
||||
self.inner.host.invocations.lock().await.insert(
|
||||
id.clone(),
|
||||
Arc::new(InvocationState {
|
||||
cancellation: request_cancellation.clone(),
|
||||
recorder,
|
||||
recorder_claimed: AtomicBool::new(false),
|
||||
}),
|
||||
);
|
||||
let (sender, receiver) = mpsc::unbounded_channel();
|
||||
self.inner
|
||||
.pending
|
||||
@@ -182,10 +207,10 @@ impl PluginWorker {
|
||||
}
|
||||
let _ = sender.send(WorkerStreamItem::Result(Err(Error::Cancelled)));
|
||||
inner.pending.lock().await.remove(&request_id);
|
||||
inner.host.cancellations.lock().await.remove(&request_id);
|
||||
inner.host.invocations.lock().await.remove(&request_id);
|
||||
}
|
||||
_ = sender.closed() => {
|
||||
inner.host.cancellations.lock().await.remove(&request_id);
|
||||
inner.host.invocations.lock().await.remove(&request_id);
|
||||
}
|
||||
}
|
||||
});
|
||||
@@ -201,7 +226,7 @@ impl PluginWorker {
|
||||
|
||||
async fn cleanup(&self, id: &str) {
|
||||
self.inner.pending.lock().await.remove(id);
|
||||
self.inner.host.cancellations.lock().await.remove(id);
|
||||
self.inner.host.invocations.lock().await.remove(id);
|
||||
}
|
||||
|
||||
async fn stdin(&self) -> Result<Arc<Mutex<ChildStdin>>> {
|
||||
@@ -393,6 +418,28 @@ async fn fail_pending(pending: &Pending, message: &str) {
|
||||
}
|
||||
}
|
||||
|
||||
fn recorded_network_request(
|
||||
params: &serde_json::Value,
|
||||
) -> Result<(serde_json::Value, serde_json::Value)> {
|
||||
let mut recorded_headers = serde_json::Map::new();
|
||||
if let Some(headers) = params.get("headers").and_then(serde_json::Value::as_object) {
|
||||
for (name, value) in headers {
|
||||
let value = value.as_str().ok_or_else(|| {
|
||||
Error::Config(format!("plugin HTTP header '{name}' must be a string"))
|
||||
})?;
|
||||
if !crate::model::is_sensitive_header(name) {
|
||||
recorded_headers.insert(name.clone(), value.into());
|
||||
}
|
||||
}
|
||||
}
|
||||
let body = params
|
||||
.get("body")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(|body| serde_json::from_str(body).unwrap_or_else(|_| body.into()))
|
||||
.unwrap_or(serde_json::Value::Null);
|
||||
Ok((serde_json::Value::Object(recorded_headers), body))
|
||||
}
|
||||
|
||||
impl HostContext {
|
||||
async fn call(
|
||||
&self,
|
||||
@@ -421,7 +468,11 @@ impl HostContext {
|
||||
&self,
|
||||
request_id: &str,
|
||||
params: &serde_json::Value,
|
||||
) -> Result<(reqwest::RequestBuilder, CancellationToken)> {
|
||||
) -> Result<(
|
||||
reqwest::RequestBuilder,
|
||||
CancellationToken,
|
||||
Option<CallRecorder>,
|
||||
)> {
|
||||
let raw_url = required_string(params, "url")?;
|
||||
let url = url::Url::parse(raw_url)
|
||||
.map_err(|error| Error::Config(format!("invalid plugin network URL: {error}")))?;
|
||||
@@ -463,14 +514,17 @@ impl HostContext {
|
||||
if let Some(body) = params.get("body").and_then(serde_json::Value::as_str) {
|
||||
request = request.body(body.to_owned());
|
||||
}
|
||||
let cancellation = self
|
||||
.cancellations
|
||||
.lock()
|
||||
.await
|
||||
.get(request_id)
|
||||
.cloned()
|
||||
let invocation = self.invocations.lock().await.get(request_id).cloned();
|
||||
let cancellation = invocation
|
||||
.as_ref()
|
||||
.map(|state| state.cancellation.clone())
|
||||
.unwrap_or_default();
|
||||
Ok((request, cancellation))
|
||||
let recorder = invocation.and_then(|state| state.claim_recorder());
|
||||
if let Some(recorder) = &recorder {
|
||||
let (headers, body) = recorded_network_request(params)?;
|
||||
recorder.request(headers, &body).await?;
|
||||
}
|
||||
Ok((request, cancellation, recorder))
|
||||
}
|
||||
|
||||
async fn fetch(
|
||||
@@ -478,13 +532,16 @@ impl HostContext {
|
||||
request_id: &str,
|
||||
params: serde_json::Value,
|
||||
) -> Result<serde_json::Value> {
|
||||
let (request, cancellation) = self.request(request_id, ¶ms).await?;
|
||||
let (request, cancellation, recorder) = self.request(request_id, ¶ms).await?;
|
||||
let request = request.timeout(Duration::from_secs(60));
|
||||
let response = tokio::select! {
|
||||
_ = cancellation.cancelled() => return Err(Error::Cancelled),
|
||||
response = request.send() => response?,
|
||||
};
|
||||
let status = response.status().as_u16();
|
||||
if let Some(recorder) = &recorder {
|
||||
recorder.response_headers(status).await?;
|
||||
}
|
||||
if response
|
||||
.content_length()
|
||||
.is_some_and(|size| size > MAX_NETWORK_RESPONSE_BYTES)
|
||||
@@ -503,6 +560,9 @@ impl HostContext {
|
||||
"plugin network response is larger than allowed".into(),
|
||||
));
|
||||
}
|
||||
if let Some(recorder) = &recorder {
|
||||
recorder.response_chunk(&body).await?;
|
||||
}
|
||||
Ok(
|
||||
serde_json::json!({ "status": status, "headers": headers, "body": String::from_utf8_lossy(&body) }),
|
||||
)
|
||||
@@ -514,12 +574,15 @@ impl HostContext {
|
||||
request_id: &str,
|
||||
params: serde_json::Value,
|
||||
) -> Result<serde_json::Value> {
|
||||
let (request, cancellation) = self.request(request_id, ¶ms).await?;
|
||||
let (request, cancellation, recorder) = self.request(request_id, ¶ms).await?;
|
||||
let response = tokio::select! {
|
||||
_ = cancellation.cancelled() => return Err(Error::Cancelled),
|
||||
response = request.send() => response?,
|
||||
};
|
||||
let status = response.status().as_u16();
|
||||
if let Some(recorder) = &recorder {
|
||||
recorder.response_headers(status).await?;
|
||||
}
|
||||
let headers = header_map(&response);
|
||||
let (sender, receiver) = mpsc::channel::<Result<String>>(256);
|
||||
tokio::spawn(async move {
|
||||
@@ -552,6 +615,12 @@ impl HostContext {
|
||||
.await;
|
||||
return;
|
||||
}
|
||||
if let Some(recorder) = &recorder {
|
||||
if let Err(error) = recorder.response_chunk(&chunk).await {
|
||||
let _ = sender.send(Err(error)).await;
|
||||
return;
|
||||
}
|
||||
}
|
||||
buffered.extend_from_slice(&chunk);
|
||||
while let Some(position) = buffered.iter().position(|byte| *byte == b'\n') {
|
||||
let mut line = buffered.drain(..=position).collect::<Vec<u8>>();
|
||||
@@ -645,3 +714,179 @@ fn required_string<'a>(params: &'a serde_json::Value, key: &str) -> Result<&'a s
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.ok_or_else(|| Error::Protocol(format!("plugin host call requires string '{key}'")))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use crate::{
|
||||
model::{NewLlmCall, ProviderType},
|
||||
provider::{CallRecorder, FinishReason},
|
||||
store::Store,
|
||||
};
|
||||
|
||||
use super::*;
|
||||
|
||||
async fn recorder(detailed: bool, call_id: &str) -> (tempfile::TempDir, Store, CallRecorder) {
|
||||
let directory = tempfile::tempdir().unwrap();
|
||||
let store = Store::connect(&format!(
|
||||
"sqlite://{}",
|
||||
directory.path().join("test.db").display()
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
store.set_detailed_logging(detailed).await.unwrap();
|
||||
let recorder = CallRecorder::start(
|
||||
store.clone(),
|
||||
NewLlmCall {
|
||||
call_id: call_id.into(),
|
||||
run_id: "run".into(),
|
||||
conversation_id: "conversation".into(),
|
||||
provider_call_index: 0,
|
||||
model_hash: "plugin:test/provider/model".into(),
|
||||
provider_type: ProviderType::Plugin,
|
||||
provider_url: "plugin://test/provider".into(),
|
||||
request_type: ProviderType::Plugin,
|
||||
request_url: "plugin://test/provider".into(),
|
||||
model_id: "model".into(),
|
||||
display_name: "Model".into(),
|
||||
reasoning_effort: None,
|
||||
fast: false,
|
||||
message_count: 1,
|
||||
tool_count: 0,
|
||||
detailed: false,
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
(directory, store, recorder)
|
||||
}
|
||||
|
||||
fn network_params() -> serde_json::Value {
|
||||
serde_json::json!({
|
||||
"url": "https://example.com/v1/responses",
|
||||
"method": "POST",
|
||||
"headers": {
|
||||
"Authorization": "Bearer secret",
|
||||
"X-Api-Key": "secret-key",
|
||||
"Cookie": "session=secret",
|
||||
"content-type": "application/json",
|
||||
"x-client-request-id": "request-1"
|
||||
},
|
||||
"body": "{\"model\":\"test\",\"stream\":true}"
|
||||
})
|
||||
}
|
||||
|
||||
async fn host_with_recorder(store: Store, recorder: CallRecorder) -> HostContext {
|
||||
let invocations = Arc::new(Mutex::new(HashMap::new()));
|
||||
invocations.lock().await.insert(
|
||||
"invocation".into(),
|
||||
Arc::new(InvocationState {
|
||||
cancellation: CancellationToken::new(),
|
||||
recorder: Some(recorder),
|
||||
recorder_claimed: AtomicBool::new(false),
|
||||
}),
|
||||
);
|
||||
HostContext {
|
||||
plugin_id: "test".into(),
|
||||
network_hosts: Arc::new(HashSet::from(["example.com".into()])),
|
||||
store,
|
||||
invocations,
|
||||
streams: Arc::new(Mutex::new(HashMap::new())),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn recorded_plugin_request_omits_sensitive_headers() {
|
||||
let (headers, body) = recorded_network_request(&network_params()).unwrap();
|
||||
|
||||
assert_eq!(
|
||||
headers,
|
||||
serde_json::json!({
|
||||
"content-type": "application/json",
|
||||
"x-client-request-id": "request-1"
|
||||
})
|
||||
);
|
||||
assert_eq!(body, serde_json::json!({ "model": "test", "stream": true }));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn detailed_plugin_network_recording_persists_request_and_raw_response() {
|
||||
let (_directory, store, recorder) = recorder(true, "detailed-plugin").await;
|
||||
let host = host_with_recorder(store.clone(), recorder.clone()).await;
|
||||
let params = network_params();
|
||||
let (_, _, first_recorder) = host.request("invocation", ¶ms).await.unwrap();
|
||||
let (_, _, second_recorder) = host.request("invocation", ¶ms).await.unwrap();
|
||||
let (_, body) = recorded_network_request(¶ms).unwrap();
|
||||
|
||||
assert!(first_recorder.is_some());
|
||||
assert!(second_recorder.is_none());
|
||||
recorder.response_headers(200).await.unwrap();
|
||||
recorder
|
||||
.response_chunk(b"data: {\"type\":\"response.created\"}\n\n")
|
||||
.await
|
||||
.unwrap();
|
||||
recorder.response_chunk(b"data: [DONE]\n\n").await.unwrap();
|
||||
recorder.completed(FinishReason::Stop).await.unwrap();
|
||||
|
||||
let request = store
|
||||
.llm_call_request("detailed-plugin")
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
request.headers,
|
||||
serde_json::json!({
|
||||
"content-type": "application/json",
|
||||
"x-client-request-id": "request-1"
|
||||
})
|
||||
);
|
||||
assert_eq!(request.body, body);
|
||||
let chunks = store.llm_call_chunks("detailed-plugin").await.unwrap();
|
||||
let expected_response = "data: {\"type\":\"response.created\"}\n\ndata: [DONE]\n\n";
|
||||
assert_eq!(chunks.len(), 2);
|
||||
assert_eq!(
|
||||
chunks
|
||||
.iter()
|
||||
.map(|chunk| chunk.data.as_str())
|
||||
.collect::<String>(),
|
||||
expected_response
|
||||
);
|
||||
let summary = store.llm_call("detailed-plugin").await.unwrap().unwrap();
|
||||
assert_eq!(summary.http_status, Some(200));
|
||||
assert_eq!(summary.stream_event_count, 2);
|
||||
assert_eq!(summary.response_bytes, expected_response.len() as i64);
|
||||
assert!(summary.detailed);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn standard_plugin_network_recording_keeps_metrics_without_payloads() {
|
||||
let (_directory, store, recorder) = recorder(false, "standard-plugin").await;
|
||||
let host = host_with_recorder(store.clone(), recorder.clone()).await;
|
||||
let params = network_params();
|
||||
let (_, body) = recorded_network_request(¶ms).unwrap();
|
||||
let request_bytes = serde_json::to_string(&body).unwrap().len() as i64;
|
||||
let response = b"data: [DONE]\n\n";
|
||||
|
||||
let (_, _, observed) = host.request("invocation", ¶ms).await.unwrap();
|
||||
assert!(observed.is_some());
|
||||
recorder.response_headers(204).await.unwrap();
|
||||
recorder.response_chunk(response).await.unwrap();
|
||||
recorder.completed(FinishReason::Stop).await.unwrap();
|
||||
|
||||
assert!(store
|
||||
.llm_call_request("standard-plugin")
|
||||
.await
|
||||
.unwrap()
|
||||
.is_none());
|
||||
assert!(store
|
||||
.llm_call_chunks("standard-plugin")
|
||||
.await
|
||||
.unwrap()
|
||||
.is_empty());
|
||||
let summary = store.llm_call("standard-plugin").await.unwrap().unwrap();
|
||||
assert_eq!(summary.http_status, Some(204));
|
||||
assert_eq!(summary.request_bytes, Some(request_bytes));
|
||||
assert_eq!(summary.response_bytes, response.len() as i64);
|
||||
assert_eq!(summary.stream_event_count, 1);
|
||||
assert!(!summary.detailed);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -12,9 +12,9 @@ use crate::{
|
||||
};
|
||||
|
||||
use super::{
|
||||
attempt::{send_once, Attempt},
|
||||
map_sse_error, merge_extra_params, provider_event_error,
|
||||
recorder::recorded_headers,
|
||||
retry::{send_with_retry, Attempt, RetryPolicy},
|
||||
CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream,
|
||||
};
|
||||
|
||||
@@ -90,17 +90,14 @@ impl Provider for AnthropicProvider {
|
||||
if let Some(recorder) = &recorder {
|
||||
recorder.request(request_headers.clone(), &body).await?;
|
||||
}
|
||||
let attempt = send_with_retry(
|
||||
let attempt = send_once(
|
||||
"Anthropic",
|
||||
|| client.post(&config.request_url)
|
||||
.header("x-api-key", &config.api_key).header("anthropic-version", "2023-06-01")
|
||||
.headers(config.custom_headers.clone())
|
||||
.json(&body),
|
||||
RetryPolicy { retries: config.retry_count, ..RetryPolicy::default() },
|
||||
&cancellation,
|
||||
recorder.as_ref(),
|
||||
request_headers,
|
||||
&body,
|
||||
).await?;
|
||||
let Attempt::Response(response) = attempt else { return };
|
||||
yield ModelEvent::Start { model_call_id: call_id };
|
||||
@@ -321,10 +318,19 @@ fn apply_model(body: &mut Value, model: &crate::model::ModelSpec) -> Result<()>
|
||||
|
||||
fn merge_usage(total: &mut Usage, update: Usage) {
|
||||
merge_usage_field(&mut total.input_tokens, update.input_tokens);
|
||||
merge_usage_field(&mut total.context_input_tokens, update.context_input_tokens);
|
||||
merge_usage_field(&mut total.output_tokens, update.output_tokens);
|
||||
merge_usage_field(&mut total.total_tokens, update.total_tokens);
|
||||
merge_usage_field(&mut total.cache_read_tokens, update.cache_read_tokens);
|
||||
merge_usage_field(&mut total.cache_write_tokens, update.cache_write_tokens);
|
||||
merge_usage_field(&mut total.reasoning_tokens, update.reasoning_tokens);
|
||||
if let Some(request_tokens) = total
|
||||
.context_input_tokens
|
||||
.zip(total.output_tokens)
|
||||
.and_then(|(input, output)| input.checked_add(output))
|
||||
{
|
||||
total.total_tokens = Some(request_tokens);
|
||||
}
|
||||
}
|
||||
|
||||
fn merge_usage_field(total: &mut Option<u64>, update: Option<u64>) {
|
||||
@@ -475,14 +481,67 @@ fn required_u64(value: &Value, name: &str) -> Result<u64> {
|
||||
}
|
||||
|
||||
fn anthropic_usage(value: &Value) -> Usage {
|
||||
let input_tokens = value.get("input_tokens").and_then(Value::as_u64);
|
||||
let cache_read_tokens = value.get("cache_read_input_tokens").and_then(Value::as_u64);
|
||||
let cache_write_tokens = value
|
||||
.get("cache_creation_input_tokens")
|
||||
.and_then(Value::as_u64);
|
||||
let context_input_tokens = input_tokens.map(|input| {
|
||||
input
|
||||
.saturating_add(cache_read_tokens.unwrap_or_default())
|
||||
.saturating_add(cache_write_tokens.unwrap_or_default())
|
||||
});
|
||||
Usage {
|
||||
input_tokens: value.get("input_tokens").and_then(Value::as_u64),
|
||||
input_tokens,
|
||||
context_input_tokens,
|
||||
output_tokens: value.get("output_tokens").and_then(Value::as_u64),
|
||||
total_tokens: value.get("total_tokens").and_then(Value::as_u64),
|
||||
cache_read_tokens: value.get("cache_read_input_tokens").and_then(Value::as_u64),
|
||||
cache_write_tokens: value
|
||||
.get("cache_creation_input_tokens")
|
||||
.and_then(Value::as_u64),
|
||||
cache_read_tokens,
|
||||
cache_write_tokens,
|
||||
reasoning_tokens: None,
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn cached_tokens_are_included_once_in_anthropic_context_input() {
|
||||
let usage = anthropic_usage(&serde_json::json!({
|
||||
"input_tokens": 10,
|
||||
"output_tokens": 5,
|
||||
"cache_read_input_tokens": 20,
|
||||
"cache_creation_input_tokens": 30
|
||||
}));
|
||||
|
||||
assert_eq!(usage.input_tokens, Some(10));
|
||||
assert_eq!(usage.context_input_tokens, Some(60));
|
||||
assert_eq!(usage.output_tokens, Some(5));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn streamed_anthropic_usage_includes_cached_input_in_total() {
|
||||
let mut usage = Usage::default();
|
||||
merge_usage(
|
||||
&mut usage,
|
||||
anthropic_usage(&serde_json::json!({
|
||||
"input_tokens": 414,
|
||||
"output_tokens": 0,
|
||||
"cache_read_input_tokens": 100_352,
|
||||
"cache_creation_input_tokens": 0
|
||||
})),
|
||||
);
|
||||
merge_usage(
|
||||
&mut usage,
|
||||
anthropic_usage(&serde_json::json!({"output_tokens": 191})),
|
||||
);
|
||||
|
||||
assert_eq!(usage.input_tokens, Some(414));
|
||||
assert_eq!(usage.cache_read_tokens, Some(100_352));
|
||||
assert_eq!(usage.cache_write_tokens, Some(0));
|
||||
assert_eq!(usage.context_input_tokens, Some(100_766));
|
||||
assert_eq!(usage.output_tokens, Some(191));
|
||||
assert_eq!(usage.total_tokens, Some(100_957));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,100 @@
|
||||
//! Sends one Provider HTTP attempt without applying retry policy.
|
||||
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
use crate::{Error, Result};
|
||||
|
||||
use super::CallRecorder;
|
||||
|
||||
#[derive(Debug)]
|
||||
pub(crate) enum Attempt {
|
||||
Response(reqwest::Response),
|
||||
Cancelled,
|
||||
}
|
||||
|
||||
pub(crate) async fn send_once<F>(
|
||||
label: &str,
|
||||
build: F,
|
||||
cancellation: &CancellationToken,
|
||||
recorder: Option<&CallRecorder>,
|
||||
) -> Result<Attempt>
|
||||
where
|
||||
F: FnOnce() -> reqwest::RequestBuilder,
|
||||
{
|
||||
let response = tokio::select! {
|
||||
_ = cancellation.cancelled() => return Ok(Attempt::Cancelled),
|
||||
response = build().send() => response,
|
||||
}?;
|
||||
if let Some(recorder) = recorder {
|
||||
recorder
|
||||
.response_headers(response.status().as_u16())
|
||||
.await?;
|
||||
}
|
||||
if response.status().is_success() {
|
||||
return Ok(Attempt::Response(response));
|
||||
}
|
||||
let status = response.status();
|
||||
let bytes = tokio::select! {
|
||||
_ = cancellation.cancelled() => return Ok(Attempt::Cancelled),
|
||||
bytes = response.bytes() => bytes,
|
||||
}?;
|
||||
Err(Error::Provider(format!(
|
||||
"{label} {status}: {}",
|
||||
String::from_utf8_lossy(&bytes)
|
||||
)))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
||||
|
||||
async fn server(response: &'static [u8]) -> String {
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let address = listener.local_addr().unwrap();
|
||||
tokio::spawn(async move {
|
||||
let (mut socket, _) = listener.accept().await.unwrap();
|
||||
let mut request = [0_u8; 1024];
|
||||
let _ = socket.read(&mut request).await;
|
||||
socket.write_all(response).await.unwrap();
|
||||
});
|
||||
format!("http://{address}")
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn non_success_status_is_one_failed_attempt() {
|
||||
let url =
|
||||
server(b"HTTP/1.1 503 Service Unavailable\r\nContent-Length: 4\r\n\r\ndown").await;
|
||||
let client = reqwest::Client::new();
|
||||
let error = send_once("test", || client.get(&url), &CancellationToken::new(), None)
|
||||
.await
|
||||
.unwrap_err();
|
||||
assert!(
|
||||
matches!(error, Error::Provider(message) if message.contains("503") && message.contains("down"))
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn response_body_transport_failure_is_one_failed_attempt() {
|
||||
let url =
|
||||
server(b"HTTP/1.1 500 Internal Server Error\r\nContent-Length: 100\r\n\r\nshort").await;
|
||||
let client = reqwest::Client::new();
|
||||
let error = send_once("test", || client.get(&url), &CancellationToken::new(), None)
|
||||
.await
|
||||
.unwrap_err();
|
||||
assert!(matches!(error, Error::Http(_)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn request_transport_failure_is_one_failed_attempt() {
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let address = listener.local_addr().unwrap();
|
||||
drop(listener);
|
||||
let url = format!("http://{address}");
|
||||
let client = reqwest::Client::new();
|
||||
let error = send_once("test", || client.get(&url), &CancellationToken::new(), None)
|
||||
.await
|
||||
.unwrap_err();
|
||||
assert!(matches!(error, Error::Http(_)));
|
||||
}
|
||||
}
|
||||
@@ -1,11 +1,11 @@
|
||||
//! Defines the provider interface and exports provider implementations.
|
||||
mod anthropic;
|
||||
mod attempt;
|
||||
mod event;
|
||||
mod normalize;
|
||||
mod openai_chat;
|
||||
mod openai_responses;
|
||||
mod recorder;
|
||||
mod retry;
|
||||
mod router;
|
||||
|
||||
use std::pin::Pin;
|
||||
|
||||
@@ -17,10 +17,10 @@ use crate::{
|
||||
};
|
||||
|
||||
use super::{
|
||||
apply_body_allowlist, apply_openai_prompt_cache_key, map_sse_error, merge_extra_params,
|
||||
provider_event_error,
|
||||
apply_body_allowlist, apply_openai_prompt_cache_key,
|
||||
attempt::{send_once, Attempt},
|
||||
map_sse_error, merge_extra_params, provider_event_error,
|
||||
recorder::recorded_headers,
|
||||
retry::{send_with_retry, Attempt, RetryPolicy},
|
||||
CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream,
|
||||
};
|
||||
|
||||
@@ -92,15 +92,12 @@ impl Provider for OpenAiChatProvider {
|
||||
if let Some(recorder) = &recorder {
|
||||
recorder.request(request_headers.clone(), &body).await?;
|
||||
}
|
||||
let attempt = send_with_retry(
|
||||
let attempt = send_once(
|
||||
"OpenAI Chat",
|
||||
|| client.post(&config.request_url)
|
||||
.bearer_auth(&config.api_key).headers(config.custom_headers.clone()).json(&body),
|
||||
RetryPolicy { retries: config.retry_count, ..RetryPolicy::default() },
|
||||
&cancellation,
|
||||
recorder.as_ref(),
|
||||
request_headers,
|
||||
&body,
|
||||
).await?;
|
||||
let Attempt::Response(response) = attempt else { return };
|
||||
yield ModelEvent::Start { model_call_id: call_id };
|
||||
@@ -421,8 +418,10 @@ fn merge_chat_fragment(target: &mut String, fragment: &str) {
|
||||
}
|
||||
|
||||
pub(crate) fn openai_usage(value: &Value) -> Usage {
|
||||
let input_tokens = value.get("prompt_tokens").and_then(Value::as_u64);
|
||||
Usage {
|
||||
input_tokens: value.get("prompt_tokens").and_then(Value::as_u64),
|
||||
input_tokens,
|
||||
context_input_tokens: input_tokens,
|
||||
output_tokens: value.get("completion_tokens").and_then(Value::as_u64),
|
||||
total_tokens: value.get("total_tokens").and_then(Value::as_u64),
|
||||
cache_read_tokens: value
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user