use crate::provider::{Message, Provider, ProviderError, Role};
use serde::{Deserialize, Serialize};
pub fn estimate_tokens(text: &str) -> usize {
text.len().div_ceil(3)
}
pub const IMAGE_TOKEN_COST: usize = 1200;
pub fn estimate_image_tokens() -> usize {
IMAGE_TOKEN_COST
}
pub fn estimate_messages(messages: &[Message]) -> usize {
const STRUCTURAL_OVERHEAD: usize = 22;
const TOOL_CALL_OVERHEAD: usize = 18;
let mut bytes: usize = 0;
for m in messages {
bytes += m.role.to_string().len();
bytes += m.content.len();
if let Some(tid) = &m.tool_call_id {
bytes += tid.len() + TOOL_CALL_OVERHEAD;
}
bytes += STRUCTURAL_OVERHEAD;
}
bytes.div_ceil(3)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum Severity {
Critical,
Important,
Informational,
}
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)]
pub enum CompactionMarker {
Task,
FileReference,
Decision,
ToolOutput,
UserCorrection,
SystemNote,
}
impl CompactionMarker {
pub fn severity(&self) -> Severity {
match self {
Self::Task => Severity::Critical,
Self::UserCorrection => Severity::Critical,
Self::Decision => Severity::Important,
Self::FileReference => Severity::Important,
Self::ToolOutput => Severity::Informational,
Self::SystemNote => Severity::Informational,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CompactionConfig {
pub context_window: usize,
pub reserve: usize,
pub keep_recent: usize,
pub trigger_threshold: usize,
}
impl CompactionConfig {
pub const DEFAULT_CONTEXT_WINDOW: usize = 128_000;
pub const DEFAULT_RESERVE: usize = 10_240;
pub const DEFAULT_KEEP_RECENT: usize = 12_800;
pub fn new(context_window: usize, reserve: usize, keep_recent: usize) -> Self {
Self {
context_window,
reserve,
keep_recent,
trigger_threshold: context_window.saturating_sub(reserve),
}
}
}
impl Default for CompactionConfig {
fn default() -> Self {
Self::new(
Self::DEFAULT_CONTEXT_WINDOW,
Self::DEFAULT_RESERVE,
Self::DEFAULT_KEEP_RECENT,
)
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct CompactionResult {
pub summary: String,
pub removed_count: usize,
pub removed_tokens: usize,
pub remaining_tokens: usize,
pub markers_preserved: Vec<CompactionMarker>,
}
pub fn compact_messages(messages: &[Message], config: &CompactionConfig) -> CompactionResult {
let total = estimate_messages(messages);
if total <= config.trigger_threshold {
return CompactionResult {
summary: String::new(),
removed_count: 0,
removed_tokens: 0,
remaining_tokens: total,
markers_preserved: Vec::new(),
};
}
let system_end = messages
.iter()
.position(|m| m.role != Role::System)
.unwrap_or(messages.len());
let mut preserved_tokens = 0usize;
let mut tail_start = messages.len();
for i in (system_end..messages.len()).rev() {
let cost = estimate_messages(std::slice::from_ref(&messages[i]));
if preserved_tokens + cost > config.keep_recent {
break;
}
preserved_tokens += cost;
tail_start = i;
}
let removable_end = tail_start.max(system_end);
let removed = &messages[system_end..removable_end];
let removed_tokens = estimate_messages(removed);
let summary = summarize_removed(removed);
let remaining_tokens = total.saturating_sub(removed_tokens);
let mut markers_preserved = Vec::new();
for m in &messages[..system_end] {
collect_markers(&m.content, &mut markers_preserved);
}
for m in &messages[tail_start..] {
collect_markers(&m.content, &mut markers_preserved);
}
markers_preserved.sort();
markers_preserved.dedup();
CompactionResult {
summary,
removed_count: removed.len(),
removed_tokens,
remaining_tokens,
markers_preserved,
}
}
pub fn apply_compaction(
messages: &mut Vec<Message>,
config: &CompactionConfig,
) -> CompactionResult {
let result = compact_messages(messages, config);
if result.removed_count == 0 {
return result;
}
let system_end = messages
.iter()
.position(|m| m.role != Role::System)
.unwrap_or(messages.len());
let mut preserved_tokens = 0usize;
let mut tail_start = messages.len();
for i in (system_end..messages.len()).rev() {
let cost = estimate_messages(std::slice::from_ref(&messages[i]));
if preserved_tokens + cost > config.keep_recent {
break;
}
preserved_tokens += cost;
tail_start = i;
}
let tail: Vec<Message> = messages[tail_start..].to_vec();
messages.truncate(system_end);
if !result.summary.is_empty() {
messages.push(Message::system(format!(
"[context compacted] {} Markers preserved: {:?}",
result.summary, result.markers_preserved
)));
}
messages.extend(tail);
result
}
pub async fn compact_messages_semantically(
messages: &[Message],
config: &CompactionConfig,
provider: &dyn Provider,
model: &str,
) -> Result<CompactionResult, ProviderError> {
let mut result = compact_messages(messages, config);
if result.removed_count == 0 {
return Ok(result);
}
let system_end = messages
.iter()
.position(|m| m.role != Role::System)
.unwrap_or(messages.len());
let removed_end = system_end + result.removed_count;
let system = Some(
"Summarize the conversation into a continuation-grade checkpoint. Preserve the user's objective and corrections, decisions and rationale, files changed or inspected, commands and test results, failures, and unfinished work. Be concise, factual, and specific. Do not continue the task."
.to_string(),
);
let transcript = messages[system_end..removed_end]
.iter()
.map(|message| {
let tool_calls = message
.tool_calls
.iter()
.map(|call| format!("\ntool call {} {}: {}", call.id, call.name, call.arguments))
.collect::<String>();
format!("{}: {}{}", message.role, message.content, tool_calls)
})
.collect::<Vec<_>>()
.join("\n\n");
result.summary = provider
.generate(&[Message::user(transcript)], &system, model, &[])
.await?;
if result.summary.trim().is_empty() {
return Err(ProviderError::Api(
"compaction provider returned an empty summary".to_string(),
));
}
Ok(result)
}
pub(crate) fn apply_compaction_result(
messages: &mut Vec<Message>,
original: &[Message],
result: &CompactionResult,
) -> bool {
if result.removed_count == 0 {
return false;
}
if !messages.starts_with(original) {
return false;
}
let system_end = messages
.iter()
.position(|m| m.role != Role::System)
.unwrap_or(messages.len());
let removed_end = system_end + result.removed_count;
messages.drain(system_end..removed_end);
messages.insert(
system_end,
Message::system(format!(
"[context compacted] {} Markers preserved: {:?}",
result.summary, result.markers_preserved
)),
);
true
}
fn summarize_removed(removed: &[Message]) -> String {
if removed.is_empty() {
return String::new();
}
let mut user_turns = 0usize;
let mut assistant_turns = 0usize;
let mut tool_turns = 0usize;
let mut chars: usize = 0;
for m in removed {
chars += m.content.len();
match m.role {
Role::User => user_turns += 1,
Role::Assistant => assistant_turns += 1,
Role::Tool => tool_turns += 1,
Role::System => {}
}
}
format!(
"Compacted {} messages ({} user, {} assistant, {} tool, ~{} chars).",
removed.len(),
user_turns,
assistant_turns,
tool_turns,
chars,
)
}
fn collect_markers(content: &str, out: &mut Vec<CompactionMarker>) {
let lower = content.to_ascii_lowercase();
if lower.contains("task:") || lower.contains("objective:") {
out.push(CompactionMarker::Task);
}
if lower.contains(".rs") || lower.contains("file:") || lower.contains("path:") {
out.push(CompactionMarker::FileReference);
}
if lower.contains("decided") || lower.contains("decision:") {
out.push(CompactionMarker::Decision);
}
if lower.contains("tool output") || lower.contains("tool_result") {
out.push(CompactionMarker::ToolOutput);
}
if lower.contains("correction") || lower.contains("actually,") {
out.push(CompactionMarker::UserCorrection);
}
if lower.contains("system note") {
out.push(CompactionMarker::SystemNote);
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::provider::Message;
struct SummaryProvider {
summary: Option<&'static str>,
}
#[async_trait::async_trait]
impl Provider for SummaryProvider {
fn id(&self) -> &str {
"summary"
}
fn name(&self) -> &str {
"Summary"
}
#[cfg(feature = "providers")]
async fn stream(
&self,
_messages: &[Message],
_system: &Option<String>,
_model: &str,
_tools: &[serde_json::Value],
_reasoning_effort: Option<&str>,
) -> Result<crate::provider::StreamResult, ProviderError> {
unreachable!()
}
async fn generate(
&self,
_messages: &[Message],
_system: &Option<String>,
_model: &str,
_tools: &[serde_json::Value],
) -> Result<String, ProviderError> {
self.summary
.map(str::to_string)
.ok_or_else(|| ProviderError::Api("summary failed".to_string()))
}
}
#[test]
fn estimate_tokens_three_chars_per_token() {
assert_eq!(estimate_tokens(""), 0);
assert_eq!(estimate_tokens("abc"), 1);
assert_eq!(estimate_tokens("abcd"), 2);
assert_eq!(estimate_tokens("abcdef"), 2);
assert_eq!(estimate_tokens("abcdefg"), 3);
}
#[test]
fn test_estimate_image_tokens_prevents_regressions() {
let tokens = estimate_image_tokens();
assert_eq!(
tokens, IMAGE_TOKEN_COST,
"Should return the constant IMAGE_TOKEN_COST"
);
assert_eq!(
tokens, 1200,
"Image token cost must remain exactly 1200 to prevent regressions"
);
}
#[test]
fn estimate_messages_grows_with_content() {
let one = vec![Message::user("hello world")];
let two = vec![Message::user("hello world"), Message::assistant("bye")];
assert!(estimate_messages(&two) > estimate_messages(&one));
assert!(estimate_messages(&one) > 0);
}
#[test]
fn estimate_messages_includes_tool_call_id() {
let plain = vec![Message::user("hello")];
let with_tool = vec![Message::tool("call_1", "hello")];
assert!(estimate_messages(&with_tool) > estimate_messages(&plain));
}
#[test]
fn no_compaction_under_threshold() {
let config = CompactionConfig::new(1_000, 100, 200);
let messages = vec![
Message::system("system prompt"),
Message::user("short message"),
];
let result = compact_messages(&messages, &config);
assert_eq!(result.removed_count, 0);
assert_eq!(result.removed_tokens, 0);
assert_eq!(result.remaining_tokens, estimate_messages(&messages));
assert!(result.summary.is_empty());
}
#[test]
fn apply_compaction_mutates_message_list() {
let config = CompactionConfig::new(100, 30, 20);
let mut messages = vec![
Message::system("system prompt"),
Message::user("old ".repeat(50)),
Message::assistant("mid ".repeat(50)),
Message::user("new ".repeat(50)),
];
let before_len = messages.len();
let result = apply_compaction(&mut messages, &config);
assert!(result.removed_count > 0);
assert!(messages.len() < before_len);
assert_eq!(messages.first().unwrap().role, Role::System);
assert!(messages
.iter()
.any(|m| m.content.contains("context compacted")));
}
#[tokio::test]
async fn semantic_compaction_uses_provider_summary() {
let config = CompactionConfig::new(100, 30, 20);
let messages = vec![
Message::system("system prompt"),
Message::user("old objective ".repeat(50)),
Message::assistant("old work ".repeat(50)),
Message::user("recent tail"),
];
let provider = SummaryProvider {
summary: Some("Objective retained; tests still need to run."),
};
let result = compact_messages_semantically(&messages, &config, &provider, "test")
.await
.unwrap();
assert_eq!(
result.summary,
"Objective retained; tests still need to run."
);
}
#[tokio::test]
async fn semantic_compaction_failure_leaves_messages_untouched() {
let config = CompactionConfig::new(100, 30, 20);
let messages = vec![
Message::system("system prompt"),
Message::user("old objective ".repeat(50)),
Message::assistant("old work ".repeat(50)),
Message::user("recent tail"),
];
let original = messages.clone();
let provider = SummaryProvider { summary: None };
assert!(
compact_messages_semantically(&messages, &config, &provider, "test")
.await
.is_err()
);
assert_eq!(messages, original);
}
#[test]
fn applying_semantic_result_preserves_messages_appended_after_snapshot() {
let config = CompactionConfig::new(100, 30, 20);
let snapshot = vec![
Message::system("system prompt"),
Message::user("old objective ".repeat(50)),
Message::assistant("old work ".repeat(50)),
Message::user("recent tail"),
];
let mut result = compact_messages(&snapshot, &config);
result.summary = "checkpoint".to_string();
let mut messages = snapshot.clone();
messages.push(Message::user("appended while summarizing"));
assert!(apply_compaction_result(&mut messages, &snapshot, &result));
assert_eq!(
messages.last().unwrap().content,
"appended while summarizing"
);
}
#[test]
fn applying_semantic_result_rejects_divergent_prefix() {
let config = CompactionConfig::new(100, 30, 20);
let snapshot = vec![
Message::system("system prompt"),
Message::user("old objective ".repeat(50)),
Message::assistant("old work ".repeat(50)),
Message::user("recent tail"),
];
let mut result = compact_messages(&snapshot, &config);
result.summary = "checkpoint".to_string();
let mut messages = snapshot.clone();
messages[1] = Message::user("changed while summarizing");
let divergent = messages.clone();
assert!(!apply_compaction_result(&mut messages, &snapshot, &result));
assert_eq!(messages, divergent);
}
#[test]
fn compaction_removes_oldest_messages() {
let config = CompactionConfig::new(100, 30, 20);
let messages = vec![
Message::system("system prompt"),
Message::user("old ".repeat(50)),
Message::assistant("mid ".repeat(50)),
Message::user("new ".repeat(50)),
];
let result = compact_messages(&messages, &config);
assert!(result.removed_count > 0);
assert!(result.removed_tokens > 0);
assert!(result.remaining_tokens < estimate_messages(&messages));
assert!(!result.summary.is_empty());
}
#[test]
fn system_prompt_is_preserved() {
let config = CompactionConfig::new(100, 30, 20);
let system_content = "important system prompt";
let messages = vec![
Message::system(system_content),
Message::user("a".repeat(200)),
Message::assistant("b".repeat(200)),
Message::user("recent"),
];
let result = compact_messages(&messages, &config);
assert!(result.removed_count > 0);
assert!(result.remaining_tokens >= estimate_tokens(system_content));
}
#[test]
fn keep_recent_is_respected() {
let config = CompactionConfig::new(200, 60, 30);
let messages = vec![
Message::system("sys"),
Message::user("a".repeat(300)),
Message::assistant("b".repeat(300)),
Message::user("c".repeat(60)),
];
let result = compact_messages(&messages, &config);
let total = estimate_messages(&messages);
assert!(
result.remaining_tokens <= total,
"remaining {} should not exceed total {}",
result.remaining_tokens,
total
);
assert!(
result.remaining_tokens
<= estimate_messages(std::slice::from_ref(&messages[0])) + config.keep_recent,
"remaining {} should not exceed system + keep_recent {}",
result.remaining_tokens,
config.keep_recent
);
}
#[test]
fn trigger_threshold_is_context_minus_reserve() {
let config = CompactionConfig::new(128_000, 10_240, 12_800);
assert_eq!(config.trigger_threshold, 128_000 - 10_240);
}
#[test]
fn default_config_matches_spec() {
let config = CompactionConfig::default();
assert_eq!(config.context_window, 128_000);
assert_eq!(config.reserve, 10_240);
assert_eq!(config.keep_recent, 12_800);
assert_eq!(config.trigger_threshold, 128_000 - 10_240);
}
#[test]
fn marker_severity_classification() {
assert_eq!(CompactionMarker::Task.severity(), Severity::Critical);
assert_eq!(
CompactionMarker::UserCorrection.severity(),
Severity::Critical
);
assert_eq!(CompactionMarker::Decision.severity(), Severity::Important);
assert_eq!(
CompactionMarker::FileReference.severity(),
Severity::Important
);
assert_eq!(
CompactionMarker::ToolOutput.severity(),
Severity::Informational
);
assert_eq!(
CompactionMarker::SystemNote.severity(),
Severity::Informational
);
}
#[test]
fn markers_collected_from_preserved_messages() {
let config = CompactionConfig::new(100, 30, 20);
let messages = vec![
Message::system("Task: do the thing"),
Message::user("a".repeat(200)),
Message::assistant("b".repeat(200)),
Message::user("Decision: keep it simple"),
];
let result = compact_messages(&messages, &config);
assert!(result.markers_preserved.contains(&CompactionMarker::Task));
assert!(result
.markers_preserved
.contains(&CompactionMarker::Decision));
}
#[test]
fn compaction_result_serializes() {
let result = CompactionResult {
summary: "test".to_string(),
removed_count: 1,
removed_tokens: 10,
remaining_tokens: 20,
markers_preserved: vec![CompactionMarker::Task],
};
let json = serde_json::to_string(&result).unwrap();
let back: CompactionResult = serde_json::from_str(&json).unwrap();
assert_eq!(back, result);
}
}