mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-07 06:04:53 +08:00
feat: implement context usage anchor for improved token estimation
- Introduced `ContextUsageAnchor` struct to track context input tokens and message count for conversations. - Updated token estimation functions to utilize the context usage anchor, enhancing accuracy in estimating tokens for projected messages. - Refactored compaction logic to incorporate context usage anchor, allowing for more efficient management of token budgets during model runs. - Added tests to validate the behavior of the context usage anchor across different scenarios, including model switching and message additions.
This commit is contained in:
@@ -11,10 +11,15 @@ pub(crate) fn estimate_context_tokens(prompt: &PromptSpec, messages: &[Projected
|
|||||||
let tools = prompt.tools.iter().fold(0_u64, |total, tool| {
|
let tools = prompt.tools.iter().fold(0_u64, |total, tool| {
|
||||||
total.saturating_add(estimate_json_tokens(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))
|
total.saturating_add(estimate_message_tokens(message))
|
||||||
});
|
})
|
||||||
instructions.saturating_add(tools).saturating_add(messages)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn estimate_message_tokens(message: &ProjectedMessage) -> u64 {
|
fn estimate_message_tokens(message: &ProjectedMessage) -> u64 {
|
||||||
@@ -23,7 +28,7 @@ fn estimate_message_tokens(message: &ProjectedMessage) -> u64 {
|
|||||||
ProjectedContent::Assistant {
|
ProjectedContent::Assistant {
|
||||||
text,
|
text,
|
||||||
thinking,
|
thinking,
|
||||||
replay_state,
|
replay_state: _,
|
||||||
calls,
|
calls,
|
||||||
} => {
|
} => {
|
||||||
let calls = calls.iter().fold(0_u64, |total, call| {
|
let calls = calls.iter().fold(0_u64, |total, call| {
|
||||||
@@ -35,12 +40,6 @@ fn estimate_message_tokens(message: &ProjectedMessage) -> u64 {
|
|||||||
});
|
});
|
||||||
estimate_text_tokens(text)
|
estimate_text_tokens(text)
|
||||||
.saturating_add(estimate_text_tokens(thinking))
|
.saturating_add(estimate_text_tokens(thinking))
|
||||||
.saturating_add(
|
|
||||||
replay_state
|
|
||||||
.as_ref()
|
|
||||||
.map(estimate_json_tokens)
|
|
||||||
.unwrap_or_default(),
|
|
||||||
)
|
|
||||||
.saturating_add(calls)
|
.saturating_add(calls)
|
||||||
}
|
}
|
||||||
ProjectedContent::ToolResult(result) => {
|
ProjectedContent::ToolResult(result) => {
|
||||||
@@ -118,7 +117,9 @@ pub(crate) fn format_token_count(tokens: u64) -> String {
|
|||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
use super::*;
|
use super::*;
|
||||||
use crate::model::{Role, ToolCallContent, ToolDefinition, ToolResultContent};
|
use crate::model::{
|
||||||
|
ProviderReplayState, Role, ToolCallContent, ToolDefinition, ToolResultContent,
|
||||||
|
};
|
||||||
|
|
||||||
fn prompt() -> PromptSpec {
|
fn prompt() -> PromptSpec {
|
||||||
PromptSpec {
|
PromptSpec {
|
||||||
@@ -214,4 +215,34 @@ mod tests {
|
|||||||
assert!(with_text > with_call);
|
assert!(with_text > with_call);
|
||||||
assert!(with_image > with_text);
|
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])
|
||||||
|
);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -2,7 +2,13 @@
|
|||||||
|
|
||||||
use std::collections::HashSet;
|
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;
|
const FALLBACK_CHARS: usize = 12_000;
|
||||||
|
|
||||||
@@ -20,25 +26,36 @@ pub(super) fn input_budget(prepared: &PreparedRun) -> Option<u64> {
|
|||||||
pub(super) fn estimated_tokens(
|
pub(super) fn estimated_tokens(
|
||||||
prepared: &PreparedRun,
|
prepared: &PreparedRun,
|
||||||
projected_messages: &[ProjectedMessage],
|
projected_messages: &[ProjectedMessage],
|
||||||
|
anchor: Option<ContextUsageAnchor>,
|
||||||
) -> u64 {
|
) -> 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(
|
pub(super) fn should_compact(
|
||||||
prepared: &PreparedRun,
|
prepared: &PreparedRun,
|
||||||
projected_messages: &[ProjectedMessage],
|
projected_messages: &[ProjectedMessage],
|
||||||
|
anchor: Option<ContextUsageAnchor>,
|
||||||
) -> bool {
|
) -> bool {
|
||||||
let Some(budget) = input_budget(prepared) else {
|
let Some(budget) = input_budget(prepared) else {
|
||||||
return false;
|
return false;
|
||||||
};
|
};
|
||||||
estimated_tokens(prepared, projected_messages) > budget
|
estimated_tokens(prepared, projected_messages, anchor) > budget
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(super) fn validate_compacted(
|
pub(super) fn validate_compacted(
|
||||||
prepared: &PreparedRun,
|
prepared: &PreparedRun,
|
||||||
projected_messages: &[ProjectedMessage],
|
projected_messages: &[ProjectedMessage],
|
||||||
) -> std::result::Result<u64, String> {
|
) -> std::result::Result<u64, String> {
|
||||||
let estimated = estimated_tokens(prepared, projected_messages);
|
let estimated = estimate_context_tokens(&prepared.prompt, projected_messages);
|
||||||
let Some(budget) = input_budget(prepared) else {
|
let Some(budget) = input_budget(prepared) else {
|
||||||
return Ok(estimated);
|
return Ok(estimated);
|
||||||
};
|
};
|
||||||
@@ -125,14 +142,97 @@ mod tests {
|
|||||||
let estimated = estimate_context_tokens(&prepared(1).prompt, &projected);
|
let estimated = estimate_context_tokens(&prepared(1).prompt, &projected);
|
||||||
let mut prepared = prepared(estimated + RESERVE_TOKENS);
|
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);
|
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 {
|
prepared.action = RunAction::Resume {
|
||||||
pending_tool_round: None,
|
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]
|
#[test]
|
||||||
|
|||||||
@@ -10,7 +10,7 @@ use crate::{
|
|||||||
ToolRoundId, Usage,
|
ToolRoundId, Usage,
|
||||||
},
|
},
|
||||||
provider::Provider,
|
provider::Provider,
|
||||||
store::{RunStatus, Store},
|
store::{ContextUsageAnchor, RunStatus, Store},
|
||||||
};
|
};
|
||||||
|
|
||||||
use super::{
|
use super::{
|
||||||
@@ -92,6 +92,14 @@ impl RunEngine {
|
|||||||
cancellation: &CancellationToken,
|
cancellation: &CancellationToken,
|
||||||
) -> (RunOutcome, Option<Usage>) {
|
) -> (RunOutcome, Option<Usage>) {
|
||||||
let mut usage = None;
|
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!(
|
tracing::info!(
|
||||||
checkpoint_id = checkpoint.0,
|
checkpoint_id = checkpoint.0,
|
||||||
"Run claimed conversation ownership"
|
"Run claimed conversation ownership"
|
||||||
@@ -176,7 +184,7 @@ impl RunEngine {
|
|||||||
Err(error) => return (RunOutcome::Failed(error.into()), usage),
|
Err(error) => return (RunOutcome::Failed(error.into()), usage),
|
||||||
};
|
};
|
||||||
if prepared.action != RunAction::Compact
|
if prepared.action != RunAction::Compact
|
||||||
&& super::compaction::should_compact(prepared, &history)
|
&& super::compaction::should_compact(prepared, &history, context_usage_anchor)
|
||||||
{
|
{
|
||||||
match self
|
match self
|
||||||
.auto_compact(prepared, checkpoint, &messages, client, cancellation)
|
.auto_compact(prepared, checkpoint, &messages, client, cancellation)
|
||||||
@@ -184,6 +192,7 @@ impl RunEngine {
|
|||||||
{
|
{
|
||||||
Ok((next_checkpoint, compaction_usage)) => {
|
Ok((next_checkpoint, compaction_usage)) => {
|
||||||
checkpoint = next_checkpoint;
|
checkpoint = next_checkpoint;
|
||||||
|
context_usage_anchor = None;
|
||||||
if let Some(compaction_usage) = compaction_usage {
|
if let Some(compaction_usage) = compaction_usage {
|
||||||
accumulate_usage(&mut usage, compaction_usage);
|
accumulate_usage(&mut usage, compaction_usage);
|
||||||
}
|
}
|
||||||
@@ -272,11 +281,21 @@ impl RunEngine {
|
|||||||
match interrupted {
|
match interrupted {
|
||||||
Ok(cycle) => {
|
Ok(cycle) => {
|
||||||
if let Some(cycle_usage) = cycle.usage {
|
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);
|
accumulate_usage(&mut usage, cycle_usage);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
Err(failure) => {
|
Err(failure) => {
|
||||||
if let Some(cycle_usage) = failure.usage {
|
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);
|
accumulate_usage(&mut usage, cycle_usage);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -319,6 +338,11 @@ impl RunEngine {
|
|||||||
Ok(cycle) => break 'attempt cycle,
|
Ok(cycle) => break 'attempt cycle,
|
||||||
Err(cycle_failure) => {
|
Err(cycle_failure) => {
|
||||||
if let Some(cycle_usage) = cycle_failure.usage {
|
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);
|
accumulate_usage(&mut usage, cycle_usage);
|
||||||
}
|
}
|
||||||
if cancellation.is_cancelled() {
|
if cancellation.is_cancelled() {
|
||||||
@@ -420,6 +444,11 @@ impl RunEngine {
|
|||||||
}
|
}
|
||||||
};
|
};
|
||||||
if let Some(cycle_usage) = cycle.usage {
|
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);
|
accumulate_usage(&mut usage, cycle_usage);
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -879,6 +908,19 @@ async fn hydrate_tool_images(
|
|||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn update_context_usage_anchor(
|
||||||
|
anchor: &mut Option<ContextUsageAnchor>,
|
||||||
|
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: Usage) {
|
fn accumulate_usage(total: &mut Option<Usage>, usage: Usage) {
|
||||||
match total {
|
match total {
|
||||||
Some(total) => *total += usage,
|
Some(total) => *total += usage,
|
||||||
|
|||||||
@@ -8,6 +8,12 @@ use crate::{
|
|||||||
|
|
||||||
use super::{now_ms, Store};
|
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)]
|
#[derive(Clone, Debug)]
|
||||||
pub(crate) struct BufferedLlmChunk {
|
pub(crate) struct BufferedLlmChunk {
|
||||||
pub(crate) seq: i64,
|
pub(crate) seq: i64,
|
||||||
@@ -266,6 +272,37 @@ impl Store {
|
|||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub(crate) async fn latest_context_usage(
|
||||||
|
&self,
|
||||||
|
conversation_id: &str,
|
||||||
|
) -> Result<Option<ContextUsageAnchor>> {
|
||||||
|
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::<i64, _>("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<Vec<LlmCallSummary>> {
|
pub async fn llm_calls(&self, limit: i64) -> Result<Vec<LlmCallSummary>> {
|
||||||
let rows = sqlx::query("SELECT * FROM llm_calls ORDER BY created_at_ms DESC LIMIT ?")
|
let rows = sqlx::query("SELECT * FROM llm_calls ORDER BY created_at_ms DESC LIMIT ?")
|
||||||
.bind(limit.clamp(1, 500))
|
.bind(limit.clamp(1, 500))
|
||||||
@@ -421,4 +458,65 @@ mod tests {
|
|||||||
assert_eq!(overview.metrics.llm_calls, 1);
|
assert_eq!(overview.metrics.llm_calls, 1);
|
||||||
assert_eq!(overview.metrics.successful_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);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -19,7 +19,7 @@ mod writer;
|
|||||||
|
|
||||||
pub use cas::*;
|
pub use cas::*;
|
||||||
pub(crate) use cursor_traces::BufferedCursorTraceChunk;
|
pub(crate) use cursor_traces::BufferedCursorTraceChunk;
|
||||||
pub(crate) use llm_calls::BufferedLlmChunk;
|
pub(crate) use llm_calls::{BufferedLlmChunk, ContextUsageAnchor};
|
||||||
pub use runs::*;
|
pub use runs::*;
|
||||||
pub use settings::*;
|
pub use settings::*;
|
||||||
pub(crate) use sqlite::now_ms;
|
pub(crate) use sqlite::now_ms;
|
||||||
|
|||||||
@@ -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]
|
#[tokio::test]
|
||||||
async fn irreducibly_oversized_current_input_fails_before_provider_dispatch() {
|
async fn irreducibly_oversized_current_input_fails_before_provider_dispatch() {
|
||||||
let (_directory, store) = fixtures::temp_store().await;
|
let (_directory, store) = fixtures::temp_store().await;
|
||||||
|
|||||||
Reference in New Issue
Block a user