From 29fde7d7c797fd71c2f19774502028ef335b2f4a Mon Sep 17 00:00:00 2001 From: leokun Date: Tue, 1 Sep 2026 10:10:09 +0800 Subject: [PATCH] feat: enhance context token estimation and compaction logic - Added `estimate_context_tokens` function to calculate provider-visible context size based on prompt specifications and projected messages. - Updated `CheckpointBuilder` to record estimated context tokens during message processing. - Refactored compaction logic to utilize the new token estimation, ensuring proper context management during model runs. - Introduced tests to validate context estimation and compaction behavior under various scenarios. --- server/src/cursor/checkpoint/summary.rs | 8 +- server/src/model/observability.rs | 27 +-- server/src/model/token_count.rs | 199 ++++++++++++++++++- server/src/run/compaction.rs | 153 +++++++-------- server/src/run/engine.rs | 39 ++-- server/src/store/llm_calls.rs | 43 +---- server/tests/compaction.rs | 166 ++++++++++++++++ server/tests/interrupt.rs | 32 +++- server/tests/tool_round.rs | 244 +++++++++++++++++++++++- 9 files changed, 723 insertions(+), 188 deletions(-) diff --git a/server/src/cursor/checkpoint/summary.rs b/server/src/cursor/checkpoint/summary.rs index 6f84f6e..c391d84 100644 --- a/server/src/cursor/checkpoint/summary.rs +++ b/server/src/cursor/checkpoint/summary.rs @@ -3,7 +3,7 @@ use prost::Message; use crate::{ cursor::{checkpoint::PendingSteps, protocol::proto::agent::v1 as pb}, - model::CanonicalMessage, + model::{estimate_context_tokens, project_messages, CanonicalMessage, PromptSpec}, store::{BlobEdge, BlobId}, Error, Result, }; @@ -85,6 +85,12 @@ impl CheckpointBuilder { .push(archive_id.as_bytes().to_vec()); } self.base.self_summary_count = self.base.self_summary_count.saturating_add(1); + let projected = project_messages(messages)?; + let prompt = PromptSpec { + instructions: self.instructions.clone(), + tools: self.tool_definitions.clone(), + }; + self.record_context_tokens(Some(estimate_context_tokens(&prompt, &projected))); if let Some(details) = self.base.token_details.as_mut() { details.breakdown = Some(crate::cursor::services::usage::breakdown( details.used_tokens, diff --git a/server/src/model/observability.rs b/server/src/model/observability.rs index 7fd8fb9..ec5f443 100644 --- a/server/src/model/observability.rs +++ b/server/src/model/observability.rs @@ -6,8 +6,6 @@ mod usage { use serde::{Deserialize, Serialize}; - use super::ProviderType; - #[derive(Clone, Copy, Debug, Default, Serialize, Deserialize, PartialEq, Eq)] pub struct Usage { pub input_tokens: Option, @@ -18,21 +16,6 @@ mod usage { pub reasoning_tokens: Option, } - impl Usage { - /// Returns the provider-visible input context without counting cached tokens twice. - pub(crate) fn context_input_tokens(self, provider: ProviderType) -> Option { - let input = self.input_tokens?; - match provider { - ProviderType::OpenAiChat | ProviderType::OpenAiResponses | ProviderType::Plugin => { - Some(input) - } - ProviderType::Anthropic => input - .checked_add(self.cache_read_tokens.unwrap_or_default())? - .checked_add(self.cache_write_tokens.unwrap_or_default()), - } - } - } - impl AddAssign for Usage { fn add_assign(&mut self, rhs: Self) { self.input_tokens = sum(self.input_tokens, rhs.input_tokens); @@ -53,7 +36,7 @@ pub use usage::*; mod llm_call { use serde::Serialize; - use super::{ProviderType, Usage}; + use super::ProviderType; #[derive(Clone, Debug)] pub struct NewLlmCall { @@ -75,14 +58,6 @@ mod llm_call { pub detailed: bool, } - #[derive(Clone, Copy, Debug, PartialEq, Eq)] - pub(crate) struct LlmCallUsageAnchor { - pub request_type: ProviderType, - pub usage: Usage, - pub message_count: usize, - pub tool_count: usize, - } - #[derive(Clone, Debug, Serialize)] pub struct LlmCallSummary { pub call_id: String, diff --git a/server/src/model/token_count.rs b/server/src/model/token_count.rs index 14e499c..e1ae538 100644 --- a/server/src/model/token_count.rs +++ b/server/src/model/token_count.rs @@ -1,4 +1,100 @@ -//! Estimates and records model token usage. +//! Estimates provider-visible context size and formats configured token counts. + +use super::{ContentPart, ProjectedContent, ProjectedMessage, PromptSpec}; + +const TOKENS_PER_MESSAGE_OVERHEAD: u64 = 8; +const TOKENS_PER_TOOL_CALL_OVERHEAD: u64 = 6; +const TOKENS_PER_IMAGE: u64 = 1_024; + +pub(crate) fn estimate_context_tokens(prompt: &PromptSpec, messages: &[ProjectedMessage]) -> u64 { + let instructions = estimate_text_tokens(&prompt.instructions); + let tools = prompt.tools.iter().fold(0_u64, |total, tool| { + total.saturating_add(estimate_json_tokens(tool)) + }); + let messages = 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 { + let content = match &message.content { + ProjectedContent::Parts(parts) => estimate_parts_tokens(parts), + ProjectedContent::Assistant { + text, + thinking, + replay_state, + calls, + } => { + let calls = calls.iter().fold(0_u64, |total, call| { + total + .saturating_add(TOKENS_PER_TOOL_CALL_OVERHEAD) + .saturating_add(estimate_text_tokens(&call.call_id)) + .saturating_add(estimate_text_tokens(&call.name)) + .saturating_add(estimate_json_tokens(&call.arguments)) + }); + estimate_text_tokens(text) + .saturating_add(estimate_text_tokens(thinking)) + .saturating_add( + replay_state + .as_ref() + .map(estimate_json_tokens) + .unwrap_or_default(), + ) + .saturating_add(calls) + } + ProjectedContent::ToolResult(result) => { + let content = if result.provider_parts.is_empty() { + estimate_text_tokens(&result.content).saturating_add( + result + .image + .as_ref() + .map(|image| { + TOKENS_PER_IMAGE.saturating_add(estimate_text_tokens(&image.mime_type)) + }) + .unwrap_or_default(), + ) + } else { + estimate_parts_tokens(&result.provider_parts) + }; + estimate_text_tokens(&result.call_id) + .saturating_add(estimate_text_tokens(&result.name)) + .saturating_add(content) + } + }; + TOKENS_PER_MESSAGE_OVERHEAD.saturating_add(content) +} + +fn estimate_parts_tokens(parts: &[ContentPart]) -> u64 { + parts.iter().fold(0_u64, |total, part| { + let tokens = match part { + ContentPart::Text { text } => estimate_text_tokens(text), + ContentPart::Image { mime_type, .. } => { + TOKENS_PER_IMAGE.saturating_add(estimate_text_tokens(mime_type)) + } + }; + total.saturating_add(tokens) + }) +} + +fn estimate_json_tokens(value: &impl serde::Serialize) -> u64 { + serde_json::to_string(value) + .map(|value| estimate_text_tokens(&value)) + .unwrap_or_default() +} + +fn estimate_text_tokens(text: &str) -> u64 { + let text = text.trim(); + if text.is_empty() { + return 0; + } + let characters = text.chars().count() as u64; + characters + .div_ceil(4) + .saturating_add(text.bytes().filter(|byte| *byte == b'\n').count() as u64) + .max(1) +} + pub(crate) fn parse_token_count(value: &str) -> Option { let value = value.trim().to_ascii_lowercase(); let (number, multiplier) = match value.chars().last()? { @@ -18,3 +114,104 @@ pub(crate) fn format_token_count(tokens: u64) -> String { tokens.to_string() } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::model::{Role, ToolCallContent, ToolDefinition, ToolResultContent}; + + fn prompt() -> PromptSpec { + PromptSpec { + instructions: "system instructions".into(), + tools: vec![ToolDefinition { + name: "Read".into(), + description: "Read a file".into(), + parameters: serde_json::json!({"type": "object", "properties": {"path": {"type": "string"}}}), + }], + } + } + + #[test] + fn context_estimate_grows_with_provider_visible_text_and_tools() { + let short = vec![ProjectedMessage { + message_id: "short".into(), + role: Role::User, + content: ProjectedContent::Parts(vec![ContentPart::Text { + text: "hello".into(), + }]), + }]; + let long = vec![ProjectedMessage { + message_id: "long".into(), + role: Role::User, + content: ProjectedContent::Parts(vec![ContentPart::Text { + text: "x".repeat(40_000), + }]), + }]; + let without_tools = PromptSpec { + instructions: prompt().instructions, + tools: Vec::new(), + }; + + assert!( + estimate_context_tokens(&prompt(), &short) + > estimate_context_tokens(&without_tools, &short) + ); + assert!( + estimate_context_tokens(&prompt(), &long) > estimate_context_tokens(&prompt(), &short) + ); + } + + #[test] + fn context_estimate_counts_tool_calls_results_and_images() { + let assistant = ProjectedMessage { + message_id: "assistant".into(), + role: Role::Assistant, + content: ProjectedContent::Assistant { + text: String::new(), + thinking: "reasoning".into(), + replay_state: None, + calls: vec![ToolCallContent { + index: 0, + call_id: "call-1".into(), + name: "Read".into(), + arguments: serde_json::json!({"path": "/tmp/file"}), + }], + }, + }; + let text_result = ProjectedMessage { + message_id: "result-text".into(), + role: Role::Tool, + content: ProjectedContent::ToolResult(ToolResultContent { + call_id: "call-1".into(), + name: "Read".into(), + content: "file contents".into(), + is_error: false, + image: None, + provider_parts: Vec::new(), + }), + }; + let image_result = ProjectedMessage { + message_id: "result-image".into(), + role: Role::Tool, + content: ProjectedContent::ToolResult(ToolResultContent { + call_id: "call-1".into(), + name: "Read".into(), + content: "file contents".into(), + is_error: false, + image: None, + provider_parts: vec![ContentPart::Image { + mime_type: "image/png".into(), + data: vec![0; 32], + }], + }), + }; + + let base = estimate_context_tokens(&prompt(), &[]); + let with_call = estimate_context_tokens(&prompt(), std::slice::from_ref(&assistant)); + let with_text = estimate_context_tokens(&prompt(), &[assistant.clone(), text_result]); + let with_image = estimate_context_tokens(&prompt(), &[assistant, image_result]); + assert!(with_call > base); + assert!(with_text > with_call); + assert!(with_image > with_text); + } +} diff --git a/server/src/run/compaction.rs b/server/src/run/compaction.rs index e97e415..6fbe1a7 100644 --- a/server/src/run/compaction.rs +++ b/server/src/run/compaction.rs @@ -1,62 +1,53 @@ -//! Decides when to compact context and builds a stable fallback summary. +//! Decides when to compact provider-visible context and builds a stable fallback summary. use std::collections::HashSet; -use crate::model::{CanonicalMessage, LlmCallUsageAnchor, PreparedRun, ProjectedMessage}; +use crate::model::{estimate_context_tokens, CanonicalMessage, PreparedRun, ProjectedMessage}; const FALLBACK_CHARS: usize = 12_000; +pub(super) const RESERVE_TOKENS: u64 = 10_000; pub(super) const OUTPUT_TOKENS: u64 = 4_096; pub(super) const INSTRUCTIONS: &str = "Summarize the conversation for the next model turn. Preserve goals, constraints, decisions, files, commands, errors, results, and unfinished work. Do not call tools. Return only the concise durable summary."; -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub(super) struct ContextUsageAnchor { - input_tokens: u64, - message_count: usize, - tool_count: usize, +pub(super) fn input_budget(prepared: &PreparedRun) -> Option { + prepared + .model + .context_window_tokens + .map(|window| window.saturating_sub(RESERVE_TOKENS)) } -impl ContextUsageAnchor { - pub(super) fn from_llm_call(anchor: LlmCallUsageAnchor) -> Option { - Some(Self { - input_tokens: anchor.usage.context_input_tokens(anchor.request_type)?, - message_count: anchor.message_count, - tool_count: anchor.tool_count, - }) - } +pub(super) fn estimated_tokens( + prepared: &PreparedRun, + projected_messages: &[ProjectedMessage], +) -> u64 { + estimate_context_tokens(&prepared.prompt, projected_messages) } pub(super) fn should_compact( prepared: &PreparedRun, - messages: &[CanonicalMessage], projected_messages: &[ProjectedMessage], - anchor: Option, ) -> bool { - let Some(context_window) = prepared.model.context_window_tokens else { + let Some(budget) = input_budget(prepared) else { return false; }; - if context_window == 0 || messages.len() <= prepared.initial_messages.len() { - return false; + estimated_tokens(prepared, projected_messages) > budget +} + +pub(super) fn validate_compacted( + prepared: &PreparedRun, + projected_messages: &[ProjectedMessage], +) -> std::result::Result { + let estimated = estimated_tokens(prepared, projected_messages); + let Some(budget) = input_budget(prepared) else { + return Ok(estimated); + }; + if estimated <= budget { + return Ok(estimated); } - let estimated_input = anchor - .filter(|anchor| { - anchor.message_count <= projected_messages.len() - && anchor.tool_count == prepared.prompt.tools.len() - }) - .map(|anchor| { - anchor - .input_tokens - .saturating_add(estimate_serialized_tokens( - &serde_json::to_string(&projected_messages[anchor.message_count..]) - .unwrap_or_default(), - )) - }) - .unwrap_or_else(|| { - estimate_serialized_tokens( - &serde_json::to_string(&(&prepared.prompt, messages)).unwrap_or_default(), - ) - }); - estimated_input > context_window + Err(format!( + "context overflow after compaction: estimated input {estimated} tokens exceeds budget {budget} tokens" + )) } pub(super) fn partition( @@ -95,15 +86,6 @@ pub(super) fn fallback_summary(messages: &[CanonicalMessage]) -> String { ) } -fn estimate_serialized_tokens(serialized: &str) -> u64 { - serialized - .chars() - .fold(0_u64, |units, character| { - units.saturating_add(if character.is_ascii() { 273 } else { 550 }) - }) - .div_ceil(1_000) -} - #[cfg(test)] mod tests { use super::*; @@ -112,11 +94,10 @@ mod tests { RunAction, RunId, RunKind, }; - #[test] - fn automatic_compaction_runs_for_start_and_resume_actions_after_the_limit() { + fn prepared(context_window_tokens: u64) -> PreparedRun { let mut model = ModelSpec::new("model"); - model.context_window_tokens = Some(200_000); - let mut prepared = PreparedRun { + model.context_window_tokens = Some(context_window_tokens); + PreparedRun { run_id: RunId::new("run"), cursor_request_id: None, conversation_id: ConversationId::new("conversation"), @@ -129,50 +110,50 @@ mod tests { initial_messages: Vec::new(), action: RunAction::Start, base_checkpoint_id: CheckpointId(1), - }; + } + } + + #[test] + fn automatic_compaction_uses_fixed_reserve_for_every_action() { let messages = vec![CanonicalMessage::text( "user", Role::User, Origin::Runtime, - "hello", + "x".repeat(40_000), )]; let projected = project_messages(&messages).unwrap(); - let tail_tokens = estimate_serialized_tokens(&serde_json::to_string(&projected).unwrap()); - let anchor = |estimated_input| { - Some(ContextUsageAnchor { - input_tokens: estimated_input - tail_tokens, - message_count: 0, - tool_count: 0, - }) - }; + let estimated = estimate_context_tokens(&prepared(1).prompt, &projected); + let mut prepared = prepared(estimated + RESERVE_TOKENS); - assert!(!should_compact( - &prepared, - &messages, - &projected, - anchor(199_999) - )); - assert!(!should_compact( - &prepared, - &messages, - &projected, - anchor(200_000) - )); - assert!(should_compact( - &prepared, - &messages, - &projected, - anchor(200_001) - )); + assert!(!should_compact(&prepared, &projected)); + prepared.model.context_window_tokens = Some(estimated + RESERVE_TOKENS - 1); + assert!(should_compact(&prepared, &projected)); prepared.action = RunAction::Resume { pending_tool_round: None, }; - assert!(should_compact( - &prepared, - &messages, - &projected, - anchor(200_001) - )); + assert!(should_compact(&prepared, &projected)); + } + + #[test] + fn compacted_history_is_validated_against_the_same_budget() { + let messages = vec![CanonicalMessage::text( + "user", + Role::User, + Origin::Runtime, + "x".repeat(40_000), + )]; + let projected = project_messages(&messages).unwrap(); + let estimated = estimate_context_tokens(&prepared(1).prompt, &projected); + + assert_eq!( + validate_compacted(&prepared(estimated + RESERVE_TOKENS), &projected), + Ok(estimated) + ); + assert!( + validate_compacted(&prepared(estimated + RESERVE_TOKENS - 1), &projected) + .unwrap_err() + .contains("context overflow after compaction") + ); } } diff --git a/server/src/run/engine.rs b/server/src/run/engine.rs index 54ca793..1c75334 100644 --- a/server/src/run/engine.rs +++ b/server/src/run/engine.rs @@ -161,7 +161,6 @@ impl RunEngine { }; } - let mut auto_compacted = prepared.action == RunAction::Compact; 'model: loop { if cancellation.is_cancelled() { return (RunOutcome::Cancelled, usage); @@ -170,31 +169,13 @@ impl RunEngine { Ok(messages) => messages, Err(error) => return (RunOutcome::Failed(error.into()), usage), }; - let context_anchor = if !auto_compacted { - match self - .store - .latest_llm_call_usage_anchor( - &prepared.conversation_id, - &prepared.model.model_id, - ) - .await - { - Ok(anchor) => { - anchor.and_then(super::compaction::ContextUsageAnchor::from_llm_call) - } - Err(error) => return (RunOutcome::Failed(error.into()), usage), - } - } else { - None - }; let history = match crate::model::project_messages(&messages) { Ok(history) => history, Err(error) => return (RunOutcome::Failed(error.into()), usage), }; - if !auto_compacted - && super::compaction::should_compact(prepared, &messages, &history, context_anchor) + if prepared.action != RunAction::Compact + && super::compaction::should_compact(prepared, &history) { - auto_compacted = true; match self .auto_compact(prepared, checkpoint, &messages, client, cancellation) .await @@ -583,7 +564,15 @@ impl RunEngine { let (compactable, retained_request_context) = super::compaction::partition(messages, ¤t_ids); if compactable.is_empty() { - return Ok((checkpoint, None)); + let projected = crate::model::project_messages(messages) + .map_err(|error| RunOutcome::Failed(error.into()))?; + let message = super::compaction::validate_compacted(prepared, &projected) + .err() + .unwrap_or_else(|| { + "context overflow after compaction: no conversation history can be compacted" + .into() + }); + return Err(RunOutcome::Failed(RunFailure::Protocol(message))); } emit(client, RunEvent::AutoCompactionStarted) @@ -689,7 +678,7 @@ impl RunEngine { ) } }; - let event_id = format!("summary:auto:{}", prepared.run_id); + let event_id = format!("summary:auto:{}:{provider_call_index}", prepared.run_id); let summary_message = CanonicalMessage { message_id: format!("runtime:{event_id}"), role: Role::User, @@ -704,6 +693,10 @@ impl RunEngine { let mut replacement = retained_request_context.into_iter().collect::>(); replacement.push(summary_message); replacement.extend(prepared.initial_messages.iter().cloned()); + let projected_replacement = crate::model::project_messages(&replacement) + .map_err(|error| RunOutcome::Failed(error.into()))?; + super::compaction::validate_compacted(prepared, &projected_replacement) + .map_err(|message| RunOutcome::Failed(RunFailure::Protocol(message)))?; let mut checkpoint = self .store .replace_checkpoint( diff --git a/server/src/store/llm_calls.rs b/server/src/store/llm_calls.rs index e3063f4..b882071 100644 --- a/server/src/store/llm_calls.rs +++ b/server/src/store/llm_calls.rs @@ -1,13 +1,8 @@ //! Persists provider call payloads, timing, and usage. -use std::str::FromStr; - use sqlx::Row; use crate::{ - model::{ - ConversationId, LlmCallRequest, LlmCallResponseChunk, LlmCallSummary, LlmCallUsageAnchor, - NewLlmCall, ProviderType, Usage, - }, + model::{LlmCallRequest, LlmCallResponseChunk, LlmCallSummary, NewLlmCall, Usage}, Result, }; @@ -288,41 +283,6 @@ impl Store { .transpose() } - pub(crate) async fn latest_llm_call_usage_anchor( - &self, - conversation_id: &ConversationId, - model_hash: &str, - ) -> Result> { - let row = sqlx::query( - r#"SELECT request_type, usage_json, message_count, tool_count - FROM llm_calls - WHERE conversation_id = ? - AND model_hash = ? - AND status = 'completed' - AND input_tokens IS NOT NULL - AND usage_json IS NOT NULL - ORDER BY rowid DESC - LIMIT 1"#, - ) - .bind(conversation_id.as_str()) - .bind(model_hash) - .fetch_optional(&self.pool) - .await?; - row.map(|row| { - let message_count = - usize::try_from(row.try_get::("message_count")?).unwrap_or(usize::MAX); - let tool_count = - usize::try_from(row.try_get::("tool_count")?).unwrap_or(usize::MAX); - Ok(LlmCallUsageAnchor { - request_type: ProviderType::from_str(row.try_get("request_type")?)?, - usage: serde_json::from_str(row.try_get("usage_json")?)?, - message_count, - tool_count, - }) - }) - .transpose() - } - pub async fn llm_call_request(&self, call_id: &str) -> Result> { let row = sqlx::query( "SELECT headers_json, body_json, byte_count FROM llm_call_requests WHERE call_id = ?", @@ -416,6 +376,7 @@ fn summary_from_row(row: sqlx::sqlite::SqliteRow) -> Result { #[cfg(test)] mod tests { use super::*; + use crate::model::ProviderType; /// 插件模型不在 model_configs 中,调用记录必须照常落库并可按其稳定 ID 筛选。 #[tokio::test] diff --git a/server/tests/compaction.rs b/server/tests/compaction.rs index 46b8ea2..3a4ab8f 100644 --- a/server/tests/compaction.rs +++ b/server/tests/compaction.rs @@ -191,6 +191,172 @@ async fn summarize_replaces_model_history_and_preserves_cursor_history() { ); } +#[tokio::test] +async fn automatic_compaction_preflights_provider_input_and_records_rebuilt_tokens() { + let (_directory, store) = fixtures::temp_store().await; + let model = store + .create_model(&ModelConfigInput { + sort_order: 0, + display_name: "Auto Compact Model".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: "Auto Compact Model".into(), + model_id: "auto-compact-model".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(100_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(&"x".repeat(400_000), 150_000, 1_000)); + provider.push(text_response("automatic durable summary", 120_000, 20)); + provider.push(text_response("continued after compaction", 20_000, 20)); + 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, + "auto-first", + user_request( + "auto-conversation", + "auto-user-1", + "start", + &model.model_hash, + None, + ), + ) + .await; + let first_state = first.checkpoints.last().unwrap().clone(); + assert!(first_state.token_details.as_ref().unwrap().used_tokens > 100_000); + + let second = run( + ®istry, + "auto-second", + user_request( + "auto-conversation", + "auto-user-2", + "continue", + &model.model_hash, + Some(first_state), + ), + ) + .await; + assert_eq!(second.summary_started, 1); + assert_eq!(second.summary_completed, 1); + let compacted_tokens = second + .checkpoints + .iter() + .filter_map(|state| state.token_details.as_ref()) + .map(|details| details.used_tokens) + .find(|tokens| *tokens > 0 && *tokens < 100_000) + .expect("compacted checkpoint must record rebuilt context tokens"); + assert!(compacted_tokens < 90_000); + + let requests = provider.requests(); + assert_eq!(requests.len(), 3); + assert!(!requests[0].prompt.tools.is_empty()); + assert!(requests[1].prompt.tools.is_empty()); + assert!(!requests[2].prompt.tools.is_empty()); + assert!(requests[1] + .history + .iter() + .any(|message| match &message.content { + ProjectedContent::Assistant { text, .. } => text.len() == 400_000, + _ => false, + })); + assert!(requests[2] + .history + .iter() + .all(|message| match &message.content { + ProjectedContent::Assistant { text, .. } => text.len() != 400_000, + _ => true, + })); +} + +#[tokio::test] +async fn irreducibly_oversized_current_input_fails_before_provider_dispatch() { + let (_directory, store) = fixtures::temp_store().await; + let model = store + .create_model(&ModelConfigInput { + sort_order: 0, + display_name: "Overflow Model".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: "Overflow Model".into(), + model_id: "overflow-model".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(100_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(); + 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 output = run( + ®istry, + "overflow-request", + user_request( + "overflow-conversation", + "overflow-user", + &"x".repeat(400_000), + &model.model_hash, + None, + ), + ) + .await; + + assert!(provider.requests().is_empty()); + assert_eq!(output.summary_started, 0); + assert_eq!(output.summary_completed, 0); +} + #[derive(Default)] struct Output { checkpoints: Vec, diff --git a/server/tests/interrupt.rs b/server/tests/interrupt.rs index 0573e43..c2db150 100644 --- a/server/tests/interrupt.rs +++ b/server/tests/interrupt.rs @@ -1021,7 +1021,7 @@ async fn injected_user_context_interrupts_automatic_compaction() { custom_headers: serde_json::json!({}), anthropic_extra_params_enabled: false, anthropic_extra_params: serde_json::json!({}), - context_window_tokens: Some(10_001), + context_window_tokens: Some(100_000), max_completion_tokens: None, anthropic_max_tokens: None, anthropic_thinking_effort: None, @@ -1030,7 +1030,8 @@ async fn injected_user_context_interrupts_automatic_compaction() { .await .unwrap(); let provider = fake_provider::FakeProvider::default(); - provider.push(text_response("seed answer")); + let seed_answer = "x".repeat(400_000); + provider.push(text_response(&seed_answer)); provider.push_pending(); provider.push(text_response("continued after compacting injection")); let assets = PromptAssets::load( @@ -1045,7 +1046,7 @@ async fn injected_user_context_interrupts_automatic_compaction() { PromptCompiler::new(assets), ); - let seed_state = run_to_end( + run_to_end( ®istry, "seed-request", client_run_for_model( @@ -1065,7 +1066,7 @@ async fn injected_user_context_interrupts_automatic_compaction() { "inject-during-compaction", "compaction-injection-conversation", &model.model_hash, - Some(seed_state), + None, ); let Some(pb::agent_client_message::Message::RunRequest(request)) = compacting_request.message.as_mut() @@ -1075,9 +1076,17 @@ async fn injected_user_context_interrupts_automatic_compaction() { request.requested_model.as_mut().unwrap().parameters.push( pb::requested_model::ModelParameterValue { id: "context".into(), - value: "10001".into(), + value: "100000".into(), }, ); + let Some(pb::conversation_action::Action::UserMessageAction(action)) = request + .action + .as_mut() + .and_then(|action| action.action.as_mut()) + else { + panic!("expected UserMessageAction") + }; + action.user_message.as_mut().unwrap().message_id = "compaction-user".into(); handle .command(TransportCommand::Append { seqno: 0, @@ -1133,10 +1142,15 @@ async fn injected_user_context_interrupts_automatic_compaction() { let requests = provider.requests(); assert_eq!(requests.len(), 3); - assert!(requests[1] - .prompt - .instructions - .starts_with("Summarize the conversation for the next model turn.")); + assert!( + requests[1] + .prompt + .instructions + .starts_with("Summarize the conversation for the next model turn."), + "second request was not compaction: instructions={:?}, history={:?}", + requests[1].prompt.instructions, + requests[1].history + ); assert!(!serde_json::to_string(&requests[1].history) .unwrap() .contains("injected follow-up")); diff --git a/server/tests/tool_round.rs b/server/tests/tool_round.rs index 243090b..bc8d1e2 100644 --- a/server/tests/tool_round.rs +++ b/server/tests/tool_round.rs @@ -20,7 +20,10 @@ use cursor_server::{ }, }, cursor::{TransportCommand, TransportRegistry}, - model::{MessageContent, ToolCall}, + model::{ + MessageContent, ModelConfigInput, ModelType, ProjectedContent, ToolCall, + OPENAI_CHAT_ENDPOINT, + }, provider::{FinishReason, ModelEvent}, }; use prost::Message; @@ -745,6 +748,178 @@ async fn unknown_exec_id_is_a_protocol_error() { )); } +#[tokio::test] +async fn one_run_can_auto_compact_again_after_more_tool_output() { + let (_directory, store) = fixtures::temp_store().await; + let model = store + .create_model(&ModelConfigInput { + sort_order: 0, + display_name: "Repeated compaction".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: "Repeated compaction".into(), + model_id: "repeated-compaction-model".into(), + reasoning_effort: None, + openai_endpoint: OPENAI_CHAT_ENDPOINT.into(), + openai_extra_params_enabled: false, + openai_extra_params: json!({}), + custom_headers_enabled: false, + custom_headers: json!({}), + anthropic_extra_params_enabled: false, + anthropic_extra_params: json!({}), + context_window_tokens: Some(25_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(tool_call_response("repeat-call-1")); + provider.push(text_events("first summary")); + provider.push(tool_call_response("repeat-call-2")); + provider.push(text_events("second summary")); + provider.push(text_events("done")); + 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("repeated-compaction-request") + .await + .unwrap(); + let mut output = handle.subscribe(); + handle + .command(TransportCommand::Append { + seqno: 0, + message: Box::new(client_run_for_model( + "repeated-compaction-conversation", + "repeated-compaction-request", + &model.model_hash, + )), + }) + .await + .unwrap(); + let mut seqno = 1; + let oversized = format!("HEAD{}TAIL", "x".repeat(4 * 1024 * 1024)); + loop { + let frame = tokio::time::timeout(std::time::Duration::from_secs(10), output.recv()) + .await + .unwrap() + .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)) => { + 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/large.txt".into(), + total_lines: 1, + file_size: oversized.len() as i64, + output: Some( + pb::read_success::Output::Content( + oversized.clone(), + ), + ), + ..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; + } + _ => {} + } + } + + let requests = provider.requests(); + let shapes = requests + .iter() + .map(|request| { + ( + request.prompt.tools.len(), + request.history.len(), + request + .history + .iter() + .filter_map(|message| match &message.content { + ProjectedContent::ToolResult(result) => Some(result.content.len()), + _ => None, + }) + .collect::>(), + ) + }) + .collect::>(); + assert_eq!(requests.len(), 5, "provider requests: {shapes:?}"); + assert!(!requests[0].prompt.tools.is_empty()); + assert!(requests[1].prompt.tools.is_empty()); + assert!(!requests[2].prompt.tools.is_empty()); + assert!(requests[3].prompt.tools.is_empty()); + assert!(!requests[4].prompt.tools.is_empty()); +} + #[tokio::test] async fn provider_tool_use_waits_for_client_result_then_calls_provider_again() { let (directory, store) = fixtures::temp_store().await; @@ -935,6 +1110,73 @@ async fn provider_tool_use_waits_for_client_result_then_calls_provider_again() { assert_eq!(tool_calls[0].call_id, "call-1"); } +fn tool_call_response(call_id: &str) -> Vec { + vec![ + ModelEvent::Start { + model_call_id: format!("model-{call_id}"), + }, + ModelEvent::ToolCallStart { + index: 0, + call_id: call_id.into(), + name: "Read".into(), + }, + ModelEvent::ToolCallArgumentsDelta { + index: 0, + delta: "{\"path\":\"/tmp/large.txt\"}".into(), + }, + ModelEvent::ToolCallEnd { index: 0 }, + ModelEvent::Done(FinishReason::ToolUse), + ] +} + +fn text_events(text: &str) -> Vec { + vec![ + ModelEvent::Start { + model_call_id: format!("model-{text}"), + }, + ModelEvent::TextStart, + ModelEvent::TextDelta(text.into()), + ModelEvent::TextEnd, + ModelEvent::Done(FinishReason::Stop), + ] +} + +fn client_run_for_model( + conversation_id: &str, + run_id: &str, + model_id: &str, +) -> pb::AgentClientMessage { + let user = pb::UserMessage { + text: "read it".into(), + message_id: "user".into(), + mode: pb::AgentMode::Agent as i32, + ..Default::default() + }; + pb::AgentClientMessage { + message: Some(pb::agent_client_message::Message::RunRequest( + pb::AgentRunRequest { + action: Some(pb::ConversationAction { + action: Some(pb::conversation_action::Action::UserMessageAction( + pb::UserMessageAction { + user_message: Some(user), + request_context: Some(pb::RequestContext::default()), + ..Default::default() + }, + )), + ..Default::default() + }), + conversation_id: Some(conversation_id.into()), + run_id: Some(run_id.into()), + requested_model: Some(pb::RequestedModel { + model_id: model_id.into(), + ..Default::default() + }), + ..Default::default() + }, + )), + } +} + fn client_run() -> pb::AgentClientMessage { let user = pb::UserMessage { text: "read it".into(),