mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-08 15:43:10 +08:00
feat: enhance context token estimation and compaction logic
- Added `estimate_context_tokens` function to calculate provider-visible context size based on prompt specifications and projected messages. - Updated `CheckpointBuilder` to record estimated context tokens during message processing. - Refactored compaction logic to utilize the new token estimation, ensuring proper context management during model runs. - Introduced tests to validate context estimation and compaction behavior under various scenarios.
This commit is contained in:
@@ -3,7 +3,7 @@ use prost::Message;
|
|||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
cursor::{checkpoint::PendingSteps, protocol::proto::agent::v1 as pb},
|
cursor::{checkpoint::PendingSteps, protocol::proto::agent::v1 as pb},
|
||||||
model::CanonicalMessage,
|
model::{estimate_context_tokens, project_messages, CanonicalMessage, PromptSpec},
|
||||||
store::{BlobEdge, BlobId},
|
store::{BlobEdge, BlobId},
|
||||||
Error, Result,
|
Error, Result,
|
||||||
};
|
};
|
||||||
@@ -85,6 +85,12 @@ impl CheckpointBuilder {
|
|||||||
.push(archive_id.as_bytes().to_vec());
|
.push(archive_id.as_bytes().to_vec());
|
||||||
}
|
}
|
||||||
self.base.self_summary_count = self.base.self_summary_count.saturating_add(1);
|
self.base.self_summary_count = self.base.self_summary_count.saturating_add(1);
|
||||||
|
let projected = project_messages(messages)?;
|
||||||
|
let prompt = PromptSpec {
|
||||||
|
instructions: self.instructions.clone(),
|
||||||
|
tools: self.tool_definitions.clone(),
|
||||||
|
};
|
||||||
|
self.record_context_tokens(Some(estimate_context_tokens(&prompt, &projected)));
|
||||||
if let Some(details) = self.base.token_details.as_mut() {
|
if let Some(details) = self.base.token_details.as_mut() {
|
||||||
details.breakdown = Some(crate::cursor::services::usage::breakdown(
|
details.breakdown = Some(crate::cursor::services::usage::breakdown(
|
||||||
details.used_tokens,
|
details.used_tokens,
|
||||||
|
|||||||
@@ -6,8 +6,6 @@ mod usage {
|
|||||||
|
|
||||||
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>,
|
||||||
@@ -18,21 +16,6 @@ mod 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 | ProviderType::Plugin => {
|
|
||||||
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);
|
||||||
@@ -53,7 +36,7 @@ pub use usage::*;
|
|||||||
mod llm_call {
|
mod llm_call {
|
||||||
use serde::Serialize;
|
use serde::Serialize;
|
||||||
|
|
||||||
use super::{ProviderType, Usage};
|
use super::ProviderType;
|
||||||
|
|
||||||
#[derive(Clone, Debug)]
|
#[derive(Clone, Debug)]
|
||||||
pub struct NewLlmCall {
|
pub struct NewLlmCall {
|
||||||
@@ -75,14 +58,6 @@ mod llm_call {
|
|||||||
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,
|
||||||
|
|||||||
@@ -1,4 +1,100 @@
|
|||||||
//! Estimates and records model token usage.
|
//! Estimates provider-visible context size and formats configured token counts.
|
||||||
|
|
||||||
|
use super::{ContentPart, ProjectedContent, ProjectedMessage, PromptSpec};
|
||||||
|
|
||||||
|
const TOKENS_PER_MESSAGE_OVERHEAD: u64 = 8;
|
||||||
|
const TOKENS_PER_TOOL_CALL_OVERHEAD: u64 = 6;
|
||||||
|
const TOKENS_PER_IMAGE: u64 = 1_024;
|
||||||
|
|
||||||
|
pub(crate) fn estimate_context_tokens(prompt: &PromptSpec, messages: &[ProjectedMessage]) -> u64 {
|
||||||
|
let instructions = estimate_text_tokens(&prompt.instructions);
|
||||||
|
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| {
|
||||||
|
total.saturating_add(estimate_message_tokens(message))
|
||||||
|
});
|
||||||
|
instructions.saturating_add(tools).saturating_add(messages)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn estimate_message_tokens(message: &ProjectedMessage) -> u64 {
|
||||||
|
let content = match &message.content {
|
||||||
|
ProjectedContent::Parts(parts) => estimate_parts_tokens(parts),
|
||||||
|
ProjectedContent::Assistant {
|
||||||
|
text,
|
||||||
|
thinking,
|
||||||
|
replay_state,
|
||||||
|
calls,
|
||||||
|
} => {
|
||||||
|
let calls = calls.iter().fold(0_u64, |total, call| {
|
||||||
|
total
|
||||||
|
.saturating_add(TOKENS_PER_TOOL_CALL_OVERHEAD)
|
||||||
|
.saturating_add(estimate_text_tokens(&call.call_id))
|
||||||
|
.saturating_add(estimate_text_tokens(&call.name))
|
||||||
|
.saturating_add(estimate_json_tokens(&call.arguments))
|
||||||
|
});
|
||||||
|
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) => {
|
||||||
|
let content = if result.provider_parts.is_empty() {
|
||||||
|
estimate_text_tokens(&result.content).saturating_add(
|
||||||
|
result
|
||||||
|
.image
|
||||||
|
.as_ref()
|
||||||
|
.map(|image| {
|
||||||
|
TOKENS_PER_IMAGE.saturating_add(estimate_text_tokens(&image.mime_type))
|
||||||
|
})
|
||||||
|
.unwrap_or_default(),
|
||||||
|
)
|
||||||
|
} else {
|
||||||
|
estimate_parts_tokens(&result.provider_parts)
|
||||||
|
};
|
||||||
|
estimate_text_tokens(&result.call_id)
|
||||||
|
.saturating_add(estimate_text_tokens(&result.name))
|
||||||
|
.saturating_add(content)
|
||||||
|
}
|
||||||
|
};
|
||||||
|
TOKENS_PER_MESSAGE_OVERHEAD.saturating_add(content)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn estimate_parts_tokens(parts: &[ContentPart]) -> u64 {
|
||||||
|
parts.iter().fold(0_u64, |total, part| {
|
||||||
|
let tokens = match part {
|
||||||
|
ContentPart::Text { text } => estimate_text_tokens(text),
|
||||||
|
ContentPart::Image { mime_type, .. } => {
|
||||||
|
TOKENS_PER_IMAGE.saturating_add(estimate_text_tokens(mime_type))
|
||||||
|
}
|
||||||
|
};
|
||||||
|
total.saturating_add(tokens)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
fn estimate_json_tokens(value: &impl serde::Serialize) -> u64 {
|
||||||
|
serde_json::to_string(value)
|
||||||
|
.map(|value| estimate_text_tokens(&value))
|
||||||
|
.unwrap_or_default()
|
||||||
|
}
|
||||||
|
|
||||||
|
fn estimate_text_tokens(text: &str) -> u64 {
|
||||||
|
let text = text.trim();
|
||||||
|
if text.is_empty() {
|
||||||
|
return 0;
|
||||||
|
}
|
||||||
|
let characters = text.chars().count() as u64;
|
||||||
|
characters
|
||||||
|
.div_ceil(4)
|
||||||
|
.saturating_add(text.bytes().filter(|byte| *byte == b'\n').count() as u64)
|
||||||
|
.max(1)
|
||||||
|
}
|
||||||
|
|
||||||
pub(crate) fn parse_token_count(value: &str) -> Option<u64> {
|
pub(crate) fn parse_token_count(value: &str) -> Option<u64> {
|
||||||
let value = value.trim().to_ascii_lowercase();
|
let value = value.trim().to_ascii_lowercase();
|
||||||
let (number, multiplier) = match value.chars().last()? {
|
let (number, multiplier) = match value.chars().last()? {
|
||||||
@@ -18,3 +114,104 @@ pub(crate) fn format_token_count(tokens: u64) -> String {
|
|||||||
tokens.to_string()
|
tokens.to_string()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
use crate::model::{Role, ToolCallContent, ToolDefinition, ToolResultContent};
|
||||||
|
|
||||||
|
fn prompt() -> PromptSpec {
|
||||||
|
PromptSpec {
|
||||||
|
instructions: "system instructions".into(),
|
||||||
|
tools: vec![ToolDefinition {
|
||||||
|
name: "Read".into(),
|
||||||
|
description: "Read a file".into(),
|
||||||
|
parameters: serde_json::json!({"type": "object", "properties": {"path": {"type": "string"}}}),
|
||||||
|
}],
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn context_estimate_grows_with_provider_visible_text_and_tools() {
|
||||||
|
let short = vec![ProjectedMessage {
|
||||||
|
message_id: "short".into(),
|
||||||
|
role: Role::User,
|
||||||
|
content: ProjectedContent::Parts(vec![ContentPart::Text {
|
||||||
|
text: "hello".into(),
|
||||||
|
}]),
|
||||||
|
}];
|
||||||
|
let long = vec![ProjectedMessage {
|
||||||
|
message_id: "long".into(),
|
||||||
|
role: Role::User,
|
||||||
|
content: ProjectedContent::Parts(vec![ContentPart::Text {
|
||||||
|
text: "x".repeat(40_000),
|
||||||
|
}]),
|
||||||
|
}];
|
||||||
|
let without_tools = PromptSpec {
|
||||||
|
instructions: prompt().instructions,
|
||||||
|
tools: Vec::new(),
|
||||||
|
};
|
||||||
|
|
||||||
|
assert!(
|
||||||
|
estimate_context_tokens(&prompt(), &short)
|
||||||
|
> estimate_context_tokens(&without_tools, &short)
|
||||||
|
);
|
||||||
|
assert!(
|
||||||
|
estimate_context_tokens(&prompt(), &long) > estimate_context_tokens(&prompt(), &short)
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn context_estimate_counts_tool_calls_results_and_images() {
|
||||||
|
let assistant = ProjectedMessage {
|
||||||
|
message_id: "assistant".into(),
|
||||||
|
role: Role::Assistant,
|
||||||
|
content: ProjectedContent::Assistant {
|
||||||
|
text: String::new(),
|
||||||
|
thinking: "reasoning".into(),
|
||||||
|
replay_state: None,
|
||||||
|
calls: vec![ToolCallContent {
|
||||||
|
index: 0,
|
||||||
|
call_id: "call-1".into(),
|
||||||
|
name: "Read".into(),
|
||||||
|
arguments: serde_json::json!({"path": "/tmp/file"}),
|
||||||
|
}],
|
||||||
|
},
|
||||||
|
};
|
||||||
|
let text_result = ProjectedMessage {
|
||||||
|
message_id: "result-text".into(),
|
||||||
|
role: Role::Tool,
|
||||||
|
content: ProjectedContent::ToolResult(ToolResultContent {
|
||||||
|
call_id: "call-1".into(),
|
||||||
|
name: "Read".into(),
|
||||||
|
content: "file contents".into(),
|
||||||
|
is_error: false,
|
||||||
|
image: None,
|
||||||
|
provider_parts: Vec::new(),
|
||||||
|
}),
|
||||||
|
};
|
||||||
|
let image_result = ProjectedMessage {
|
||||||
|
message_id: "result-image".into(),
|
||||||
|
role: Role::Tool,
|
||||||
|
content: ProjectedContent::ToolResult(ToolResultContent {
|
||||||
|
call_id: "call-1".into(),
|
||||||
|
name: "Read".into(),
|
||||||
|
content: "file contents".into(),
|
||||||
|
is_error: false,
|
||||||
|
image: None,
|
||||||
|
provider_parts: vec![ContentPart::Image {
|
||||||
|
mime_type: "image/png".into(),
|
||||||
|
data: vec![0; 32],
|
||||||
|
}],
|
||||||
|
}),
|
||||||
|
};
|
||||||
|
|
||||||
|
let base = estimate_context_tokens(&prompt(), &[]);
|
||||||
|
let with_call = estimate_context_tokens(&prompt(), std::slice::from_ref(&assistant));
|
||||||
|
let with_text = estimate_context_tokens(&prompt(), &[assistant.clone(), text_result]);
|
||||||
|
let with_image = estimate_context_tokens(&prompt(), &[assistant, image_result]);
|
||||||
|
assert!(with_call > base);
|
||||||
|
assert!(with_text > with_call);
|
||||||
|
assert!(with_image > with_text);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -1,62 +1,53 @@
|
|||||||
//! Decides when to compact context and builds a stable fallback summary.
|
//! Decides when to compact provider-visible context and builds a stable fallback summary.
|
||||||
|
|
||||||
use std::collections::HashSet;
|
use std::collections::HashSet;
|
||||||
|
|
||||||
use crate::model::{CanonicalMessage, LlmCallUsageAnchor, PreparedRun, ProjectedMessage};
|
use crate::model::{estimate_context_tokens, CanonicalMessage, PreparedRun, ProjectedMessage};
|
||||||
|
|
||||||
const FALLBACK_CHARS: usize = 12_000;
|
const FALLBACK_CHARS: usize = 12_000;
|
||||||
|
|
||||||
|
pub(super) const RESERVE_TOKENS: u64 = 10_000;
|
||||||
pub(super) const OUTPUT_TOKENS: u64 = 4_096;
|
pub(super) const OUTPUT_TOKENS: u64 = 4_096;
|
||||||
pub(super) const INSTRUCTIONS: &str = "Summarize the conversation for the next model turn. Preserve goals, constraints, decisions, files, commands, errors, results, and unfinished work. Do not call tools. Return only the concise durable summary.";
|
pub(super) const INSTRUCTIONS: &str = "Summarize the conversation for the next model turn. Preserve goals, constraints, decisions, files, commands, errors, results, and unfinished work. Do not call tools. Return only the concise durable summary.";
|
||||||
|
|
||||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
pub(super) fn input_budget(prepared: &PreparedRun) -> Option<u64> {
|
||||||
pub(super) struct ContextUsageAnchor {
|
prepared
|
||||||
input_tokens: u64,
|
.model
|
||||||
message_count: usize,
|
.context_window_tokens
|
||||||
tool_count: usize,
|
.map(|window| window.saturating_sub(RESERVE_TOKENS))
|
||||||
}
|
}
|
||||||
|
|
||||||
impl ContextUsageAnchor {
|
pub(super) fn estimated_tokens(
|
||||||
pub(super) fn from_llm_call(anchor: LlmCallUsageAnchor) -> Option<Self> {
|
prepared: &PreparedRun,
|
||||||
Some(Self {
|
projected_messages: &[ProjectedMessage],
|
||||||
input_tokens: anchor.usage.context_input_tokens(anchor.request_type)?,
|
) -> u64 {
|
||||||
message_count: anchor.message_count,
|
estimate_context_tokens(&prepared.prompt, projected_messages)
|
||||||
tool_count: anchor.tool_count,
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(super) fn should_compact(
|
pub(super) fn should_compact(
|
||||||
prepared: &PreparedRun,
|
prepared: &PreparedRun,
|
||||||
messages: &[CanonicalMessage],
|
|
||||||
projected_messages: &[ProjectedMessage],
|
projected_messages: &[ProjectedMessage],
|
||||||
anchor: Option<ContextUsageAnchor>,
|
|
||||||
) -> bool {
|
) -> bool {
|
||||||
let Some(context_window) = prepared.model.context_window_tokens else {
|
let Some(budget) = input_budget(prepared) else {
|
||||||
return false;
|
return false;
|
||||||
};
|
};
|
||||||
if context_window == 0 || messages.len() <= prepared.initial_messages.len() {
|
estimated_tokens(prepared, projected_messages) > budget
|
||||||
return false;
|
}
|
||||||
|
|
||||||
|
pub(super) fn validate_compacted(
|
||||||
|
prepared: &PreparedRun,
|
||||||
|
projected_messages: &[ProjectedMessage],
|
||||||
|
) -> std::result::Result<u64, String> {
|
||||||
|
let estimated = estimated_tokens(prepared, projected_messages);
|
||||||
|
let Some(budget) = input_budget(prepared) else {
|
||||||
|
return Ok(estimated);
|
||||||
|
};
|
||||||
|
if estimated <= budget {
|
||||||
|
return Ok(estimated);
|
||||||
}
|
}
|
||||||
let estimated_input = anchor
|
Err(format!(
|
||||||
.filter(|anchor| {
|
"context overflow after compaction: estimated input {estimated} tokens exceeds budget {budget} tokens"
|
||||||
anchor.message_count <= projected_messages.len()
|
|
||||||
&& anchor.tool_count == prepared.prompt.tools.len()
|
|
||||||
})
|
|
||||||
.map(|anchor| {
|
|
||||||
anchor
|
|
||||||
.input_tokens
|
|
||||||
.saturating_add(estimate_serialized_tokens(
|
|
||||||
&serde_json::to_string(&projected_messages[anchor.message_count..])
|
|
||||||
.unwrap_or_default(),
|
|
||||||
))
|
))
|
||||||
})
|
|
||||||
.unwrap_or_else(|| {
|
|
||||||
estimate_serialized_tokens(
|
|
||||||
&serde_json::to_string(&(&prepared.prompt, messages)).unwrap_or_default(),
|
|
||||||
)
|
|
||||||
});
|
|
||||||
estimated_input > context_window
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(super) fn partition(
|
pub(super) fn partition(
|
||||||
@@ -95,15 +86,6 @@ pub(super) fn fallback_summary(messages: &[CanonicalMessage]) -> String {
|
|||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
fn estimate_serialized_tokens(serialized: &str) -> u64 {
|
|
||||||
serialized
|
|
||||||
.chars()
|
|
||||||
.fold(0_u64, |units, character| {
|
|
||||||
units.saturating_add(if character.is_ascii() { 273 } else { 550 })
|
|
||||||
})
|
|
||||||
.div_ceil(1_000)
|
|
||||||
}
|
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
use super::*;
|
use super::*;
|
||||||
@@ -112,11 +94,10 @@ mod tests {
|
|||||||
RunAction, RunId, RunKind,
|
RunAction, RunId, RunKind,
|
||||||
};
|
};
|
||||||
|
|
||||||
#[test]
|
fn prepared(context_window_tokens: u64) -> PreparedRun {
|
||||||
fn automatic_compaction_runs_for_start_and_resume_actions_after_the_limit() {
|
|
||||||
let mut model = ModelSpec::new("model");
|
let mut model = ModelSpec::new("model");
|
||||||
model.context_window_tokens = Some(200_000);
|
model.context_window_tokens = Some(context_window_tokens);
|
||||||
let mut prepared = PreparedRun {
|
PreparedRun {
|
||||||
run_id: RunId::new("run"),
|
run_id: RunId::new("run"),
|
||||||
cursor_request_id: None,
|
cursor_request_id: None,
|
||||||
conversation_id: ConversationId::new("conversation"),
|
conversation_id: ConversationId::new("conversation"),
|
||||||
@@ -129,50 +110,50 @@ mod tests {
|
|||||||
initial_messages: Vec::new(),
|
initial_messages: Vec::new(),
|
||||||
action: RunAction::Start,
|
action: RunAction::Start,
|
||||||
base_checkpoint_id: CheckpointId(1),
|
base_checkpoint_id: CheckpointId(1),
|
||||||
};
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn automatic_compaction_uses_fixed_reserve_for_every_action() {
|
||||||
let messages = vec![CanonicalMessage::text(
|
let messages = vec![CanonicalMessage::text(
|
||||||
"user",
|
"user",
|
||||||
Role::User,
|
Role::User,
|
||||||
Origin::Runtime,
|
Origin::Runtime,
|
||||||
"hello",
|
"x".repeat(40_000),
|
||||||
)];
|
)];
|
||||||
let projected = project_messages(&messages).unwrap();
|
let projected = project_messages(&messages).unwrap();
|
||||||
let tail_tokens = estimate_serialized_tokens(&serde_json::to_string(&projected).unwrap());
|
let estimated = estimate_context_tokens(&prepared(1).prompt, &projected);
|
||||||
let anchor = |estimated_input| {
|
let mut prepared = prepared(estimated + RESERVE_TOKENS);
|
||||||
Some(ContextUsageAnchor {
|
|
||||||
input_tokens: estimated_input - tail_tokens,
|
|
||||||
message_count: 0,
|
|
||||||
tool_count: 0,
|
|
||||||
})
|
|
||||||
};
|
|
||||||
|
|
||||||
assert!(!should_compact(
|
assert!(!should_compact(&prepared, &projected));
|
||||||
&prepared,
|
prepared.model.context_window_tokens = Some(estimated + RESERVE_TOKENS - 1);
|
||||||
&messages,
|
assert!(should_compact(&prepared, &projected));
|
||||||
&projected,
|
|
||||||
anchor(199_999)
|
|
||||||
));
|
|
||||||
assert!(!should_compact(
|
|
||||||
&prepared,
|
|
||||||
&messages,
|
|
||||||
&projected,
|
|
||||||
anchor(200_000)
|
|
||||||
));
|
|
||||||
assert!(should_compact(
|
|
||||||
&prepared,
|
|
||||||
&messages,
|
|
||||||
&projected,
|
|
||||||
anchor(200_001)
|
|
||||||
));
|
|
||||||
|
|
||||||
prepared.action = RunAction::Resume {
|
prepared.action = RunAction::Resume {
|
||||||
pending_tool_round: None,
|
pending_tool_round: None,
|
||||||
};
|
};
|
||||||
assert!(should_compact(
|
assert!(should_compact(&prepared, &projected));
|
||||||
&prepared,
|
}
|
||||||
&messages,
|
|
||||||
&projected,
|
#[test]
|
||||||
anchor(200_001)
|
fn compacted_history_is_validated_against_the_same_budget() {
|
||||||
));
|
let messages = vec![CanonicalMessage::text(
|
||||||
|
"user",
|
||||||
|
Role::User,
|
||||||
|
Origin::Runtime,
|
||||||
|
"x".repeat(40_000),
|
||||||
|
)];
|
||||||
|
let projected = project_messages(&messages).unwrap();
|
||||||
|
let estimated = estimate_context_tokens(&prepared(1).prompt, &projected);
|
||||||
|
|
||||||
|
assert_eq!(
|
||||||
|
validate_compacted(&prepared(estimated + RESERVE_TOKENS), &projected),
|
||||||
|
Ok(estimated)
|
||||||
|
);
|
||||||
|
assert!(
|
||||||
|
validate_compacted(&prepared(estimated + RESERVE_TOKENS - 1), &projected)
|
||||||
|
.unwrap_err()
|
||||||
|
.contains("context overflow after compaction")
|
||||||
|
);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+16
-23
@@ -161,7 +161,6 @@ impl RunEngine {
|
|||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
||||||
let mut auto_compacted = prepared.action == RunAction::Compact;
|
|
||||||
'model: loop {
|
'model: loop {
|
||||||
if cancellation.is_cancelled() {
|
if cancellation.is_cancelled() {
|
||||||
return (RunOutcome::Cancelled, usage);
|
return (RunOutcome::Cancelled, usage);
|
||||||
@@ -170,31 +169,13 @@ impl RunEngine {
|
|||||||
Ok(messages) => messages,
|
Ok(messages) => messages,
|
||||||
Err(error) => return (RunOutcome::Failed(error.into()), usage),
|
Err(error) => return (RunOutcome::Failed(error.into()), usage),
|
||||||
};
|
};
|
||||||
let context_anchor = if !auto_compacted {
|
|
||||||
match self
|
|
||||||
.store
|
|
||||||
.latest_llm_call_usage_anchor(
|
|
||||||
&prepared.conversation_id,
|
|
||||||
&prepared.model.model_id,
|
|
||||||
)
|
|
||||||
.await
|
|
||||||
{
|
|
||||||
Ok(anchor) => {
|
|
||||||
anchor.and_then(super::compaction::ContextUsageAnchor::from_llm_call)
|
|
||||||
}
|
|
||||||
Err(error) => return (RunOutcome::Failed(error.into()), usage),
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
None
|
|
||||||
};
|
|
||||||
let history = match crate::model::project_messages(&messages) {
|
let history = match crate::model::project_messages(&messages) {
|
||||||
Ok(history) => history,
|
Ok(history) => history,
|
||||||
Err(error) => return (RunOutcome::Failed(error.into()), usage),
|
Err(error) => return (RunOutcome::Failed(error.into()), usage),
|
||||||
};
|
};
|
||||||
if !auto_compacted
|
if prepared.action != RunAction::Compact
|
||||||
&& super::compaction::should_compact(prepared, &messages, &history, context_anchor)
|
&& super::compaction::should_compact(prepared, &history)
|
||||||
{
|
{
|
||||||
auto_compacted = true;
|
|
||||||
match self
|
match self
|
||||||
.auto_compact(prepared, checkpoint, &messages, client, cancellation)
|
.auto_compact(prepared, checkpoint, &messages, client, cancellation)
|
||||||
.await
|
.await
|
||||||
@@ -583,7 +564,15 @@ impl RunEngine {
|
|||||||
let (compactable, retained_request_context) =
|
let (compactable, retained_request_context) =
|
||||||
super::compaction::partition(messages, ¤t_ids);
|
super::compaction::partition(messages, ¤t_ids);
|
||||||
if compactable.is_empty() {
|
if compactable.is_empty() {
|
||||||
return Ok((checkpoint, None));
|
let projected = crate::model::project_messages(messages)
|
||||||
|
.map_err(|error| RunOutcome::Failed(error.into()))?;
|
||||||
|
let message = super::compaction::validate_compacted(prepared, &projected)
|
||||||
|
.err()
|
||||||
|
.unwrap_or_else(|| {
|
||||||
|
"context overflow after compaction: no conversation history can be compacted"
|
||||||
|
.into()
|
||||||
|
});
|
||||||
|
return Err(RunOutcome::Failed(RunFailure::Protocol(message)));
|
||||||
}
|
}
|
||||||
|
|
||||||
emit(client, RunEvent::AutoCompactionStarted)
|
emit(client, RunEvent::AutoCompactionStarted)
|
||||||
@@ -689,7 +678,7 @@ impl RunEngine {
|
|||||||
)
|
)
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
let event_id = format!("summary:auto:{}", prepared.run_id);
|
let event_id = format!("summary:auto:{}:{provider_call_index}", prepared.run_id);
|
||||||
let summary_message = CanonicalMessage {
|
let summary_message = CanonicalMessage {
|
||||||
message_id: format!("runtime:{event_id}"),
|
message_id: format!("runtime:{event_id}"),
|
||||||
role: Role::User,
|
role: Role::User,
|
||||||
@@ -704,6 +693,10 @@ impl RunEngine {
|
|||||||
let mut replacement = retained_request_context.into_iter().collect::<Vec<_>>();
|
let mut replacement = retained_request_context.into_iter().collect::<Vec<_>>();
|
||||||
replacement.push(summary_message);
|
replacement.push(summary_message);
|
||||||
replacement.extend(prepared.initial_messages.iter().cloned());
|
replacement.extend(prepared.initial_messages.iter().cloned());
|
||||||
|
let projected_replacement = crate::model::project_messages(&replacement)
|
||||||
|
.map_err(|error| RunOutcome::Failed(error.into()))?;
|
||||||
|
super::compaction::validate_compacted(prepared, &projected_replacement)
|
||||||
|
.map_err(|message| RunOutcome::Failed(RunFailure::Protocol(message)))?;
|
||||||
let mut checkpoint = self
|
let mut checkpoint = self
|
||||||
.store
|
.store
|
||||||
.replace_checkpoint(
|
.replace_checkpoint(
|
||||||
|
|||||||
@@ -1,13 +1,8 @@
|
|||||||
//! Persists provider call payloads, timing, and usage.
|
//! Persists provider call payloads, timing, and usage.
|
||||||
use std::str::FromStr;
|
|
||||||
|
|
||||||
use sqlx::Row;
|
use sqlx::Row;
|
||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
model::{
|
model::{LlmCallRequest, LlmCallResponseChunk, LlmCallSummary, NewLlmCall, Usage},
|
||||||
ConversationId, LlmCallRequest, LlmCallResponseChunk, LlmCallSummary, LlmCallUsageAnchor,
|
|
||||||
NewLlmCall, ProviderType, Usage,
|
|
||||||
},
|
|
||||||
Result,
|
Result,
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -288,41 +283,6 @@ 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 = ?",
|
||||||
@@ -416,6 +376,7 @@ fn summary_from_row(row: sqlx::sqlite::SqliteRow) -> Result<LlmCallSummary> {
|
|||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
use super::*;
|
use super::*;
|
||||||
|
use crate::model::ProviderType;
|
||||||
|
|
||||||
/// 插件模型不在 model_configs 中,调用记录必须照常落库并可按其稳定 ID 筛选。
|
/// 插件模型不在 model_configs 中,调用记录必须照常落库并可按其稳定 ID 筛选。
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
|
|||||||
@@ -191,6 +191,172 @@ async fn summarize_replaces_model_history_and_preserves_cursor_history() {
|
|||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn automatic_compaction_preflights_provider_input_and_records_rebuilt_tokens() {
|
||||||
|
let (_directory, store) = fixtures::temp_store().await;
|
||||||
|
let model = store
|
||||||
|
.create_model(&ModelConfigInput {
|
||||||
|
sort_order: 0,
|
||||||
|
display_name: "Auto Compact Model".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: "Auto Compact Model".into(),
|
||||||
|
model_id: "auto-compact-model".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(100_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(&"x".repeat(400_000), 150_000, 1_000));
|
||||||
|
provider.push(text_response("automatic durable summary", 120_000, 20));
|
||||||
|
provider.push(text_response("continued after compaction", 20_000, 20));
|
||||||
|
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,
|
||||||
|
"auto-first",
|
||||||
|
user_request(
|
||||||
|
"auto-conversation",
|
||||||
|
"auto-user-1",
|
||||||
|
"start",
|
||||||
|
&model.model_hash,
|
||||||
|
None,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
.await;
|
||||||
|
let first_state = first.checkpoints.last().unwrap().clone();
|
||||||
|
assert!(first_state.token_details.as_ref().unwrap().used_tokens > 100_000);
|
||||||
|
|
||||||
|
let second = run(
|
||||||
|
®istry,
|
||||||
|
"auto-second",
|
||||||
|
user_request(
|
||||||
|
"auto-conversation",
|
||||||
|
"auto-user-2",
|
||||||
|
"continue",
|
||||||
|
&model.model_hash,
|
||||||
|
Some(first_state),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
.await;
|
||||||
|
assert_eq!(second.summary_started, 1);
|
||||||
|
assert_eq!(second.summary_completed, 1);
|
||||||
|
let compacted_tokens = second
|
||||||
|
.checkpoints
|
||||||
|
.iter()
|
||||||
|
.filter_map(|state| state.token_details.as_ref())
|
||||||
|
.map(|details| details.used_tokens)
|
||||||
|
.find(|tokens| *tokens > 0 && *tokens < 100_000)
|
||||||
|
.expect("compacted checkpoint must record rebuilt context tokens");
|
||||||
|
assert!(compacted_tokens < 90_000);
|
||||||
|
|
||||||
|
let requests = provider.requests();
|
||||||
|
assert_eq!(requests.len(), 3);
|
||||||
|
assert!(!requests[0].prompt.tools.is_empty());
|
||||||
|
assert!(requests[1].prompt.tools.is_empty());
|
||||||
|
assert!(!requests[2].prompt.tools.is_empty());
|
||||||
|
assert!(requests[1]
|
||||||
|
.history
|
||||||
|
.iter()
|
||||||
|
.any(|message| match &message.content {
|
||||||
|
ProjectedContent::Assistant { text, .. } => text.len() == 400_000,
|
||||||
|
_ => false,
|
||||||
|
}));
|
||||||
|
assert!(requests[2]
|
||||||
|
.history
|
||||||
|
.iter()
|
||||||
|
.all(|message| match &message.content {
|
||||||
|
ProjectedContent::Assistant { text, .. } => text.len() != 400_000,
|
||||||
|
_ => true,
|
||||||
|
}));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn irreducibly_oversized_current_input_fails_before_provider_dispatch() {
|
||||||
|
let (_directory, store) = fixtures::temp_store().await;
|
||||||
|
let model = store
|
||||||
|
.create_model(&ModelConfigInput {
|
||||||
|
sort_order: 0,
|
||||||
|
display_name: "Overflow Model".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: "Overflow Model".into(),
|
||||||
|
model_id: "overflow-model".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(100_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();
|
||||||
|
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 output = run(
|
||||||
|
®istry,
|
||||||
|
"overflow-request",
|
||||||
|
user_request(
|
||||||
|
"overflow-conversation",
|
||||||
|
"overflow-user",
|
||||||
|
&"x".repeat(400_000),
|
||||||
|
&model.model_hash,
|
||||||
|
None,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
.await;
|
||||||
|
|
||||||
|
assert!(provider.requests().is_empty());
|
||||||
|
assert_eq!(output.summary_started, 0);
|
||||||
|
assert_eq!(output.summary_completed, 0);
|
||||||
|
}
|
||||||
|
|
||||||
#[derive(Default)]
|
#[derive(Default)]
|
||||||
struct Output {
|
struct Output {
|
||||||
checkpoints: Vec<pb::ConversationStateStructure>,
|
checkpoints: Vec<pb::ConversationStateStructure>,
|
||||||
|
|||||||
@@ -1021,7 +1021,7 @@ async fn injected_user_context_interrupts_automatic_compaction() {
|
|||||||
custom_headers: serde_json::json!({}),
|
custom_headers: serde_json::json!({}),
|
||||||
anthropic_extra_params_enabled: false,
|
anthropic_extra_params_enabled: false,
|
||||||
anthropic_extra_params: serde_json::json!({}),
|
anthropic_extra_params: serde_json::json!({}),
|
||||||
context_window_tokens: Some(10_001),
|
context_window_tokens: Some(100_000),
|
||||||
max_completion_tokens: None,
|
max_completion_tokens: None,
|
||||||
anthropic_max_tokens: None,
|
anthropic_max_tokens: None,
|
||||||
anthropic_thinking_effort: None,
|
anthropic_thinking_effort: None,
|
||||||
@@ -1030,7 +1030,8 @@ async fn injected_user_context_interrupts_automatic_compaction() {
|
|||||||
.await
|
.await
|
||||||
.unwrap();
|
.unwrap();
|
||||||
let provider = fake_provider::FakeProvider::default();
|
let provider = fake_provider::FakeProvider::default();
|
||||||
provider.push(text_response("seed answer"));
|
let seed_answer = "x".repeat(400_000);
|
||||||
|
provider.push(text_response(&seed_answer));
|
||||||
provider.push_pending();
|
provider.push_pending();
|
||||||
provider.push(text_response("continued after compacting injection"));
|
provider.push(text_response("continued after compacting injection"));
|
||||||
let assets = PromptAssets::load(
|
let assets = PromptAssets::load(
|
||||||
@@ -1045,7 +1046,7 @@ async fn injected_user_context_interrupts_automatic_compaction() {
|
|||||||
PromptCompiler::new(assets),
|
PromptCompiler::new(assets),
|
||||||
);
|
);
|
||||||
|
|
||||||
let seed_state = run_to_end(
|
run_to_end(
|
||||||
®istry,
|
®istry,
|
||||||
"seed-request",
|
"seed-request",
|
||||||
client_run_for_model(
|
client_run_for_model(
|
||||||
@@ -1065,7 +1066,7 @@ async fn injected_user_context_interrupts_automatic_compaction() {
|
|||||||
"inject-during-compaction",
|
"inject-during-compaction",
|
||||||
"compaction-injection-conversation",
|
"compaction-injection-conversation",
|
||||||
&model.model_hash,
|
&model.model_hash,
|
||||||
Some(seed_state),
|
None,
|
||||||
);
|
);
|
||||||
let Some(pb::agent_client_message::Message::RunRequest(request)) =
|
let Some(pb::agent_client_message::Message::RunRequest(request)) =
|
||||||
compacting_request.message.as_mut()
|
compacting_request.message.as_mut()
|
||||||
@@ -1075,9 +1076,17 @@ async fn injected_user_context_interrupts_automatic_compaction() {
|
|||||||
request.requested_model.as_mut().unwrap().parameters.push(
|
request.requested_model.as_mut().unwrap().parameters.push(
|
||||||
pb::requested_model::ModelParameterValue {
|
pb::requested_model::ModelParameterValue {
|
||||||
id: "context".into(),
|
id: "context".into(),
|
||||||
value: "10001".into(),
|
value: "100000".into(),
|
||||||
},
|
},
|
||||||
);
|
);
|
||||||
|
let Some(pb::conversation_action::Action::UserMessageAction(action)) = request
|
||||||
|
.action
|
||||||
|
.as_mut()
|
||||||
|
.and_then(|action| action.action.as_mut())
|
||||||
|
else {
|
||||||
|
panic!("expected UserMessageAction")
|
||||||
|
};
|
||||||
|
action.user_message.as_mut().unwrap().message_id = "compaction-user".into();
|
||||||
handle
|
handle
|
||||||
.command(TransportCommand::Append {
|
.command(TransportCommand::Append {
|
||||||
seqno: 0,
|
seqno: 0,
|
||||||
@@ -1133,10 +1142,15 @@ async fn injected_user_context_interrupts_automatic_compaction() {
|
|||||||
|
|
||||||
let requests = provider.requests();
|
let requests = provider.requests();
|
||||||
assert_eq!(requests.len(), 3);
|
assert_eq!(requests.len(), 3);
|
||||||
assert!(requests[1]
|
assert!(
|
||||||
|
requests[1]
|
||||||
.prompt
|
.prompt
|
||||||
.instructions
|
.instructions
|
||||||
.starts_with("Summarize the conversation for the next model turn."));
|
.starts_with("Summarize the conversation for the next model turn."),
|
||||||
|
"second request was not compaction: instructions={:?}, history={:?}",
|
||||||
|
requests[1].prompt.instructions,
|
||||||
|
requests[1].history
|
||||||
|
);
|
||||||
assert!(!serde_json::to_string(&requests[1].history)
|
assert!(!serde_json::to_string(&requests[1].history)
|
||||||
.unwrap()
|
.unwrap()
|
||||||
.contains("injected follow-up"));
|
.contains("injected follow-up"));
|
||||||
|
|||||||
+243
-1
@@ -20,7 +20,10 @@ use cursor_server::{
|
|||||||
},
|
},
|
||||||
},
|
},
|
||||||
cursor::{TransportCommand, TransportRegistry},
|
cursor::{TransportCommand, TransportRegistry},
|
||||||
model::{MessageContent, ToolCall},
|
model::{
|
||||||
|
MessageContent, ModelConfigInput, ModelType, ProjectedContent, ToolCall,
|
||||||
|
OPENAI_CHAT_ENDPOINT,
|
||||||
|
},
|
||||||
provider::{FinishReason, ModelEvent},
|
provider::{FinishReason, ModelEvent},
|
||||||
};
|
};
|
||||||
use prost::Message;
|
use prost::Message;
|
||||||
@@ -745,6 +748,178 @@ async fn unknown_exec_id_is_a_protocol_error() {
|
|||||||
));
|
));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn one_run_can_auto_compact_again_after_more_tool_output() {
|
||||||
|
let (_directory, store) = fixtures::temp_store().await;
|
||||||
|
let model = store
|
||||||
|
.create_model(&ModelConfigInput {
|
||||||
|
sort_order: 0,
|
||||||
|
display_name: "Repeated compaction".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: "Repeated compaction".into(),
|
||||||
|
model_id: "repeated-compaction-model".into(),
|
||||||
|
reasoning_effort: None,
|
||||||
|
openai_endpoint: OPENAI_CHAT_ENDPOINT.into(),
|
||||||
|
openai_extra_params_enabled: false,
|
||||||
|
openai_extra_params: json!({}),
|
||||||
|
custom_headers_enabled: false,
|
||||||
|
custom_headers: json!({}),
|
||||||
|
anthropic_extra_params_enabled: false,
|
||||||
|
anthropic_extra_params: json!({}),
|
||||||
|
context_window_tokens: Some(25_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(tool_call_response("repeat-call-1"));
|
||||||
|
provider.push(text_events("first summary"));
|
||||||
|
provider.push(tool_call_response("repeat-call-2"));
|
||||||
|
provider.push(text_events("second summary"));
|
||||||
|
provider.push(text_events("done"));
|
||||||
|
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 handle = registry
|
||||||
|
.get_or_create("repeated-compaction-request")
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
let mut output = handle.subscribe();
|
||||||
|
handle
|
||||||
|
.command(TransportCommand::Append {
|
||||||
|
seqno: 0,
|
||||||
|
message: Box::new(client_run_for_model(
|
||||||
|
"repeated-compaction-conversation",
|
||||||
|
"repeated-compaction-request",
|
||||||
|
&model.model_hash,
|
||||||
|
)),
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
let mut seqno = 1;
|
||||||
|
let oversized = format!("HEAD{}TAIL", "x".repeat(4 * 1024 * 1024));
|
||||||
|
loop {
|
||||||
|
let frame = tokio::time::timeout(std::time::Duration::from_secs(10), output.recv())
|
||||||
|
.await
|
||||||
|
.unwrap()
|
||||||
|
.unwrap();
|
||||||
|
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
|
||||||
|
if flags & connect::END_STREAM_FLAG != 0 {
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
let server = pb::AgentServerMessage::decode(payload).unwrap();
|
||||||
|
match server.message {
|
||||||
|
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
|
||||||
|
handle
|
||||||
|
.command(TransportCommand::Append {
|
||||||
|
seqno,
|
||||||
|
message: Box::new(kv_ack(kv.id)),
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
seqno += 1;
|
||||||
|
}
|
||||||
|
Some(pb::agent_server_message::Message::ExecServerMessage(exec)) => {
|
||||||
|
let exec_id = exec.id;
|
||||||
|
handle
|
||||||
|
.command(TransportCommand::Append {
|
||||||
|
seqno,
|
||||||
|
message: Box::new(pb::AgentClientMessage {
|
||||||
|
message: Some(pb::agent_client_message::Message::ExecClientMessage(
|
||||||
|
pb::ExecClientMessage {
|
||||||
|
id: exec_id,
|
||||||
|
exec_id: String::new(),
|
||||||
|
message: Some(pb::exec_client_message::Message::ReadResult(
|
||||||
|
pb::ReadResult {
|
||||||
|
result: Some(pb::read_result::Result::Success(
|
||||||
|
pb::ReadSuccess {
|
||||||
|
path: "/tmp/large.txt".into(),
|
||||||
|
total_lines: 1,
|
||||||
|
file_size: oversized.len() as i64,
|
||||||
|
output: Some(
|
||||||
|
pb::read_success::Output::Content(
|
||||||
|
oversized.clone(),
|
||||||
|
),
|
||||||
|
),
|
||||||
|
..Default::default()
|
||||||
|
},
|
||||||
|
)),
|
||||||
|
},
|
||||||
|
)),
|
||||||
|
..Default::default()
|
||||||
|
},
|
||||||
|
)),
|
||||||
|
}),
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
seqno += 1;
|
||||||
|
handle
|
||||||
|
.command(TransportCommand::Append {
|
||||||
|
seqno,
|
||||||
|
message: Box::new(pb::AgentClientMessage {
|
||||||
|
message: Some(
|
||||||
|
pb::agent_client_message::Message::ExecClientControlMessage(
|
||||||
|
pb::ExecClientControlMessage {
|
||||||
|
message: Some(
|
||||||
|
pb::exec_client_control_message::Message::StreamClose(
|
||||||
|
pb::ExecClientStreamClose { id: exec_id },
|
||||||
|
),
|
||||||
|
),
|
||||||
|
},
|
||||||
|
),
|
||||||
|
),
|
||||||
|
}),
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
.unwrap();
|
||||||
|
seqno += 1;
|
||||||
|
}
|
||||||
|
_ => {}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
let requests = provider.requests();
|
||||||
|
let shapes = requests
|
||||||
|
.iter()
|
||||||
|
.map(|request| {
|
||||||
|
(
|
||||||
|
request.prompt.tools.len(),
|
||||||
|
request.history.len(),
|
||||||
|
request
|
||||||
|
.history
|
||||||
|
.iter()
|
||||||
|
.filter_map(|message| match &message.content {
|
||||||
|
ProjectedContent::ToolResult(result) => Some(result.content.len()),
|
||||||
|
_ => None,
|
||||||
|
})
|
||||||
|
.collect::<Vec<_>>(),
|
||||||
|
)
|
||||||
|
})
|
||||||
|
.collect::<Vec<_>>();
|
||||||
|
assert_eq!(requests.len(), 5, "provider requests: {shapes:?}");
|
||||||
|
assert!(!requests[0].prompt.tools.is_empty());
|
||||||
|
assert!(requests[1].prompt.tools.is_empty());
|
||||||
|
assert!(!requests[2].prompt.tools.is_empty());
|
||||||
|
assert!(requests[3].prompt.tools.is_empty());
|
||||||
|
assert!(!requests[4].prompt.tools.is_empty());
|
||||||
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn provider_tool_use_waits_for_client_result_then_calls_provider_again() {
|
async fn provider_tool_use_waits_for_client_result_then_calls_provider_again() {
|
||||||
let (directory, store) = fixtures::temp_store().await;
|
let (directory, store) = fixtures::temp_store().await;
|
||||||
@@ -935,6 +1110,73 @@ async fn provider_tool_use_waits_for_client_result_then_calls_provider_again() {
|
|||||||
assert_eq!(tool_calls[0].call_id, "call-1");
|
assert_eq!(tool_calls[0].call_id, "call-1");
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn tool_call_response(call_id: &str) -> Vec<ModelEvent> {
|
||||||
|
vec![
|
||||||
|
ModelEvent::Start {
|
||||||
|
model_call_id: format!("model-{call_id}"),
|
||||||
|
},
|
||||||
|
ModelEvent::ToolCallStart {
|
||||||
|
index: 0,
|
||||||
|
call_id: call_id.into(),
|
||||||
|
name: "Read".into(),
|
||||||
|
},
|
||||||
|
ModelEvent::ToolCallArgumentsDelta {
|
||||||
|
index: 0,
|
||||||
|
delta: "{\"path\":\"/tmp/large.txt\"}".into(),
|
||||||
|
},
|
||||||
|
ModelEvent::ToolCallEnd { index: 0 },
|
||||||
|
ModelEvent::Done(FinishReason::ToolUse),
|
||||||
|
]
|
||||||
|
}
|
||||||
|
|
||||||
|
fn text_events(text: &str) -> Vec<ModelEvent> {
|
||||||
|
vec![
|
||||||
|
ModelEvent::Start {
|
||||||
|
model_call_id: format!("model-{text}"),
|
||||||
|
},
|
||||||
|
ModelEvent::TextStart,
|
||||||
|
ModelEvent::TextDelta(text.into()),
|
||||||
|
ModelEvent::TextEnd,
|
||||||
|
ModelEvent::Done(FinishReason::Stop),
|
||||||
|
]
|
||||||
|
}
|
||||||
|
|
||||||
|
fn client_run_for_model(
|
||||||
|
conversation_id: &str,
|
||||||
|
run_id: &str,
|
||||||
|
model_id: &str,
|
||||||
|
) -> pb::AgentClientMessage {
|
||||||
|
let user = pb::UserMessage {
|
||||||
|
text: "read it".into(),
|
||||||
|
message_id: "user".into(),
|
||||||
|
mode: pb::AgentMode::Agent as i32,
|
||||||
|
..Default::default()
|
||||||
|
};
|
||||||
|
pb::AgentClientMessage {
|
||||||
|
message: Some(pb::agent_client_message::Message::RunRequest(
|
||||||
|
pb::AgentRunRequest {
|
||||||
|
action: Some(pb::ConversationAction {
|
||||||
|
action: Some(pb::conversation_action::Action::UserMessageAction(
|
||||||
|
pb::UserMessageAction {
|
||||||
|
user_message: Some(user),
|
||||||
|
request_context: Some(pb::RequestContext::default()),
|
||||||
|
..Default::default()
|
||||||
|
},
|
||||||
|
)),
|
||||||
|
..Default::default()
|
||||||
|
}),
|
||||||
|
conversation_id: Some(conversation_id.into()),
|
||||||
|
run_id: Some(run_id.into()),
|
||||||
|
requested_model: Some(pb::RequestedModel {
|
||||||
|
model_id: model_id.into(),
|
||||||
|
..Default::default()
|
||||||
|
}),
|
||||||
|
..Default::default()
|
||||||
|
},
|
||||||
|
)),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
fn client_run() -> pb::AgentClientMessage {
|
fn client_run() -> pb::AgentClientMessage {
|
||||||
let user = pb::UserMessage {
|
let user = pb::UserMessage {
|
||||||
text: "read it".into(),
|
text: "read it".into(),
|
||||||
|
|||||||
Reference in New Issue
Block a user