Merge branch 'main' of github.com:leookun/cursor-byok

This commit is contained in:
leokun
2026-09-01 16:08:34 +08:00
6 changed files with 394 additions and 21 deletions
+42 -11
View File
@@ -11,10 +11,15 @@ pub(crate) fn estimate_context_tokens(prompt: &PromptSpec, messages: &[Projected
let tools = prompt.tools.iter().fold(0_u64, |total, 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))
});
instructions.saturating_add(tools).saturating_add(messages)
})
}
fn estimate_message_tokens(message: &ProjectedMessage) -> u64 {
@@ -23,7 +28,7 @@ fn estimate_message_tokens(message: &ProjectedMessage) -> u64 {
ProjectedContent::Assistant {
text,
thinking,
replay_state,
replay_state: _,
calls,
} => {
let calls = calls.iter().fold(0_u64, |total, call| {
@@ -35,12 +40,6 @@ fn estimate_message_tokens(message: &ProjectedMessage) -> u64 {
});
estimate_text_tokens(text)
.saturating_add(estimate_text_tokens(thinking))
.saturating_add(
replay_state
.as_ref()
.map(estimate_json_tokens)
.unwrap_or_default(),
)
.saturating_add(calls)
}
ProjectedContent::ToolResult(result) => {
@@ -118,7 +117,9 @@ pub(crate) fn format_token_count(tokens: u64) -> String {
#[cfg(test)]
mod tests {
use super::*;
use crate::model::{Role, ToolCallContent, ToolDefinition, ToolResultContent};
use crate::model::{
ProviderReplayState, Role, ToolCallContent, ToolDefinition, ToolResultContent,
};
fn prompt() -> PromptSpec {
PromptSpec {
@@ -214,4 +215,34 @@ mod tests {
assert!(with_text > with_call);
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])
);
}
}
+107 -7
View File
@@ -2,7 +2,13 @@
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;
@@ -20,25 +26,36 @@ pub(super) fn input_budget(prepared: &PreparedRun) -> Option<u64> {
pub(super) fn estimated_tokens(
prepared: &PreparedRun,
projected_messages: &[ProjectedMessage],
anchor: Option<ContextUsageAnchor>,
) -> 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(
prepared: &PreparedRun,
projected_messages: &[ProjectedMessage],
anchor: Option<ContextUsageAnchor>,
) -> bool {
let Some(budget) = input_budget(prepared) else {
return false;
};
estimated_tokens(prepared, projected_messages) > budget
estimated_tokens(prepared, projected_messages, anchor) > budget
}
pub(super) fn validate_compacted(
prepared: &PreparedRun,
projected_messages: &[ProjectedMessage],
) -> 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 {
return Ok(estimated);
};
@@ -125,14 +142,97 @@ mod tests {
let estimated = estimate_context_tokens(&prepared(1).prompt, &projected);
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);
assert!(should_compact(&prepared, &projected));
assert!(should_compact(&prepared, &projected, None));
prepared.action = RunAction::Resume {
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]
+44 -2
View File
@@ -10,7 +10,7 @@ use crate::{
ToolRoundId, Usage,
},
provider::Provider,
store::{RunStatus, Store},
store::{ContextUsageAnchor, RunStatus, Store},
};
use super::{
@@ -92,6 +92,14 @@ impl RunEngine {
cancellation: &CancellationToken,
) -> (RunOutcome, Option<Usage>) {
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!(
checkpoint_id = checkpoint.0,
"Run claimed conversation ownership"
@@ -176,7 +184,7 @@ impl RunEngine {
Err(error) => return (RunOutcome::Failed(error.into()), usage),
};
if prepared.action != RunAction::Compact
&& super::compaction::should_compact(prepared, &history)
&& super::compaction::should_compact(prepared, &history, context_usage_anchor)
{
match self
.auto_compact(prepared, checkpoint, &messages, client, cancellation)
@@ -184,6 +192,7 @@ impl RunEngine {
{
Ok((next_checkpoint, compaction_usage)) => {
checkpoint = next_checkpoint;
context_usage_anchor = None;
if let Some(compaction_usage) = compaction_usage {
accumulate_usage(&mut usage, compaction_usage);
}
@@ -272,11 +281,21 @@ impl RunEngine {
match interrupted {
Ok(cycle) => {
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);
}
}
Err(failure) => {
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);
}
}
@@ -319,6 +338,11 @@ impl RunEngine {
Ok(cycle) => break 'attempt cycle,
Err(cycle_failure) => {
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);
}
if cancellation.is_cancelled() {
@@ -420,6 +444,11 @@ impl RunEngine {
}
};
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);
}
@@ -879,6 +908,19 @@ async fn hydrate_tool_images(
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) {
match total {
Some(total) => *total += usage,
+98
View File
@@ -8,6 +8,12 @@ use crate::{
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)]
pub(crate) struct BufferedLlmChunk {
pub(crate) seq: i64,
@@ -266,6 +272,37 @@ impl Store {
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>> {
let rows = sqlx::query("SELECT * FROM llm_calls ORDER BY created_at_ms DESC LIMIT ?")
.bind(limit.clamp(1, 500))
@@ -421,4 +458,65 @@ mod tests {
assert_eq!(overview.metrics.llm_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);
}
}
+1 -1
View File
@@ -19,7 +19,7 @@ mod writer;
pub use cas::*;
pub(crate) use cursor_traces::BufferedCursorTraceChunk;
pub(crate) use llm_calls::BufferedLlmChunk;
pub(crate) use llm_calls::{BufferedLlmChunk, ContextUsageAnchor};
pub use runs::*;
pub use settings::*;
pub(crate) use sqlite::now_ms;
+102
View File
@@ -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(
&registry,
"anchor-first",
user_request(
"anchor-conversation",
"anchor-user-1",
&"x".repeat(400_000),
&model_a.model_hash,
None,
),
)
.await;
let second = run(
&registry,
"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]
async fn irreducibly_oversized_current_input_fails_before_provider_dispatch() {
let (_directory, store) = fixtures::temp_store().await;