use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use agent_base::{ChatMessage, ContextWindowManager, LlmClient, Middleware, PreLlmCtx};
use async_trait::async_trait;
use serde_json::Value;
#[derive(Clone, Debug)]
pub struct CompressionConfig {
pub enabled: bool,
pub trigger_tokens: usize,
pub keep_last_messages: usize,
pub max_transcript_chars: usize,
pub max_summary_chars: usize,
}
impl Default for CompressionConfig {
fn default() -> Self {
Self {
enabled: true,
trigger_tokens: 30_000,
keep_last_messages: 40,
max_transcript_chars: 20_000,
max_summary_chars: 2_000,
}
}
}
pub struct SummarizingMiddleware {
client: Arc<dyn LlmClient>,
config: CompressionConfig,
cache: Mutex<HashMap<(u64, u64), (usize, String)>>,
}
impl SummarizingMiddleware {
pub fn new(client: Arc<dyn LlmClient>) -> Self {
Self { client, config: CompressionConfig::default(), cache: Mutex::new(HashMap::new()) }
}
pub fn with_config(mut self, config: CompressionConfig) -> Self {
self.config = config;
self
}
pub fn with_trigger_tokens(mut self, tokens: usize) -> Self {
self.config.trigger_tokens = tokens;
self
}
pub fn with_keep_last_messages(mut self, n: usize) -> Self {
self.config.keep_last_messages = n;
self
}
pub fn with_max_summary_chars(mut self, chars: usize) -> Self {
self.config.max_summary_chars = chars;
self
}
pub fn config(&self) -> &CompressionConfig {
&self.config
}
}
#[async_trait]
impl Middleware for SummarizingMiddleware {
async fn on_pre_llm(&self, ctx: &mut PreLlmCtx) -> agent_base::AgentResult<()> {
if !self.config.enabled {
return Ok(());
}
let messages = &ctx.messages;
if messages.len() <= self.config.keep_last_messages + 2 {
return Ok(());
}
let total_tokens: usize = messages.iter().map(estimate_message_tokens).sum();
if total_tokens <= self.config.trigger_tokens {
return Ok(());
}
let keep_first = if matches!(messages.first(), Some(ChatMessage::System { .. })) { 1 } else { 0 };
let mut recent_start = messages.len().saturating_sub(self.config.keep_last_messages);
if recent_start <= keep_first + 1 {
return Ok(());
}
recent_start = safe_cut_index(messages, keep_first, recent_start);
if recent_start <= keep_first {
return Ok(());
}
let old = &messages[keep_first..recent_start];
if old.is_empty() {
return Ok(());
}
if matches!(old.first(), Some(ChatMessage::Tool { .. })) {
tracing::warn!("context compression skipped: old block starts with a tool result");
return Ok(());
}
let transcript = serialize_block(old, self.config.max_transcript_chars);
if transcript.trim().is_empty() {
return Ok(());
}
const CACHE_PREFIX_CHARS: usize = 4096;
let prefix: String = transcript.chars().take(CACHE_PREFIX_CHARS).collect();
let key = (ctx.session_id.id, transcript_hash(&prefix));
let cached = self.cache.lock().ok().and_then(|c| c.get(&key).cloned());
let summary = match cached {
Some((cached_len, s)) if cached_len <= transcript.len() => s,
_ => {
let s = match summarize(self.client.as_ref(), &transcript, self.config.max_summary_chars).await {
Ok(s) => s,
Err(e) => {
tracing::warn!(
session_id = ctx.session_id.id,
"context compression summarization failed, dropping old block: {e}"
);
String::new()
},
};
if !s.is_empty()
&& let Ok(mut cache) = self.cache.lock()
{
cache.insert(key, (transcript.len(), s.clone()));
}
s
},
};
let mut new_messages: Vec<ChatMessage> = messages[..keep_first].to_vec();
let trimmed = summary.trim();
if !trimmed.is_empty() {
new_messages.push(ChatMessage::user(format!("[Earlier conversation summary]\n{trimmed}")));
}
new_messages.extend_from_slice(&messages[recent_start..]);
tracing::info!(
session_id = ctx.session_id.id,
before = messages.len(),
after = new_messages.len(),
estimated_tokens = total_tokens,
"context compressed"
);
ctx.messages = new_messages;
Ok(())
}
}
fn safe_cut_index(messages: &[ChatMessage], keep_first: usize, mut cut: usize) -> usize {
while cut > keep_first {
let left_is_tool_call = matches!(messages[cut - 1], ChatMessage::Assistant { tool_calls: Some(_), .. });
let right_is_tool = matches!(messages[cut], ChatMessage::Tool { .. });
if !left_is_tool_call && !right_is_tool {
break;
}
cut -= 1;
}
cut
}
fn estimate_message_tokens(msg: &ChatMessage) -> usize {
match msg {
ChatMessage::System { content, .. } => ContextWindowManager::estimate_tokens(content),
ChatMessage::User { content, images, .. } => {
ContextWindowManager::estimate_tokens(content) + images.len() * 85
},
ChatMessage::Assistant { content, reasoning_content, tool_calls } => {
let mut tokens = content.as_deref().map(ContextWindowManager::estimate_tokens).unwrap_or(0);
if let Some(rc) = reasoning_content {
tokens += ContextWindowManager::estimate_tokens(rc);
}
if let Some(calls) = tool_calls {
for c in calls {
tokens += ContextWindowManager::estimate_tokens(&c.id);
tokens += ContextWindowManager::estimate_tokens(&c.name);
tokens += ContextWindowManager::estimate_tokens(&c.arguments);
}
}
tokens
},
ChatMessage::Tool { tool_call_id, content } => {
ContextWindowManager::estimate_tokens(tool_call_id) + ContextWindowManager::estimate_tokens(content)
},
}
}
fn serialize_block(messages: &[ChatMessage], max_chars: usize) -> String {
let mut parts: Vec<String> = Vec::with_capacity(messages.len());
for msg in messages {
let line = match msg {
ChatMessage::System { content, .. } => format!("[system] {}", truncate(content, 400)),
ChatMessage::User { content, .. } => format!("[user] {}", truncate(content, 400)),
ChatMessage::Assistant { content, tool_calls, .. } => match tool_calls {
Some(calls) if !calls.is_empty() => {
let calls: Vec<String> =
calls.iter().map(|c| format!("{}({})", c.name, truncate(&c.arguments, 150))).collect();
format!("[assistant tool_call] {}", calls.join("; "))
},
_ => format!("[assistant] {}", content.as_deref().map(|c| truncate(c, 400)).unwrap_or_default()),
},
ChatMessage::Tool { tool_call_id, content } => format!("[tool:{tool_call_id}] {}", truncate(content, 300)),
};
parts.push(line);
}
truncate(&parts.join("\n"), max_chars)
}
fn truncate(s: &str, max_chars: usize) -> String {
let count = s.chars().count();
if count <= max_chars {
return s.to_string();
}
let head: String = s.chars().take(max_chars).collect();
format!("{head}…")
}
fn transcript_hash(s: &str) -> u64 {
use std::hash::{DefaultHasher, Hash, Hasher};
let mut h = DefaultHasher::new();
s.hash(&mut h);
h.finish()
}
async fn summarize(client: &dyn LlmClient, transcript: &str, max_chars: usize) -> agent_base::AgentResult<String> {
let system = ChatMessage::system(
"You are a conversation summarizer for an AI agent that can call tools \
(browser, shell, search, etc.).",
);
let user = ChatMessage::user(format!(
"Compress the earlier portion of this agent conversation. Preserve:\n\
- the user's original goal and any constraints they stated;\n\
- every important fact, decision and intermediate result;\n\
- which tools were used and their key findings/returned data;\n\
- blockers, errors, and anything the agent still needs to remember to continue.\n\
Detect the conversation language and write the summary in that same language.\n\
Output ONLY the summary text, no preamble, about {max_chars} characters max.\n\n\
=== CONVERSATION ===\n{transcript}",
));
let response = client.chat(&[system, user], &[], None, None).await?;
Ok(extract_content(&response).unwrap_or_default())
}
fn extract_content(response: &Value) -> Option<String> {
if let Some(s) = response
.get("choices")
.and_then(Value::as_array)
.and_then(|arr| arr.first())
.and_then(|ch| ch.get("message"))
.and_then(|m| m.get("content"))
.and_then(Value::as_str)
{
return Some(s.to_string());
}
response.get("content").and_then(Value::as_str).map(str::to_string)
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicUsize, Ordering};
struct MockClient {
summary: String,
calls: AtomicUsize,
}
#[async_trait]
impl LlmClient for MockClient {
async fn chat(
&self,
_messages: &[ChatMessage],
_tools: &[Value],
_reasoning: Option<&agent_base::ReasoningConfig>,
_response_format: Option<&agent_base::ResponseFormat>,
) -> agent_base::AgentResult<Value> {
self.calls.fetch_add(1, Ordering::SeqCst);
Ok(serde_json::json!({
"choices": [{ "message": { "content": self.summary } }]
}))
}
async fn chat_stream(
&self,
_messages: &[ChatMessage],
_tools: &[Value],
_reasoning: Option<&agent_base::ReasoningConfig>,
_response_format: Option<&agent_base::ResponseFormat>,
) -> agent_base::AgentResult<
std::pin::Pin<
Box<dyn futures_core::Stream<Item = agent_base::AgentResult<agent_base::StreamChunk>> + Send>,
>,
> {
unreachable!("not used in tests")
}
fn capabilities(&self) -> agent_base::LlmCapabilities {
agent_base::LlmCapabilities::default()
}
}
fn sample_messages() -> Vec<ChatMessage> {
vec![
ChatMessage::system("sys"),
ChatMessage::user("q1"),
ChatMessage::assistant_tool_call("call_1", "browser_navigate", r#"{"url":"a"}"#),
ChatMessage::tool("call_1", "loaded ok"),
ChatMessage::assistant("found the page"),
ChatMessage::user("q2"),
ChatMessage::assistant("done"),
]
}
#[test]
fn test_safe_cut_never_splits_tool_pair() {
let msgs = sample_messages();
let cut = safe_cut_index(&msgs, 1, 3);
assert!(cut < 3, "must walk backward from an unsafe cut");
assert!(!matches!(msgs[cut - 1], ChatMessage::Assistant { tool_calls: Some(_), .. }));
assert!(!matches!(msgs[cut], ChatMessage::Tool { .. }));
}
#[test]
fn test_safe_cut_prefers_given_boundary_when_safe() {
let msgs = vec![
ChatMessage::system("sys"),
ChatMessage::user("q1"),
ChatMessage::assistant("a1"),
ChatMessage::user("q2"),
ChatMessage::assistant("a2"),
];
let cut = safe_cut_index(&msgs, 1, 3);
assert_eq!(cut, 3);
}
#[test]
fn test_estimate_message_tokens_cjk_and_tool() {
let sys = ChatMessage::system("中文系统提示");
let tool = ChatMessage::tool("call_9", "hello 世界".repeat(100));
assert!(estimate_message_tokens(&sys) < estimate_message_tokens(&tool));
assert!(estimate_message_tokens(&sys) > 0);
}
#[test]
fn test_serialize_block_preserves_tool_calls() {
let msgs =
vec![ChatMessage::user("short question"), ChatMessage::assistant_tool_call("c1", "browser_navigate", "{}")];
let out = serialize_block(&msgs, 1000);
assert!(out.contains("tool_call"));
assert!(out.contains("browser_navigate"));
}
#[test]
fn test_serialize_block_truncates_oversized_fields() {
let long = "x".repeat(1000);
let msgs = vec![ChatMessage::user(long.clone())];
let out = serialize_block(&msgs, 200);
assert!(out.chars().count() <= 201); assert!(!out.contains(&long), "full payload must not leak through");
}
#[test]
fn test_extract_content_shapes() {
let openai = serde_json::json!({
"choices": [{ "message": { "content": "SUMMARY" } }]
});
assert_eq!(extract_content(&openai).as_deref(), Some("SUMMARY"));
let flat = serde_json::json!({ "content": "FLAT" });
assert_eq!(extract_content(&flat).as_deref(), Some("FLAT"));
assert_eq!(extract_content(&serde_json::json!({ "nope": 1 })), None);
}
#[tokio::test]
async fn test_on_pre_llm_noop_when_under_threshold() {
let client = Arc::new(MockClient { summary: "S".to_string(), calls: AtomicUsize::new(0) });
let mw = SummarizingMiddleware::new(client);
let mut ctx = PreLlmCtx {
session_id: agent_base::SessionId { id: 1, external_id: None },
messages: vec![ChatMessage::system("sys"), ChatMessage::user("hi")],
tools: vec![],
};
mw.on_pre_llm(&mut ctx).await.unwrap();
assert_eq!(ctx.messages.len(), 2);
}
#[tokio::test]
async fn test_on_pre_llm_compresses_and_caches() {
let client = Arc::new(MockClient {
summary: "The user wanted to scrape articles.".to_string(),
calls: AtomicUsize::new(0),
});
let mw = SummarizingMiddleware::new(client.clone())
.with_trigger_tokens(1) .with_keep_last_messages(2);
let make_ctx = || PreLlmCtx {
session_id: agent_base::SessionId { id: 1, external_id: None },
messages: sample_messages(),
tools: vec![],
};
let mut ctx = make_ctx();
mw.on_pre_llm(&mut ctx).await.unwrap();
assert!(ctx.messages.len() < 7, "should have compressed, got {}", ctx.messages.len());
assert_eq!(client.calls.load(Ordering::SeqCst), 1, "summarizer called once");
assert!(matches!(ctx.messages[1], ChatMessage::User { .. }));
for m in &ctx.messages {
assert!(!matches!(m, ChatMessage::Tool { .. }), "no tool orphan after compression");
}
let mut ctx2 = make_ctx();
mw.on_pre_llm(&mut ctx2).await.unwrap();
assert_eq!(client.calls.load(Ordering::SeqCst), 1, "cache reused");
assert_eq!(ctx.messages.len(), ctx2.messages.len());
}
#[tokio::test]
async fn test_summarization_failure_drops_old_block() {
let failing = Arc::new(FailingClient);
let mw = SummarizingMiddleware::new(failing).with_trigger_tokens(1).with_keep_last_messages(2);
let mut ctx = PreLlmCtx {
session_id: agent_base::SessionId { id: 1, external_id: None },
messages: sample_messages(),
tools: vec![],
};
let result = mw.on_pre_llm(&mut ctx).await;
assert!(result.is_ok(), "must not fail the turn");
assert!(ctx.messages.len() < 7);
assert!(ctx.messages.iter().all(|m| !matches!(
m,
ChatMessage::User { content, .. } if content.contains("Earlier conversation summary")
)));
}
#[tokio::test]
async fn test_default_threshold_fires_on_long_conversation() {
let client = Arc::new(MockClient {
summary: "Earlier context summarised.".to_string(),
calls: AtomicUsize::new(0),
});
let mw = SummarizingMiddleware::new(client.clone());
let mut messages = vec![ChatMessage::system("sys")];
let chunk = "这是用于验证压缩默认阈值的中文文本。".repeat(80);
for i in 0..45 {
messages.push(ChatMessage::user(format!("q{i} {chunk}")));
messages.push(ChatMessage::assistant(format!("a{i}")));
}
let mut ctx = PreLlmCtx {
session_id: agent_base::SessionId { id: 1, external_id: None },
messages,
tools: vec![],
};
let before = ctx.messages.len();
assert!(before > 42, "test fixture must exceed the message-count gate");
mw.on_pre_llm(&mut ctx).await.unwrap();
assert!(
ctx.messages.len() < before,
"default threshold should compress a long conversation ({} → {})",
before,
ctx.messages.len()
);
assert!(
client.calls.load(Ordering::SeqCst) >= 1,
"summarizer must be invoked for a long conversation"
);
assert!(matches!(ctx.messages[1], ChatMessage::User { .. }));
}
struct FailingClient;
#[async_trait]
impl LlmClient for FailingClient {
async fn chat(
&self,
_messages: &[ChatMessage],
_tools: &[Value],
_reasoning: Option<&agent_base::ReasoningConfig>,
_response_format: Option<&agent_base::ResponseFormat>,
) -> agent_base::AgentResult<Value> {
Err(agent_base::AgentError::internal("summarization failed"))
}
async fn chat_stream(
&self,
_messages: &[ChatMessage],
_tools: &[Value],
_reasoning: Option<&agent_base::ReasoningConfig>,
_response_format: Option<&agent_base::ResponseFormat>,
) -> agent_base::AgentResult<
std::pin::Pin<
Box<dyn futures_core::Stream<Item = agent_base::AgentResult<agent_base::StreamChunk>> + Send>,
>,
> {
unreachable!("not used in tests")
}
fn capabilities(&self) -> agent_base::LlmCapabilities {
agent_base::LlmCapabilities::default()
}
}
}