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;