use agent_context_contract::ConversationRecord;
use serde_json::Value;
#[derive(Debug, Clone)]
pub struct ContextCompressionConfig {
pub max_messages: usize,
pub max_chars: usize,
pub summarize_older: bool,
}
impl Default for ContextCompressionConfig {
fn default() -> Self {
Self {
max_messages: 100,
max_chars: 100_000,
summarize_older: true,
}
}
}
pub fn compress_context(
messages: &[ConversationRecord],
config: &ContextCompressionConfig,
) -> Vec<ConversationRecord> {
if messages.len() <= config.max_messages {
let total_chars: usize = messages.iter().map(|m| estimate_chars(&m.content)).sum();
if total_chars <= config.max_chars {
return messages.to_vec();
}
}
let keep_count = config.max_messages.min(messages.len());
let mut kept: Vec<ConversationRecord> = messages[messages.len() - keep_count..].to_vec();
let mut total_chars: usize = kept.iter().map(|m| estimate_chars(&m.content)).sum();
while total_chars > config.max_chars && kept.len() > 1 {
let removed = kept.remove(0);
total_chars = total_chars.saturating_sub(estimate_chars(&removed.content));
}
if config.summarize_older && kept.len() < messages.len() {
let dropped_count = messages.len() - kept.len();
let summary = ConversationRecord {
id: format!("__compressed_{}", dropped_count),
scope: messages
.first()
.map(|m| m.scope.clone())
.unwrap_or_default(),
role: "system".to_string(),
content: Value::String(format!(
"[Context compressed: {} earlier messages were summarized to fit the context window]",
dropped_count
)),
name: None,
tool_call_id: None,
sequence: 0,
created_at_ms: 0,
metadata: Value::Null,
};
kept.insert(0, summary);
}
kept
}
fn estimate_chars(content: &Value) -> usize {
match content {
Value::String(s) => s.len(),
Value::Array(parts) => parts
.iter()
.map(|p| {
p.get("text")
.and_then(Value::as_str)
.map(|s| s.len())
.unwrap_or(0)
})
.sum(),
_ => content.to_string().len(),
}
}
#[cfg(test)]
mod tests {
use super::*;
use agent_context_contract::ConversationScope;
fn make_msg(id: &str, text: &str) -> ConversationRecord {
ConversationRecord {
id: id.to_string(),
scope: ConversationScope::conversation("conv-1"),
role: "user".to_string(),
content: Value::String(text.to_string()),
name: None,
tool_call_id: None,
sequence: 0,
created_at_ms: 0,
metadata: Value::Null,
}
}
#[test]
fn compress_under_limits_returns_all() {
let msgs: Vec<_> = (0..5)
.map(|i| make_msg(&format!("m{i}"), "short"))
.collect();
let result = compress_context(&msgs, &ContextCompressionConfig::default());
assert_eq!(result.len(), 5);
}
#[test]
fn compress_over_message_limit_drops_oldest() {
let msgs: Vec<_> = (0..200)
.map(|i| make_msg(&format!("m{i}"), "msg"))
.collect();
let config = ContextCompressionConfig {
max_messages: 50,
max_chars: 1_000_000,
summarize_older: false,
};
let result = compress_context(&msgs, &config);
assert!(result.len() <= 50);
}
#[test]
fn compress_adds_summary_when_dropping() {
let msgs: Vec<_> = (0..200)
.map(|i| make_msg(&format!("m{i}"), "msg"))
.collect();
let config = ContextCompressionConfig {
max_messages: 50,
max_chars: 1_000_000,
summarize_older: true,
};
let result = compress_context(&msgs, &config);
assert_eq!(result[0].role, "system");
assert!(result[0].content.as_str().unwrap().contains("compressed"));
}
}