use std::fmt;
use super::window::{count_message, keep_rounds, split_rounds};
use super::{Budget, MemoryError, TokenCounter, TrimResult, TrimStrategy, WindowDrop};
use crate::message::{ContentBlock, Message};
use crate::provider::{ChatRequest, ModelOptions, Provider};
const SUMMARY_PREFIX: &str = "Summary of prior context:\n";
const DEFAULT_SUMMARY_PROMPT: &str = "\
You are a conversation-history compressor. Compress the provided conversation history into a concise summary so the model can recover context in subsequent conversation.
Requirements:
- Preserve the task goal, current state, decisions made, open items, and next steps;
- Keep proper nouns such as code identifiers, file paths, URLs, and numbers verbatim;
- The history may contain a \"Summary of prior context\" (system message): first understand the prior summary, then merge it with the new history and output a single integrated summary; do not re-summarize the prior summary itself;
- Output only the summary body, with no explanations or surrounding text.";
pub struct SummarizeStrategy {
provider: Box<dyn Provider>,
prompt: String,
summary_max_tokens: u32,
}
impl fmt::Debug for SummarizeStrategy {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("SummarizeStrategy")
.field("provider", &"Box<dyn Provider>")
.field("prompt_len", &self.prompt.len())
.field("summary_max_tokens", &self.summary_max_tokens)
.finish()
}
}
impl SummarizeStrategy {
pub fn new(provider: impl Provider + 'static) -> Self {
Self {
provider: Box::new(provider),
prompt: DEFAULT_SUMMARY_PROMPT.into(),
summary_max_tokens: 1024,
}
}
pub fn with_prompt(mut self, prompt: impl Into<String>) -> Self {
self.prompt = prompt.into();
self
}
pub fn with_summary_max_tokens(mut self, max_tokens: u32) -> Self {
self.summary_max_tokens = max_tokens;
self
}
async fn trim_impl(
&self,
messages: &[Message],
counts: &[usize],
budget: &Budget,
counter: &dyn TokenCounter,
fallback: &WindowDrop,
) -> Result<TrimResult, MemoryError> {
let rounds = split_rounds(messages, counts);
let keep = keep_rounds(&rounds, budget, self.summary_max_tokens as usize);
let cut = rounds[rounds.len() - keep].0;
if cut == 0 {
return Ok(TrimResult {
messages: messages.to_vec(),
replace: false,
});
}
let request = ChatRequest {
messages: vec![
Message::system(self.prompt.clone()),
Message::user(messages_to_text(&messages[..cut])),
],
tools: Vec::new(),
options: ModelOptions {
max_tokens: Some(self.summary_max_tokens),
..Default::default()
},
};
let summary = match self.provider.chat(request).await {
Ok(response) => {
let Message::Assistant { content, .. } = &response.message else {
return fallback
.trim_with_counts(messages, counts, budget, counter)
.await;
};
if content.trim().is_empty() {
None
} else {
Some(content.trim().to_string())
}
}
Err(error) => {
tracing::warn!(%error, "summarize failed, falling back to window drop");
None
}
};
let Some(summary) = summary else {
return fallback
.trim_with_counts(messages, counts, budget, counter)
.await;
};
let mut result = Vec::with_capacity(keep + 1);
result.push(Message::system(format!("{SUMMARY_PREFIX}{summary}")));
result.extend_from_slice(&messages[cut..]);
let summary_tokens = count_message(counter, &result[0]).await?;
tracing::info!(
compressed = cut,
kept = messages.len() - cut,
summary_tokens,
"context summarized by LLM"
);
Ok(TrimResult {
messages: result,
replace: true,
})
}
}
#[async_trait::async_trait]
impl TrimStrategy for SummarizeStrategy {
async fn trim(
&self,
messages: &[Message],
budget: &Budget,
counter: &dyn TokenCounter,
) -> Result<TrimResult, MemoryError> {
let mut counts = Vec::with_capacity(messages.len());
for m in messages {
counts.push(count_message(counter, m).await?);
}
self.trim_with_counts(messages, &counts, budget, counter)
.await
}
async fn trim_with_counts(
&self,
messages: &[Message],
counts: &[usize],
budget: &Budget,
counter: &dyn TokenCounter,
) -> Result<TrimResult, MemoryError> {
if messages.is_empty() {
return Ok(TrimResult {
messages: Vec::new(),
replace: false,
});
}
debug_assert_eq!(messages.len(), counts.len());
let fallback = WindowDrop;
self.trim_impl(messages, counts, budget, counter, &fallback)
.await
}
}
fn messages_to_text(messages: &[Message]) -> String {
let mut lines = Vec::with_capacity(messages.len());
for message in messages {
match message {
Message::System(s) => lines.push(format!("system: {s}")),
Message::User(blocks) => {
let text: String = blocks
.iter()
.map(|block| match block {
ContentBlock::Text(t) => t.clone(),
ContentBlock::Image(image) => format!("[image: {}]", image.mime_type),
ContentBlock::Wire(value) => {
match value.get("type").and_then(|t| t.as_str()) {
Some(kind) => format!("[{kind}]"),
None => "[content]".to_string(),
}
}
})
.collect::<Vec<_>>()
.join(" ");
lines.push(format!("user: {text}"));
}
Message::Assistant {
content,
tool_calls,
..
} => {
if tool_calls.is_empty() {
lines.push(format!("assistant: {content}"));
} else {
let calls = tool_calls
.iter()
.map(|tc| format!("{} {}", tc.name, tc.arguments))
.collect::<Vec<_>>()
.join("; ");
lines.push(format!("assistant: {content}\ntool_calls: {calls}"));
}
}
Message::ToolResult { content, .. } => lines.push(format!("tool_result: {content}")),
}
}
lines.join("\n")
}
#[cfg(test)]
mod tests {
use super::*;
use crate::memory::{CharTokenCounter, Memory, WindowMemory};
use crate::message::ToolCall;
use crate::provider::{FakeProvider, FakeReply, ProviderError};
use std::sync::Arc;
async fn record_rounds(memory: &mut WindowMemory, start: usize, n: usize) {
for i in start..=n {
memory
.record(Message::user(format!("Question from round {i}")))
.await
.unwrap();
memory
.record(Message::assistant(format!("Answer from round {i}")))
.await
.unwrap();
}
}
fn strategy(fake: Arc<FakeProvider>) -> Arc<SummarizeStrategy> {
Arc::new(SummarizeStrategy::new(fake))
}
#[tokio::test]
async fn summarizes_old_rounds_and_keeps_recent() {
let fake = Arc::new(FakeProvider::new([FakeReply::Text(
"Key points from earlier rounds".into(),
)]));
let mut memory = WindowMemory::new(30).with_strategy(strategy(fake.clone()));
record_rounds(&mut memory, 1, 4).await;
let context = memory.context().await.unwrap();
assert_eq!(context.len(), 3);
assert_eq!(
context[0],
Message::system(format!("{SUMMARY_PREFIX}Key points from earlier rounds"))
);
assert_eq!(context[1], Message::user("Question from round 4"));
assert_eq!(context[2], Message::assistant("Answer from round 4"));
let requests = fake.requests();
assert_eq!(requests.len(), 1);
let messages = &requests[0].messages;
assert!(
matches!(&messages[0], Message::System(p) if p.contains("conversation-history compressor"))
);
let Message::User(blocks) = &messages[1] else {
panic!("expected user message");
};
let ContentBlock::Text(text) = &blocks[0] else {
panic!("expected a text block");
};
assert!(text.contains("Question from round 1") && text.contains("Question from round 3"));
assert!(!text.contains("Question from round 4"));
}
#[tokio::test]
async fn summary_budget_reserves_space_for_recent_rounds() {
let fake = Arc::new(FakeProvider::new([FakeReply::Text("Highlights".into())]));
let mut memory = WindowMemory::new(30).with_strategy(Arc::new(
SummarizeStrategy::new(fake.clone()).with_summary_max_tokens(4),
));
record_rounds(&mut memory, 1, 4).await;
let context = memory.context().await.unwrap();
assert_eq!(context.len(), 5);
assert_eq!(context[1], Message::user("Question from round 3"));
assert_eq!(context[3], Message::user("Question from round 4"));
}
#[tokio::test]
async fn incremental_summary_reuses_previous_summary() {
let fake = Arc::new(FakeProvider::new([
FakeReply::Text("first segment highlights".into()),
FakeReply::Text("merged highlights".into()),
]));
let mut memory = WindowMemory::new(30).with_strategy(Arc::new(
SummarizeStrategy::new(fake.clone()).with_summary_max_tokens(4),
));
record_rounds(&mut memory, 1, 4).await;
memory.context().await.unwrap();
record_rounds(&mut memory, 5, 6).await;
let context = memory.context().await.unwrap();
assert_eq!(
context[0],
Message::system(format!("{SUMMARY_PREFIX}merged highlights"))
);
assert_eq!(context.len(), 5);
assert_eq!(context[1], Message::user("Question from round 5"));
let requests = fake.requests();
assert_eq!(requests.len(), 2);
let Message::User(blocks) = &requests[1].messages[1] else {
panic!("expected user message");
};
let ContentBlock::Text(text) = &blocks[0] else {
panic!("expected a text block");
};
assert!(text.contains("first segment highlights"));
}
#[tokio::test]
async fn provider_failure_falls_back_to_window_drop() {
let fake = Arc::new(FakeProvider::new([
FakeReply::Error(ProviderError::Network("mock network error".into())),
FakeReply::Text("second attempt succeeded".into()),
]));
let mut memory = WindowMemory::new(30).with_strategy(strategy(fake.clone()));
record_rounds(&mut memory, 1, 4).await;
let context = memory.context().await.unwrap();
assert_eq!(
context,
vec![
Message::user("Question from round 3"),
Message::assistant("Answer from round 3"),
Message::user("Question from round 4"),
Message::assistant("Answer from round 4"),
]
);
let context = memory.context().await.unwrap();
assert_eq!(
context[0],
Message::system(format!("{SUMMARY_PREFIX}second attempt succeeded"))
);
assert_eq!(fake.requests().len(), 2);
}
#[tokio::test]
async fn empty_summary_falls_back() {
let fake = Arc::new(FakeProvider::new([FakeReply::Text("".into())]));
let mut memory = WindowMemory::new(30).with_strategy(strategy(fake.clone()));
record_rounds(&mut memory, 1, 4).await;
let context = memory.context().await.unwrap();
assert_eq!(context.len(), 4); assert!(!context.iter().any(|m| matches!(m, Message::System(_))));
}
#[tokio::test]
async fn single_round_over_budget_returns_as_is() {
let fake = Arc::new(FakeProvider::new([FakeReply::Text("summary".into())]));
let mut memory = WindowMemory::new(5).with_strategy(strategy(fake.clone()));
memory
.record(Message::user(
"An extremely long user message, over budget in a single round",
))
.await
.unwrap();
memory.record(Message::assistant("Reply")).await.unwrap();
let context = memory.context().await.unwrap();
assert_eq!(context.len(), 2);
assert!(
fake.requests().is_empty(),
"with nothing to compact the summarizer must not be called"
);
}
#[tokio::test]
async fn empty_input_returns_empty() {
let fake = Arc::new(FakeProvider::new([FakeReply::Text("summary".into())]));
let strategy = strategy(fake);
let result = strategy
.trim(&[], &Budget::tokens(10), &CharTokenCounter)
.await
.unwrap();
assert!(result.messages.is_empty());
assert!(!result.replace);
}
#[tokio::test]
async fn messages_to_text_includes_tool_calls_and_results() {
let messages = vec![
Message::system("setup"),
Message::user("calculate"),
Message::Assistant {
content: "".into(),
reasoning: Some("reasoning trace is not included".into()),
tool_calls: vec![ToolCall {
id: "t1".into(),
name: "calc".into(),
arguments: r#"{"expr":"1+1"}"#.into(),
}],
},
Message::tool_result("t1", "2"),
];
let text = messages_to_text(&messages);
assert!(text.contains("system: setup"));
assert!(text.contains("user: calculate"));
assert!(text.contains(r#"tool_calls: calc {"expr":"1+1"}"#));
assert!(text.contains("tool_result: 2"));
assert!(!text.contains("reasoning trace is not included"));
}
#[tokio::test]
async fn trim_counts_consistent_with_window() {
let fake = Arc::new(FakeProvider::new([FakeReply::Text("summary".into())]));
let strategy = strategy(fake);
let messages = vec![
Message::user("u1"),
Message::assistant("a1"),
Message::user("u2"),
Message::assistant("a2"),
];
let counter = CharTokenCounter;
let result = strategy
.trim_with_counts(&messages, &[1, 1, 1, 1], &Budget::tokens(2), &counter)
.await
.unwrap();
assert_eq!(result.messages.len(), 3);
assert!(result.replace);
}
}