merge: fix compaction input usage anchor

This commit is contained in:
leookun
2026-08-26 01:49:15 +08:00
4 changed files with 335 additions and 9 deletions
+9 -1
View File
@@ -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,
+50
View File
@@ -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<u64>,
@@ -12,6 +14,19 @@ pub struct 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 => 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<u64>, right: Option<u64>) -> Option<u64> {
#[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() {
+142 -7
View File
@@ -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<Self> {
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<ContextUsageAnchor>,
) -> 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(
+134 -1
View File
@@ -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<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 = ?",
@@ -284,3 +324,96 @@ fn summary_from_row(row: sqlx::sqlite::SqliteRow) -> Result<LlmCallSummary> {
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);
}
}