diff --git a/server/src/model/llm_call.rs b/server/src/model/llm_call.rs index 3947a21..123a80b 100644 --- a/server/src/model/llm_call.rs +++ b/server/src/model/llm_call.rs @@ -1,6 +1,6 @@ use serde::Serialize; -use super::ProviderType; +use super::{ProviderType, Usage}; #[derive(Clone, Debug)] pub struct NewLlmCall { @@ -22,6 +22,14 @@ pub struct NewLlmCall { 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/usage.rs b/server/src/model/usage.rs index d4a3c1b..e8b245e 100644 --- a/server/src/model/usage.rs +++ b/server/src/model/usage.rs @@ -2,6 +2,8 @@ use std::ops::AddAssign; use serde::{Deserialize, Serialize}; +use super::ProviderType; + #[derive(Clone, Copy, Debug, Default, Serialize, Deserialize, PartialEq, Eq)] pub struct Usage { pub input_tokens: Option, @@ -12,6 +14,19 @@ pub struct 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 => 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); @@ -30,6 +45,41 @@ fn sum(left: Option, right: Option) -> Option { #[cfg(test)] mod tests { use super::Usage; + use crate::model::ProviderType; + + #[test] + fn openai_context_input_does_not_double_count_cached_tokens() { + let usage = Usage { + input_tokens: Some(140_649), + cache_read_tokens: Some(120_000), + cache_write_tokens: Some(10_000), + ..Usage::default() + }; + + assert_eq!( + usage.context_input_tokens(ProviderType::OpenAiResponses), + Some(140_649) + ); + assert_eq!( + usage.context_input_tokens(ProviderType::OpenAiChat), + Some(140_649) + ); + } + + #[test] + fn anthropic_context_input_includes_disjoint_cache_tokens() { + let usage = Usage { + input_tokens: Some(10_649), + cache_read_tokens: Some(120_000), + cache_write_tokens: Some(10_000), + ..Usage::default() + }; + + assert_eq!( + usage.context_input_tokens(ProviderType::Anthropic), + Some(140_649) + ); + } #[test] fn turn_total_only_reports_fields_known_for_every_cycle() { diff --git a/server/src/run/engine.rs b/server/src/run/engine.rs index c98de86..09e8369 100644 --- a/server/src/run/engine.rs +++ b/server/src/run/engine.rs @@ -175,7 +175,22 @@ impl RunEngine { Ok(messages) => messages, Err(error) => return (RunOutcome::Failed(error.into()), usage), }; - if !auto_compacted && should_auto_compact(prepared, &messages) { + let context_anchor = if !auto_compacted && prepared.action == RunAction::Start { + match self + .store + .latest_llm_call_usage_anchor( + &prepared.conversation_id, + &prepared.model.model_id, + ) + .await + { + Ok(anchor) => anchor.and_then(ContextUsageAnchor::from_llm_call), + Err(error) => return (RunOutcome::Failed(error.into()), usage), + } + } else { + None + }; + if !auto_compacted && should_auto_compact(prepared, &messages, context_anchor) { auto_compacted = true; match self .auto_compact(prepared, revision, &messages, client, cancellation) @@ -586,7 +601,28 @@ fn auto_compaction_partition( (compactable, retained) } -fn should_auto_compact(prepared: &PreparedRun, messages: &[CanonicalMessage]) -> bool { +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +struct ContextUsageAnchor { + input_tokens: u64, + message_count: usize, + tool_count: usize, +} + +impl ContextUsageAnchor { + fn from_llm_call(anchor: crate::model::LlmCallUsageAnchor) -> Option { + Some(Self { + input_tokens: anchor.usage.context_input_tokens(anchor.request_type)?, + message_count: anchor.message_count, + tool_count: anchor.tool_count, + }) + } +} + +fn should_auto_compact( + prepared: &PreparedRun, + messages: &[CanonicalMessage], + anchor: Option, +) -> bool { if prepared.action != RunAction::Start { return false; } @@ -598,8 +634,18 @@ fn should_auto_compact(prepared: &PreparedRun, messages: &[CanonicalMessage]) -> { return false; } - estimate_context_tokens(&prepared.prompt, messages) - > context_window.saturating_sub(COMPACTION_RESERVE_TOKENS) + let estimated_input = anchor + .filter(|anchor| { + anchor.message_count <= messages.len() + && anchor.tool_count == prepared.prompt.tools.len() + }) + .map(|anchor| { + anchor + .input_tokens + .saturating_add(estimate_message_tokens(&messages[anchor.message_count..])) + }) + .unwrap_or_else(|| estimate_context_tokens(&prepared.prompt, messages)); + estimated_input > context_window.saturating_sub(COMPACTION_RESERVE_TOKENS) } fn estimate_context_tokens( @@ -607,6 +653,15 @@ fn estimate_context_tokens( messages: &[CanonicalMessage], ) -> u64 { let serialized = serde_json::to_string(&(prompt, messages)).unwrap_or_default(); + estimate_serialized_tokens(&serialized) +} + +fn estimate_message_tokens(messages: &[CanonicalMessage]) -> u64 { + let serialized = serde_json::to_string(messages).unwrap_or_default(); + estimate_serialized_tokens(&serialized) +} + +fn estimate_serialized_tokens(serialized: &str) -> u64 { serialized .chars() .fold(0_u64, |units, character| { @@ -745,11 +800,15 @@ fn failure_message(failure: &RunFailure) -> String { #[cfg(test)] mod tests { - use super::{auto_compaction_partition, estimate_context_tokens, hydrate_tool_images}; + use super::{ + auto_compaction_partition, estimate_context_tokens, hydrate_tool_images, + should_auto_compact, ContextUsageAnchor, + }; use crate::{ model::{ - CanonicalMessage, ContentPart, Origin, ProjectedContent, ProjectedMessage, PromptSpec, - Role, ToolImageReference, ToolResultContent, + CanonicalMessage, ContentPart, ConversationId, ModelSpec, Origin, PreparedRun, + ProjectedContent, ProjectedMessage, PromptSpec, RevisionId, Role, RunAction, RunId, + RunKind, ToolImageReference, ToolResultContent, }, store::Store, }; @@ -778,6 +837,82 @@ mod tests { assert!(estimate_context_tokens(&prompt, &long) > estimate_context_tokens(&prompt, &short)); } + #[test] + fn real_previous_input_only_estimates_messages_added_after_the_anchor() { + let old_history = + CanonicalMessage::text("old-history", Role::User, Origin::User, "x".repeat(698_641)); + let current_runtime = CanonicalMessage::text( + "runtime:current", + Role::User, + Origin::Runtime, + "current request", + ); + let messages = vec![old_history, current_runtime.clone()]; + let prepared = PreparedRun { + run_id: RunId::new("run"), + conversation_id: ConversationId::new("conversation"), + kind: RunKind::Root, + model: ModelSpec { + context_window_tokens: Some(200_000), + ..ModelSpec::new("model") + }, + prompt: PromptSpec { + instructions: "system".into(), + tools: Vec::new(), + }, + initial_messages: vec![current_runtime], + action: RunAction::Start, + base_revision_id: RevisionId(1), + }; + let anchor = ContextUsageAnchor { + input_tokens: 140_649, + message_count: 1, + tool_count: 0, + }; + + assert_eq!( + estimate_context_tokens(&prepared.prompt, &messages), + 190_813 + ); + assert!(!should_auto_compact(&prepared, &messages, Some(anchor))); + } + + #[test] + fn real_previous_input_compacts_after_the_new_message_crosses_the_reserve() { + let old_history = + CanonicalMessage::text("old-history", Role::User, Origin::User, "old history"); + let current_runtime = CanonicalMessage::text( + "runtime:current", + Role::User, + Origin::Runtime, + "x".repeat(190_000), + ); + let messages = vec![old_history, current_runtime.clone()]; + let prepared = PreparedRun { + run_id: RunId::new("run"), + conversation_id: ConversationId::new("conversation"), + kind: RunKind::Root, + model: ModelSpec { + context_window_tokens: Some(200_000), + ..ModelSpec::new("model") + }, + prompt: PromptSpec { + instructions: "system".into(), + tools: Vec::new(), + }, + initial_messages: vec![current_runtime], + action: RunAction::Start, + base_revision_id: RevisionId(1), + }; + let anchor = ContextUsageAnchor { + input_tokens: 140_649, + message_count: 1, + tool_count: 0, + }; + + assert!(should_auto_compact(&prepared, &messages, Some(anchor))); + } + #[test] fn auto_compaction_preserves_only_the_latest_request_context() { let first_context = CanonicalMessage::text( diff --git a/server/src/store/llm_calls.rs b/server/src/store/llm_calls.rs index 248b32b..573c108 100644 --- a/server/src/store/llm_calls.rs +++ b/server/src/store/llm_calls.rs @@ -1,7 +1,12 @@ +use std::str::FromStr; + use sqlx::Row; use crate::{ - model::{LlmCallRequest, LlmCallResponseChunk, LlmCallSummary, NewLlmCall, Usage}, + model::{ + ConversationId, LlmCallRequest, LlmCallResponseChunk, LlmCallSummary, LlmCallUsageAnchor, + NewLlmCall, ProviderType, Usage, + }, Result, }; @@ -197,6 +202,41 @@ 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 = ?", @@ -284,3 +324,96 @@ fn summary_from_row(row: sqlx::sqlite::SqliteRow) -> Result { detailed: row.try_get("detailed")?, }) } + +#[cfg(test)] +mod tests { + use super::*; + use crate::model::{ProviderEndpointInput, ProviderModelInput}; + + #[tokio::test] + async fn latest_usage_anchor_uses_the_latest_completed_call_for_the_same_conversation_and_model( + ) { + let store = Store::connect("sqlite::memory:").await.unwrap(); + let (provider, model) = store + .create_provider_with_model( + &ProviderEndpointInput { + name: "Test".into(), + provider_type: ProviderType::OpenAiResponses, + base_url: "https://example.com".into(), + api_key: Some("secret".into()), + custom_headers: serde_json::json!({}), + extra_params: serde_json::json!({}), + }, + &ProviderModelInput { + model_id: "model".into(), + display_name: "Model".into(), + endpoint_type: ProviderType::OpenAiResponses, + request_url: String::new(), + enabled: true, + sort_order: 0, + context_window_tokens: Some(200_000), + max_output_tokens: Some(16_000), + reasoning_enabled: false, + reasoning_effort: None, + supports_image_generation: false, + }, + ) + .await + .unwrap(); + let conversation_id = ConversationId::new("conversation"); + + for (call_id, status, input_tokens, message_count) in [ + ("completed-old", "completed", 120_000, 10), + ("failed-newer", "error", 180_000, 11), + ("completed-latest", "completed", 140_649, 12), + ] { + store + .start_llm_call(&NewLlmCall { + call_id: call_id.into(), + run_id: format!("run-{call_id}"), + conversation_id: conversation_id.to_string(), + provider_call_index: 0, + model_hash: model.model_hash.clone(), + provider_type: provider.provider_type, + provider_url: provider.base_url.clone(), + request_type: model.endpoint_type, + request_url: model.request_url.clone(), + model_id: model.model_id.clone(), + display_name: model.display_name.clone(), + reasoning_effort: None, + fast: false, + message_count, + tool_count: 7, + detailed: false, + }) + .await + .unwrap(); + store + .record_llm_usage( + call_id, + Usage { + input_tokens: Some(input_tokens), + cache_read_tokens: Some(100_000), + ..Usage::default() + }, + ) + .await + .unwrap(); + store + .finish_llm_call(call_id, status, None, 1, None, None) + .await + .unwrap(); + } + + let anchor = store + .latest_llm_call_usage_anchor(&conversation_id, &model.model_hash) + .await + .unwrap() + .unwrap(); + + assert_eq!(anchor.request_type, ProviderType::OpenAiResponses); + assert_eq!(anchor.usage.input_tokens, Some(140_649)); + assert_eq!(anchor.message_count, 12); + assert_eq!(anchor.tool_count, 7); + } +}