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.
This commit is contained in:
leokun
2026-09-01 10:10:09 +08:00
parent ee2592c469
commit 29fde7d7c7
9 changed files with 723 additions and 188 deletions
+7 -1
View File
@@ -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,
+1 -26
View File
@@ -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<u64>,
@@ -18,21 +16,6 @@ mod usage {
pub reasoning_tokens: Option<u64>,
}
impl Usage {
/// Returns the provider-visible input context without counting cached tokens twice.
pub(crate) fn context_input_tokens(self, provider: ProviderType) -> Option<u64> {
let input = self.input_tokens?;
match provider {
ProviderType::OpenAiChat | ProviderType::OpenAiResponses | ProviderType::Plugin => {
Some(input)
}
ProviderType::Anthropic => input
.checked_add(self.cache_read_tokens.unwrap_or_default())?
.checked_add(self.cache_write_tokens.unwrap_or_default()),
}
}
}
impl AddAssign for Usage {
fn add_assign(&mut self, rhs: Self) {
self.input_tokens = sum(self.input_tokens, rhs.input_tokens);
@@ -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,
+198 -1
View File
@@ -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<u64> {
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);
}
}
+67 -86
View File
@@ -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<u64> {
prepared
.model
.context_window_tokens
.map(|window| window.saturating_sub(RESERVE_TOKENS))
}
impl ContextUsageAnchor {
pub(super) fn from_llm_call(anchor: LlmCallUsageAnchor) -> Option<Self> {
Some(Self {
input_tokens: anchor.usage.context_input_tokens(anchor.request_type)?,
message_count: anchor.message_count,
tool_count: anchor.tool_count,
})
}
pub(super) fn estimated_tokens(
prepared: &PreparedRun,
projected_messages: &[ProjectedMessage],
) -> u64 {
estimate_context_tokens(&prepared.prompt, projected_messages)
}
pub(super) fn should_compact(
prepared: &PreparedRun,
messages: &[CanonicalMessage],
projected_messages: &[ProjectedMessage],
anchor: Option<ContextUsageAnchor>,
) -> 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<u64, String> {
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")
);
}
}
+16 -23
View File
@@ -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, &current_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::<Vec<_>>();
replacement.push(summary_message);
replacement.extend(prepared.initial_messages.iter().cloned());
let projected_replacement = crate::model::project_messages(&replacement)
.map_err(|error| RunOutcome::Failed(error.into()))?;
super::compaction::validate_compacted(prepared, &projected_replacement)
.map_err(|message| RunOutcome::Failed(RunFailure::Protocol(message)))?;
let mut checkpoint = self
.store
.replace_checkpoint(
+2 -41
View File
@@ -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<Option<LlmCallUsageAnchor>> {
let row = sqlx::query(
r#"SELECT request_type, usage_json, message_count, tool_count
FROM llm_calls
WHERE conversation_id = ?
AND model_hash = ?
AND status = 'completed'
AND input_tokens IS NOT NULL
AND usage_json IS NOT NULL
ORDER BY rowid DESC
LIMIT 1"#,
)
.bind(conversation_id.as_str())
.bind(model_hash)
.fetch_optional(&self.pool)
.await?;
row.map(|row| {
let message_count =
usize::try_from(row.try_get::<i64, _>("message_count")?).unwrap_or(usize::MAX);
let tool_count =
usize::try_from(row.try_get::<i64, _>("tool_count")?).unwrap_or(usize::MAX);
Ok(LlmCallUsageAnchor {
request_type: ProviderType::from_str(row.try_get("request_type")?)?,
usage: serde_json::from_str(row.try_get("usage_json")?)?,
message_count,
tool_count,
})
})
.transpose()
}
pub async fn llm_call_request(&self, call_id: &str) -> Result<Option<LlmCallRequest>> {
let row = sqlx::query(
"SELECT headers_json, body_json, byte_count FROM llm_call_requests WHERE call_id = ?",
@@ -416,6 +376,7 @@ fn summary_from_row(row: sqlx::sqlite::SqliteRow) -> Result<LlmCallSummary> {
#[cfg(test)]
mod tests {
use super::*;
use crate::model::ProviderType;
/// 插件模型不在 model_configs 中,调用记录必须照常落库并可按其稳定 ID 筛选。
#[tokio::test]
+166
View File
@@ -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(
&registry,
"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(
&registry,
"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(
&registry,
"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<pb::ConversationStateStructure>,
+23 -9
View File
@@ -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(
&registry,
"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"));
+243 -1
View File
@@ -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::<Vec<_>>(),
)
})
.collect::<Vec<_>>();
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<ModelEvent> {
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<ModelEvent> {
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(),