From e87abace8b7fc2c28a822e67f4909d1080dadb83 Mon Sep 17 00:00:00 2001 From: kevin9327 Date: Mon, 31 Aug 2026 19:20:30 +0900 Subject: [PATCH 01/32] fix(tests): acknowledge conversation Blob writes in the local rules test `local_markdown_rules_land_in_the_request_context_message` never finishes: `cargo test --workspace` fails on `main` with panicked at server\tests\local_rules_context.rs:64:14: run finishes within timeout: Elapsed(()) The Run publishes the conversation checkpoint by asking the client to write Blobs, and it does not continue until every `KvServerMessage` is answered with a `SetBlobResult`. The test drained the output stream without replying, so the Run stalled after the first frame, the provider was never invoked, and none of the assertions the test exists for were ever reached. Answer the Blob writes the way every other transport test already does (`error_lifecycle.rs`, `conversation_delivery.rs`, `interrupt.rs`). With the acknowledgement in place the Run reaches `EndStream` in ~0.3s and the original assertions run and pass, so `merge_local_rules` is now genuinely covered: exactly one `request-context:` message is projected and it carries `Always answer in haiku.`. No production code changes. Before: `cargo test -p cursor-server --test local_rules_context` -> FAILED (0 passed; 1 failed) after a 5s timeout After: `cargo test -p cursor-server --test local_rules_context` -> ok (1 passed; 0 failed) in 0.28s --- server/tests/local_rules_context.rs | 36 +++++++++++++++++++++++++---- 1 file changed, 31 insertions(+), 5 deletions(-) diff --git a/server/tests/local_rules_context.rs b/server/tests/local_rules_context.rs index 67d2bdc..fdc6e3c 100644 --- a/server/tests/local_rules_context.rs +++ b/server/tests/local_rules_context.rs @@ -16,6 +16,7 @@ use cursor_server::{ model::{ContentPart, ProjectedContent}, provider::{FinishReason, ModelEvent}, }; +use prost::Message; #[tokio::test] async fn local_markdown_rules_land_in_the_request_context_message() { @@ -58,18 +59,30 @@ async fn local_markdown_rules_land_in_the_request_context_message() { .await .unwrap(); + let mut append_seqno = 1; loop { let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv()) .await .expect("run finishes within timeout") .expect("output stays open until EndStream"); - let ended = connect::decode_frames(&frame) - .unwrap() - .iter() - .any(|(flags, _)| flags & connect::END_STREAM_FLAG != 0); - if ended { + let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap(); + if flags & connect::END_STREAM_FLAG != 0 { break; } + // The Run waits for the client to confirm every conversation Blob write, + // so the stream only advances once each KvServerMessage is acknowledged. + if let Some(pb::agent_server_message::Message::KvServerMessage(kv)) = + pb::AgentServerMessage::decode(payload).unwrap().message + { + handle + .command(TransportCommand::Append { + seqno: append_seqno, + message: Box::new(set_blob_result(kv.id)), + }) + .await + .unwrap(); + append_seqno += 1; + } } let requests = provider.requests(); @@ -102,6 +115,19 @@ async fn local_markdown_rules_land_in_the_request_context_message() { registry.shutdown().await; } +fn set_blob_result(id: u32) -> pb::AgentClientMessage { + pb::AgentClientMessage { + message: Some(pb::agent_client_message::Message::KvClientMessage( + pb::KvClientMessage { + id, + message: Some(pb::kv_client_message::Message::SetBlobResult( + pb::SetBlobResult { error: None }, + )), + }, + )), + } +} + fn user_run() -> pb::AgentClientMessage { pb::AgentClientMessage { message: Some(pb::agent_client_message::Message::RunRequest( From b6fc7326e7c3e998195006bf885ce39990320e3c Mon Sep 17 00:00:00 2001 From: kevin9327 Date: Mon, 31 Aug 2026 19:26:38 +0900 Subject: [PATCH 02/32] fix(provider): keep OpenAI Chat tool calls that finish with reason "stop" `map_finish` matched `"stop" | "content_filter"` before the `has_tools` fallback, so the fallback only ever applied to finish reasons the adapter did not recognise. When an OpenAI-compatible server streams tool calls and then reports `finish_reason: "stop"` the adapter yielded `Done(Stop)` even though `ToolCallStart` / `ToolCallEnd` had already been emitted. `run/model_cycle.rs` rejects that combination: let has_tool_calls = !calls.is_empty(); if matches!(finish_reason, FinishReason::ToolUse) != has_tool_calls { return Err(failure(RunFailure::Protocol( "finish reason and tool calls disagree".into()), ...)); } so the whole turn fails with a protocol error and the tool never runs. Servers that report `"stop"` alongside `tool_calls` are common in BYOK setups (llama.cpp, Ollama's OpenAI shim, several proxies), which makes those models unusable for anything agentic. Move the `has_tools` arm ahead of `"stop" | "content_filter"` so an observed tool call outranks the label the provider attached to the stop. This matches the two sibling adapters (`anthropic.rs` checks `_ if saw_tool` before every reason except `Length`, `openai_responses.rs` derives the reason from `saw_tool` alone) and this adapter's own `[DONE]` fallback, which already infers `ToolUse` from `!tools.is_empty()`. `"length"` still wins so a truncated response is still reported as truncated, and behaviour with no tool calls is byte-for-byte unchanged. Before (with the arm restored): cargo test -p cursor-server --lib provider::openai_chat -> observed_tool_calls_outrank_a_stop_finish_reason FAILED assertion `left == right` failed: left: Stop, right: ToolUse After: cargo test -p cursor-server --lib provider::openai_chat -> 3 passed --- server/src/provider/openai_chat.rs | 31 +++++++++++++++++++++++++++++- 1 file changed, 30 insertions(+), 1 deletion(-) diff --git a/server/src/provider/openai_chat.rs b/server/src/provider/openai_chat.rs index 32835a4..f2d962d 100644 --- a/server/src/provider/openai_chat.rs +++ b/server/src/provider/openai_chat.rs @@ -368,7 +368,9 @@ fn map_finish(value: &str, has_tools: bool) -> FinishReason { match value { "tool_calls" | "function_call" => FinishReason::ToolUse, "length" => FinishReason::Length, - "stop" | "content_filter" => FinishReason::Stop, + // Observed tool calls outrank whatever the provider labelled the stop: + // OpenAI-compatible servers routinely report "stop" while streaming tool + // calls, and the run rejects a finish reason that contradicts them. _ if has_tools => FinishReason::ToolUse, _ => FinishReason::Stop, } @@ -434,3 +436,30 @@ pub(crate) fn openai_usage(value: &Value) -> Usage { .and_then(Value::as_u64), } } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn observed_tool_calls_outrank_a_stop_finish_reason() { + assert_eq!(map_finish("stop", true), FinishReason::ToolUse); + assert_eq!(map_finish("content_filter", true), FinishReason::ToolUse); + assert_eq!(map_finish("", true), FinishReason::ToolUse); + } + + #[test] + fn finish_reason_mapping_without_tool_calls_is_unchanged() { + assert_eq!(map_finish("stop", false), FinishReason::Stop); + assert_eq!(map_finish("content_filter", false), FinishReason::Stop); + assert_eq!(map_finish("", false), FinishReason::Stop); + assert_eq!(map_finish("tool_calls", false), FinishReason::ToolUse); + assert_eq!(map_finish("function_call", false), FinishReason::ToolUse); + } + + #[test] + fn a_truncated_response_stays_truncated_even_with_tool_calls() { + assert_eq!(map_finish("length", true), FinishReason::Length); + assert_eq!(map_finish("length", false), FinishReason::Length); + } +} From 24c41d0347d013e5f2a3fd1e24ed0811836d0050 Mon Sep 17 00:00:00 2001 From: kevin9327 Date: Mon, 31 Aug 2026 19:35:39 +0900 Subject: [PATCH 03/32] fix(tools): keep result truncation inside its byte budget and terminating `truncate_text` looks for a fixed point where the kept prefix length equals the byte count printed in its own truncation notice. Two things go wrong when the limit is small, and both are reachable because two call sites pass a *remaining* budget rather than a constant: * `gate_grep_content` -> `truncate_text("Grep", .., budget.content_bytes)` * `gate_mcp` -> `truncate_text("MCP text", .., remaining_text)` 1. The loop can spin forever. `notice.len()` grows with the decimal digit count of `shown`, so `kept.len()` alternates between two values across a power-of- ten boundary and never equals `shown`. With `tool_name = "Grep"` this happens at `limit` 78 and 170; with `"MCP text"` at 82 and 174. The conversation task then spins at 100% CPU and the turn never completes. `truncate_middle` in `model/tool_result_replay.rs` already guards against exactly this. 2. When the notice itself does not fit, `available` saturates to 0 and the function returns the ~68 byte notice alone, i.e. *more* than `limit`. `gate_grep_content` then evaluates `budget.content_bytes -= <68 bytes>` on a budget of at most 68, which panics with "attempt to subtract with overflow" in debug/test builds and wraps in release, silently disabling the 32 KiB content cap for the rest of the result. Both are ordinary Grep results away: 16 matches of ~2 KiB leave a two-digit remainder of the 32 KiB budget, and the next match then hits the small-limit path. Fix: return a plain prefix when the notice cannot fit, and stop as soon as the kept length repeats a previous value. The reported byte count is then off by one at most in that rare oscillating case, and the result is guaranteed to be at most `limit` bytes. Every constant-limit call site is unaffected: their notices always fit and their limits do not oscillate. `budget.content_bytes` now uses `saturating_sub`, matching every other subtraction in this file. Verified against the unfixed function: cargo test -p cursor-server --lib gate::tests::truncate_text_never_exceeds_its_limit -> FAILED: limit 1 produced 67 bytes cargo test -p cursor-server --lib gate::tests::grep_content_gate_survives -> FAILED: panicked at gate.rs:261: attempt to subtract with overflow cargo test -p cursor-server --lib gate::tests::truncate_text_terminates -> never returns (killed after 45s) cargo test -p cursor-server --lib gate::tests::grep_content_gate_terminates -> never returns (killed after 60s) After: cargo test -p cursor-server --lib tools::tool_call_result::gate -> 4 passed --- .../src/cursor/tools/tool_call_result/gate.rs | 100 +++++++++++++++++- 1 file changed, 96 insertions(+), 4 deletions(-) diff --git a/server/src/cursor/tools/tool_call_result/gate.rs b/server/src/cursor/tools/tool_call_result/gate.rs index c174a5c..44d4d9b 100644 --- a/server/src/cursor/tools/tool_call_result/gate.rs +++ b/server/src/cursor/tools/tool_call_result/gate.rs @@ -258,7 +258,9 @@ fn gate_grep_content(content: &mut pb::GrepContentResult, budget: &mut GrepBudge truncated = true; break; } - budget.content_bytes -= next_match.content.len(); + budget.content_bytes = budget + .content_bytes + .saturating_sub(next_match.content.len()); budget.matches -= 1; next.matches.push(next_match); } @@ -630,15 +632,25 @@ fn truncate_text(tool_name: &str, content: &str, limit: usize) -> String { } let original = content.len(); let mut shown = limit; + let mut previous = None; loop { let notice = format!( "\n\n[truncated: {tool_name} result exceeded {limit} bytes; showing {shown} of {original} bytes]" ); - let available = limit.saturating_sub(notice.len()); - let kept = utf8_prefix(content, available); - if kept.len() == shown { + if notice.len() >= limit { + // The notice alone would blow the budget; keep a plain prefix so the + // result never costs more than `limit` bytes. + return utf8_prefix(content, limit).to_string(); + } + let kept = utf8_prefix(content, limit - notice.len()); + // `notice.len()` grows with the digit count of `shown`, so `kept.len()` + // can alternate between two values across a power-of-ten boundary + // instead of reaching a fixed point. Settle on the first repeat; the + // reported count is then off by one at most and the result still fits. + if kept.len() == shown || previous == Some(kept.len()) { return format!("{}{notice}", kept.trim_end_matches('\n')); } + previous = Some(shown); shown = kept.len(); } } @@ -685,3 +697,83 @@ fn utf8_suffix(value: &str, limit: usize) -> &str { } &value[start..] } + +#[cfg(test)] +mod tests { + use super::*; + + fn grep_tool(matches: Vec) -> pb::tool_call::Tool { + pb::tool_call::Tool::GrepToolCall(pb::GrepToolCall { + args: None, + result: Some(pb::GrepResult { + result: Some(pb::grep_result::Result::Success(pb::GrepSuccess { + active_editor_result: Some(pb::GrepUnionResult { + result: Some(pb::grep_union_result::Result::Content( + pb::GrepContentResult { + matches: vec![pb::GrepFileMatch { + file: "src/lib.rs".into(), + matches: matches + .into_iter() + .enumerate() + .map(|(index, content)| pb::GrepContentMatch { + line_number: index as i32 + 1, + content, + ..Default::default() + }) + .collect(), + }], + ..Default::default() + }, + )), + }), + ..Default::default() + })), + }), + }) + } + + #[test] + fn truncate_text_never_exceeds_its_limit() { + let content = "b".repeat(200); + for limit in 1..=250 { + let output = truncate_text("Grep", &content, limit); + assert!( + output.len() <= limit, + "limit {limit} produced {} bytes", + output.len() + ); + } + } + + #[test] + fn truncate_text_terminates_when_the_notice_length_oscillates() { + // `limit` values where the notice grows and shrinks with the digit count + // of the reported byte count, so the fixed point is never reached. + assert!(truncate_text("Grep", &"b".repeat(200), 78).len() <= 78); + assert!(truncate_text("Grep", &"b".repeat(200), 170).len() <= 170); + assert!(truncate_text("MCP text", &"b".repeat(500), 82).len() <= 82); + } + + #[test] + fn grep_content_gate_survives_a_nearly_exhausted_byte_budget() { + // 16 matches leave 16 bytes of the 32 KiB content budget, which is less + // than the truncation notice for the 17th match. + let mut matches = vec!["a".repeat(2047); 16]; + matches.push("b".repeat(100)); + let mut tool = grep_tool(matches); + let mut content = String::new(); + tool_completion("Grep", &mut tool, &mut content); + } + + #[test] + fn grep_content_gate_terminates_on_an_oscillating_remaining_budget() { + // The same path, tuned so the remaining budget lands on a `limit` where + // the truncation notice length oscillates. + let mut matches = vec!["a".repeat(2043); 15]; + matches.push("a".repeat(2045)); + matches.push("b".repeat(200)); + let mut tool = grep_tool(matches); + let mut content = String::new(); + tool_completion("Grep", &mut tool, &mut content); + } +} From ab4d3adee5e16ad277de5eba72f4a883d1c34aed Mon Sep 17 00:00:00 2001 From: kevin9327 Date: Mon, 31 Aug 2026 19:39:02 +0900 Subject: [PATCH 04/32] fix(model): stop token accumulation from erasing counts a later call omits `Usage::add_assign` folds every field through fn sum(left: Option, right: Option) -> Option { left?.checked_add(right?) } so `sum(Some(900), None)` is `None`, not `Some(900)`. Adding a record that does not report a field therefore does not leave the running total alone - it wipes it. Both accumulators are per-turn and run over every provider call in the turn: `run/engine.rs::accumulate_usage` (the compaction call plus each model cycle) and `cursor/conversation/output.rs::turn_usage`, which is what `turn_ended` reports to Cursor. Providers report these fields inconsistently *across calls of one turn*, which is all it takes. `openai_usage` reads `cache_read_tokens` from `prompt_tokens_details`, an object most OpenAI-compatible gateways omit until the prompt cache warms: call 1 (cold) yields `None`, call 2 yields `Some(1024)`, and the turn reports `None`. The same happens to `reasoning_tokens` when only some cycles reason, and to `total_tokens` on gateways that omit it from the streaming usage chunk. Any tool-using turn makes more than one call, so this is the normal case rather than an edge case. Treat an unreported count as zero and keep the field unknown only when neither side reported it. `checked_add` also becomes `saturating_add`: with the new rule, silently turning an overflow into `None` would be the same erasure by another route, and token counts never approach `u64::MAX` anyway. `context_input_tokens` is untouched. Before (with `left?.checked_add(right?)` restored): cargo test -p cursor-server --lib model::observability -> 2 failed: assertion `left == right` failed: left: None, right: Some(900) After: cargo test -p cursor-server --lib model::observability -> 3 passed --- server/src/model/observability.rs | 66 ++++++++++++++++++++++++++++++- 1 file changed, 65 insertions(+), 1 deletion(-) diff --git a/server/src/model/observability.rs b/server/src/model/observability.rs index 7fd8fb9..f50d541 100644 --- a/server/src/model/observability.rs +++ b/server/src/model/observability.rs @@ -44,8 +44,72 @@ mod usage { } } + /// Adds two optional counts, treating an unreported count as zero. + /// The result is only unknown when neither side reported the field. fn sum(left: Option, right: Option) -> Option { - left?.checked_add(right?) + if left.is_none() && right.is_none() { + return None; + } + Some( + left.unwrap_or_default() + .saturating_add(right.unwrap_or_default()), + ) + } + + #[cfg(test)] + mod tests { + use super::*; + + #[test] + fn a_record_without_cache_details_keeps_the_counts_already_accumulated() { + let mut total = Usage { + input_tokens: Some(1_000), + output_tokens: Some(20), + total_tokens: Some(1_020), + cache_read_tokens: Some(900), + ..Default::default() + }; + // A second provider call that omits `prompt_tokens_details`. + total += Usage { + input_tokens: Some(1_200), + output_tokens: Some(30), + total_tokens: Some(1_230), + ..Default::default() + }; + assert_eq!(total.input_tokens, Some(2_200)); + assert_eq!(total.output_tokens, Some(50)); + assert_eq!(total.total_tokens, Some(2_250)); + assert_eq!(total.cache_read_tokens, Some(900)); + } + + #[test] + fn a_late_reported_count_is_not_discarded() { + let mut total = Usage { + input_tokens: Some(1_000), + ..Default::default() + }; + total += Usage { + input_tokens: Some(1_200), + cache_read_tokens: Some(900), + reasoning_tokens: Some(64), + ..Default::default() + }; + assert_eq!(total.cache_read_tokens, Some(900)); + assert_eq!(total.reasoning_tokens, Some(64)); + } + + #[test] + fn a_field_no_call_reported_stays_unknown() { + let mut total = Usage { + input_tokens: Some(1_000), + ..Default::default() + }; + total += Usage { + input_tokens: Some(1_200), + ..Default::default() + }; + assert_eq!(total.cache_write_tokens, None); + } } } pub use usage::*; From 38a4c3f516f7ef0a9483ee35ea565017ab1ef306 Mon Sep 17 00:00:00 2001 From: kevin9327 Date: Mon, 31 Aug 2026 19:59:28 +0900 Subject: [PATCH 05/32] chore: restore a green `make check` on main MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit `make check` currently fails on main before any change is made: one `cargo fmt --all -- --check` diff and four `cargo clippy --workspace --all-targets -- -D warnings` errors. All five are pre-existing and none of them change behaviour. - `server/tests/knowledge_rules.rs:125` — rustfmt wants the long `assert!` split across lines. Applied `cargo fmt --all` verbatim. - `server/src/plugin/data.rs:206,215` — `path` is only read under `#[cfg(unix)]`, so every other target sees an unused binding. Added a `#[cfg(not(unix))] { let _ = path; }` arm, matching the `let _ = error;` idiom already used at line 189 of the same file. Windows behaviour is unchanged: these helpers stay no-ops there. - `server/src/provider/openai_responses.rs:158` — `collapsible_match`. Applied clippy's own suggestion (move `thinking_open` into a match guard). The match ends in `_ => {}`, so a failed guard falls through to a no-op exactly as the inner `if` did. - `server/src/store/models.rs:327` — `items_after_test_module`. Moved `optional_u64` and `to_i64` above `mod tests`; the bodies are untouched. Co-Authored-By: Claude Opus 4.8 --- server/src/plugin/data.rs | 8 ++++++++ server/src/provider/openai_responses.rs | 5 ++--- server/src/store/models.rs | 24 ++++++++++++------------ server/tests/knowledge_rules.rs | 5 ++++- 4 files changed, 26 insertions(+), 16 deletions(-) diff --git a/server/src/plugin/data.rs b/server/src/plugin/data.rs index f9205b1..2e0d1c1 100644 --- a/server/src/plugin/data.rs +++ b/server/src/plugin/data.rs @@ -209,6 +209,10 @@ fn set_directory_permissions(path: &Path) -> Result<()> { use std::os::unix::fs::PermissionsExt; std::fs::set_permissions(path, std::fs::Permissions::from_mode(0o700))?; } + #[cfg(not(unix))] + { + let _ = path; + } Ok(()) } @@ -218,6 +222,10 @@ fn set_file_permissions(path: &Path) -> Result<()> { use std::os::unix::fs::PermissionsExt; std::fs::set_permissions(path, std::fs::Permissions::from_mode(0o600))?; } + #[cfg(not(unix))] + { + let _ = path; + } Ok(()) } diff --git a/server/src/provider/openai_responses.rs b/server/src/provider/openai_responses.rs index 993f936..304267f 100644 --- a/server/src/provider/openai_responses.rs +++ b/server/src/provider/openai_responses.rs @@ -154,9 +154,8 @@ impl Provider for OpenAiResponsesProvider { if !thinking_open { thinking_open = true; yield ModelEvent::ThinkingStart; } if let Some(delta) = value.get("delta").and_then(Value::as_str) { yield ModelEvent::ThinkingDelta(delta.into()); } } - "response.reasoning_summary_text.done" | "response.reasoning_text.done" => { - if thinking_open { thinking_open = false; yield ModelEvent::ThinkingEnd; } - } + "response.reasoning_summary_text.done" | "response.reasoning_text.done" + if thinking_open => { thinking_open = false; yield ModelEvent::ThinkingEnd; } "response.output_item.added" => { let item = value.get("item").unwrap_or(&Value::Null); if item.get("type").and_then(Value::as_str) == Some("function_call") { diff --git a/server/src/store/models.rs b/server/src/store/models.rs index 0e6eeef..6d29d3c 100644 --- a/server/src/store/models.rs +++ b/server/src/store/models.rs @@ -323,6 +323,18 @@ fn model_from_row(row: sqlx::sqlite::SqliteRow) -> Result { }) } +fn optional_u64(row: &sqlx::sqlite::SqliteRow, column: &str) -> Result> { + row.try_get::, _>(column)? + .map(|value| { + u64::try_from(value).map_err(|_| Error::Config(format!("{column} cannot be negative"))) + }) + .transpose() +} + +fn to_i64(value: u64) -> Result { + i64::try_from(value).map_err(|_| Error::Config("token value is too large".into())) +} + #[cfg(test)] mod tests { use super::*; @@ -387,15 +399,3 @@ mod tests { assert_eq!(cleared.group_name, None); } } - -fn optional_u64(row: &sqlx::sqlite::SqliteRow, column: &str) -> Result> { - row.try_get::, _>(column)? - .map(|value| { - u64::try_from(value).map_err(|_| Error::Config(format!("{column} cannot be negative"))) - }) - .transpose() -} - -fn to_i64(value: u64) -> Result { - i64::try_from(value).map_err(|_| Error::Config("token value is too large".into())) -} diff --git a/server/tests/knowledge_rules.rs b/server/tests/knowledge_rules.rs index 1dc4bc5..14505d1 100644 --- a/server/tests/knowledge_rules.rs +++ b/server/tests/knowledge_rules.rs @@ -125,7 +125,10 @@ async fn offline_crud_round_trip_persists_markdown() { .unwrap(); let added: AddResponse = decode(response).await; assert!(added.success); - assert!(added.id.starts_with("local-"), "offline add uses a local id"); + assert!( + added.id.starts_with("local-"), + "offline add uses a local id" + ); let markdown = rules_root.join(format!("{}.md", added.id)); assert_eq!( std::fs::read_to_string(&markdown).unwrap(), From 86c3899016b5be04700435a98d2bde0a8ba486e7 Mon Sep 17 00:00:00 2001 From: kevin9327 Date: Mon, 31 Aug 2026 20:02:14 +0900 Subject: [PATCH 06/32] ci: run the Makefile check gates on pull requests `release.yml` only runs on `v*` tags, so nothing verifies a commit before it lands on `main`. `main` is currently red on three of the five gates `make check` defines, which is the drift this is meant to catch. Adds a single job running the three Rust gates verbatim as the Makefile spells them: cargo fmt --all -- --check cargo clippy --workspace --all-targets -- -D warnings cargo test --workspace --all-targets Conventions follow release.yml so the two workflows stay consistent: ubuntu-22.04, the same apt packages (the workspace includes apps/desktop/src-tauri, so even `cargo check` needs webkit2gtk), dtolnay/rust-toolchain@stable and Swatinem/rust-cache@v2 with the same `. -> target` workspace key. Co-Authored-By: Claude Opus 4.8 --- .github/workflows/ci.yml | 45 ++++++++++++++++++++++++++++++++++++++++ 1 file changed, 45 insertions(+) create mode 100644 .github/workflows/ci.yml diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml new file mode 100644 index 0000000..d839e05 --- /dev/null +++ b/.github/workflows/ci.yml @@ -0,0 +1,45 @@ +name: CI + +on: + pull_request: + push: + branches: + - main + +permissions: + contents: read + +concurrency: + group: ci-${{ github.ref }} + cancel-in-progress: true + +jobs: + rust: + name: Rust + runs-on: ubuntu-22.04 + steps: + - uses: actions/checkout@v4 + + # The workspace includes apps/desktop/src-tauri, so even `cargo check` + # needs the same system libraries the release workflow installs. + - name: Install Linux system dependencies + run: | + sudo apt-get update + sudo apt-get install -y libwebkit2gtk-4.1-dev libayatana-appindicator3-dev librsvg2-dev patchelf + + - uses: dtolnay/rust-toolchain@stable + with: + components: rustfmt, clippy + + - uses: Swatinem/rust-cache@v2 + with: + workspaces: ". -> target" + + - name: cargo fmt --all -- --check + run: cargo fmt --all -- --check + + - name: cargo clippy --workspace --all-targets -- -D warnings + run: cargo clippy --workspace --all-targets -- -D warnings + + - name: cargo test --workspace --all-targets + run: cargo test --workspace --all-targets From 1269110615c3d01fc230dc93f12ef4840008d84e Mon Sep 17 00:00:00 2001 From: kevin9327 Date: Mon, 31 Aug 2026 20:13:22 +0900 Subject: [PATCH 07/32] fix(tools): report the ReadLints paths that were never checked ReadLints accepts an array of paths, but codec::request encodes only paths[0] into the DiagnosticsArgs exec, and the exec protocol cannot carry more than one path. Nothing downstream mentions the drop: for {"paths": ["a.ts", "b.ts", "c.ts"]} the model is handed "No diagnostics found in a.ts" with is_error false, so it concludes b.ts and c.ts are clean when neither was ever opened. The existing truncation notice cannot catch this, since it compares total_diagnostics against a per-file diagnostics count. Name the unread paths in the model-facing result, using the bracketed notice convention already in this file and the ToolCall that output() is already given (as task() and render::read() already use it). Single-path calls are unchanged. This does not close the capability gap. A real fix fans out one DiagnosticsArgs exec per path and aggregates the results into the repeated FileDiagnostics that ReadLintsToolSuccess already defines, which needs multi-exec reservation and a completion barrier because take_exec drops the pending entry on the first result. That is a separate change; this one only stops the silent misinformation. Co-Authored-By: Claude Opus 4.8 --- .../tools/tool_call_result/exec/output.rs | 121 +++++++++++++++++- 1 file changed, 116 insertions(+), 5 deletions(-) diff --git a/server/src/cursor/tools/tool_call_result/exec/output.rs b/server/src/cursor/tools/tool_call_result/exec/output.rs index ab6dc19..316085c 100644 --- a/server/src/cursor/tools/tool_call_result/exec/output.rs +++ b/server/src/cursor/tools/tool_call_result/exec/output.rs @@ -12,7 +12,7 @@ pub(super) fn output( Message::WriteResult(value) => write(value), Message::DeleteResult(value) => delete(value), Message::GrepResult(value) => grep(value), - Message::DiagnosticsResult(value) => diagnostics(value), + Message::DiagnosticsResult(value) => diagnostics(value, call), Message::McpResult(value) => mcp(value), Message::ReadMcpResourceExecResult(value) => read_mcp(value), Message::SubagentResult(value) => task(value, call), @@ -212,14 +212,14 @@ fn grep_truncation( } } -fn diagnostics(value: &pb::DiagnosticsResult) -> Result<(String, bool)> { +fn diagnostics(value: &pb::DiagnosticsResult, call: &ToolCall) -> Result<(String, bool)> { use pb::diagnostics_result::Result as R; match value .result .as_ref() .ok_or_else(|| missing("diagnostics"))? { - R::Success(value) => Ok((diagnostics_success(value), false)), + R::Success(value) => Ok((diagnostics_success(value, call), false)), R::Error(value) => Ok((value.error.clone(), true)), R::Rejected(value) => Ok((value.reason.clone(), true)), R::FileNotFound(value) => Ok((format!("file not found: {}", value.path), true)), @@ -227,9 +227,40 @@ fn diagnostics(value: &pb::DiagnosticsResult) -> Result<(String, bool)> { } } -fn diagnostics_success(value: &pb::DiagnosticsSuccess) -> String { +/// `codec::request` encodes only `paths[0]` into the `DiagnosticsArgs` exec, and +/// `DiagnosticsSuccess` carries a single `path`, so a multi-path ReadLints call +/// only ever inspects the first entry. Name the rest instead of letting a clean +/// result for one file read as a clean bill of health for all of them. +fn unchecked_lint_paths(call: &ToolCall) -> Vec<&str> { + call.arguments + .get("paths") + .and_then(serde_json::Value::as_array) + .map(|paths| { + paths + .iter() + .skip(1) + .filter_map(serde_json::Value::as_str) + .filter(|path| !path.is_empty()) + .collect() + }) + .unwrap_or_default() +} + +fn diagnostics_success(value: &pb::DiagnosticsSuccess, call: &ToolCall) -> String { + let unchecked = unchecked_lint_paths(call); + let notice = (!unchecked.is_empty()).then(|| { + format!( + "[Only {} was checked; ReadLints reads one path per call. Not checked: {}]", + value.path, + unchecked.join(", ") + ) + }); if value.diagnostics.is_empty() { - return format!("No diagnostics found in {}", value.path); + let clean = format!("No diagnostics found in {}", value.path); + return match notice { + Some(notice) => format!("{clean}\n{notice}"), + None => clean, + }; } let mut lines = value .diagnostics @@ -261,6 +292,7 @@ fn diagnostics_success(value: &pb::DiagnosticsSuccess) -> String { value.diagnostics.len() )); } + lines.extend(notice); lines.join("\n") } @@ -413,3 +445,82 @@ fn creates_subagent(call: &ToolCall) -> bool { fn missing(name: &str) -> Error { Error::Protocol(format!("{name} returned no result")) } + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + fn read_lints(paths: serde_json::Value) -> ToolCall { + ToolCall { + index: 0, + call_id: "lints".into(), + model_call_id: "model:0".into(), + name: "ReadLints".into(), + arguments_text: String::new(), + arguments: json!({ "paths": paths }), + } + } + + fn clean_result(path: &str) -> pb::exec_client_message::Message { + pb::exec_client_message::Message::DiagnosticsResult(pb::DiagnosticsResult { + result: Some(pb::diagnostics_result::Result::Success( + pb::DiagnosticsSuccess { + path: path.into(), + diagnostics: Vec::new(), + total_diagnostics: 0, + }, + )), + }) + } + + /// The codec encodes only `paths[0]` into the DiagnosticsArgs exec, so a + /// clean result for that one path must not read as "these files are all + /// clean" for the paths that were never looked at. + #[test] + fn read_lints_names_the_paths_it_did_not_check() { + let call = read_lints(json!(["a.ts", "b.ts", "c.ts"])); + let (content, is_error) = output(&clean_result("a.ts"), &call).unwrap(); + assert!(!is_error); + assert!(content.contains("a.ts")); + assert!( + content.contains("b.ts") && content.contains("c.ts"), + "unchecked paths must be reported, got: {content}" + ); + } + + /// The overwhelmingly common single-path call must be byte-for-byte + /// unchanged. + #[test] + fn read_lints_single_path_result_is_unchanged() { + let call = read_lints(json!(["a.ts"])); + let (content, _) = output(&clean_result("a.ts"), &call).unwrap(); + assert_eq!(content, "No diagnostics found in a.ts"); + } + + /// The notice belongs on a result that did report diagnostics too: those + /// diagnostics are still only a.ts's. + #[test] + fn read_lints_reports_unchecked_paths_alongside_diagnostics() { + let call = read_lints(json!(["a.ts", "b.ts"])); + let message = pb::exec_client_message::Message::DiagnosticsResult(pb::DiagnosticsResult { + result: Some(pb::diagnostics_result::Result::Success( + pb::DiagnosticsSuccess { + path: "a.ts".into(), + diagnostics: vec![pb::Diagnostic { + message: "unused import".into(), + severity: pb::DiagnosticSeverity::Warning as i32, + ..Default::default() + }], + total_diagnostics: 1, + }, + )), + }); + let (content, _) = output(&message, &call).unwrap(); + assert!(content.contains("unused import")); + assert!( + content.contains("Not checked: b.ts"), + "unchecked paths must be reported, got: {content}" + ); + } +} From d004139526ce0b1db0580e06a42c6f320c6f8d7c Mon Sep 17 00:00:00 2001 From: leookun Date: Tue, 1 Sep 2026 16:07:51 +0800 Subject: [PATCH 08/32] feat: implement context usage anchor for improved token estimation - Introduced `ContextUsageAnchor` struct to track context input tokens and message count for conversations. - Updated token estimation functions to utilize the context usage anchor, enhancing accuracy in estimating tokens for projected messages. - Refactored compaction logic to incorporate context usage anchor, allowing for more efficient management of token budgets during model runs. - Added tests to validate the behavior of the context usage anchor across different scenarios, including model switching and message additions. --- server/src/model/token_count.rs | 53 ++++++++++++--- server/src/run/compaction.rs | 114 ++++++++++++++++++++++++++++++-- server/src/run/engine.rs | 46 ++++++++++++- server/src/store/llm_calls.rs | 98 +++++++++++++++++++++++++++ server/src/store/mod.rs | 2 +- server/tests/compaction.rs | 102 ++++++++++++++++++++++++++++ 6 files changed, 394 insertions(+), 21 deletions(-) diff --git a/server/src/model/token_count.rs b/server/src/model/token_count.rs index e1ae538..20a7e59 100644 --- a/server/src/model/token_count.rs +++ b/server/src/model/token_count.rs @@ -11,10 +11,15 @@ pub(crate) fn estimate_context_tokens(prompt: &PromptSpec, messages: &[Projected let tools = prompt.tools.iter().fold(0_u64, |total, tool| { total.saturating_add(estimate_json_tokens(tool)) }); - let messages = messages.iter().fold(0_u64, |total, message| { + 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)) - }); - instructions.saturating_add(tools).saturating_add(messages) + }) } fn estimate_message_tokens(message: &ProjectedMessage) -> u64 { @@ -23,7 +28,7 @@ fn estimate_message_tokens(message: &ProjectedMessage) -> u64 { ProjectedContent::Assistant { text, thinking, - replay_state, + replay_state: _, calls, } => { let calls = calls.iter().fold(0_u64, |total, call| { @@ -35,12 +40,6 @@ fn estimate_message_tokens(message: &ProjectedMessage) -> u64 { }); estimate_text_tokens(text) .saturating_add(estimate_text_tokens(thinking)) - .saturating_add( - replay_state - .as_ref() - .map(estimate_json_tokens) - .unwrap_or_default(), - ) .saturating_add(calls) } ProjectedContent::ToolResult(result) => { @@ -118,7 +117,9 @@ pub(crate) fn format_token_count(tokens: u64) -> String { #[cfg(test)] mod tests { use super::*; - use crate::model::{Role, ToolCallContent, ToolDefinition, ToolResultContent}; + use crate::model::{ + ProviderReplayState, Role, ToolCallContent, ToolDefinition, ToolResultContent, + }; fn prompt() -> PromptSpec { PromptSpec { @@ -214,4 +215,34 @@ mod tests { 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]) + ); + } } diff --git a/server/src/run/compaction.rs b/server/src/run/compaction.rs index 6fbe1a7..ec41441 100644 --- a/server/src/run/compaction.rs +++ b/server/src/run/compaction.rs @@ -2,7 +2,13 @@ use std::collections::HashSet; -use crate::model::{estimate_context_tokens, CanonicalMessage, PreparedRun, ProjectedMessage}; +use crate::{ + model::{ + estimate_context_tokens, estimate_projected_messages_tokens, CanonicalMessage, PreparedRun, + ProjectedMessage, + }, + store::ContextUsageAnchor, +}; const FALLBACK_CHARS: usize = 12_000; @@ -20,25 +26,36 @@ pub(super) fn input_budget(prepared: &PreparedRun) -> Option { pub(super) fn estimated_tokens( prepared: &PreparedRun, projected_messages: &[ProjectedMessage], + anchor: Option, ) -> u64 { - estimate_context_tokens(&prepared.prompt, projected_messages) + anchor + .filter(|anchor| anchor.message_count <= projected_messages.len()) + .map(|anchor| { + anchor + .context_input_tokens + .saturating_add(estimate_projected_messages_tokens( + &projected_messages[anchor.message_count..], + )) + }) + .unwrap_or_else(|| estimate_context_tokens(&prepared.prompt, projected_messages)) } pub(super) fn should_compact( prepared: &PreparedRun, projected_messages: &[ProjectedMessage], + anchor: Option, ) -> bool { let Some(budget) = input_budget(prepared) else { return false; }; - estimated_tokens(prepared, projected_messages) > budget + estimated_tokens(prepared, projected_messages, anchor) > budget } pub(super) fn validate_compacted( prepared: &PreparedRun, projected_messages: &[ProjectedMessage], ) -> std::result::Result { - let estimated = estimated_tokens(prepared, projected_messages); + let estimated = estimate_context_tokens(&prepared.prompt, projected_messages); let Some(budget) = input_budget(prepared) else { return Ok(estimated); }; @@ -125,14 +142,97 @@ mod tests { let estimated = estimate_context_tokens(&prepared(1).prompt, &projected); let mut prepared = prepared(estimated + RESERVE_TOKENS); - assert!(!should_compact(&prepared, &projected)); + assert!(!should_compact(&prepared, &projected, None)); prepared.model.context_window_tokens = Some(estimated + RESERVE_TOKENS - 1); - assert!(should_compact(&prepared, &projected)); + assert!(should_compact(&prepared, &projected, None)); prepared.action = RunAction::Resume { pending_tool_round: None, }; - assert!(should_compact(&prepared, &projected)); + assert!(should_compact(&prepared, &projected, None)); + } + + #[test] + fn provider_usage_anchor_only_estimates_messages_added_after_last_request() { + let messages = vec![ + CanonicalMessage::text("old", Role::User, Origin::Runtime, "x".repeat(400_000)), + CanonicalMessage::text("new", Role::User, Origin::Runtime, "short follow-up"), + ]; + let projected = project_messages(&messages).unwrap(); + let anchor = ContextUsageAnchor { + context_input_tokens: 103_904, + message_count: 1, + }; + let expected = 103_904 + estimate_projected_messages_tokens(&projected[1..]); + + assert_eq!( + estimated_tokens(&prepared(200_000), &projected, Some(anchor)), + expected + ); + assert!(!should_compact( + &prepared(200_000), + &projected, + Some(anchor) + )); + } + + #[test] + fn provider_usage_anchor_triggers_after_new_messages_cross_budget() { + let messages = vec![ + CanonicalMessage::text("old", Role::User, Origin::Runtime, "old"), + CanonicalMessage::text("new", Role::User, Origin::Runtime, "x".repeat(80_000)), + ]; + let projected = project_messages(&messages).unwrap(); + + assert!(should_compact( + &prepared(200_000), + &projected, + Some(ContextUsageAnchor { + context_input_tokens: 180_000, + message_count: 1, + }) + )); + } + + #[test] + fn missing_anchor_uses_full_fallback() { + let messages = vec![CanonicalMessage::text( + "user", + Role::User, + Origin::Runtime, + "x".repeat(40_000), + )]; + let projected = project_messages(&messages).unwrap(); + let prepared = prepared(200_000); + + assert_eq!( + estimated_tokens(&prepared, &projected, None), + estimate_context_tokens(&prepared.prompt, &projected) + ); + } + + #[test] + fn invalid_anchor_message_count_uses_full_fallback() { + let messages = vec![CanonicalMessage::text( + "user", + Role::User, + Origin::Runtime, + "x".repeat(40_000), + )]; + let projected = project_messages(&messages).unwrap(); + let expected = estimate_context_tokens(&prepared(200_000).prompt, &projected); + + assert_eq!( + estimated_tokens( + &prepared(200_000), + &projected, + Some(ContextUsageAnchor { + context_input_tokens: 1, + message_count: 2, + }) + ), + expected + ); } #[test] diff --git a/server/src/run/engine.rs b/server/src/run/engine.rs index 1195fd6..b83c695 100644 --- a/server/src/run/engine.rs +++ b/server/src/run/engine.rs @@ -10,7 +10,7 @@ use crate::{ ToolRoundId, Usage, }, provider::Provider, - store::{RunStatus, Store}, + store::{ContextUsageAnchor, RunStatus, Store}, }; use super::{ @@ -92,6 +92,14 @@ impl RunEngine { cancellation: &CancellationToken, ) -> (RunOutcome, Option) { let mut usage = None; + let mut context_usage_anchor = match self + .store + .latest_context_usage(prepared.conversation_id.as_str()) + .await + { + Ok(anchor) => anchor, + Err(error) => return (RunOutcome::Failed(error.into()), usage), + }; tracing::info!( checkpoint_id = checkpoint.0, "Run claimed conversation ownership" @@ -176,7 +184,7 @@ impl RunEngine { Err(error) => return (RunOutcome::Failed(error.into()), usage), }; if prepared.action != RunAction::Compact - && super::compaction::should_compact(prepared, &history) + && super::compaction::should_compact(prepared, &history, context_usage_anchor) { match self .auto_compact(prepared, checkpoint, &messages, client, cancellation) @@ -184,6 +192,7 @@ impl RunEngine { { Ok((next_checkpoint, compaction_usage)) => { checkpoint = next_checkpoint; + context_usage_anchor = None; if let Some(compaction_usage) = compaction_usage { accumulate_usage(&mut usage, compaction_usage); } @@ -272,11 +281,21 @@ impl RunEngine { match interrupted { Ok(cycle) => { if let Some(cycle_usage) = cycle.usage { + update_context_usage_anchor( + &mut context_usage_anchor, + cycle_usage, + request.history.len(), + ); accumulate_usage(&mut usage, cycle_usage); } } Err(failure) => { if let Some(cycle_usage) = failure.usage { + update_context_usage_anchor( + &mut context_usage_anchor, + cycle_usage, + request.history.len(), + ); accumulate_usage(&mut usage, cycle_usage); } } @@ -319,6 +338,11 @@ impl RunEngine { Ok(cycle) => break 'attempt cycle, Err(cycle_failure) => { if let Some(cycle_usage) = cycle_failure.usage { + update_context_usage_anchor( + &mut context_usage_anchor, + cycle_usage, + request.history.len(), + ); accumulate_usage(&mut usage, cycle_usage); } if cancellation.is_cancelled() { @@ -420,6 +444,11 @@ impl RunEngine { } }; if let Some(cycle_usage) = cycle.usage { + update_context_usage_anchor( + &mut context_usage_anchor, + cycle_usage, + request.history.len(), + ); accumulate_usage(&mut usage, cycle_usage); } @@ -879,6 +908,19 @@ async fn hydrate_tool_images( Ok(()) } +fn update_context_usage_anchor( + anchor: &mut Option, + usage: Usage, + message_count: usize, +) { + if let Some(context_input_tokens) = usage.context_input_tokens { + *anchor = Some(ContextUsageAnchor { + context_input_tokens, + message_count, + }); + } +} + fn accumulate_usage(total: &mut Option, usage: Usage) { match total { Some(total) => *total += usage, diff --git a/server/src/store/llm_calls.rs b/server/src/store/llm_calls.rs index b882071..b9db421 100644 --- a/server/src/store/llm_calls.rs +++ b/server/src/store/llm_calls.rs @@ -8,6 +8,12 @@ use crate::{ use super::{now_ms, Store}; +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) struct ContextUsageAnchor { + pub(crate) context_input_tokens: u64, + pub(crate) message_count: usize, +} + #[derive(Clone, Debug)] pub(crate) struct BufferedLlmChunk { pub(crate) seq: i64, @@ -266,6 +272,37 @@ impl Store { Ok(()) } + pub(crate) async fn latest_context_usage( + &self, + conversation_id: &str, + ) -> Result> { + let row = sqlx::query( + "SELECT usage_json, message_count FROM llm_calls + WHERE conversation_id = ? + AND json_extract(usage_json, '$.context_input_tokens') IS NOT NULL + ORDER BY created_at_ms DESC, rowid DESC + LIMIT 1", + ) + .bind(conversation_id) + .fetch_optional(&self.pool) + .await?; + let Some(row) = row else { + return Ok(None); + }; + let usage: Usage = serde_json::from_str(row.try_get("usage_json")?)?; + let Some(context_input_tokens) = usage.context_input_tokens else { + return Ok(None); + }; + let message_count = row.try_get::("message_count")?; + let Ok(message_count) = usize::try_from(message_count) else { + return Ok(None); + }; + Ok(Some(ContextUsageAnchor { + context_input_tokens, + message_count, + })) + } + pub async fn llm_calls(&self, limit: i64) -> Result> { let rows = sqlx::query("SELECT * FROM llm_calls ORDER BY created_at_ms DESC LIMIT ?") .bind(limit.clamp(1, 500)) @@ -421,4 +458,65 @@ mod tests { assert_eq!(overview.metrics.llm_calls, 1); assert_eq!(overview.metrics.successful_calls, 1); } + + #[tokio::test] + async fn latest_context_usage_follows_conversation_chronology() { + let directory = tempfile::tempdir().unwrap(); + let store = Store::connect(&format!( + "sqlite://{}", + directory.path().join("test.db").display() + )) + .await + .unwrap(); + + for (call_id, model_id, context_input_tokens, message_count) in [ + ("call-a-1", "model-a", 100_u64, 3_usize), + ("call-b", "model-b", 200_u64, 5_usize), + ("call-a-2", "model-a", 300_u64, 7_usize), + ] { + store + .start_llm_call(&NewLlmCall { + call_id: call_id.into(), + run_id: format!("run-{call_id}"), + conversation_id: "conversation".into(), + provider_call_index: 0, + model_hash: model_id.into(), + provider_type: ProviderType::Plugin, + provider_url: "plugin://test".into(), + request_type: ProviderType::Plugin, + request_url: "plugin://test".into(), + model_id: model_id.into(), + display_name: model_id.into(), + reasoning_effort: None, + fast: false, + message_count, + tool_count: 0, + detailed: false, + }) + .await + .unwrap(); + store + .record_llm_usage( + call_id, + Usage { + input_tokens: Some(context_input_tokens), + context_input_tokens: Some(context_input_tokens), + output_tokens: Some(10), + total_tokens: Some(context_input_tokens + 10), + ..Default::default() + }, + ) + .await + .unwrap(); + } + + assert_eq!( + store.latest_context_usage("conversation").await.unwrap(), + Some(ContextUsageAnchor { + context_input_tokens: 300, + message_count: 7, + }) + ); + assert_eq!(store.latest_context_usage("other").await.unwrap(), None); + } } diff --git a/server/src/store/mod.rs b/server/src/store/mod.rs index b9a0bf5..4043c42 100644 --- a/server/src/store/mod.rs +++ b/server/src/store/mod.rs @@ -19,7 +19,7 @@ mod writer; pub use cas::*; pub(crate) use cursor_traces::BufferedCursorTraceChunk; -pub(crate) use llm_calls::BufferedLlmChunk; +pub(crate) use llm_calls::{BufferedLlmChunk, ContextUsageAnchor}; pub use runs::*; pub use settings::*; pub(crate) use sqlite::now_ms; diff --git a/server/tests/compaction.rs b/server/tests/compaction.rs index d778a66..0476772 100644 --- a/server/tests/compaction.rs +++ b/server/tests/compaction.rs @@ -297,6 +297,108 @@ async fn automatic_compaction_preflights_provider_input_and_records_rebuilt_toke })); } +#[tokio::test] +async fn incremental_preflight_uses_conversation_anchor_across_model_switch() { + let (_directory, store) = fixtures::temp_store().await; + let model_a = store + .create_model(&ModelConfigInput { + sort_order: 0, + display_name: "Anchor Model A".into(), + group_name: None, + model_type: ModelType::OpenAi, + base_url: "https://example.com/v1/chat/completions".into(), + use_full_url: true, + api_key: "test-key".into(), + tooltip_data: "Anchor Model A".into(), + model_id: "anchor-model-a".into(), + reasoning_effort: None, + openai_endpoint: OPENAI_CHAT_ENDPOINT.into(), + openai_extra_params_enabled: false, + openai_extra_params: serde_json::json!({}), + custom_headers_enabled: false, + custom_headers: serde_json::json!({}), + anthropic_extra_params_enabled: false, + anthropic_extra_params: serde_json::json!({}), + context_window_tokens: None, + max_completion_tokens: None, + anthropic_max_tokens: None, + anthropic_thinking_effort: None, + thinking_budget_tokens: None, + }) + .await + .unwrap(); + let model_b = store + .create_model(&ModelConfigInput { + sort_order: 1, + display_name: "Anchor Model B".into(), + group_name: None, + model_type: ModelType::OpenAi, + base_url: "https://example.com/v1/chat/completions".into(), + use_full_url: true, + api_key: "test-key".into(), + tooltip_data: "Anchor Model B".into(), + model_id: "anchor-model-b".into(), + reasoning_effort: None, + openai_endpoint: OPENAI_CHAT_ENDPOINT.into(), + openai_extra_params_enabled: false, + openai_extra_params: serde_json::json!({}), + custom_headers_enabled: false, + custom_headers: serde_json::json!({}), + anthropic_extra_params_enabled: false, + anthropic_extra_params: serde_json::json!({}), + context_window_tokens: Some(200_000), + max_completion_tokens: None, + anthropic_max_tokens: None, + anthropic_thinking_effort: None, + thinking_budget_tokens: None, + }) + .await + .unwrap(); + let provider = fake_provider::FakeProvider::default(); + provider.push(text_response("old answer", 103_904, 12)); + provider.push(text_response("new answer", 104_000, 12)); + let assets = PromptAssets::load( + std::path::Path::new(env!("CARGO_MANIFEST_DIR")) + .join("prompt/cursor") + .as_path(), + ) + .unwrap(); + let registry = TransportRegistry::new( + store, + Arc::new(provider.clone()), + PromptCompiler::new(assets), + ); + + let first = run( + ®istry, + "anchor-first", + user_request( + "anchor-conversation", + "anchor-user-1", + &"x".repeat(400_000), + &model_a.model_hash, + None, + ), + ) + .await; + let second = run( + ®istry, + "anchor-second", + user_request( + "anchor-conversation", + "anchor-user-2", + "short follow-up", + &model_b.model_hash, + first.checkpoints.last().cloned(), + ), + ) + .await; + + assert_eq!(second.summary_started, 0); + assert_eq!(second.summary_completed, 0); + assert_eq!(provider.requests().len(), 2); +} + #[tokio::test] async fn irreducibly_oversized_current_input_fails_before_provider_dispatch() { let (_directory, store) = fixtures::temp_store().await; From 2c63bd845a64f50fde2c2bf57420e3e436d2e6ab Mon Sep 17 00:00:00 2001 From: leokun Date: Tue, 1 Sep 2026 16:54:45 +0800 Subject: [PATCH 09/32] feat: track interaction events during automatic compaction - Added tracking for interaction events in the `Output` struct, including `summary_started` and `token_delta`. - Updated the `run` function to push relevant interaction events to the `interaction_events` vector. - Enhanced the automatic compaction test to verify the immediate reset of cursor usage and the correct logging of interaction events. --- server/src/run/engine.rs | 14 ++++++++++++++ server/tests/compaction.rs | 16 ++++++++++++++-- 2 files changed, 28 insertions(+), 2 deletions(-) diff --git a/server/src/run/engine.rs b/server/src/run/engine.rs index b83c695..42b3808 100644 --- a/server/src/run/engine.rs +++ b/server/src/run/engine.rs @@ -703,6 +703,20 @@ impl RunEngine { emit(client, RunEvent::AutoCompactionStarted) .await .map_err(|_| client_failure())?; + emit( + client, + RunEvent::Usage(Usage { + input_tokens: Some(0), + context_input_tokens: Some(0), + output_tokens: Some(0), + total_tokens: Some(0), + cache_read_tokens: Some(0), + cache_write_tokens: Some(0), + reasoning_tokens: Some(0), + }), + ) + .await + .map_err(|_| client_failure())?; let provider_call_index = self .store .begin_provider_call(&prepared.run_id) diff --git a/server/tests/compaction.rs b/server/tests/compaction.rs index 0476772..2781544 100644 --- a/server/tests/compaction.rs +++ b/server/tests/compaction.rs @@ -267,6 +267,11 @@ async fn automatic_compaction_preflights_provider_input_and_records_rebuilt_toke .await; assert_eq!(second.summary_started, 1); assert_eq!(second.summary_completed, 1); + assert_eq!( + &second.interaction_events[..2], + &["summary_started", "token_delta:0"], + "automatic compaction must immediately reset Cursor usage" + ); let compacted_tokens = second .checkpoints .iter() @@ -469,6 +474,7 @@ struct Output { summary_completed: usize, turn_ended: usize, token_delta: usize, + interaction_events: Vec, } async fn run( @@ -517,7 +523,8 @@ async fn run( Some(pb::agent_server_message::Message::InteractionUpdate(update)) => { match update.message { Some(pb::interaction_update::Message::SummaryStarted(_)) => { - output.summary_started += 1 + output.summary_started += 1; + output.interaction_events.push("summary_started".into()); } Some(pb::interaction_update::Message::Summary(delta)) => { output.summary.push_str(&delta.summary) @@ -526,7 +533,12 @@ async fn run( output.summary_completed += 1 } Some(pb::interaction_update::Message::TurnEnded(_)) => output.turn_ended += 1, - Some(pb::interaction_update::Message::TokenDelta(_)) => output.token_delta += 1, + Some(pb::interaction_update::Message::TokenDelta(delta)) => { + output.token_delta += 1; + output + .interaction_events + .push(format!("token_delta:{}", delta.tokens)); + } _ => {} } } From e768980dadd96c8e96a6af483d53be6e82c2f591 Mon Sep 17 00:00:00 2001 From: leokun Date: Tue, 1 Sep 2026 17:43:36 +0800 Subject: [PATCH 10/32] feat: enhance reasoning replay functionality and integrate call recording - Added a new test to validate the projection of reasoning response items to valid input items in the Codex API. - Introduced `CallRecorder` to track network requests and responses during plugin interactions. - Updated the `PluginRegistry` and `PluginWorker` to support call recording, ensuring that reasoning items are correctly processed and recorded. - Refactored the `responses_input` function to handle reasoning items more effectively, improving the overall response handling logic. --- .../plugins/build-in/codex-auth/codex_test.ts | 38 +++ server/src/plugin/mod.rs | 1 - server/src/plugin/registry.rs | 10 +- .../plugin/sdk/protocol/openai_responses.ts | 12 +- server/src/plugin/worker.rs | 293 ++++++++++++++++-- server/src/provider/openai_responses.rs | 68 +++- server/src/provider/router.rs | 6 +- 7 files changed, 396 insertions(+), 32 deletions(-) diff --git a/server/plugins/build-in/codex-auth/codex_test.ts b/server/plugins/build-in/codex-auth/codex_test.ts index b572cdd..1149537 100644 --- a/server/plugins/build-in/codex-auth/codex_test.ts +++ b/server/plugins/build-in/codex-auth/codex_test.ts @@ -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, @@ -305,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({ diff --git a/server/src/plugin/mod.rs b/server/src/plugin/mod.rs index ce7faa3..7789c40 100644 --- a/server/src/plugin/mod.rs +++ b/server/src/plugin/mod.rs @@ -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)] diff --git a/server/src/plugin/registry.rs b/server/src/plugin/registry.rs index e107caf..29e22bc 100644 --- a/server/src/plugin/registry.rs +++ b/server/src/plugin/registry.rs @@ -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 { diff --git a/server/src/plugin/sdk/protocol/openai_responses.ts b/server/src/plugin/sdk/protocol/openai_responses.ts index b8e479a..b04e037 100644 --- a/server/src/plugin/sdk/protocol/openai_responses.ts +++ b/server/src/plugin/sdk/protocol/openai_responses.ts @@ -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 = { 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 { diff --git a/server/src/plugin/worker.rs b/server/src/plugin/worker.rs index 2863b04..b56c867 100644 --- a/server/src/plugin/worker.rs +++ b/server/src/plugin/worker.rs @@ -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>, } +struct InvocationState { + cancellation: CancellationToken, + recorder: Option, + recorder_claimed: AtomicBool, +} + +impl InvocationState { + fn claim_recorder(&self) -> Option { + 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>, store: Store, - cancellations: Arc>>, + invocations: Arc>>>, streams: Arc>>, } @@ -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 { - 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, ) -> Result> { 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>> { @@ -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, + )> { 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 { - 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 { - 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::>(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::>(); @@ -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::(), + 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); + } +} diff --git a/server/src/provider/openai_responses.rs b/server/src/provider/openai_responses.rs index 99eb7a7..6439ea4 100644 --- a/server/src/provider/openai_responses.rs +++ b/server/src/provider/openai_responses.rs @@ -420,7 +420,12 @@ fn responses_input(messages: &[ProjectedMessage]) -> Result> { .ok_or_else(|| { Error::Protocol("OpenAI Responses replay state is missing items".into()) })?; - input.extend(items.iter().cloned()); + input.extend( + items + .iter() + .map(response_reasoning_input) + .collect::>>()?, + ); } push_responses_text(&mut input, &message.role, text); for call in calls { @@ -437,6 +442,23 @@ fn responses_input(messages: &[ProjectedMessage]) -> Result> { Ok(input) } +fn response_reasoning_input(item: &Value) -> Result { + let source = item + .as_object() + .filter(|object| object.get("type").and_then(Value::as_str) == Some("reasoning")) + .ok_or_else(|| { + Error::Protocol("OpenAI Responses replay state contains a non-reasoning item".into()) + })?; + let mut projected = Map::new(); + projected.insert("type".into(), json!("reasoning")); + for field in ["id", "summary", "content", "encrypted_content"] { + if let Some(value) = source.get(field) { + projected.insert(field.into(), value.clone()); + } + } + Ok(Value::Object(projected)) +} + fn push_responses_parts(input: &mut Vec, role: &Role, parts: &[ContentPart]) -> Result<()> { let text_type = if *role == Role::Assistant { "output_text" @@ -517,3 +539,47 @@ fn responses_usage(value: &Value) -> Usage { .and_then(Value::as_u64), } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::model::ProviderReplayState; + + #[test] + fn reasoning_replay_projects_response_items_to_valid_input_items() { + let messages = [ProjectedMessage { + message_id: "assistant-1".into(), + role: Role::Assistant, + content: ProjectedContent::Assistant { + text: String::new(), + thinking: String::new(), + replay_state: Some(ProviderReplayState { + provider_kind: "openai_responses".into(), + value: json!({ + "items": [{ + "type": "reasoning", + "id": "item-1", + "status": "completed", + "summary": [{"type": "summary_text", "text": "why"}], + "content": [], + "encrypted_content": "opaque", + "output_only": true + }] + }), + }), + calls: Vec::new(), + }, + }]; + + assert_eq!( + responses_input(&messages).unwrap(), + vec![json!({ + "type": "reasoning", + "id": "item-1", + "summary": [{"type": "summary_text", "text": "why"}], + "content": [], + "encrypted_content": "opaque" + })] + ); + } +} diff --git a/server/src/provider/router.rs b/server/src/provider/router.rs index d16f122..7eabe64 100644 --- a/server/src/provider/router.rs +++ b/server/src/provider/router.rs @@ -62,7 +62,6 @@ impl Provider for ProviderRouter { let plan = plugins.plan_model(&selected).await?; let recorder = start_recorder(&store, &invocation, &selected, &plan.model.display_name, ProviderType::Plugin, &plan.request_url, &plan.model.model_id).await?; let guard = recorder.cancel_on_drop(); - recorder.request(serde_json::json!({}), &crate::plugin::plugin_llm_request(&invocation)?).await?; let mut routed = invocation.clone(); routed.request.model.display_name = Some(plan.model.display_name.clone()); if let Some(tokens) = plan.model.max_output_tokens { @@ -70,6 +69,7 @@ impl Provider for ProviderRouter { } let provider: Arc = Arc::new(NormalizedProvider::new(Arc::new(PluginModelProvider { registry: plugins.clone(), + recorder: recorder.clone(), }))); (recorder, guard, provider.stream(routed, cancellation.clone())) } else { @@ -227,6 +227,7 @@ async fn finish_stream(recorder: &CallRecorder, cancellation: &CancellationToken /// 插件模型的 Provider 实现;对路由与规范化层完全等同于内置 Provider。 struct PluginModelProvider { registry: PluginRegistry, + recorder: CallRecorder, } impl Provider for PluginModelProvider { @@ -235,7 +236,8 @@ impl Provider for PluginModelProvider { invocation: ModelInvocation, cancellation: CancellationToken, ) -> ProviderStream { - self.registry.stream_model(invocation, cancellation) + self.registry + .stream_model(invocation, cancellation, self.recorder.clone()) } } From 76417e005b4e0b96fbbcf1fee46b4681e56de72b Mon Sep 17 00:00:00 2001 From: leokun Date: Tue, 1 Sep 2026 20:04:35 +0800 Subject: [PATCH 11/32] feat: add usage snapshot event and enhance compaction logic - Introduced `UsageSnapshot` event to track token usage during conversation runs. - Updated `RunEngine` to emit usage snapshots, providing better visibility into token consumption. - Refactored compaction logic to utilize a new `compaction_estimate` function for improved token budget management. - Added tests to validate timeout constants for blob synchronization and ensure correct behavior of usage tracking during compaction. --- server/src/cursor/conversation/output.rs | 8 ++++ server/src/cursor/services/blob_sync.rs | 22 ++++++++++- server/src/run/compaction.rs | 16 ++++++-- server/src/run/engine.rs | 47 +++++++++++++++--------- server/src/run/event.rs | 1 + server/tests/compaction.rs | 14 +++++-- 6 files changed, 81 insertions(+), 27 deletions(-) diff --git a/server/src/cursor/conversation/output.rs b/server/src/cursor/conversation/output.rs index d66c54d..654937b 100644 --- a/server/src/cursor/conversation/output.rs +++ b/server/src/cursor/conversation/output.rs @@ -399,6 +399,14 @@ impl ConversationOutput { .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 { if let Some(output_tokens) = usage.output_tokens { diff --git a/server/src/cursor/services/blob_sync.rs b/server/src/cursor/services/blob_sync.rs index d9fa6c6..e88947a 100644 --- a/server/src/cursor/services/blob_sync.rs +++ b/server/src/cursor/services/blob_sync.rs @@ -20,6 +20,9 @@ use crate::{ type BlobSetSender = oneshot::Sender>; +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, @@ -130,7 +133,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); @@ -183,7 +186,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); @@ -318,3 +321,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)); + } +} diff --git a/server/src/run/compaction.rs b/server/src/run/compaction.rs index ec41441..9cb45b6 100644 --- a/server/src/run/compaction.rs +++ b/server/src/run/compaction.rs @@ -40,15 +40,23 @@ pub(super) fn estimated_tokens( .unwrap_or_else(|| estimate_context_tokens(&prepared.prompt, projected_messages)) } +pub(super) fn compaction_estimate( + prepared: &PreparedRun, + projected_messages: &[ProjectedMessage], + anchor: Option, +) -> Option { + let budget = input_budget(prepared)?; + let estimated = estimated_tokens(prepared, projected_messages, anchor); + (estimated > budget).then_some(estimated) +} + +#[cfg(test)] pub(super) fn should_compact( prepared: &PreparedRun, projected_messages: &[ProjectedMessage], anchor: Option, ) -> bool { - let Some(budget) = input_budget(prepared) else { - return false; - }; - estimated_tokens(prepared, projected_messages, anchor) > budget + compaction_estimate(prepared, projected_messages, anchor).is_some() } pub(super) fn validate_compacted( diff --git a/server/src/run/engine.rs b/server/src/run/engine.rs index 42b3808..4694f36 100644 --- a/server/src/run/engine.rs +++ b/server/src/run/engine.rs @@ -183,9 +183,21 @@ impl RunEngine { Ok(history) => history, Err(error) => return (RunOutcome::Failed(error.into()), usage), }; - if prepared.action != RunAction::Compact - && super::compaction::should_compact(prepared, &history, context_usage_anchor) - { + let compaction_estimate = (prepared.action != RunAction::Compact) + .then(|| { + super::compaction::compaction_estimate(prepared, &history, context_usage_anchor) + }) + .flatten(); + if let Some(estimated_tokens) = compaction_estimate { + if emit( + client, + RunEvent::UsageSnapshot(context_usage_snapshot(estimated_tokens)), + ) + .await + .is_err() + { + return (client_failure(), usage); + } match self .auto_compact(prepared, checkpoint, &messages, client, cancellation) .await @@ -703,20 +715,6 @@ impl RunEngine { emit(client, RunEvent::AutoCompactionStarted) .await .map_err(|_| client_failure())?; - emit( - client, - RunEvent::Usage(Usage { - input_tokens: Some(0), - context_input_tokens: Some(0), - output_tokens: Some(0), - total_tokens: Some(0), - cache_read_tokens: Some(0), - cache_write_tokens: Some(0), - reasoning_tokens: Some(0), - }), - ) - .await - .map_err(|_| client_failure())?; let provider_call_index = self .store .begin_provider_call(&prepared.run_id) @@ -862,6 +860,9 @@ impl RunEngine { emit(client, RunEvent::AutoCompactionCompleted) .await .map_err(|_| client_failure())?; + emit(client, RunEvent::UsageSnapshot(context_usage_snapshot(0))) + .await + .map_err(|_| client_failure())?; checkpoint = super::messages::append_batches( &self.store, prepared, @@ -922,6 +923,18 @@ async fn hydrate_tool_images( Ok(()) } +fn context_usage_snapshot(tokens: u64) -> Usage { + Usage { + input_tokens: Some(tokens), + context_input_tokens: Some(tokens), + output_tokens: Some(0), + total_tokens: Some(tokens), + cache_read_tokens: Some(0), + cache_write_tokens: Some(0), + reasoning_tokens: Some(0), + } +} + fn update_context_usage_anchor( anchor: &mut Option, usage: Usage, diff --git a/server/src/run/event.rs b/server/src/run/event.rs index 345be46..01918aa 100644 --- a/server/src/run/event.rs +++ b/server/src/run/event.rs @@ -129,6 +129,7 @@ pub enum RunEvent { ToolCallEnd { index: usize, }, + UsageSnapshot(Usage), Usage(Usage), ExecuteToolRound { round_id: ToolRoundId, diff --git a/server/tests/compaction.rs b/server/tests/compaction.rs index 2781544..4eeb8ee 100644 --- a/server/tests/compaction.rs +++ b/server/tests/compaction.rs @@ -268,9 +268,14 @@ async fn automatic_compaction_preflights_provider_input_and_records_rebuilt_toke assert_eq!(second.summary_started, 1); assert_eq!(second.summary_completed, 1); assert_eq!( - &second.interaction_events[..2], - &["summary_started", "token_delta:0"], - "automatic compaction must immediately reset Cursor usage" + &second.interaction_events[..4], + &[ + "token_delta:0", + "summary_started", + "summary_completed", + "token_delta:0", + ], + "automatic compaction must publish estimated usage before summarizing and zero usage after" ); let compacted_tokens = second .checkpoints @@ -530,7 +535,8 @@ async fn run( output.summary.push_str(&delta.summary) } Some(pb::interaction_update::Message::SummaryCompleted(_)) => { - output.summary_completed += 1 + output.summary_completed += 1; + output.interaction_events.push("summary_completed".into()); } Some(pb::interaction_update::Message::TurnEnded(_)) => output.turn_ended += 1, Some(pb::interaction_update::Message::TokenDelta(delta)) => { From 2dad5932639530e0c2adc9c6b368b0bce848938d Mon Sep 17 00:00:00 2001 From: leokun Date: Tue, 1 Sep 2026 21:03:02 +0800 Subject: [PATCH 12/32] feat: add resource limits management and network client integration - Introduced a new module for managing process resource limits, specifically for raising the open file limit on Unix systems. - Added a `NetworkClients` struct to handle reusable outbound HTTP clients, improving network request management. - Updated various components, including `ControlService` and `CursorProxy`, to utilize the new network client structure for better client handling. - Enhanced the API router to accept network clients, ensuring consistent client usage across different services. - Added tests to validate the integration of network clients and resource limits functionality. --- Cargo.lock | 1 + apps/desktop/src-tauri/Cargo.toml | 1 + apps/desktop/src-tauri/src/desktop.rs | 17 ++++ apps/desktop/src-tauri/src/lib.rs | 1 + apps/desktop/src-tauri/src/resource_limits.rs | 56 ++++++++++++ server/src/api/cursor/handlers.rs | 7 +- server/src/api/cursor/proxy.rs | 21 ++--- server/src/api/router.rs | 6 +- server/src/app.rs | 13 ++- server/src/control/service.rs | 13 ++- server/src/network.rs | 89 ++++++++++++++++++- server/src/provider/router.rs | 6 +- server/tests/knowledge_rules.rs | 4 +- 13 files changed, 203 insertions(+), 32 deletions(-) create mode 100644 apps/desktop/src-tauri/src/resource_limits.rs diff --git a/Cargo.lock b/Cargo.lock index e6d0ba7..effc024 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1176,6 +1176,7 @@ version = "0.1.5" dependencies = [ "axum", "cursor-server", + "libc", "rfd", "serde", "serde_json", diff --git a/apps/desktop/src-tauri/Cargo.toml b/apps/desktop/src-tauri/Cargo.toml index 8ec2faa..7f6328e 100644 --- a/apps/desktop/src-tauri/Cargo.toml +++ b/apps/desktop/src-tauri/Cargo.toml @@ -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" diff --git a/apps/desktop/src-tauri/src/desktop.rs b/apps/desktop/src-tauri/src/desktop.rs index 31f2d85..6afa44c 100644 --- a/apps/desktop/src-tauri/src/desktop.rs +++ b/apps/desktop/src-tauri/src/desktop.rs @@ -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, diff --git a/apps/desktop/src-tauri/src/lib.rs b/apps/desktop/src-tauri/src/lib.rs index ec2df31..14e5164 100644 --- a/apps/desktop/src-tauri/src/lib.rs +++ b/apps/desktop/src-tauri/src/lib.rs @@ -1,6 +1,7 @@ mod desktop; #[cfg(not(dev))] mod frontend; +mod resource_limits; mod startup; mod tray; diff --git a/apps/desktop/src-tauri/src/resource_limits.rs b/apps/desktop/src-tauri/src/resource_limits.rs new file mode 100644 index 0000000..b553f66 --- /dev/null +++ b/apps/desktop/src-tauri/src/resource_limits.rs @@ -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 { + 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, + }) +} diff --git a/server/src/api/cursor/handlers.rs b/server/src/api/cursor/handlers.rs index 6a608d2..aa0c0d3 100644 --- a/server/src/api/cursor/handlers.rs +++ b/server/src/api/cursor/handlers.rs @@ -27,8 +27,11 @@ use crate::{ Result, }; -pub fn router(registry: TransportRegistry) -> Result { - let proxy = CursorProxy::cursor(registry.store().clone())?; +pub fn router( + registry: TransportRegistry, + clients: crate::network::NetworkClients, +) -> Result { + let proxy = CursorProxy::cursor(clients); let knowledge = knowledge::KnowledgeService::managed()?; Ok(router_with_proxy(registry, proxy, knowledge)) } diff --git a/server/src/api/cursor/proxy.rs b/server/src/api/cursor/proxy.rs index df1dff6..6463fe3 100644 --- a/server/src/api/cursor/proxy.rs +++ b/server/src/api/cursor/proxy.rs @@ -14,8 +14,7 @@ pub const UPSTREAM_URL_HEADER: &str = "x-server-upstream-url"; #[derive(Clone)] pub struct CursorProxy { - client: Option, - store: Option, + clients: crate::network::NetworkClients, upstream: String, } @@ -47,23 +46,15 @@ impl BufferedResponse { } impl CursorProxy { - pub fn cursor(store: crate::store::Store) -> Result { - 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 { - 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 } } diff --git a/server/src/api/router.rs b/server/src/api/router.rs index 4a7e8e9..377831c 100644 --- a/server/src/api/router.rs +++ b/server/src/api/router.rs @@ -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 { - super::cursor::router(registry) +pub fn router(registry: TransportRegistry, clients: NetworkClients) -> Result { + super::cursor::router(registry, clients) } diff --git a/server/src/app.rs b/server/src/app.rs index d4ec4b1..62cb015 100644 --- a/server/src/app.rs +++ b/server/src/app.rs @@ -44,9 +44,11 @@ impl App { plugin_runtime.clone(), config.app_version.clone(), )?; + let clients = crate::network::NetworkClients::new(store.clone()); let provider = std::sync::Arc::new(ProviderRouter::new( store.clone(), plugins.clone(), + clients.clone(), config.provider_request_timeout, config.provider_stream_idle_timeout, )); @@ -58,10 +60,15 @@ impl App { plugins.clone(), crate::config::managed_data_dir()?.join("rules"), ); - let control = - control::ControlService::new(store.clone(), provider, plugin_runtime, plugins)?; + let control = control::ControlService::new( + store.clone(), + provider, + plugin_runtime, + plugins, + clients.clone(), + )?; 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)) diff --git a/server/src/control/service.rs b/server/src/control/service.rs index 5c3fa02..28e9e09 100644 --- a/server/src/control/service.rs +++ b/server/src/control/service.rs @@ -41,6 +41,7 @@ pub struct ControlService { provider: Arc, plugin_runtime: PluginRuntime, plugins: PluginRegistry, + clients: crate::network::NetworkClients, model_tests: Arc>>, } @@ -151,6 +152,7 @@ impl ControlService { provider: Arc, plugin_runtime: PluginRuntime, plugins: PluginRegistry, + clients: crate::network::NetworkClients, ) -> Result { Ok(Self { cursor_harness: CursorHarness::new(store.clone())?, @@ -158,6 +160,7 @@ impl ControlService { provider, plugin_runtime, plugins, + clients, model_tests: Arc::new(Mutex::new(BTreeMap::new())), }) } @@ -256,7 +259,7 @@ impl ControlService { disabled_ad_ids: Option<&str>, language: &str, ) -> 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 request = client .get(ADS_ENDPOINT) @@ -281,7 +284,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}")) @@ -510,7 +513,7 @@ impl ControlService { } pub async fn discover_models(&self, input: &ModelDiscoveryInput) -> Result { - 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, @@ -677,7 +680,9 @@ impl ControlService { } pub async fn set_proxy_settings(&self, settings: ProxySettingsInput) -> Result { - 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 { diff --git a/server/src/network.rs b/server/src/network.rs index d31f02c..b790dae 100644 --- a/server/src/network.rs +++ b/server/src/network.rs @@ -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>, +} + +#[derive(Default)] +struct ClientCache { + default: Option, + cursor: Option, + 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 { + 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 { + 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 { + 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 { let settings = store.proxy_settings_secret().await?; // Use the platform TLS stack for compatibility with provider gateways that diff --git a/server/src/provider/router.rs b/server/src/provider/router.rs index 7eabe64..577c0eb 100644 --- a/server/src/provider/router.rs +++ b/server/src/provider/router.rs @@ -21,6 +21,7 @@ use super::{ pub struct ProviderRouter { store: Store, plugins: PluginRegistry, + clients: crate::network::NetworkClients, request_timeout: Duration, stream_idle_timeout: Duration, } @@ -29,12 +30,14 @@ impl ProviderRouter { pub fn new( store: Store, plugins: PluginRegistry, + clients: crate::network::NetworkClients, request_timeout: Duration, stream_idle_timeout: Duration, ) -> Self { Self { store, plugins, + clients, request_timeout, stream_idle_timeout, } @@ -49,6 +52,7 @@ impl Provider for ProviderRouter { ) -> ProviderStream { let store = self.store.clone(); let plugins = self.plugins.clone(); + let clients = self.clients.clone(); let request_timeout = self.request_timeout; let stream_idle_timeout = self.stream_idle_timeout; Box::pin(try_stream! { @@ -91,7 +95,7 @@ impl Provider for ProviderRouter { request_timeout, allowed_body_fields: None, }; - let client = crate::network::client_builder(&store).await?.timeout(request_timeout).build()?; + let client = clients.provider_client(request_timeout).await?; let provider = build_observed(&config, recorder.clone(), client)?; (recorder, guard, provider.stream(routed, cancellation.clone())) }; diff --git a/server/tests/knowledge_rules.rs b/server/tests/knowledge_rules.rs index 1dc4bc5..05687b0 100644 --- a/server/tests/knowledge_rules.rs +++ b/server/tests/knowledge_rules.rs @@ -106,7 +106,7 @@ async fn decode(response: Response) -> M { #[tokio::test] async fn offline_crud_round_trip_persists_markdown() { let (_store_dir, store) = fixtures::temp_store().await; - let upstream = CursorProxy::cursor(store).unwrap(); + let upstream = CursorProxy::cursor(cursor_server::network::NetworkClients::new(store)); let rules_dir = tempfile::tempdir().unwrap(); let rules_root = rules_dir.path().join("rules"); let service = KnowledgeService::with_root(rules_root.clone()).unwrap(); @@ -193,7 +193,7 @@ async fn offline_crud_round_trip_persists_markdown() { #[tokio::test] async fn updating_missing_rule_reports_failure() { let (_store_dir, store) = fixtures::temp_store().await; - let upstream = CursorProxy::cursor(store).unwrap(); + let upstream = CursorProxy::cursor(cursor_server::network::NetworkClients::new(store)); let rules_dir = tempfile::tempdir().unwrap(); let service = KnowledgeService::with_root(rules_dir.path().join("rules")).unwrap(); From 5cdf642dd1530c026cc45ba4e2a3c63542688cfd Mon Sep 17 00:00:00 2001 From: leookun Date: Wed, 2 Sep 2026 01:04:54 +0800 Subject: [PATCH 13/32] feat: enhance bidi request handling and observability tracing - Updated the `append` function to include a flag for replacing closing requests, improving request handling. - Refactored the `run_sse_handler` and `bidi_handler` functions to utilize a new tracing mechanism, enhancing observability. - Introduced a new `trace_outcome` function to standardize tracing outcomes for requests. - Removed the `CursorTraceRecorder` in favor of a new `CursorTraceService` for better performance and non-blocking behavior. - Added tests to validate the new tracing functionality and ensure correct behavior during request processing. --- server/src/api/cursor/bidi.rs | 6 +- server/src/api/cursor/handlers.rs | 110 ++++-- server/src/api/cursor/run_sse.rs | 10 +- server/src/cursor/checkpoint/builder.rs | 24 +- server/src/cursor/compile/context.rs | 16 +- server/src/cursor/compile/run.rs | 4 +- server/src/cursor/conversation/command.rs | 14 +- server/src/cursor/conversation/output.rs | 22 +- server/src/cursor/conversation/runtime.rs | 163 ++++++-- server/src/cursor/services/blob_sync.rs | 116 +++--- server/src/cursor/services/observability.rs | 225 ----------- .../cursor/services/observability/event.rs | 71 ++++ .../src/cursor/services/observability/mod.rs | 157 ++++++++ .../cursor/services/observability/worker.rs | 349 ++++++++++++++++++ server/src/cursor/transport/handle.rs | 50 ++- server/src/cursor/transport/lifecycle.rs | 170 +++++++++ server/src/cursor/transport/mod.rs | 2 + server/src/cursor/transport/output.rs | 14 +- server/src/cursor/transport/registry.rs | 95 ++++- server/src/model/projection.rs | 61 ++- server/src/model/tool.rs | 18 + server/src/run/model_cycle.rs | 31 +- server/src/store/cursor_traces.rs | 51 ++- server/tests/cursor_trace_queue.rs | 129 +++++++ server/tests/cursor_transport_lifecycle.rs | 72 ++++ 25 files changed, 1522 insertions(+), 458 deletions(-) delete mode 100644 server/src/cursor/services/observability.rs create mode 100644 server/src/cursor/services/observability/event.rs create mode 100644 server/src/cursor/services/observability/mod.rs create mode 100644 server/src/cursor/services/observability/worker.rs create mode 100644 server/src/cursor/transport/lifecycle.rs create mode 100644 server/tests/cursor_trace_queue.rs create mode 100644 server/tests/cursor_transport_lifecycle.rs diff --git a/server/src/api/cursor/bidi.rs b/server/src/api/cursor/bidi.rs index 8592c66..f069981 100644 --- a/server/src/api/cursor/bidi.rs +++ b/server/src/api/cursor/bidi.rs @@ -184,7 +184,11 @@ pub async fn append( request: DecodedAppend, parent: Option, ) -> Result { - 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)?; } diff --git a/server/src/api/cursor/handlers.rs b/server/src/api/cursor/handlers.rs index aa0c0d3..5d78278 100644 --- a/server/src/api/cursor/handlers.rs +++ b/server/src/api/cursor/handlers.rs @@ -19,9 +19,7 @@ use crate::{ connect, proto::{agent::v1 as agent, aiserver::v1 as ai}, }, - services::{ - account, analytics, knowledge, model_catalog, observability::CursorTraceRecorder, tab, - }, + services::{account, analytics, knowledge, model_catalog, tab}, transport::{TransportParent, TransportRegistry}, }, Result, @@ -123,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 @@ -143,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) } } } @@ -159,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) @@ -183,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" @@ -198,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( @@ -227,6 +269,22 @@ async fn bidi_handler( Ok(response) } +fn trace_outcome( + mut metadata: serde_json::Value, + accepted: bool, + route_outcome: &str, + error: Option, +) -> 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) -> Result<(axum::http::request::Parts, Bytes)> { let (parts, body) = request.into_parts(); let body = to_bytes(body, usize::MAX) diff --git a/server/src/api/cursor/run_sse.rs b/server/src/api/cursor/run_sse.rs index b81b97f..6da2b52 100644 --- a/server/src/api/cursor/run_sse.rs +++ b/server/src/api/cursor/run_sse.rs @@ -22,7 +22,7 @@ pub async fn stream(registry: &TransportRegistry, request_id: &str) -> Result Response { 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), diff --git a/server/src/cursor/checkpoint/builder.rs b/server/src/cursor/checkpoint/builder.rs index 28c1c29..704833a 100644 --- a/server/src/cursor/checkpoint/builder.rs +++ b/server/src/cursor/checkpoint/builder.rs @@ -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 } diff --git a/server/src/cursor/compile/context.rs b/server/src/cursor/compile/context.rs index e29e634..859ecc5 100644 --- a/server/src/cursor/compile/context.rs +++ b/server/src/cursor/compile/context.rs @@ -12,7 +12,7 @@ use crate::{ protocol::proto::agent::v1 as pb, services::context_sync::RequestContextSynchronizer, tools::runtime::McpRoute, }, - model::ToolDefinition, + model::{normalize_tool_name, ToolDefinition}, store::BlobId, Error, Result, }; @@ -493,7 +493,7 @@ pub fn dynamic_mcp( })?), }; let parameters = normalize_mcp_parameters(&wire.name, parameters)?; - let name = model_tool_name(&wire.name); + let name = normalize_tool_name(&wire.name); let definition = ToolDefinition { name: name.clone(), description: wire.description.clone(), @@ -553,18 +553,6 @@ fn invalid_mcp_parameters(tool_name: &str) -> Error { )) } -fn model_tool_name(name: &str) -> String { - name.chars() - .map(|character| { - if character.is_ascii_alphanumeric() || matches!(character, '_' | '-') { - character - } else { - '_' - } - }) - .collect() -} - fn prost_value(value: &prost_types::Value) -> Value { use prost_types::value::Kind; match value.kind.as_ref() { diff --git a/server/src/cursor/compile/run.rs b/server/src/cursor/compile/run.rs index f351ead..b3712ce 100644 --- a/server/src/cursor/compile/run.rs +++ b/server/src/cursor/compile/run.rs @@ -117,9 +117,7 @@ pub(crate) async fn prepare( "selected_source": "root_prompt_messages_json", }); let encoded = serde_json::to_vec(&summary)?; - trace - .artifact("history_projection", "byok_server", &encoded, summary) - .await; + trace.artifact("history_projection", "byok_server", &encoded, summary); } let mut request_context = context::hydrate(request, context_sync).await?; if let Some(rules_dir) = local_rules_dir { diff --git a/server/src/cursor/conversation/command.rs b/server/src/cursor/conversation/command.rs index 3cdd2f7..fa97fe4 100644 --- a/server/src/cursor/conversation/command.rs +++ b/server/src/cursor/conversation/command.rs @@ -1,6 +1,13 @@ //! 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 TransportFinish { + Success, + Failed(Error), + Cancelled, +} #[derive(Debug)] pub enum TransportCommand { @@ -8,6 +15,9 @@ pub enum TransportCommand { seqno: i64, message: Box, }, + RunFinished { + generation: u64, + finish: TransportFinish, + }, Disconnect, - Close, } diff --git a/server/src/cursor/conversation/output.rs b/server/src/cursor/conversation/output.rs index 654937b..86765e8 100644 --- a/server/src/cursor/conversation/output.rs +++ b/server/src/cursor/conversation/output.rs @@ -34,7 +34,7 @@ use crate::{ Error, Result, }; -use super::{CompiledMessages, ConversationRegistry, MessageDelivery}; +use super::{CompiledMessages, ConversationRegistry, MessageDelivery, 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 { 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 { 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(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(TransportFinish::Cancelled); } action = self.runtime_actions.recv() => Input::RuntimeAction(action.map(Box::new)), event = self.core.events.recv() => Input::Event(event), @@ -719,7 +719,7 @@ impl ConversationOutput { if self.superseded.is_cancelled() { worker.abort(); self.abort_execs().await; - return Ok(()); + return Ok(TransportFinish::Cancelled); } return match outcome { RunOutcome::Completed => { @@ -735,8 +735,7 @@ impl ConversationOutput { for _ in 0..3 { self.checkpoint.publish(&self.handle, &checkpoint).await?; } - finish_success(&self.handle); - return Ok(()); + return Ok(TransportFinish::Success); } let checkpoints = final_checkpoint.take().ok_or_else(|| { Error::Protocol("Completed without final state".into()) @@ -752,18 +751,17 @@ impl ConversationOutput { ttft_breakdown: None, message: Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(checkpoints.settled)), })?; - finish_success(&self.handle); - Ok(()) + Ok(TransportFinish::Success) } RunOutcome::Cancelled => { worker.abort(); self.abort_execs().await; - finish_cancelled(&self.handle) + Ok(TransportFinish::Cancelled) } RunOutcome::Failed(failure) => { worker.abort(); self.abort_execs().await; - finish_failed(&self.handle, &cursor_error(failure)) + Ok(TransportFinish::Failed(cursor_error(failure))) } }; } diff --git a/server/src/cursor/conversation/runtime.rs b/server/src/cursor/conversation/runtime.rs index dc3f1ef..4f6d3f6 100644 --- a/server/src/cursor/conversation/runtime.rs +++ b/server/src/cursor/conversation/runtime.rs @@ -23,13 +23,14 @@ use crate::{ use super::{ CompiledMessages, ConversationDependencies, ConversationOutput, ConversationOutputDependencies, - ConversationRegistry, MessageDelivery, TransportCommand, + ConversationRegistry, MessageDelivery, TransportCommand, TransportFinish, }; pub struct ConversationRuntime; #[derive(Clone)] struct RunGeneration { + id: u64, superseded: CancellationToken, finished: CancellationToken, run: Arc>>, @@ -41,6 +42,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 +63,7 @@ impl ConversationRuntime { mut receiver: mpsc::Receiver, ) { 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 +75,47 @@ impl ConversationRuntime { let context_sync = RequestContextSynchronizer::new(handle.clone(), dependencies.store.clone()); let mut current = None::; + let mut next_generation = 1_u64; + let mut pending_finish = None::<(u64, TransportFinish)>; + let mut draining = 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 { + 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 { @@ -95,8 +133,16 @@ impl ConversationRuntime { 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; + } + handle.begin_close(); + pending_finish = Some((generation, finish)); + draining = true; } TransportCommand::Append { seqno, message } => { for (_seqno, message) in inbox.push(seqno, *message) { @@ -105,6 +151,11 @@ impl ConversationRuntime { Some(pb::agent_client_message::Message::RunRequest( request, )) => { + if draining { + handle.reopen(); + draining = false; + pending_finish = None; + } if let Some(conversation_id) = request.conversation_id.as_deref() { @@ -117,8 +168,6 @@ impl ConversationRuntime { "invalid Cursor conversation id" ); let _ = super::finish_failed(&handle, &error); - let _ = - handle.command(TransportCommand::Close).await; return; } } @@ -150,6 +199,7 @@ impl ConversationRuntime { dependencies.web_cache.clone(), ); let generation = RunGeneration { + id: next_generation, superseded: CancellationToken::new(), finished: CancellationToken::new(), run: Arc::new(parking_lot::Mutex::new(None)), @@ -158,6 +208,7 @@ impl ConversationRuntime { tool_runtime, tools, }; + next_generation = next_generation.saturating_add(1); current = Some(generation.clone()); spawn_run_request( registry.clone(), @@ -428,6 +479,34 @@ impl ConversationRuntime { } } +fn finish_pending( + handle: &TransportHandle, + current: &Option, + 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, @@ -487,8 +566,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: TransportFinish::Failed(error), + }) + .await; return; } }; @@ -520,8 +603,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: TransportFinish::Success, + }) + .await; } return; } @@ -545,8 +632,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: TransportFinish::Success, + }) + .await; } return; } @@ -605,16 +696,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() { + TransportFinish::Cancelled + } else { + tracing::error!( + request_id = handle.request_id(), + %error, + "Cursor session failed" + ); + TransportFinish::Failed(error) + } } - } + }; let _ = core_run.await; registry.release(&conversation_id, &run_id).await; if generation @@ -626,7 +722,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; } }); } diff --git a/server/src/cursor/services/blob_sync.rs b/server/src/cursor/services/blob_sync.rs index e88947a..316ce7d 100644 --- a/server/src/cursor/services/blob_sync.rs +++ b/server/src/cursor/services/blob_sync.rs @@ -76,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::>(), - }), - ) - .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::>(), + }), + ); } result?; Ok(id) @@ -144,18 +142,16 @@ impl BlobSynchronizer { pub async fn get(&self, blob_id: &BlobId) -> Result>> { 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)); } @@ -194,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(), + }), + ); } } } diff --git a/server/src/cursor/services/observability.rs b/server/src/cursor/services/observability.rs deleted file mode 100644 index 56b9859..0000000 --- a/server/src/cursor/services/observability.rs +++ /dev/null @@ -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>, - finished: Arc, -} - -#[derive(Default)] -struct TraceChunkBuffer { - chunks: Vec, - bytes: usize, - first_chunk_at: Option, - 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 { - 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 { - 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(()) - } -} diff --git a/server/src/cursor/services/observability/event.rs b/server/src/cursor/services/observability/event.rs new file mode 100644 index 0000000..9491b7a --- /dev/null +++ b/server/src/cursor/services/observability/event.rs @@ -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, + conversation_id: Option, + route: String, + model_id: Option, + }, + Resume { + request_id: String, + activation: Arc, + }, + 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, + }, +} + +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, + } + } +} diff --git a/server/src/cursor/services/observability/mod.rs b/server/src/cursor/services/observability/mod.rs new file mode 100644 index 0000000..76727cf --- /dev/null +++ b/server/src/cursor/services/observability/mod.rs @@ -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, +} + +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, + sender: mpsc::Sender, + finished: Arc, + activation: Arc, +} + +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" + ); + } + } +} diff --git a/server/src/cursor/services/observability/worker.rs b/server/src/cursor/services/observability/worker.rs new file mode 100644 index 0000000..62e6046 --- /dev/null +++ b/server/src/cursor/services/observability/worker.rs @@ -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, + bytes: usize, +} + +struct BufferedRequest { + artifact_type: String, + data: bytes::Bytes, + metadata: serde_json::Value, +} + +#[derive(Default)] +struct RequestOrder { + next: i64, + pending: BTreeMap>, +} + +pub(super) async fn run(store: Store, mut receiver: mpsc::Receiver) { + let mut states = HashMap::::new(); + let mut buffers = HashMap::::new(); + let mut request_orders = HashMap::::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, + buffers: &mut HashMap, + request_orders: &mut HashMap, + 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, + 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, + 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, + 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, 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) { + let request_ids = buffers.keys().cloned().collect::>(); + for request_id in request_ids { + flush_one(store, buffers, &request_id).await; + } +} diff --git a/server/src/cursor/transport/handle.rs b/server/src/cursor/transport/handle.rs index 4b72b5b..87a0a2e 100644 --- a/server/src/cursor/transport/handle.rs +++ b/server/src/cursor/transport/handle.rs @@ -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, conversation_id: Arc>, parent: Arc>, - trace: Option, + trace: CursorTraceRecorder, + lifecycle: TransportLifecycle, disconnect: CancellationToken, } @@ -39,7 +40,7 @@ impl TransportHandle { request_id: String, commands: mpsc::Sender, output: Arc, - trace: Option, + 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 { + 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 { diff --git a/server/src/cursor/transport/lifecycle.rs b/server/src/cursor/transport/lifecycle.rs new file mode 100644 index 0000000..b1faeeb --- /dev/null +++ b/server/src/cursor/transport/lifecycle.rs @@ -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, +} + +struct LifecycleInner { + state: parking_lot::Mutex, + admissions_drained: Notify, + closed: Notify, +} + +struct LifecycleState { + state: TransportState, + admissions: usize, +} + +pub(crate) struct TransportAdmission { + inner: Arc, +} + +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 { + 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()); + } +} diff --git a/server/src/cursor/transport/mod.rs b/server/src/cursor/transport/mod.rs index 20815bd..263e02c 100644 --- a/server/src/cursor/transport/mod.rs +++ b/server/src/cursor/transport/mod.rs @@ -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::*; diff --git a/server/src/cursor/transport/output.rs b/server/src/cursor/transport/output.rs index 0f513d9..bbd55ef 100644 --- a/server/src/cursor/transport/output.rs +++ b/server/src/cursor/transport/output.rs @@ -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, - 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; - } - } } diff --git a/server/src/cursor/transport/registry.rs b/server/src/cursor/transport/registry.rs index db57032..687987c 100644 --- a/server/src/cursor/transport/registry.rs +++ b/server/src/cursor/transport/registry.rs @@ -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>, + local: Mutex>, + next_local_generation: AtomicU64, upstream: Mutex>, route_changed: Notify, store: Store, + traces: CursorTraceService, web_cache: WebCache, plugins: Option, conversations: ConversationRegistry, } +#[derive(Clone)] +struct LocalTransport { + generation: u64, + handle: TransportHandle, +} + #[derive(Clone, Copy, Debug, PartialEq, Eq)] pub enum TransportRoute { Local, @@ -99,8 +113,10 @@ impl TransportRegistry { Self { inner: Arc::new(RegistryInner { local: Mutex::new(HashMap::new()), + next_local_generation: AtomicU64::new(1), upstream: Mutex::new(HashMap::new()), route_changed: Notify::new(), + traces: CursorTraceService::new(store.clone()), conversations: ConversationRegistry::new( store.clone(), provider, @@ -119,6 +135,13 @@ impl TransportRegistry { &self.inner.store } + pub fn trace( + &self, + request_id: &str, + ) -> crate::cursor::services::observability::CursorTraceRecorder { + self.inner.traces.recorder(request_id) + } + pub fn web_cache(&self) -> &WebCache { &self.inner.web_cache } @@ -132,18 +155,37 @@ impl TransportRegistry { } pub async fn get_or_create(&self, request_id: &str) -> Result { - 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 { + let mut local = self.inner.local.lock().await; + if let Some(transport) = local.get(request_id) { + if transport.handle.accepting_appends() || !replace_closing { + return Ok(transport.handle.clone()); + } } + local.remove(request_id); let (commands, receiver) = mpsc::channel(128); let output = Arc::new(OutputHub::default()); - let trace = CursorTraceRecorder::resume(self.inner.store.clone(), request_id).await; - let handle = TransportHandle::new(request_id.into(), commands, output.clone(), trace); - let mut local = self.inner.local.lock().await; - if let Some(existing) = local.get(request_id).cloned() { - return Ok(existing); - } - local.insert(request_id.into(), handle.clone()); + let trace = self.inner.traces.recorder(request_id); + trace.resume(); + let handle = TransportHandle::new(request_id.into(), commands, output, trace); + let generation = self + .inner + .next_local_generation + .fetch_add(1, Ordering::Relaxed); + local.insert( + request_id.into(), + LocalTransport { + generation, + handle: handle.clone(), + }, + ); drop(local); self.inner.route_changed.notify_waiters(); self.inner @@ -152,17 +194,29 @@ impl TransportRegistry { let registry = Arc::downgrade(&self.inner); let request_id = request_id.to_string(); + let lifecycle = handle.clone(); tokio::spawn(async move { - output.wait_closed().await; + lifecycle.wait_transport_closed().await; if let Some(registry) = registry.upgrade() { - registry.local.lock().await.remove(&request_id); + let mut local = registry.local.lock().await; + if local + .get(&request_id) + .is_some_and(|transport| transport.generation == generation) + { + local.remove(&request_id); + } } }); Ok(handle) } pub async fn local(&self, request_id: &str) -> Option { - self.inner.local.lock().await.get(request_id).cloned() + self.inner + .local + .lock() + .await + .get(request_id) + .map(|transport| transport.handle.clone()) } pub async fn mark_upstream(&self, request_id: &str) { @@ -206,10 +260,13 @@ impl TransportRegistry { self.inner.conversations.shutdown().await; let handles = std::mem::take(&mut *self.inner.local.lock().await); self.inner.upstream.lock().await.clear(); - for handle in handles.into_values() { - handle.disconnect().await; - let _ = - tokio::time::timeout(std::time::Duration::from_secs(2), handle.wait_closed()).await; + for transport in handles.into_values() { + transport.handle.disconnect().await; + let _ = tokio::time::timeout( + std::time::Duration::from_secs(2), + transport.handle.wait_transport_closed(), + ) + .await; } } } diff --git a/server/src/model/projection.rs b/server/src/model/projection.rs index 05dfaee..c6a7774 100644 --- a/server/src/model/projection.rs +++ b/server/src/model/projection.rs @@ -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"); + } +} diff --git a/server/src/model/tool.rs b/server/src/model/tool.rs index c522765..c35f115 100644 --- a/server/src/model/tool.rs +++ b/server/src/model/tool.rs @@ -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::(); + if normalized.is_empty() { + "_".into() + } else { + normalized + } +} + #[derive(Clone, Debug, Serialize, Deserialize, PartialEq)] pub struct ToolDefinition { pub name: String, diff --git a/server/src/run/model_cycle.rs b/server/src/run/model_cycle.rs index e6a24ba..f9d3b7a 100644 --- a/server/src/run/model_cycle.rs +++ b/server/src/run/model_cycle.rs @@ -6,7 +6,7 @@ use tokio::sync::mpsc; use tokio_util::sync::CancellationToken; use crate::{ - model::{ProviderReplayState, ToolCall, Usage}, + model::{normalize_tool_name, ProviderReplayState, ToolCall, Usage}, provider::{FinishReason, ModelEvent, ProviderStream}, }; @@ -165,6 +165,7 @@ pub async fn consume_model_cycle( call_id, name, } => { + let name = normalize_tool_name(&name); let Some(model_call_id) = model_call_id.as_ref() else { return Err(failure( RunFailure::Protocol("provider emitted content before Start".into()), @@ -406,6 +407,34 @@ mod tests { }; use tokio_stream::wrappers::ReceiverStream; + #[tokio::test] + async fn provider_tool_names_are_normalized_when_received() { + let events = vec![ + Ok(ModelEvent::Start { + model_call_id: "call".into(), + }), + Ok(ModelEvent::ToolCallStart { + index: 0, + call_id: "tool-call".into(), + name: "multi_tool_use.parallel".into(), + }), + Ok(ModelEvent::ToolCallEnd { index: 0 }), + Ok(ModelEvent::Done(FinishReason::ToolUse)), + ]; + let stream = Box::pin(tokio_stream::iter(events)); + let (event_tx, mut event_rx) = tokio::sync::mpsc::channel(4); + + let result = consume_model_cycle(stream, &event_tx, &CancellationToken::new()) + .await + .unwrap(); + + assert_eq!(result.calls[0].name, "multi_tool_use_parallel"); + assert!(matches!( + event_rx.recv().await, + Some(RunEvent::ToolCallStart { name, .. }) if name == "multi_tool_use_parallel" + )); + } + #[tokio::test] async fn usage_is_forwarded_before_the_provider_call_finishes() { let (provider_tx, provider_rx) = tokio::sync::mpsc::channel(4); diff --git a/server/src/store/cursor_traces.rs b/server/src/store/cursor_traces.rs index e91b2fe..55f8a33 100644 --- a/server/src/store/cursor_traces.rs +++ b/server/src/store/cursor_traces.rs @@ -62,6 +62,40 @@ impl Store { .await?) } + pub async fn append_cursor_trace_request( + &self, + request_id: &str, + artifact_type: &str, + source: &str, + data: &[u8], + metadata: &serde_json::Value, + ) -> Result<()> { + let metadata_json = serde_json::to_string(metadata)?; + let blob_id = BlobId::digest(data); + let _write = self.writes.lock().await; + let mut tx = self.pool.begin_with("BEGIN IMMEDIATE").await?; + Self::put_blob_tx(&mut tx, &blob_id, data, &[]).await?; + Self::link_cursor_trace_artifact_tx( + &mut tx, + request_id, + artifact_type, + source, + &blob_id, + &metadata_json, + ) + .await?; + sqlx::query( + "UPDATE cursor_run_traces + SET request_bytes = request_bytes + ? WHERE request_id = ?", + ) + .bind(as_i64(data.len())) + .bind(request_id) + .execute(&mut *tx) + .await?; + tx.commit().await?; + Ok(()) + } + pub async fn append_cursor_trace_artifact( &self, request_id: &str, @@ -144,23 +178,6 @@ impl Store { Ok(()) } - pub async fn add_cursor_trace_request_bytes( - &self, - request_id: &str, - bytes: usize, - ) -> Result<()> { - let _write = self.writes.lock().await; - sqlx::query( - "UPDATE cursor_run_traces - SET request_bytes = request_bytes + ? WHERE request_id = ?", - ) - .bind(as_i64(bytes)) - .bind(request_id) - .execute(&self.pool) - .await?; - Ok(()) - } - pub async fn start_cursor_trace_response(&self, request_id: &str, status: u16) -> Result<()> { let now = now_ms(); let _write = self.writes.lock().await; diff --git a/server/tests/cursor_trace_queue.rs b/server/tests/cursor_trace_queue.rs new file mode 100644 index 0000000..7debc58 --- /dev/null +++ b/server/tests/cursor_trace_queue.rs @@ -0,0 +1,129 @@ +//! Verifies that Cursor trace persistence is ordered and detached from producers. + +#[path = "support/fixtures.rs"] +mod fixtures; + +use std::time::{Duration, Instant}; + +use bytes::Bytes; +use cursor_server::{cursor::services::observability::CursorTraceService, store::Store}; +use sqlx::{Connection, SqliteConnection}; + +#[tokio::test] +async fn trace_producers_do_not_wait_for_sqlite_and_artifacts_stay_ordered() { + let directory = tempfile::tempdir().unwrap(); + let url = format!("sqlite://{}", directory.path().join("test.db").display()); + let store = Store::connect(&url).await.unwrap(); + store.set_detailed_logging(true).await.unwrap(); + let traces = CursorTraceService::new(store.clone()); + let recorder = traces.recorder("trace-queue-order"); + recorder.begin(Some("conversation-1"), "local_byok", Some("model-1")); + + tokio::time::timeout(Duration::from_secs(2), async { + loop { + if store + .cursor_trace("trace-queue-order") + .await + .unwrap() + .is_some() + { + break; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .unwrap(); + + let mut write_lock = SqliteConnection::connect(&url).await.unwrap(); + sqlx::query("BEGIN IMMEDIATE") + .execute(&mut write_lock) + .await + .unwrap(); + recorder.request( + "bidi_request", + Bytes::from_static(b"request-0"), + serde_json::json!({ + "append_seqno": 0, + "accepted": true, + "route_outcome": "local" + }), + ); + tokio::time::sleep(Duration::from_millis(25)).await; + + let started = Instant::now(); + let mut seqnos = (1..64).collect::>(); + for pair in seqnos.chunks_mut(2) { + pair.reverse(); + } + for seqno in seqnos { + recorder.request( + "bidi_request", + Bytes::from(format!("request-{seqno}")), + serde_json::json!({ + "append_seqno": seqno, + "accepted": true, + "route_outcome": "local" + }), + ); + } + recorder.finish(None); + assert!(started.elapsed() < Duration::from_millis(100)); + sqlx::query("ROLLBACK") + .execute(&mut write_lock) + .await + .unwrap(); + + let artifacts = tokio::time::timeout(Duration::from_secs(5), async { + loop { + let artifacts = store + .cursor_trace_artifacts("trace-queue-order") + .await + .unwrap(); + if artifacts.len() == 64 { + break artifacts; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .unwrap(); + + for (expected, artifact) in artifacts.iter().enumerate() { + assert_eq!(artifact.seq, expected as i64); + assert_eq!(artifact.metadata["append_seqno"], expected as i64); + } + let trace = store + .cursor_trace("trace-queue-order") + .await + .unwrap() + .unwrap(); + assert_eq!(trace.status, "completed"); + assert_eq!( + trace.request_bytes, + (0..64) + .map(|seqno| format!("request-{seqno}").len() as i64) + .sum::() + ); +} + +#[tokio::test] +async fn events_for_disabled_detailed_logging_are_discarded_off_path() { + let (_directory, store) = fixtures::temp_store().await; + let traces = CursorTraceService::new(store.clone()); + let recorder = traces.recorder("trace-disabled"); + + recorder.begin(None, "local_byok", Some("model-1")); + recorder.request( + "bidi_request", + Bytes::from_static(b"body"), + serde_json::json!({"append_seqno": 0}), + ); + + tokio::time::sleep(Duration::from_millis(50)).await; + assert!(store + .cursor_trace("trace-disabled") + .await + .unwrap() + .is_none()); +} diff --git a/server/tests/cursor_transport_lifecycle.rs b/server/tests/cursor_transport_lifecycle.rs new file mode 100644 index 0000000..258ed6a --- /dev/null +++ b/server/tests/cursor_transport_lifecycle.rs @@ -0,0 +1,72 @@ +//! Verifies registry ownership follows the transport actor rather than output subscriptions. + +#[path = "support/fake_provider.rs"] +mod fake_provider; +#[path = "support/fixtures.rs"] +mod fixtures; + +use std::{sync::Arc, time::Duration}; + +use cursor_server::cursor::{ + conversation::TransportCommand, + prompting::{PromptAssets, PromptCompiler}, + transport::TransportRegistry, +}; + +async fn registry() -> (tempfile::TempDir, TransportRegistry) { + let (directory, store) = fixtures::temp_store().await; + let assets = PromptAssets::load( + std::path::Path::new(env!("CARGO_MANIFEST_DIR")) + .join("prompt/cursor") + .as_path(), + ) + .unwrap(); + ( + directory, + TransportRegistry::new( + store, + Arc::new(fake_provider::FakeProvider::default()), + PromptCompiler::new(assets), + ), + ) +} + +#[tokio::test] +async fn actor_exit_removes_the_matching_transport_and_allows_a_new_generation() { + let (_directory, registry) = registry().await; + let first = registry.get_or_create("lifecycle-request").await.unwrap(); + assert!(registry.local("lifecycle-request").await.is_some()); + + first.command(TransportCommand::Disconnect).await.unwrap(); + tokio::time::timeout(Duration::from_secs(2), async { + loop { + if registry.local("lifecycle-request").await.is_none() { + break; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .unwrap(); + + let second = registry.get_or_create("lifecycle-request").await.unwrap(); + assert_eq!(second.request_id(), "lifecycle-request"); + assert!(registry.local("lifecycle-request").await.is_some()); + second.command(TransportCommand::Disconnect).await.unwrap(); +} + +#[tokio::test] +async fn dropping_an_output_subscription_does_not_remove_the_transport() { + let (_directory, registry) = registry().await; + let handle = registry + .get_or_create("subscription-request") + .await + .unwrap(); + let subscription = handle.subscribe(); + drop(subscription); + + tokio::time::sleep(Duration::from_millis(25)).await; + assert!(registry.local("subscription-request").await.is_some()); + + handle.command(TransportCommand::Disconnect).await.unwrap(); +} From f22c7b6680d87fc1fd4661646f961a451e256c3c Mon Sep 17 00:00:00 2001 From: leookun Date: Wed, 2 Sep 2026 10:23:50 +0800 Subject: [PATCH 14/32] feat: integrate app version into control service and update ads endpoint - Added `app_version` field to `ControlService` and updated its initialization to include the app version. - Modified the `ADS_ENDPOINT` to point to a local server for development purposes. - Refactored conversation command and output handling to utilize a new `RunFinish` enum for better state management. - Enhanced the conversation runtime to handle queued user messages after a turn has ended, ensuring smooth transitions between turns. - Added tests to validate the new behavior of queued messages and transport handling. --- server/src/app.rs | 1 + server/src/control/ads.rs | 3 +- server/src/control/service.rs | 7 +- server/src/cursor/conversation/command.rs | 8 +- server/src/cursor/conversation/output.rs | 22 +- server/src/cursor/conversation/runtime.rs | 240 +++++++++++++++------- server/tests/interrupt.rs | 128 ++++++++++++ 7 files changed, 321 insertions(+), 88 deletions(-) diff --git a/server/src/app.rs b/server/src/app.rs index 62cb015..6f97c86 100644 --- a/server/src/app.rs +++ b/server/src/app.rs @@ -66,6 +66,7 @@ impl App { plugin_runtime, plugins, clients.clone(), + config.app_version.clone(), )?; let harness = control.cursor_harness().clone(); let mut router = api::router(registry.clone(), clients)?; diff --git a/server/src/control/ads.rs b/server/src/control/ads.rs index 7749f65..bd5ffe9 100644 --- a/server/src/control/ads.rs +++ b/server/src/control/ads.rs @@ -16,7 +16,8 @@ use super::ControlService; // 此广告拉取不涉及用户隐私,用户id随机产生 // 开源项目广告为作者唯一收入来源,PR时请勿去除此广告 -pub(super) const ADS_ENDPOINT: &str = "https://tab.leokun.cn/api/v1/ads?placement=menu"; +// pub(super) const ADS_ENDPOINT: &str = "https://tab.leokun.cn/api/v1/ads?placement=menu"; +pub(super) const ADS_ENDPOINT: &str = "http://localhost:8080/api/v1/ads?placement=menu"; pub(super) const DEVICE_ID_HEADER: &str = "X-Cursor-Assistant-Device-ID"; pub(super) const OS_HEADER: &str = "X-Cursor-Assistant-OS"; pub(super) const APP_VERSION_HEADER: &str = "X-Cursor-Assistant-Version"; diff --git a/server/src/control/service.rs b/server/src/control/service.rs index 28e9e09..7a9d827 100644 --- a/server/src/control/service.rs +++ b/server/src/control/service.rs @@ -42,6 +42,7 @@ pub struct ControlService { plugin_runtime: PluginRuntime, plugins: PluginRegistry, clients: crate::network::NetworkClients, + app_version: String, model_tests: Arc>>, } @@ -153,6 +154,7 @@ impl ControlService { plugin_runtime: PluginRuntime, plugins: PluginRegistry, clients: crate::network::NetworkClients, + app_version: String, ) -> Result { Ok(Self { cursor_harness: CursorHarness::new(store.clone())?, @@ -161,6 +163,7 @@ impl ControlService { plugin_runtime, plugins, clients, + app_version, model_tests: Arc::new(Mutex::new(BTreeMap::new())), }) } @@ -265,7 +268,7 @@ impl ControlService { .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()) { @@ -299,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() diff --git a/server/src/cursor/conversation/command.rs b/server/src/cursor/conversation/command.rs index fa97fe4..4ba5f2b 100644 --- a/server/src/cursor/conversation/command.rs +++ b/server/src/cursor/conversation/command.rs @@ -2,6 +2,12 @@ use crate::{cursor::protocol::proto::agent::v1 as pb, Error}; +#[derive(Debug)] +pub enum RunFinish { + TurnCompleted, + Transport(TransportFinish), +} + #[derive(Debug)] pub enum TransportFinish { Success, @@ -17,7 +23,7 @@ pub enum TransportCommand { }, RunFinished { generation: u64, - finish: TransportFinish, + finish: RunFinish, }, Disconnect, } diff --git a/server/src/cursor/conversation/output.rs b/server/src/cursor/conversation/output.rs index 86765e8..349fccc 100644 --- a/server/src/cursor/conversation/output.rs +++ b/server/src/cursor/conversation/output.rs @@ -34,7 +34,7 @@ use crate::{ Error, Result, }; -use super::{CompiledMessages, ConversationRegistry, MessageDelivery, TransportFinish}; +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 { 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 { 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(TransportFinish::Cancelled); + 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(TransportFinish::Cancelled); + 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), @@ -719,7 +719,7 @@ impl ConversationOutput { if self.superseded.is_cancelled() { worker.abort(); self.abort_execs().await; - return Ok(TransportFinish::Cancelled); + return Ok(RunFinish::Transport(TransportFinish::Cancelled)); } return match outcome { RunOutcome::Completed => { @@ -735,7 +735,7 @@ impl ConversationOutput { for _ in 0..3 { self.checkpoint.publish(&self.handle, &checkpoint).await?; } - return Ok(TransportFinish::Success); + return Ok(RunFinish::TurnCompleted); } let checkpoints = final_checkpoint.take().ok_or_else(|| { Error::Protocol("Completed without final state".into()) @@ -751,17 +751,19 @@ impl ConversationOutput { ttft_breakdown: None, message: Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(checkpoints.settled)), })?; - Ok(TransportFinish::Success) + Ok(RunFinish::TurnCompleted) } RunOutcome::Cancelled => { worker.abort(); self.abort_execs().await; - Ok(TransportFinish::Cancelled) + Ok(RunFinish::Transport(TransportFinish::Cancelled)) } RunOutcome::Failed(failure) => { worker.abort(); self.abort_execs().await; - Ok(TransportFinish::Failed(cursor_error(failure))) + Ok(RunFinish::Transport(TransportFinish::Failed(cursor_error( + failure, + )))) } }; } diff --git a/server/src/cursor/conversation/runtime.rs b/server/src/cursor/conversation/runtime.rs index 4f6d3f6..3a1edec 100644 --- a/server/src/cursor/conversation/runtime.rs +++ b/server/src/cursor/conversation/runtime.rs @@ -18,19 +18,22 @@ use crate::{ }, transport::{OrderedInbox, TransportHandle}, }, - run::{CommandResult, RunEngine, RunHandle}, + run::{CommandResult, RunEngine, RunHandle, RunPhase}, }; use super::{ CompiledMessages, ConversationDependencies, ConversationOutput, ConversationOutputDependencies, - ConversationRegistry, MessageDelivery, TransportCommand, TransportFinish, + 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>>, @@ -78,6 +81,7 @@ impl ConversationRuntime { 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 = if draining { if !handle.admissions_drained() { @@ -102,6 +106,28 @@ impl ConversationRuntime { } } } + } 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, @@ -130,7 +156,19 @@ 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::RunFinished { generation, finish } => { @@ -140,9 +178,19 @@ impl ConversationRuntime { { continue; } - handle.begin_close(); - pending_finish = Some((generation, finish)); - draining = true; + 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) { @@ -151,6 +199,7 @@ impl ConversationRuntime { Some(pb::agent_client_message::Message::RunRequest( request, )) => { + waiting_for_action = false; if draining { handle.reopen(); draining = false; @@ -171,57 +220,18 @@ impl ConversationRuntime { 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::(); - 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, - 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(), + 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, @@ -382,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() { @@ -479,6 +510,67 @@ 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, + 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::(); + 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, @@ -569,7 +661,7 @@ fn spawn_run_request( let _ = handle .command(TransportCommand::RunFinished { generation: generation.id, - finish: TransportFinish::Failed(error), + finish: RunFinish::Transport(TransportFinish::Failed(error)), }) .await; return; @@ -606,7 +698,7 @@ fn spawn_run_request( let _ = handle .command(TransportCommand::RunFinished { generation: generation.id, - finish: TransportFinish::Success, + finish: RunFinish::Transport(TransportFinish::Success), }) .await; } @@ -635,7 +727,7 @@ fn spawn_run_request( let _ = handle .command(TransportCommand::RunFinished { generation: generation.id, - finish: TransportFinish::Success, + finish: RunFinish::Transport(TransportFinish::Success), }) .await; } @@ -700,14 +792,14 @@ fn spawn_run_request( Ok(finish) => finish, Err(error) => { if generation.superseded.is_cancelled() { - TransportFinish::Cancelled + RunFinish::Transport(TransportFinish::Cancelled) } else { tracing::error!( request_id = handle.request_id(), %error, "Cursor session failed" ); - TransportFinish::Failed(error) + RunFinish::Transport(TransportFinish::Failed(error)) } } }; diff --git a/server/tests/interrupt.rs b/server/tests/interrupt.rs index 1cfd882..9c50c8f 100644 --- a/server/tests/interrupt.rs +++ b/server/tests/interrupt.rs @@ -422,6 +422,85 @@ async fn runtime_cancel_action_aborts_active_exec_before_canceled_end_stream() { assert_eq!(output.recv().await, None); } +#[tokio::test] +async fn queued_user_message_after_turn_ended_starts_the_next_turn() { + let (_directory, store) = fixtures::temp_store().await; + let provider = fake_provider::FakeProvider::default(); + provider.push(text_response("first turn")); + provider.push(text_response("queued turn")); + let assets = PromptAssets::load( + std::path::Path::new(env!("CARGO_MANIFEST_DIR")) + .join("prompt/cursor") + .as_path(), + ) + .unwrap(); + let registry = TransportRegistry::new( + store, + Arc::new(provider.clone()), + PromptCompiler::new(assets), + ); + let handle = registry.get_or_create("queued-after-turn").await.unwrap(); + let mut output = handle.subscribe(); + handle + .command(TransportCommand::Append { + seqno: 0, + message: Box::new(client_run_for( + "queued-after-turn", + "queued-after-turn-conversation", + )), + }) + .await + .unwrap(); + + let mut append_seqno = 1; + wait_for_turn_ended(&handle, &mut output, &mut append_seqno).await; + assert_transport_remains_open(&handle, &mut output, &mut append_seqno).await; + + cursor_server::api::cursor::bidi::append( + ®istry, + cursor_server::api::cursor::bidi::DecodedAppend { + request_id: "queued-after-turn".into(), + seqno: append_seqno, + message: runtime_user_message(), + }, + None, + ) + .await + .unwrap(); + append_seqno += 1; + + let mut text = String::new(); + loop { + let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv()) + .await + .unwrap() + .expect("queued turn closed without EndStream"); + let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap(); + if flags & connect::END_STREAM_FLAG != 0 { + assert_eq!(payload.as_ref(), b"{}"); + break; + } + let server = pb::AgentServerMessage::decode(payload).unwrap(); + if let Some(pb::agent_server_message::Message::InteractionUpdate(update)) = server.message { + if let Some(pb::interaction_update::Message::TextDelta(delta)) = update.message { + text.push_str(&delta.text); + } + } + acknowledge_kv(&handle, &mut append_seqno, &frame).await; + } + + assert!(text.contains("queued turn")); + let requests = provider.requests(); + assert_eq!(requests.len(), 2); + assert_eq!( + &requests[1].history[..requests[0].history.len()], + requests[0].history.as_slice(), + "queued continuation must preserve the first provider request as a prefix" + ); + let history = serde_json::to_string(&requests[1].history).unwrap(); + assert!(history.contains("queued follow-up")); +} + #[tokio::test] async fn runtime_user_message_action_interrupts_and_continues_with_new_message() { let (_directory, store) = fixtures::temp_store().await; @@ -1732,6 +1811,55 @@ async fn run_to_end( } } +async fn wait_for_turn_ended( + handle: &cursor_server::cursor::TransportHandle, + output: &mut tokio::sync::mpsc::UnboundedReceiver, + append_seqno: &mut i64, +) { + loop { + let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv()) + .await + .unwrap() + .expect("RunSSE closed before turnEnded"); + let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap(); + assert_eq!(flags & connect::END_STREAM_FLAG, 0); + let server = pb::AgentServerMessage::decode(payload).unwrap(); + let turn_ended = matches!( + server.message, + Some(pb::agent_server_message::Message::InteractionUpdate( + pb::InteractionUpdate { + message: Some(pb::interaction_update::Message::TurnEnded(_)), + } + )) + ); + acknowledge_kv(handle, append_seqno, &frame).await; + if turn_ended { + return; + } + } +} + +async fn assert_transport_remains_open( + handle: &cursor_server::cursor::TransportHandle, + output: &mut tokio::sync::mpsc::UnboundedReceiver, + append_seqno: &mut i64, +) { + let deadline = tokio::time::Instant::now() + std::time::Duration::from_millis(100); + loop { + let remaining = deadline.saturating_duration_since(tokio::time::Instant::now()); + let Ok(Some(frame)) = tokio::time::timeout(remaining, output.recv()).await else { + return; + }; + let (flags, _) = connect::decode_frames(&frame).unwrap().pop().unwrap(); + assert_eq!( + flags & connect::END_STREAM_FLAG, + 0, + "turnEnded closed the transport before the queued action arrived" + ); + acknowledge_kv(handle, append_seqno, &frame).await; + } +} + async fn acknowledge_kv( handle: &cursor_server::cursor::TransportHandle, append_seqno: &mut i64, From 919c1d80325bf126aaf0bd403eb5e499987c17a8 Mon Sep 17 00:00:00 2001 From: leookun Date: Wed, 2 Sep 2026 12:49:57 +0800 Subject: [PATCH 15/32] fix: update ADS_ENDPOINT to production URL - Changed the `ADS_ENDPOINT` from a local server URL to the production URL for ads. - Commented out the local server URL for clarity and future reference. --- server/src/control/ads.rs | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/server/src/control/ads.rs b/server/src/control/ads.rs index bd5ffe9..6d17814 100644 --- a/server/src/control/ads.rs +++ b/server/src/control/ads.rs @@ -16,8 +16,8 @@ use super::ControlService; // 此广告拉取不涉及用户隐私,用户id随机产生 // 开源项目广告为作者唯一收入来源,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 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"; From ddaa61c827c5072eefd9a8439c2336e73ffe0bbc Mon Sep 17 00:00:00 2001 From: leookun Date: Wed, 2 Sep 2026 13:02:11 +0800 Subject: [PATCH 16/32] fix(desktop): update Windows executable in place --- .github/scripts/generate-portable-update.mjs | 71 +++++ .../scripts/generate-portable-update.test.mjs | 41 +++ .github/workflows/release.yml | 31 +- Cargo.lock | 3 + apps/desktop/package.json | 2 +- apps/desktop/src-tauri/Cargo.toml | 5 + apps/desktop/src-tauri/build.rs | 6 +- .../src-tauri/capabilities/default.json | 2 + .../autogenerated/check_portable_update.toml | 11 + .../install_portable_update.toml | 11 + apps/desktop/src-tauri/src/desktop.rs | 7 +- apps/desktop/src-tauri/src/lib.rs | 8 +- apps/desktop/src-tauri/src/update/mod.rs | 269 ++++++++++++++++ .../src-tauri/src/update/replacement.rs | 296 ++++++++++++++++++ .../settings/AppLifecycleSettingsCard.tsx | 6 +- apps/desktop/src/i18n/generated/catalog.json | 68 ++-- apps/desktop/src/i18n/locales/en-US.json | 2 + apps/desktop/src/i18n/locales/zh-CN.json | 2 + .../desktop/src/shared/native/appLifecycle.ts | 40 ++- apps/desktop/src/shared/store/updateStore.ts | 6 +- 20 files changed, 850 insertions(+), 37 deletions(-) create mode 100644 .github/scripts/generate-portable-update.mjs create mode 100644 .github/scripts/generate-portable-update.test.mjs create mode 100644 apps/desktop/src-tauri/permissions/autogenerated/check_portable_update.toml create mode 100644 apps/desktop/src-tauri/permissions/autogenerated/install_portable_update.toml create mode 100644 apps/desktop/src-tauri/src/update/mod.rs create mode 100644 apps/desktop/src-tauri/src/update/replacement.rs diff --git a/.github/scripts/generate-portable-update.mjs b/.github/scripts/generate-portable-update.mjs new file mode 100644 index 0000000..b7980db --- /dev/null +++ b/.github/scripts/generate-portable-update.mjs @@ -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; + }); +} diff --git a/.github/scripts/generate-portable-update.test.mjs b/.github/scripts/generate-portable-update.test.mjs new file mode 100644 index 0000000..1bcd223 --- /dev/null +++ b/.github/scripts/generate-portable-update.test.mjs @@ -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/); +}); diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index a55b8e8..92fd50f 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -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 }} diff --git a/Cargo.lock b/Cargo.lock index effc024..9096523 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1188,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]] diff --git a/apps/desktop/package.json b/apps/desktop/package.json index b691673..4972ca3 100644 --- a/apps/desktop/package.json +++ b/apps/desktop/package.json @@ -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", diff --git a/apps/desktop/src-tauri/Cargo.toml b/apps/desktop/src-tauri/Cargo.toml index 7f6328e..9d172fc 100644 --- a/apps/desktop/src-tauri/Cargo.toml +++ b/apps/desktop/src-tauri/Cargo.toml @@ -25,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"] } diff --git a/apps/desktop/src-tauri/build.rs b/apps/desktop/src-tauri/build.rs index da09dc1..dcb6da1 100644 --- a/apps/desktop/src-tauri/build.rs +++ b/apps/desktop/src-tauri/build.rs @@ -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") } diff --git a/apps/desktop/src-tauri/capabilities/default.json b/apps/desktop/src-tauri/capabilities/default.json index 1954d26..c34b30e 100644 --- a/apps/desktop/src-tauri/capabilities/default.json +++ b/apps/desktop/src-tauri/capabilities/default.json @@ -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", diff --git a/apps/desktop/src-tauri/permissions/autogenerated/check_portable_update.toml b/apps/desktop/src-tauri/permissions/autogenerated/check_portable_update.toml new file mode 100644 index 0000000..aea0675 --- /dev/null +++ b/apps/desktop/src-tauri/permissions/autogenerated/check_portable_update.toml @@ -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"] diff --git a/apps/desktop/src-tauri/permissions/autogenerated/install_portable_update.toml b/apps/desktop/src-tauri/permissions/autogenerated/install_portable_update.toml new file mode 100644 index 0000000..edf77fb --- /dev/null +++ b/apps/desktop/src-tauri/permissions/autogenerated/install_portable_update.toml @@ -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"] diff --git a/apps/desktop/src-tauri/src/desktop.rs b/apps/desktop/src-tauri/src/desktop.rs index 6afa44c..c700bea 100644 --- a/apps/desktop/src-tauri/src/desktop.rs +++ b/apps/desktop/src-tauri/src/desktop.rs @@ -177,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); @@ -245,6 +249,7 @@ pub fn run() -> ExitCode { window.set_focus()?; } tray::create(app)?; + crate::update::signal_ready_if_requested()?; Ok(()) }) .build(tauri::generate_context!()); diff --git a/apps/desktop/src-tauri/src/lib.rs b/apps/desktop/src-tauri/src/lib.rs index 14e5164..cef6375 100644 --- a/apps/desktop/src-tauri/src/lib.rs +++ b/apps/desktop/src-tauri/src/lib.rs @@ -4,5 +4,11 @@ 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() +} diff --git a/apps/desktop/src-tauri/src/update/mod.rs b/apps/desktop/src-tauri/src/update/mod.rs new file mode 100644 index 0000000..a6f0e62 --- /dev/null +++ b/apps/desktop/src-tauri/src/update/mod.rs @@ -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 { + #[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, 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, 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, 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 { + 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 { + 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()); + } +} diff --git a/apps/desktop/src-tauri/src/update/replacement.rs b/apps/desktop/src-tauri/src/update/replacement.rs new file mode 100644 index 0000000..cf2cb78 --- /dev/null +++ b/apps/desktop/src-tauri/src/update/replacement.rs @@ -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, String> { + let args = std::env::args_os().collect::>(); + 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::() + .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 { + let args = std::env::args_os().collect::>(); + 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()); + } +} diff --git a/apps/desktop/src/features/settings/AppLifecycleSettingsCard.tsx b/apps/desktop/src/features/settings/AppLifecycleSettingsCard.tsx index 7c14ea3..eff6f94 100644 --- a/apps/desktop/src/features/settings/AppLifecycleSettingsCard.tsx +++ b/apps/desktop/src/features/settings/AppLifecycleSettingsCard.tsx @@ -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 })); } }; diff --git a/apps/desktop/src/i18n/generated/catalog.json b/apps/desktop/src/i18n/generated/catalog.json index 5f06481..c294ccf 100644 --- a/apps/desktop/src/i18n/generated/catalog.json +++ b/apps/desktop/src/i18n/generated/catalog.json @@ -606,7 +606,7 @@ "refs": [ { "file": "features/settings/AppLifecycleSettingsCard.tsx", - "line": 156, + "line": 158, "column": 27 } ] @@ -916,7 +916,7 @@ "refs": [ { "file": "features/settings/AppLifecycleSettingsCard.tsx", - "line": 152, + "line": 154, "column": 13 } ] @@ -1673,7 +1673,7 @@ "refs": [ { "file": "features/settings/AppLifecycleSettingsCard.tsx", - "line": 126, + "line": 128, "column": 17 } ] @@ -1725,7 +1725,7 @@ "refs": [ { "file": "features/settings/AppLifecycleSettingsCard.tsx", - "line": 151, + "line": 153, "column": 13 } ] @@ -2353,7 +2353,7 @@ "refs": [ { "file": "shared/api.ts", - "line": 488, + "line": 486, "column": 43 } ] @@ -2365,7 +2365,7 @@ "refs": [ { "file": "features/settings/AppLifecycleSettingsCard.tsx", - "line": 149, + "line": 151, "column": 18 } ] @@ -2451,12 +2451,12 @@ "refs": [ { "file": "features/settings/AppLifecycleSettingsCard.tsx", - "line": 137, + "line": 139, "column": 18 }, { "file": "features/settings/AppLifecycleSettingsCard.tsx", - "line": 143, + "line": 145, "column": 16 } ] @@ -2680,7 +2680,7 @@ "refs": [ { "file": "shared/api.ts", - "line": 483, + "line": 481, "column": 43 } ] @@ -2795,7 +2795,7 @@ "refs": [ { "file": "features/settings/AppLifecycleSettingsCard.tsx", - "line": 160, + "line": 162, "column": 37 } ] @@ -2834,7 +2834,7 @@ "refs": [ { "file": "shared/api.ts", - "line": 420, + "line": 418, "column": 21 } ] @@ -2900,12 +2900,12 @@ "refs": [ { "file": "features/settings/AppLifecycleSettingsCard.tsx", - "line": 113, + "line": 115, "column": 18 }, { "file": "features/settings/AppLifecycleSettingsCard.tsx", - "line": 119, + "line": 121, "column": 16 } ] @@ -3194,6 +3194,20 @@ } ] }, + "92e26b27d5ea8f0e": { + "source": "检查更新失败:{error}", + "kind": "template", + "placeholders": [ + "error" + ], + "refs": [ + { + "file": "features/settings/AppLifecycleSettingsCard.tsx", + "line": 99, + "column": 15 + } + ] + }, "940a168911ade998": { "source": "每页条数", "kind": "text", @@ -3520,7 +3534,7 @@ "refs": [ { "file": "features/settings/AppLifecycleSettingsCard.tsx", - "line": 110, + "line": 112, "column": 29 } ] @@ -3532,12 +3546,12 @@ "refs": [ { "file": "features/settings/AppLifecycleSettingsCard.tsx", - "line": 125, + "line": 127, "column": 18 }, { "file": "features/settings/AppLifecycleSettingsCard.tsx", - "line": 131, + "line": 133, "column": 16 } ] @@ -3749,6 +3763,20 @@ } ] }, + "a80b53f8848e6d27": { + "source": "安装更新失败:{error}", + "kind": "template", + "placeholders": [ + "error" + ], + "refs": [ + { + "file": "features/settings/AppLifecycleSettingsCard.tsx", + "line": 108, + "column": 15 + } + ] + }, "a98585871c5313ff": { "source": "显示名称", "kind": "text", @@ -3814,7 +3842,7 @@ "refs": [ { "file": "features/settings/AppLifecycleSettingsCard.tsx", - "line": 156, + "line": 158, "column": 39 } ] @@ -3854,7 +3882,7 @@ "refs": [ { "file": "features/settings/AppLifecycleSettingsCard.tsx", - "line": 138, + "line": 140, "column": 17 } ] @@ -4577,7 +4605,7 @@ "refs": [ { "file": "features/settings/AppLifecycleSettingsCard.tsx", - "line": 114, + "line": 116, "column": 17 } ] @@ -5522,7 +5550,7 @@ }, { "file": "features/settings/AppLifecycleSettingsCard.tsx", - "line": 160, + "line": 162, "column": 25 } ] diff --git a/apps/desktop/src/i18n/locales/en-US.json b/apps/desktop/src/i18n/locales/en-US.json index d352612..83eaa21 100644 --- a/apps/desktop/src/i18n/locales/en-US.json +++ b/apps/desktop/src/i18n/locales/en-US.json @@ -221,6 +221,7 @@ "91aaf184cfc17ffd": "Overview", "91af6e57e7453fbe": "Add account", "92156a483d4ba248": "Only request, response, and trace attachments are deleted; call summaries, metrics, and configuration are kept.", + "92e26b27d5ea8f0e": "Failed to check for updates: {error}", "940a168911ade998": "Items per page", "945fb1c67eca8493": "Installing the plugin runtime", "946b3ffc02f026c0": "Delete this model?", @@ -259,6 +260,7 @@ "a748cc074f78de00": "View details", "a7617f42f898b2bf": "Use complete request URL", "a8036485f9227f2c": "Drag to reorder", + "a80b53f8848e6d27": "Failed to install update: {error}", "a98585871c5313ff": "Display name", "ab9084a640fbb864": "Deselect all", "abecab6701177721": "Launch at login enabled", diff --git a/apps/desktop/src/i18n/locales/zh-CN.json b/apps/desktop/src/i18n/locales/zh-CN.json index 53ed05e..2c5a221 100644 --- a/apps/desktop/src/i18n/locales/zh-CN.json +++ b/apps/desktop/src/i18n/locales/zh-CN.json @@ -221,6 +221,7 @@ "91aaf184cfc17ffd": "数据概览", "91af6e57e7453fbe": "添加账号", "92156a483d4ba248": "仅删除请求、响应和追踪附件等详细内容,保留调用汇总、统计指标和配置。", + "92e26b27d5ea8f0e": "检查更新失败:{error}", "940a168911ade998": "每页条数", "945fb1c67eca8493": "正在安装插件运行时", "946b3ffc02f026c0": "确定删除这个模型吗?", @@ -259,6 +260,7 @@ "a748cc074f78de00": "查看详情", "a7617f42f898b2bf": "使用完整请求地址", "a8036485f9227f2c": "拖动排序", + "a80b53f8848e6d27": "安装更新失败:{error}", "a98585871c5313ff": "显示名称", "ab9084a640fbb864": "全不选", "abecab6701177721": "已开启开机启动", diff --git a/apps/desktop/src/shared/native/appLifecycle.ts b/apps/desktop/src/shared/native/appLifecycle.ts index da7379c..e959ac6 100644 --- a/apps/desktop/src/shared/native/appLifecycle.ts +++ b/apps/desktop/src/shared/native/appLifecycle.ts @@ -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 { } } -export async function checkForUpdate(): Promise { - return check(); +export type AppUpdate = { + version: string; + install(): Promise; + close(): Promise; +}; + +type PortableUpdateInfo = { + version: string; +}; + +export async function checkForUpdate(): Promise { + if (desktopPlatform() === "windows") { + const update = await invoke("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 { - await update.downloadAndInstall(); - await relaunch(); +export async function installUpdate(update: AppUpdate): Promise { + await update.install(); } diff --git a/apps/desktop/src/shared/store/updateStore.ts b/apps/desktop/src/shared/store/updateStore.ts index a5dfed9..97fe4f1 100644 --- a/apps/desktop/src/shared/store/updateStore.ts +++ b/apps/desktop/src/shared/store/updateStore.ts @@ -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 | null = null; const listeners = new Set<() => void>(); @@ -26,7 +26,7 @@ function update(patch: Partial) { 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 }); From 0fd5e9d6f2382f303c7a55d5d2898384f40b191c Mon Sep 17 00:00:00 2001 From: leokun Date: Wed, 2 Sep 2026 13:56:28 +0800 Subject: [PATCH 17/32] fix(desktop): update content security policy to allow HTTPS images - Modified the content security policy in `tauri.conf.json` to include `https:` in the `img-src` directive, enhancing security by allowing images from secure sources. --- apps/desktop/src-tauri/tauri.conf.json | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/apps/desktop/src-tauri/tauri.conf.json b/apps/desktop/src-tauri/tauri.conf.json index 7d6f768..809440f 100644 --- a/apps/desktop/src-tauri/tauri.conf.json +++ b/apps/desktop/src-tauri/tauri.conf.json @@ -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" ] From 75b7ea9cc885760f66b8bb1b6186abb5eaab9328 Mon Sep 17 00:00:00 2001 From: leokun Date: Wed, 2 Sep 2026 14:03:35 +0800 Subject: [PATCH 18/32] chore(release): bump desktop to 0.1.6 --- Cargo.lock | 2 +- apps/desktop/package-lock.json | 4 ++-- apps/desktop/package.json | 2 +- apps/desktop/src-tauri/Cargo.toml | 2 +- apps/desktop/src-tauri/tauri.conf.json | 2 +- 5 files changed, 6 insertions(+), 6 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 9096523..2e4661a 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1172,7 +1172,7 @@ checksum = "52560adf09603e58c9a7ee1fe1dcb95a16927b17c127f0ac02d6e768a0e25bc1" [[package]] name = "cursor-byok-desktop" -version = "0.1.5" +version = "0.1.6" dependencies = [ "axum", "cursor-server", diff --git a/apps/desktop/package-lock.json b/apps/desktop/package-lock.json index f0b9ab5..cac9448 100644 --- a/apps/desktop/package-lock.json +++ b/apps/desktop/package-lock.json @@ -1,12 +1,12 @@ { "name": "cursor-byok-desktop", - "version": "0.1.5", + "version": "0.1.6", "lockfileVersion": 3, "requires": true, "packages": { "": { "name": "cursor-byok-desktop", - "version": "0.1.5", + "version": "0.1.6", "license": "MIT", "dependencies": { "@floating-ui/dom": "^1.8.0", diff --git a/apps/desktop/package.json b/apps/desktop/package.json index 4972ca3..50a1009 100644 --- a/apps/desktop/package.json +++ b/apps/desktop/package.json @@ -1,6 +1,6 @@ { "name": "cursor-byok-desktop", - "version": "0.1.5", + "version": "0.1.6", "description": "Cursor BYOK desktop management application", "type": "module", "scripts": { diff --git a/apps/desktop/src-tauri/Cargo.toml b/apps/desktop/src-tauri/Cargo.toml index 9d172fc..10c01dd 100644 --- a/apps/desktop/src-tauri/Cargo.toml +++ b/apps/desktop/src-tauri/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "cursor-byok-desktop" -version = "0.1.5" +version = "0.1.6" edition = "2021" publish = false diff --git a/apps/desktop/src-tauri/tauri.conf.json b/apps/desktop/src-tauri/tauri.conf.json index 809440f..3ab75d8 100644 --- a/apps/desktop/src-tauri/tauri.conf.json +++ b/apps/desktop/src-tauri/tauri.conf.json @@ -1,7 +1,7 @@ { "$schema": "https://schema.tauri.app/config/2", "productName": "Cursor BYOK", - "version": "0.1.5", + "version": "0.1.6", "identifier": "dev.cursorbyok.desktop", "build": { "beforeDevCommand": "npm run dev", From a42f84cfe78b03b94fa07e9caa04a8588c67f4b3 Mon Sep 17 00:00:00 2001 From: kevin9327 Date: Mon, 31 Aug 2026 20:07:52 +0900 Subject: [PATCH 19/32] fix(conversation): scope completed tool call ids to their round A provider that reuses a tool call id across two rounds of one run wedges the run permanently. `ToolDispatcher::start_batch` skips any call whose id is in `ToolBatchState::completed`, so the second call is never dispatched and never produces a `ToolCompletion`, while `tool_round::execute` blocks waiting for `calls.len()` results with no timeout on that path. The client sees the tool call appear and then nothing: no completion, no further output, no end-stream frame. The `completed` set is built once per run and never cleared, so it is run-scoped. A tool call id is only unique within a round, which the schema already states as `UNIQUE (round_id, call_id)`; the sibling runtime completed-map is likewise already cleared per round at output.rs:615. Clear `completed` when `ExecuteToolRound` begins a new round, tracked independently of `active_round` so it does not depend on the order in which ToolRoundStarted and ExecuteToolRound are observed. Replaying a round still skips the calls that round already committed. The practical trigger is openai_chat.rs:196, which synthesizes `call-{index}` from a per-stream index when a provider omits tool call ids, so `call-0` recurs on every model call. That file is left alone here. Co-Authored-By: Claude Opus 4.8 --- server/src/cursor/conversation/output.rs | 11 ++ server/tests/tool_round.rs | 145 +++++++++++++++++++++++ 2 files changed, 156 insertions(+) diff --git a/server/src/cursor/conversation/output.rs b/server/src/cursor/conversation/output.rs index 349fccc..55dd9b6 100644 --- a/server/src/cursor/conversation/output.rs +++ b/server/src/cursor/conversation/output.rs @@ -158,6 +158,7 @@ impl ConversationOutput { let mut streams = BTreeMap::::new(); let mut completions = HashMap::::new(); let mut completed = HashSet::::new(); + let mut completed_round = None::; let mut response_text = String::new(); let mut response_thinking = String::new(); let mut active_round = None::; @@ -425,6 +426,16 @@ impl ConversationOutput { round_id, calls: round_calls, } => { + // `completed` exists so that replaying ExecuteToolRound for a round + // does not dispatch a call this round already committed. Tool call ids + // are only unique *within* a round -- the schema says as much with + // `UNIQUE (round_id, call_id)` -- so an id retained from an earlier + // round would make start_batch skip a fresh call, and tool_round wait + // forever for a result nothing will ever produce. + if completed_round.as_ref() != Some(&round_id) { + completed.clear(); + completed_round = Some(round_id.clone()); + } active_round = Some(round_id.clone()); active_tool_calls = round_calls .iter() diff --git a/server/tests/tool_round.rs b/server/tests/tool_round.rs index 309ed1d..bfea8f8 100644 --- a/server/tests/tool_round.rs +++ b/server/tests/tool_round.rs @@ -1253,6 +1253,151 @@ fn client_run_for_model( } } +/// A provider that reuses a tool call id across two rounds of the same run +/// must not wedge the run. `ToolDispatcher::start_batch` skips any call whose +/// id is already in `ToolBatchState::completed`, and that set is never cleared +/// for the lifetime of the run, so the skipped call produced no completion and +/// `tool_round::execute` waited forever for a result that could never arrive. +#[tokio::test] +async fn duplicate_tool_call_id_across_rounds_does_not_wedge_the_run() { + let (_directory, store) = fixtures::temp_store().await; + let provider = fake_provider::FakeProvider::default(); + for _ in 0..2 { + provider.push(vec![ + ModelEvent::Start { + model_call_id: "ignored".into(), + }, + ModelEvent::ToolCallStart { + index: 0, + call_id: "call-1".into(), + name: "Read".into(), + }, + ModelEvent::ToolCallArgumentsDelta { + index: 0, + delta: "{\"path\":\"/tmp/a\"}".into(), + }, + ModelEvent::ToolCallEnd { index: 0 }, + ModelEvent::Done(FinishReason::ToolUse), + ]); + } + provider.push(vec![ + ModelEvent::Start { + model_call_id: "ignored".into(), + }, + ModelEvent::TextStart, + ModelEvent::TextDelta("done".into()), + ModelEvent::TextEnd, + ModelEvent::Done(FinishReason::Stop), + ]); + let assets = PromptAssets::load( + std::path::Path::new(env!("CARGO_MANIFEST_DIR")) + .join("prompt/cursor") + .as_path(), + ) + .unwrap(); + let registry = TransportRegistry::new( + store.clone(), + Arc::new(provider.clone()), + PromptCompiler::new(assets), + ); + let handle = registry.get_or_create("tool-request").await.unwrap(); + let mut output = handle.subscribe(); + handle + .command(TransportCommand::Append { + seqno: 0, + message: Box::new(client_run()), + }) + .await + .unwrap(); + let mut seqno = 1; + let mut execs = 0; + loop { + let frame = tokio::time::timeout(std::time::Duration::from_secs(10), output.recv()) + .await + .expect("run must not hang on a reused tool call id") + .unwrap(); + let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap(); + if flags & connect::END_STREAM_FLAG != 0 { + break; + } + let server = pb::AgentServerMessage::decode(payload).unwrap(); + match server.message { + Some(pb::agent_server_message::Message::KvServerMessage(kv)) => { + handle + .command(TransportCommand::Append { + seqno, + message: Box::new(kv_ack(kv.id)), + }) + .await + .unwrap(); + seqno += 1; + } + Some(pb::agent_server_message::Message::ExecServerMessage(exec)) => { + execs += 1; + let exec_id = exec.id; + handle + .command(TransportCommand::Append { + seqno, + message: Box::new(pb::AgentClientMessage { + message: Some(pb::agent_client_message::Message::ExecClientMessage( + pb::ExecClientMessage { + id: exec_id, + exec_id: String::new(), + message: Some(pb::exec_client_message::Message::ReadResult( + pb::ReadResult { + result: Some(pb::read_result::Result::Success( + pb::ReadSuccess { + path: "/tmp/a".into(), + total_lines: 1, + file_size: 1, + output: Some( + pb::read_success::Output::Content( + "x".into(), + ), + ), + ..Default::default() + }, + )), + }, + )), + ..Default::default() + }, + )), + }), + }) + .await + .unwrap(); + seqno += 1; + handle + .command(TransportCommand::Append { + seqno, + message: Box::new(pb::AgentClientMessage { + message: Some( + pb::agent_client_message::Message::ExecClientControlMessage( + pb::ExecClientControlMessage { + message: Some( + pb::exec_client_control_message::Message::StreamClose( + pb::ExecClientStreamClose { id: exec_id }, + ), + ), + }, + ), + ), + }), + }) + .await + .unwrap(); + seqno += 1; + } + _ => {} + } + } + // Both rounds have to reach the client, and the run has to get far enough + // to ask the provider a third time and finish. + assert_eq!(execs, 2); + assert_eq!(provider.requests().len(), 3); +} + fn client_run() -> pb::AgentClientMessage { let user = pb::UserMessage { text: "read it".into(), From e4b5e132d66ec1d916e882f23aa06ce31bdb2fd3 Mon Sep 17 00:00:00 2001 From: kevin9327 Date: Thu, 3 Sep 2026 07:55:51 +0900 Subject: [PATCH 20/32] chore: restore a green make check on main - output.rs: the read_lints test helper (#398) builds ToolCall without the argument_error field added in d83e14a, so the lib-test target does not compile - connect_wire.rs: cursor::router takes NetworkClients since 2dad593 - runtime.rs: clippy::nonminimal_bool from 5cdf642 - update/mod.rs: clippy::needless_return in the Windows arm, ddaa61c --- apps/desktop/src-tauri/src/update/mod.rs | 4 ++-- server/src/cursor/conversation/runtime.rs | 4 ++-- server/src/cursor/tools/tool_call_result/exec/output.rs | 1 + server/tests/connect_wire.rs | 4 +++- 4 files changed, 8 insertions(+), 5 deletions(-) diff --git a/apps/desktop/src-tauri/src/update/mod.rs b/apps/desktop/src-tauri/src/update/mod.rs index a6f0e62..3a83f6c 100644 --- a/apps/desktop/src-tauri/src/update/mod.rs +++ b/apps/desktop/src-tauri/src/update/mod.rs @@ -57,9 +57,9 @@ pub(crate) async fn check_portable_update( #[cfg(target_os = "windows")] { let update = portable_update(&app).await?; - return Ok(update.map(|update| PortableUpdateInfo { + Ok(update.map(|update| PortableUpdateInfo { version: update.version, - })); + })) } #[cfg(not(target_os = "windows"))] { diff --git a/server/src/cursor/conversation/runtime.rs b/server/src/cursor/conversation/runtime.rs index 3a1edec..dae5cf5 100644 --- a/server/src/cursor/conversation/runtime.rs +++ b/server/src/cursor/conversation/runtime.rs @@ -172,9 +172,9 @@ impl ConversationRuntime { break; } TransportCommand::RunFinished { generation, finish } => { - if !current + if current .as_ref() - .is_some_and(|current| current.id == generation) + .is_none_or(|current| current.id != generation) { continue; } diff --git a/server/src/cursor/tools/tool_call_result/exec/output.rs b/server/src/cursor/tools/tool_call_result/exec/output.rs index 316085c..2441a5a 100644 --- a/server/src/cursor/tools/tool_call_result/exec/output.rs +++ b/server/src/cursor/tools/tool_call_result/exec/output.rs @@ -459,6 +459,7 @@ mod tests { name: "ReadLints".into(), arguments_text: String::new(), arguments: json!({ "paths": paths }), + argument_error: None, } } diff --git a/server/tests/connect_wire.rs b/server/tests/connect_wire.rs index a370f36..3a4f0a6 100644 --- a/server/tests/connect_wire.rs +++ b/server/tests/connect_wire.rs @@ -21,6 +21,7 @@ use cursor_server::{ proto::{agent::v1 as pb, aiserver::v1 as ai}, }, cursor::transport::TransportRegistry, + network::NetworkClients, }; use flate2::{write::GzEncoder, Compression}; use prost::Message; @@ -102,6 +103,7 @@ async fn bidi_append_gzip_body_is_decompressed_before_protobuf_decode() { .as_path(), ) .unwrap(); + let clients = NetworkClients::new(store.clone()); let registry = TransportRegistry::new( store, Arc::new(fake_provider::FakeProvider::default()), @@ -118,7 +120,7 @@ async fn bidi_append_gzip_body_is_decompressed_before_protobuf_decode() { encoder.write_all(&wire).unwrap(); let compressed = encoder.finish().unwrap(); - let response = cursor::router(registry) + let response = cursor::router(registry, clients) .unwrap() .oneshot( Request::post("/aiserver.v1.BidiService/BidiAppend") From 42811a27f59df99a692d68a1ee9b23f4d66f3ccd Mon Sep 17 00:00:00 2001 From: ProtectCookies <1439080921@qq.com> Date: Thu, 3 Sep 2026 10:31:41 +0800 Subject: [PATCH 21/32] =?UTF-8?q?feat:=20=E6=96=B0=E5=A2=9E=20Commit=20?= =?UTF-8?q?=E8=AE=BE=E7=BD=AE=E5=8F=8A=E6=9C=AC=E5=9C=B0=E6=8F=90=E4=BA=A4?= =?UTF-8?q?=E4=BF=A1=E6=81=AF=E7=94=9F=E6=88=90?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 新增 CommitSettingsCard 组件,支持选择生成模型与编辑提示词 - 新增 /settings/commit GET/PUT 接口及 CommitSettings 持久化 - 实现 WriteGitCommitMessage RPC 本地生成,空 model_id 时直连转发 - 新增 NetworkService/IsConnected 探针响应,防止流式生成被中断 - 添加 commit prompt 模板及 proto 消息定义 - 补充 zh-CN / en-US 国际化词条 --- .../settings/CommitSettingsCard.module.scss | 53 +++ .../features/settings/CommitSettingsCard.tsx | 153 +++++++ .../src/features/settings/SettingsPage.tsx | 2 + apps/desktop/src/i18n/generated/catalog.json | 219 +++++++--- apps/desktop/src/i18n/locales/en-US.json | 7 + apps/desktop/src/i18n/locales/zh-CN.json | 7 + apps/desktop/src/shared/api.ts | 11 + server/prompt/cursor/commit/prompt.md | 66 +++ server/src/api/cursor/handlers.rs | 27 +- server/src/control/mod.rs | 4 + server/src/control/service.rs | 12 +- server/src/control/settings.rs | 39 +- server/src/cursor/protocol/proto.rs | 35 ++ server/src/cursor/services/commit_message.rs | 406 ++++++++++++++++++ server/src/cursor/services/mod.rs | 1 + server/src/cursor/services/tab.rs | 3 +- server/src/local_app/proxy.rs | 2 + server/src/store/settings.rs | 62 +++ server/tests/commit_message.rs | 246 +++++++++++ 19 files changed, 1292 insertions(+), 63 deletions(-) create mode 100644 apps/desktop/src/features/settings/CommitSettingsCard.module.scss create mode 100644 apps/desktop/src/features/settings/CommitSettingsCard.tsx create mode 100644 server/prompt/cursor/commit/prompt.md create mode 100644 server/src/cursor/services/commit_message.rs create mode 100644 server/tests/commit_message.rs diff --git a/apps/desktop/src/features/settings/CommitSettingsCard.module.scss b/apps/desktop/src/features/settings/CommitSettingsCard.module.scss new file mode 100644 index 0000000..17a8905 --- /dev/null +++ b/apps/desktop/src/features/settings/CommitSettingsCard.module.scss @@ -0,0 +1,53 @@ +@use "../../styles/typography" as type; + +.row { + display: flex; + align-items: center; + justify-content: space-between; + gap: 20px; + padding: 16px; + + > div:first-child { + display: grid; + gap: 5px; + } + + small { + color: var(--vscode-descriptionForeground); + font-size: type.$font-size-xs; + } +} + +.select { + min-width: 280px; +} + +.textButton { + padding: 4px; + color: var(--vscode-textLink-foreground); + background: transparent; + border: 0; + + &:hover { color: var(--vscode-textLink-activeForeground); } + &:disabled { opacity: 0.5; cursor: not-allowed; } +} + +.promptEditor { + width: 100%; + flex: 1; + min-height: 320px; + box-sizing: border-box; + padding: 10px 12px; + color: var(--vscode-input-foreground); + background: var(--vscode-input-background); + border: 1px solid var(--vscode-input-border, var(--vscode-sideBar-border)); + border-radius: 6px; + font-family: var(--vscode-editor-font-family, monospace); + font-size: type.$font-size-xs; + line-height: 1.6; + resize: none; + + &:focus { + outline: 1px solid var(--vscode-focusBorder); + } +} diff --git a/apps/desktop/src/features/settings/CommitSettingsCard.tsx b/apps/desktop/src/features/settings/CommitSettingsCard.tsx new file mode 100644 index 0000000..2702173 --- /dev/null +++ b/apps/desktop/src/features/settings/CommitSettingsCard.tsx @@ -0,0 +1,153 @@ +import { useCallback, useEffect, useMemo, useState } from "react"; +import { api, type CommitSettingsView } from "../../shared/api"; +import { useAppStore } from "../../shared/store/appStore"; +import { Modal } from "../../shared/ui/Modal"; +import { Select } from "../../shared/ui/Select"; +import { TitledCard } from "../../shared/ui/TitledCard"; +import { useMessage } from "../../shared/ui/message"; +import controls from "../../shared/ui/Controls.module.scss"; +import styles from "./CommitSettingsCard.module.scss"; + +function errorText(cause: unknown) { + return cause instanceof Error ? cause.message : String(cause); +} + +export function CommitSettingsCard() { + const { models } = useAppStore(); + const message = useMessage(); + const [view, setView] = useState(null); + const [savingModel, setSavingModel] = useState(false); + const [promptOpen, setPromptOpen] = useState(false); + const [promptDraft, setPromptDraft] = useState(""); + const [savingPrompt, setSavingPrompt] = useState(false); + + useEffect(() => { + void api.commitSettings().then(setView).catch((cause) => message(errorText(cause))); + }, [message]); + + const modelOptions = useMemo(() => { + const options: Array<{ value: string; label: string }> = [{ value: "", label: t("直连") }]; + const seen = new Set(); + for (const model of models) { + seen.add(model.model_hash); + options.push({ + value: model.model_hash, + label: model.display_name && model.display_name !== model.model_id + ? `${model.display_name}(${model.model_id})` + : model.display_name || model.model_id, + }); + } + if (view?.model_id && !seen.has(view.model_id)) { + options.push({ value: view.model_id, label: view.model_id }); + } + return options; + }, [models, view]); + + const selectedModelId = view?.model_id ?? ""; + + const persist = useCallback( + async (modelId: string, prompt: string) => { + if (!view) return null; + const normalizedPrompt = + prompt.trim() === view.default_prompt.trim() ? "" : prompt.trim(); + return api.setCommitSettings({ model_id: modelId, prompt: normalizedPrompt }); + }, + [view], + ); + + const changeModel = useCallback( + async (modelId: string) => { + if (!view) return; + setSavingModel(true); + try { + const saved = await persist(modelId, view.prompt); + if (saved) setView(saved); + } catch (cause) { + message(errorText(cause)); + } finally { + setSavingModel(false); + } + }, + [view, persist, message], + ); + + const openPrompt = useCallback(() => { + if (!view) return; + setPromptDraft(view.prompt || view.default_prompt); + setPromptOpen(true); + }, [view]); + + const savePrompt = useCallback(async () => { + if (!view) return; + setSavingPrompt(true); + try { + const saved = await persist(view.model_id, promptDraft); + if (saved) setView(saved); + setPromptOpen(false); + message(t("提示词设置已保存")); + } catch (cause) { + message(errorText(cause)); + } finally { + setSavingPrompt(false); + } + }, [view, promptDraft, persist, message]); + + const resetPrompt = useCallback(() => { + if (!view) return; + setPromptDraft(view.default_prompt); + }, [view]); + + return ( + <> + + {t("提示词设置")} + + } + > +
+
+ {t("生成模型")} + {t("直连走 Cursor 原有通道;选择 Cursor 页已配置的模型后由本地生成。")} +
+
+