mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-03 18:23:51 +08:00
fix: anchor compaction to provider input usage
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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
@@ -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(
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user