mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-04 02:52:55 +08:00
merge: fix compaction input usage anchor
This commit is contained in:
@@ -1,6 +1,6 @@
|
|||||||
use serde::Serialize;
|
use serde::Serialize;
|
||||||
|
|
||||||
use super::ProviderType;
|
use super::{ProviderType, Usage};
|
||||||
|
|
||||||
#[derive(Clone, Debug)]
|
#[derive(Clone, Debug)]
|
||||||
pub struct NewLlmCall {
|
pub struct NewLlmCall {
|
||||||
@@ -22,6 +22,14 @@ pub struct NewLlmCall {
|
|||||||
pub detailed: bool,
|
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)]
|
#[derive(Clone, Debug, Serialize)]
|
||||||
pub struct LlmCallSummary {
|
pub struct LlmCallSummary {
|
||||||
pub call_id: String,
|
pub call_id: String,
|
||||||
|
|||||||
@@ -2,6 +2,8 @@ use std::ops::AddAssign;
|
|||||||
|
|
||||||
use serde::{Deserialize, Serialize};
|
use serde::{Deserialize, Serialize};
|
||||||
|
|
||||||
|
use super::ProviderType;
|
||||||
|
|
||||||
#[derive(Clone, Copy, Debug, Default, Serialize, Deserialize, PartialEq, Eq)]
|
#[derive(Clone, Copy, Debug, Default, Serialize, Deserialize, PartialEq, Eq)]
|
||||||
pub struct Usage {
|
pub struct Usage {
|
||||||
pub input_tokens: Option<u64>,
|
pub input_tokens: Option<u64>,
|
||||||
@@ -12,6 +14,19 @@ pub struct Usage {
|
|||||||
pub reasoning_tokens: Option<u64>,
|
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 {
|
impl AddAssign for Usage {
|
||||||
fn add_assign(&mut self, rhs: Self) {
|
fn add_assign(&mut self, rhs: Self) {
|
||||||
self.input_tokens = sum(self.input_tokens, rhs.input_tokens);
|
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)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
use super::Usage;
|
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]
|
#[test]
|
||||||
fn turn_total_only_reports_fields_known_for_every_cycle() {
|
fn turn_total_only_reports_fields_known_for_every_cycle() {
|
||||||
|
|||||||
+142
-7
@@ -175,7 +175,22 @@ impl RunEngine {
|
|||||||
Ok(messages) => messages,
|
Ok(messages) => messages,
|
||||||
Err(error) => return (RunOutcome::Failed(error.into()), usage),
|
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;
|
auto_compacted = true;
|
||||||
match self
|
match self
|
||||||
.auto_compact(prepared, revision, &messages, client, cancellation)
|
.auto_compact(prepared, revision, &messages, client, cancellation)
|
||||||
@@ -586,7 +601,28 @@ fn auto_compaction_partition(
|
|||||||
(compactable, retained)
|
(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 {
|
if prepared.action != RunAction::Start {
|
||||||
return false;
|
return false;
|
||||||
}
|
}
|
||||||
@@ -598,8 +634,18 @@ fn should_auto_compact(prepared: &PreparedRun, messages: &[CanonicalMessage]) ->
|
|||||||
{
|
{
|
||||||
return false;
|
return false;
|
||||||
}
|
}
|
||||||
estimate_context_tokens(&prepared.prompt, messages)
|
let estimated_input = anchor
|
||||||
> context_window.saturating_sub(COMPACTION_RESERVE_TOKENS)
|
.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(
|
fn estimate_context_tokens(
|
||||||
@@ -607,6 +653,15 @@ fn estimate_context_tokens(
|
|||||||
messages: &[CanonicalMessage],
|
messages: &[CanonicalMessage],
|
||||||
) -> u64 {
|
) -> u64 {
|
||||||
let serialized = serde_json::to_string(&(prompt, messages)).unwrap_or_default();
|
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
|
serialized
|
||||||
.chars()
|
.chars()
|
||||||
.fold(0_u64, |units, character| {
|
.fold(0_u64, |units, character| {
|
||||||
@@ -745,11 +800,15 @@ fn failure_message(failure: &RunFailure) -> String {
|
|||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
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::{
|
use crate::{
|
||||||
model::{
|
model::{
|
||||||
CanonicalMessage, ContentPart, Origin, ProjectedContent, ProjectedMessage, PromptSpec,
|
CanonicalMessage, ContentPart, ConversationId, ModelSpec, Origin, PreparedRun,
|
||||||
Role, ToolImageReference, ToolResultContent,
|
ProjectedContent, ProjectedMessage, PromptSpec, RevisionId, Role, RunAction, RunId,
|
||||||
|
RunKind, ToolImageReference, ToolResultContent,
|
||||||
},
|
},
|
||||||
store::Store,
|
store::Store,
|
||||||
};
|
};
|
||||||
@@ -778,6 +837,82 @@ mod tests {
|
|||||||
assert!(estimate_context_tokens(&prompt, &long) > estimate_context_tokens(&prompt, &short));
|
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]
|
#[test]
|
||||||
fn auto_compaction_preserves_only_the_latest_request_context() {
|
fn auto_compaction_preserves_only_the_latest_request_context() {
|
||||||
let first_context = CanonicalMessage::text(
|
let first_context = CanonicalMessage::text(
|
||||||
|
|||||||
@@ -1,7 +1,12 @@
|
|||||||
|
use std::str::FromStr;
|
||||||
|
|
||||||
use sqlx::Row;
|
use sqlx::Row;
|
||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
model::{LlmCallRequest, LlmCallResponseChunk, LlmCallSummary, NewLlmCall, Usage},
|
model::{
|
||||||
|
ConversationId, LlmCallRequest, LlmCallResponseChunk, LlmCallSummary, LlmCallUsageAnchor,
|
||||||
|
NewLlmCall, ProviderType, Usage,
|
||||||
|
},
|
||||||
Result,
|
Result,
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -197,6 +202,41 @@ impl Store {
|
|||||||
.transpose()
|
.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>> {
|
pub async fn llm_call_request(&self, call_id: &str) -> Result<Option<LlmCallRequest>> {
|
||||||
let row = sqlx::query(
|
let row = sqlx::query(
|
||||||
"SELECT headers_json, body_json, byte_count FROM llm_call_requests WHERE call_id = ?",
|
"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")?,
|
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