use std::collections::hash_map::DefaultHasher;
use std::hash::{Hash, Hasher};
use std::sync::{Arc, Mutex};
use agent_base::llm_trait::LlmProvider;
use agent_base::{AgentResult, AgentRuntime, ChatMessage, ContextWindowManager, SessionId};
use crate::compression::config::CompressionConfig;
use crate::compression::events::{CompressionEvent, CompressionTrigger};
use crate::compression::filter::{SUMMARY_PREFIX, is_summary_message, split_system_prompt};
use crate::compression::summarizer::summarize;
type CompactionCache = std::collections::HashMap<(u64, u64), (usize, String)>;
pub struct ContextCompactor {
client: Arc<dyn LlmProvider>,
config: CompressionConfig,
cache: Arc<Mutex<CompactionCache>>,
}
#[allow(missing_docs)]
impl ContextCompactor {
pub fn new(client: Arc<dyn LlmProvider>, config: CompressionConfig) -> Self {
Self {
client,
config,
cache: Arc::new(Mutex::new(std::collections::HashMap::new())),
}
}
pub fn clone_handle(&self) -> Self {
Self {
client: self.client.clone(),
config: self.config.clone(),
cache: self.cache.clone(),
}
}
pub fn config(&self) -> &CompressionConfig {
&self.config
}
const MAX_PRESERVED_USER_TOKENS: usize = 20_000;
pub async fn compact(
&self,
session_id: u64,
messages: &[ChatMessage],
trigger: CompressionTrigger,
emit_fn: Option<&(dyn Fn(CompressionEvent) + Sync)>,
) -> AgentResult<Option<Vec<ChatMessage>>> {
let t0 = std::time::Instant::now();
if !self.config.enabled {
return Ok(None);
}
let total_tokens: usize = messages.iter().map(estimate_message_tokens).sum();
tracing::info!(
total_tokens,
trigger = self.config.trigger_tokens,
"[compact-timing] enter compact()"
);
if total_tokens <= self.config.trigger_tokens {
return Ok(None);
}
let (system_msgs, conversation) = split_system_prompt(messages);
tracing::info!(
elapsed_ms = t0.elapsed().as_millis() as u64,
conv_len = conversation.len(),
keep = self.config.keep_recent_messages,
"[compact-timing] after split_system_prompt"
);
if conversation.len() <= self.config.keep_recent_messages + 1 {
return Ok(None);
}
let mut recent_start = conversation.len() - self.config.keep_recent_messages;
recent_start = safe_cut_index(conversation, 0, recent_start);
tracing::info!(
elapsed_ms = t0.elapsed().as_millis() as u64,
recent_start,
"[compact-timing] after split old/recent"
);
if recent_start == 0 {
return Ok(None);
}
let old = &conversation[..recent_start];
let recent = &conversation[recent_start..];
if matches!(old.first(), Some(ChatMessage::Tool { .. })) {
tracing::warn!("context compression skipped: old block starts with a tool result");
return Ok(None);
}
let old: Vec<&ChatMessage> = old.iter().filter(|m| !is_summary_message(m)).collect();
let recent: Vec<&ChatMessage> = recent.iter().filter(|m| !is_summary_message(m)).collect();
let original_goal = old
.iter()
.find_map(|m| match m {
ChatMessage::User { content, .. } if !content.starts_with(SUMMARY_PREFIX) => {
Some(content.as_str())
}
_ => None,
})
.unwrap_or("(unknown goal)");
let original_goal = truncate_str(original_goal, 400);
let preserved_user_msgs: Vec<&ChatMessage> = old
.iter()
.filter(|m| matches!(m, ChatMessage::User { .. }))
.copied()
.collect();
let assistant_tool_msgs: Vec<&ChatMessage> = old
.iter()
.filter(|m| !matches!(m, ChatMessage::User { .. }))
.copied()
.collect();
let transcript = truncate_to_chars(&assistant_tool_msgs, self.config.max_transcript_chars);
if transcript.trim().is_empty() && preserved_user_msgs.is_empty() {
return Ok(None);
}
let msg_count = messages.len();
if let Some(emit) = emit_fn {
emit(CompressionEvent::Preparing {
session_id,
tokens_before: total_tokens,
msg_count,
trigger: trigger.clone(),
});
emit(CompressionEvent::Started {
session_id,
tokens_before: total_tokens,
msg_count,
trigger: trigger.clone(),
});
}
const CACHE_PREFIX_CHARS: usize = 4096;
let prefix: String = transcript.chars().take(CACHE_PREFIX_CHARS).collect();
let key = (session_id, hash_str(&prefix));
let cached = self.cache.lock().ok().and_then(|c| c.get(&key).cloned());
let cache_effective = matches!(cached, Some((cl, _)) if cl == transcript.len() || transcript.len() >= self.config.max_transcript_chars);
let summary = if transcript.trim().is_empty() {
String::new()
} else {
match cached {
Some((_, ref s)) if cache_effective => s.clone(),
_ => {
if let Some(emit) = emit_fn {
emit(CompressionEvent::Progress {
session_id,
chars: 0,
});
}
tracing::info!(
elapsed_ms = t0.elapsed().as_millis() as u64,
"[compact-timing] before summarize() LLM call"
);
let s = if let Some(emit) = emit_fn {
let progress_fn = |chars: usize| {
emit(CompressionEvent::Progress { session_id, chars });
};
summarize(
self.client.as_ref(),
&transcript,
&original_goal,
self.config.max_summary_chars,
Some(&progress_fn),
)
.await
} else {
summarize(
self.client.as_ref(),
&transcript,
&original_goal,
self.config.max_summary_chars,
None,
)
.await
};
let s = match s {
Ok(s) => {
tracing::info!(
elapsed_ms = t0.elapsed().as_millis() as u64,
"[compact-timing] summarize() returned Ok"
);
s
}
Err(e) => {
tracing::warn!(
session_id,
"summarisation failed, falling back to 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> = system_msgs.to_vec();
let truncated_user =
truncate_user_messages(&preserved_user_msgs, Self::MAX_PRESERVED_USER_TOKENS);
let trimmed = summary.trim();
let old_block_chars: usize = old.iter().map(|m| message_content_len(m)).sum();
let replacement_chars: usize = truncated_user
.iter()
.map(|m| message_content_len(m))
.sum::<usize>()
+ if trimmed.is_empty() {
0
} else {
SUMMARY_PREFIX.len() + 1 + trimmed.len()
};
if replacement_chars >= old_block_chars {
tracing::info!(
session_id,
old_block_chars,
replacement_chars,
"[compact-timing] replacement larger than original, skipping compression"
);
return Ok(None);
}
new_messages.extend(truncated_user.into_iter().cloned());
if !trimmed.is_empty() {
new_messages.push(ChatMessage::user(format!("{SUMMARY_PREFIX}\n{trimmed}")));
}
new_messages.extend(recent.into_iter().cloned());
let compressed_tokens: usize = new_messages.iter().map(estimate_message_tokens).sum();
tracing::info!(
session_id,
tokens_before = total_tokens,
tokens_after = compressed_tokens,
kept_recent = self.config.keep_recent_messages,
summary_len = trimmed.len(),
preserved_user_count = preserved_user_msgs.len(),
old_block_chars,
replacement_chars,
cache_hit = cache_effective,
elapsed_ms = t0.elapsed().as_millis() as u64,
"[compact-timing] compression complete"
);
Ok(Some(new_messages))
}
pub fn clear_cache(&self) {
if let Ok(mut cache) = self.cache.lock() {
cache.clear();
}
}
pub async fn compact_session(
&self,
runtime: &AgentRuntime,
session_id: &SessionId,
emit_fn: Option<std::sync::Arc<dyn Fn(agent_base::UserEvent) + Send + Sync>>,
) -> AgentResult<bool> {
let sid = session_id.id;
let trigger = CompressionTrigger::Manual;
let messages = runtime
.with_session_mut(session_id, |session| session.chat_messages().to_vec())
.await?;
let msg_count_before = messages.len();
let tokens_before: usize = messages.iter().map(estimate_message_tokens).sum();
let filtered: Vec<ChatMessage> = messages
.iter()
.filter(|m| !is_summary_message(m))
.cloned()
.collect();
let compressed = if let Some(ref ef) = emit_fn {
let ef = ef.clone();
self.compact(
sid,
&filtered,
trigger.clone(),
Some(&move |ev: CompressionEvent| {
ef(ev.into_user_event());
}),
)
.await?
} else {
self.compact(sid, &filtered, trigger.clone(), None).await?
};
let compressed = match compressed {
Some(msgs) => msgs,
None => {
return Ok(false);
}
};
let tokens_after: usize = compressed.iter().map(estimate_message_tokens).sum();
let msg_count_after = compressed.len();
runtime
.with_session_mut(session_id, |session| {
let msg_count_now = session.chat_messages().len();
if msg_count_now != msg_count_before {
return Err(agent_base::AgentError::internal(format!(
"session modified concurrently ({} → {} messages), aborting write-back",
msg_count_before, msg_count_now
)));
}
session.set_chat_messages(compressed).map_err(|e| {
agent_base::AgentError::internal(format!(
"set_chat_messages validation failed: {e}"
))
})
})
.await??;
let reduction_pct = if tokens_before > 0 {
((tokens_before as f64 - tokens_after as f64) / tokens_before as f64 * 100.0).round()
as i32
} else {
0
};
if let Some(ref f) = emit_fn {
f(CompressionEvent::Completed {
session_id: sid,
tokens_before,
tokens_after,
reduction_pct,
msg_count_before,
msg_count_after,
trigger,
}
.into_user_event());
}
Ok(true)
}
}
fn message_content_len(msg: &ChatMessage) -> usize {
match msg {
ChatMessage::System { content, .. } => content.len(),
ChatMessage::User { content, .. } => content.len(),
ChatMessage::Assistant {
content,
reasoning_content,
tool_calls,
thinking_signature: _,
} => {
let mut len = content.as_deref().map(|c| c.len()).unwrap_or(0);
if let Some(rc) = reasoning_content {
len += rc.len();
}
if let Some(calls) = tool_calls {
for c in calls {
len += c.id.len() + c.name.len() + c.arguments.len();
}
}
len
}
ChatMessage::Tool {
tool_call_id,
content,
..
} => tool_call_id.len() + content.len(),
ChatMessage::Custom { role, data } => role.len() + data.to_string().len(),
}
}
pub fn safe_cut_index(messages: &[ChatMessage], start: usize, mut cut: usize) -> usize {
while cut > start {
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
}
pub(crate) 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,
thinking_signature: _,
} => {
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)
}
ChatMessage::Custom { role, data } => {
ContextWindowManager::estimate_tokens(role)
+ ContextWindowManager::estimate_tokens(&data.to_string())
}
}
}
#[allow(dead_code)]
fn estimate_total_tokens(messages: &[ChatMessage]) -> usize {
messages.iter().map(estimate_message_tokens).sum()
}
fn hash_str(s: &str) -> u64 {
let mut h = DefaultHasher::new();
s.hash(&mut h);
h.finish()
}
pub 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_str(content, 400))
}
ChatMessage::User { content, .. } => format!("[user] {}", truncate_str(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_str(&c.arguments, 150)))
.collect();
format!("[assistant tool_call] {}", calls.join("; "))
}
_ => format!(
"[assistant] {}",
content
.as_deref()
.map(|c| truncate_str(c, 400))
.unwrap_or_default()
),
},
ChatMessage::Tool {
tool_call_id,
content,
..
} => format!(
"[tool:{}] {}",
truncate_str(tool_call_id, 50),
truncate_str(content, 300)
),
ChatMessage::Custom { role, data } => {
format!("[custom:{}] {}", role, truncate_str(&data.to_string(), 400))
}
};
parts.push(line);
}
let joined = parts.join("\n");
truncate_str(&joined, max_chars)
}
pub fn truncate_str(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 truncate_to_chars(messages: &[&ChatMessage], max_chars: usize) -> String {
serialize_block(messages, max_chars)
}
fn truncate_user_messages<'a>(
messages: &[&'a ChatMessage],
max_tokens: usize,
) -> Vec<&'a ChatMessage> {
if max_tokens == 0 {
return Vec::new();
}
let mut result: Vec<&ChatMessage> = Vec::with_capacity(messages.len());
let mut remaining = max_tokens;
for msg in messages.iter().rev() {
let tokens = estimate_message_tokens(msg);
if tokens <= remaining {
result.push(msg);
remaining -= tokens;
}
}
result.reverse();
result
}
#[async_trait::async_trait]
impl agent_base::ContextCompaction for ContextCompactor {
async fn compact(
&self,
session_id: &SessionId,
messages: &[ChatMessage],
) -> Option<Vec<ChatMessage>> {
match self
.compact(
session_id.id,
messages,
CompressionTrigger::InlineCompaction,
None,
)
.await
{
Ok(result) => result,
Err(e) => {
tracing::warn!(
session_id = session_id.id,
error = %e,
"inline compaction failed"
);
None
}
}
}
fn token_count_hint(&self, _session_id: &SessionId) -> Option<usize> {
None
}
}
#[cfg(test)]
mod tests {
use super::*;
use agent_base::ToolCallMessage;
use agent_base::llm_trait::response::FinishReason;
use agent_base::llm_trait::types::UsageInfo;
use agent_base::llm_trait::{
Capabilities, ChatRequest, ChatResponse, ChatStream, LlmError, LlmProvider, ProviderInfo,
};
struct MockClient(&'static str);
#[async_trait::async_trait]
impl LlmProvider for MockClient {
async fn stream(&self, _request: ChatRequest) -> Result<ChatStream, LlmError> {
let response = self.0.to_string();
Ok(ChatStream::new(Box::pin(futures_util::stream::once(
async move { Ok(agent_base::StreamChunk::Text(response)) },
))))
}
async fn chat(&self, _request: ChatRequest) -> Result<ChatResponse, LlmError> {
Ok(ChatResponse {
content: self.0.to_string(),
tool_calls: vec![],
usage: UsageInfo::default(),
finish_reason: FinishReason::Stop,
raw: None,
reasoning_content: None,
thinking_signature: None,
})
}
fn capabilities(&self) -> Capabilities {
Capabilities::default()
}
fn info(&self) -> ProviderInfo {
ProviderInfo {
name: "stub".to_string(),
model: "stub-model".to_string(),
version: None,
}
}
}
struct CountingClient {
response: &'static str,
calls: std::sync::Arc<std::sync::atomic::AtomicUsize>,
}
#[async_trait::async_trait]
impl LlmProvider for CountingClient {
async fn stream(&self, _request: ChatRequest) -> Result<ChatStream, LlmError> {
self.calls.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
let response = self.response.to_string();
Ok(ChatStream::new(Box::pin(futures_util::stream::once(
async move { Ok(agent_base::StreamChunk::Text(response)) },
))))
}
async fn chat(&self, _request: ChatRequest) -> Result<ChatResponse, LlmError> {
self.calls.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
Ok(ChatResponse {
content: self.response.to_string(),
tool_calls: vec![],
usage: UsageInfo::default(),
finish_reason: FinishReason::Stop,
raw: None,
reasoning_content: None,
thinking_signature: None,
})
}
fn capabilities(&self) -> Capabilities {
Capabilities::default()
}
fn info(&self) -> ProviderInfo {
ProviderInfo {
name: "stub".to_string(),
model: "stub-model".to_string(),
version: None,
}
}
}
struct FailingClient;
#[async_trait::async_trait]
impl LlmProvider for FailingClient {
async fn stream(&self, _request: ChatRequest) -> Result<ChatStream, LlmError> {
Err(LlmError::llm("summarisation failed"))
}
async fn chat(&self, _request: ChatRequest) -> Result<ChatResponse, LlmError> {
Err(LlmError::llm("summarisation failed"))
}
fn capabilities(&self) -> Capabilities {
Capabilities::default()
}
fn info(&self) -> ProviderInfo {
ProviderInfo {
name: "stub".to_string(),
model: "stub-model".to_string(),
version: None,
}
}
}
fn make_messages(count: usize) -> Vec<ChatMessage> {
let mut msgs = vec![ChatMessage::system("You are a test agent.")];
for i in 0..count {
msgs.push(ChatMessage::user(format!("question {i}")));
msgs.push(ChatMessage::assistant(format!(
"answer {i} with some extra content to make it longer"
)));
}
msgs
}
#[test]
fn test_safe_cut_never_splits_tool_pair() {
let msgs = vec![
ChatMessage::user("hi"),
ChatMessage::assistant(""),
ChatMessage::assistant_tool_call("call_1", "bash", "{}"),
ChatMessage::tool("call_1", "result"),
ChatMessage::user("next"),
];
let cut = super::safe_cut_index(&msgs, 0, 3);
assert!(
cut <= 2,
"cut should walk back before assistant{{tool_call}}"
);
assert!(
!matches!(
&msgs[cut - 1],
ChatMessage::Assistant {
tool_calls: Some(_),
..
}
),
"left boundary must not be assistant with pending tool calls"
);
}
#[test]
fn test_safe_cut_prefers_given_boundary_when_safe() {
let msgs = vec![
ChatMessage::user("hi"),
ChatMessage::assistant("hello"),
ChatMessage::user("next"),
ChatMessage::assistant("world"),
];
assert_eq!(super::safe_cut_index(&msgs, 0, 2), 2);
}
#[test]
fn test_safe_cut_walks_back_from_orphan_tool() {
let msgs = vec![
ChatMessage::user("a"),
ChatMessage::assistant("b"),
ChatMessage::tool("c1", "r1"),
ChatMessage::user("d"),
];
assert_eq!(super::safe_cut_index(&msgs, 0, 1), 1);
}
#[test]
fn test_estimate_tokens_cjk_and_tool() {
let msg = ChatMessage::tool("t1", "你好世界");
let tokens = super::estimate_message_tokens(&msg);
assert!((2..=10).contains(&tokens));
}
#[test]
fn test_estimate_tokens_user_with_images() {
let msg = ChatMessage::user_with_images(
"describe this",
vec![agent_base::ImageAttachment::Url {
url: "http://example.com/img.png".into(),
detail: Some(agent_base::ImageDetail::Auto),
}],
);
let tokens = super::estimate_message_tokens(&msg);
assert!(tokens >= 85);
}
#[test]
fn test_serialize_block_preserves_tool_calls() {
let msgs = [
ChatMessage::assistant(""),
ChatMessage::tool("c1", "result"),
];
let msg_refs: Vec<&ChatMessage> = msgs.iter().collect();
let out = super::serialize_block(&msg_refs, 2000);
assert!(out.contains("[tool:c1]"), "got: {out}");
assert!(out.contains("result"), "got: {out}");
}
#[test]
fn test_serialize_block_truncates_oversized_fields() {
let long_arg = "x".repeat(500);
let msgs = [ChatMessage::Assistant {
content: None,
reasoning_content: None,
tool_calls: Some(vec![ToolCallMessage {
id: "t1".into(),
name: "bash".into(),
arguments: long_arg,
}]),
thinking_signature: None,
}];
let msg_refs: Vec<&ChatMessage> = msgs.iter().collect();
let out = super::serialize_block(&msg_refs, 2000);
assert!(out.contains("bash("));
assert!(out.contains("…"));
assert!(out.len() < 400);
}
#[tokio::test]
async fn test_compact_noop_when_disabled() {
let config = CompressionConfig::default().with_enabled(false);
let client = std::sync::Arc::new(MockClient("summary"));
let compactor = ContextCompactor::new(client, config);
let msgs = make_messages(100);
let result = compactor
.compact(1, &msgs, CompressionTrigger::Auto, None)
.await
.unwrap();
assert!(result.is_none());
}
#[tokio::test]
async fn test_compact_noop_when_below_threshold() {
let config = CompressionConfig::default().with_trigger_tokens(999_999);
let client = std::sync::Arc::new(MockClient("summary"));
let compactor = ContextCompactor::new(client, config);
let msgs = make_messages(10);
let result = compactor
.compact(1, &msgs, CompressionTrigger::Auto, None)
.await
.unwrap();
assert!(result.is_none());
}
#[tokio::test]
async fn test_compact_produces_valid_output() {
let config = CompressionConfig::default()
.with_trigger_tokens(1) .with_keep_recent_messages(4);
let client = std::sync::Arc::new(MockClient("test summary text"));
let compactor = ContextCompactor::new(client, config);
let msgs = make_messages(20);
let result = compactor
.compact(1, &msgs, CompressionTrigger::Auto, None)
.await
.unwrap()
.unwrap();
assert!(matches!(&result[0], ChatMessage::System { .. }));
let has_summary = result.iter().any(|m| match m {
ChatMessage::User { content, .. } => content.starts_with(SUMMARY_PREFIX),
_ => false,
});
assert!(has_summary, "expected summary message in output");
let last = result.last().unwrap();
assert!(matches!(
last,
ChatMessage::User { .. } | ChatMessage::Assistant { .. }
));
for (i, msg) in result.iter().enumerate() {
assert!(
!matches!(msg, ChatMessage::Tool { .. }) || i > 0,
"orphan Tool at index 0"
);
}
}
#[tokio::test]
async fn test_compact_fallback_on_summarisation_failure() {
let config = CompressionConfig::default()
.with_trigger_tokens(1)
.with_keep_recent_messages(4);
let client = std::sync::Arc::new(FailingClient);
let compactor = ContextCompactor::new(client, config);
let msgs = make_messages(20);
let result = compactor
.compact(1, &msgs, CompressionTrigger::Auto, None)
.await
.unwrap()
.unwrap();
assert!(matches!(&result[0], ChatMessage::System { .. }));
let has_summary = result.iter().any(|m| match m {
ChatMessage::User { content, .. } => content.starts_with(SUMMARY_PREFIX),
_ => false,
});
assert!(!has_summary, "should have no summary on failure");
}
#[tokio::test]
async fn test_compact_caches_summary() {
let config = CompressionConfig::default()
.with_trigger_tokens(1)
.with_keep_recent_messages(4);
let calls = std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0));
let client = std::sync::Arc::new(CountingClient {
response: "cached summary",
calls: calls.clone(),
});
let compactor = ContextCompactor::new(client, config);
let msgs = make_messages(20);
let _ = compactor
.compact(1, &msgs, CompressionTrigger::Auto, None)
.await
.unwrap();
assert_eq!(calls.load(std::sync::atomic::Ordering::SeqCst), 1);
let _ = compactor
.compact(1, &msgs, CompressionTrigger::Auto, None)
.await
.unwrap();
assert_eq!(calls.load(std::sync::atomic::Ordering::SeqCst), 1);
let mut grown = msgs.clone();
grown.push(ChatMessage::user("new question"));
grown.push(ChatMessage::assistant("new answer"));
let _ = compactor
.compact(1, &grown, CompressionTrigger::Auto, None)
.await
.unwrap();
assert_eq!(calls.load(std::sync::atomic::Ordering::SeqCst), 2);
}
#[tokio::test]
async fn test_compact_cache_length_guard_re_summarises_on_growth() {
let config = CompressionConfig::default()
.with_trigger_tokens(1)
.with_keep_recent_messages(4)
.with_max_transcript_chars(60_000); let calls = std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0));
let client = std::sync::Arc::new(CountingClient {
response: "summary",
calls: calls.clone(),
});
let compactor = ContextCompactor::new(client, config);
let mut msgs = vec![ChatMessage::system("sys")];
for i in 0..200 {
msgs.push(ChatMessage::user(format!("question {i:04}")));
msgs.push(ChatMessage::assistant(format!("answer {i:04}")));
}
let _ = compactor
.compact(1, &msgs, CompressionTrigger::Auto, None)
.await
.unwrap();
assert_eq!(calls.load(std::sync::atomic::Ordering::SeqCst), 1);
let _ = compactor
.compact(1, &msgs, CompressionTrigger::Auto, None)
.await
.unwrap();
assert_eq!(calls.load(std::sync::atomic::Ordering::SeqCst), 1);
msgs.push(ChatMessage::user("extra question"));
msgs.push(ChatMessage::assistant("extra answer"));
let _ = compactor
.compact(1, &msgs, CompressionTrigger::Auto, None)
.await
.unwrap();
assert_eq!(
calls.load(std::sync::atomic::Ordering::SeqCst),
2,
"length guard should have triggered re-summarisation"
);
}
#[tokio::test]
async fn test_compact_no_orphan_tool_messages() {
let config = CompressionConfig::default()
.with_trigger_tokens(1)
.with_keep_recent_messages(4);
let client = std::sync::Arc::new(MockClient("summary"));
let compactor = ContextCompactor::new(client, config);
let mut msgs = vec![ChatMessage::system("sys")];
for i in 0..10 {
msgs.push(ChatMessage::user(format!("q{i}")));
msgs.push(ChatMessage::assistant_tool_call(
format!("tc{i}"),
"bash",
format!("{{\"cmd\":\"echo {i} with some extra arguments to make it longer\"}}"),
));
msgs.push(ChatMessage::tool(
format!("tc{i}"),
format!("result {i} with some extra output content"),
));
}
let result = compactor
.compact(1, &msgs, CompressionTrigger::Auto, None)
.await
.unwrap()
.unwrap();
for (i, msg) in result.iter().enumerate() {
if let ChatMessage::Tool { tool_call_id, .. } = msg {
assert!(i > 0, "Tool message at index 0 has no preceding assistant");
let prev = &result[i - 1];
match prev {
ChatMessage::Assistant {
tool_calls: Some(calls),
..
} => {
let ids: Vec<&str> = calls.iter().map(|c| c.id.as_str()).collect();
assert!(
ids.contains(&tool_call_id.as_str()),
"Tool {tool_call_id} not referenced by preceding assistant: {ids:?}"
);
}
other => panic!(
"Tool {tool_call_id} at index {i} not preceded by \
Assistant{{tool_calls}}, got: {other:?}"
),
}
}
}
}
#[tokio::test]
async fn test_compact_noop_empty_messages() {
let config = CompressionConfig::default().with_trigger_tokens(1);
let client = std::sync::Arc::new(MockClient("summary"));
let compactor = ContextCompactor::new(client, config);
let result = compactor
.compact(1, &[], CompressionTrigger::Auto, None)
.await
.unwrap();
assert!(result.is_none());
}
#[tokio::test]
async fn test_compact_noop_system_only() {
let config = CompressionConfig::default().with_trigger_tokens(1);
let client = std::sync::Arc::new(MockClient("summary"));
let compactor = ContextCompactor::new(client, config);
let msgs = vec![ChatMessage::system("You are helpful.")];
let result = compactor
.compact(1, &msgs, CompressionTrigger::Auto, None)
.await
.unwrap();
assert!(result.is_none());
}
#[tokio::test]
async fn test_compact_noop_conversation_equals_keep() {
let config = CompressionConfig::default()
.with_trigger_tokens(1)
.with_keep_recent_messages(100);
let client = std::sync::Arc::new(MockClient("summary"));
let compactor = ContextCompactor::new(client, config);
let msgs = make_messages(3); let result = compactor
.compact(1, &msgs, CompressionTrigger::Auto, None)
.await
.unwrap();
assert!(result.is_none());
}
#[tokio::test]
async fn test_compact_filters_old_summary_in_old_block() {
let config = CompressionConfig::default()
.with_trigger_tokens(1)
.with_keep_recent_messages(4);
let client = std::sync::Arc::new(MockClient("new summary"));
let compactor = ContextCompactor::new(client, config);
let mut msgs = vec![ChatMessage::system("sys")];
msgs.push(ChatMessage::user(format!(
"{SUMMARY_PREFIX}\nprevious summary text"
)));
for i in 0..10 {
msgs.push(ChatMessage::user(format!("question {i}")));
msgs.push(ChatMessage::assistant(format!(
"answer {i} with some extra content to make the old block large enough"
)));
}
let result = compactor
.compact(1, &msgs, CompressionTrigger::Auto, None)
.await
.unwrap()
.unwrap();
let summary_count = result
.iter()
.filter(|m| match m {
ChatMessage::User { content, .. } => content.starts_with(SUMMARY_PREFIX),
_ => false,
})
.count();
assert_eq!(
summary_count, 1,
"should have exactly 1 summary (the new one), old one must be filtered"
);
}
#[tokio::test]
async fn test_compact_filters_old_summary_in_recent_block() {
let config = CompressionConfig::default()
.with_trigger_tokens(1)
.with_keep_recent_messages(4);
let client = std::sync::Arc::new(MockClient("new summary"));
let compactor = ContextCompactor::new(client, config);
let mut msgs = vec![ChatMessage::system("sys")];
for i in 0..5 {
msgs.push(ChatMessage::user(format!("question {i}")));
msgs.push(ChatMessage::assistant(format!(
"answer {i} with some extra content to make the old block large enough"
)));
}
msgs.push(ChatMessage::user(format!(
"{SUMMARY_PREFIX}\nold summary in recent"
)));
for i in 5..10 {
msgs.push(ChatMessage::user(format!("question {i}")));
msgs.push(ChatMessage::assistant(format!(
"answer {i} with some extra content to make the old block large enough"
)));
}
let result = compactor
.compact(1, &msgs, CompressionTrigger::Auto, None)
.await
.unwrap()
.unwrap();
let summary_count = result
.iter()
.filter(|m| match m {
ChatMessage::User { content, .. } => content.starts_with(SUMMARY_PREFIX),
_ => false,
})
.count();
assert_eq!(
summary_count, 1,
"old summary in recent block must be filtered"
);
}
#[tokio::test]
#[ignore] async fn test_compact_with_real_deepseek_api() {
let api_key = std::env::var("DEEPSEEK_API_KEY").unwrap_or_default();
if api_key.is_empty() {
eprintln!("Skipping test: DEEPSEEK_API_KEY not set");
return;
}
let _base_url = std::env::var("DEEPSEEK_BASE_URL")
.unwrap_or_else(|_| "https://api.deepseek.com".to_string());
return;
}
mod proptest_tests {
use super::*;
use proptest::prelude::*;
proptest! {
#[test]
fn truncate_str_chars_count_bounded(s in ".*", max in 0usize..500) {
let result = truncate_str(&s, max);
let char_count = result.chars().count();
let original_count = s.chars().count();
if original_count <= max {
assert_eq!(char_count, original_count);
} else {
assert!(char_count <= max + 1,
"truncated {} chars > max {} + 1", char_count, max);
}
}
#[test]
fn truncate_str_short_string_unchanged(s in "[a-zA-Z\u{4e00}-\u{9fff}]{0,50}", max in 50usize..200) {
let result = truncate_str(&s, max);
assert_eq!(result, s, "short string should be unchanged");
}
#[test]
fn truncate_str_cjk_char_boundary_safe(s in "[\u{4e00}-\u{9fff}]{1,200}", max in 1usize..100) {
let result = truncate_str(&s, max);
assert!(result.chars().count() <= max + 1);
}
#[test]
fn truncate_str_result_is_valid_utf8(s in ".*", max in 0usize..500) {
let result = truncate_str(&s, max);
assert!(std::str::from_utf8(result.as_bytes()).is_ok());
let _ = result.len(); let _ = result.chars().count(); }
}
fn arb_messages() -> impl Strategy<Value = Vec<ChatMessage>> {
prop::collection::vec(
prop_oneof![
"[a-z ]{0,50}".prop_map(|s| ChatMessage::user(&s)),
"[a-z ]{0,50}".prop_map(|s| ChatMessage::assistant(&s)),
("[a-z]{1,10}", "[a-z ]{0,50}")
.prop_map(|(id, content)| ChatMessage::tool(&id, &content)),
],
0..15,
)
}
proptest! {
#[test]
fn safe_cut_index_never_panics(
messages in arb_messages(),
start in 0usize..15,
) {
if messages.is_empty() || start >= messages.len() {
return Ok(());
}
let cut = start + (messages.len() - start) / 2;
if cut > messages.len() {
return Ok(());
}
let _result = super::safe_cut_index(&messages, start, cut);
}
#[test]
fn safe_cut_index_result_in_bounds(
messages in arb_messages(),
start in 0usize..15,
) {
if messages.is_empty() || start >= messages.len() {
return Ok(());
}
let cut = start + (messages.len() - start) / 2;
if cut > messages.len() || cut < start {
return Ok(());
}
let result = super::safe_cut_index(&messages, start, cut);
assert!(result >= start, "result {} < start {}", result, start);
assert!(result <= messages.len(), "result {} > len {}", result, messages.len());
}
#[test]
fn safe_cut_index_never_ends_on_tool_call(
messages in arb_messages(),
start in 0usize..15,
) {
if messages.is_empty() || start >= messages.len() {
return Ok(());
}
let cut = start + (messages.len() - start) / 2;
if cut > messages.len() || cut < start {
return Ok(());
}
let result = super::safe_cut_index(&messages, start, cut);
if result > 0 && result < messages.len() {
assert!(
!matches!(&messages[result - 1], ChatMessage::Assistant { tool_calls: Some(_), .. }),
"safe_cut left boundary should not be assistant with tool_calls"
);
}
}
}
}
}