use std::sync::atomic::{AtomicUsize, Ordering};
use talos_core::message::{AgentEvent, Message, MessageToolResult};
use talos_core::provider::LanguageModel;
use crate::token::TokenEstimator;
use super::constants::{
CIRCUIT_BREAKER_THRESHOLD, COLLAPSE_TURN_THRESHOLD, MAX_TOOL_RESULT_CHARS, PRESERVED_TURNS,
TRIM_TURN_THRESHOLD, TRUNCATION_SUFFIX,
};
use super::{CompactionError, CompactionResult, CompactionStatus};
pub struct Compactor {
token_estimator: TokenEstimator,
model_limit: u32,
trigger_threshold: f32,
consecutive_failures: AtomicUsize,
}
impl Compactor {
#[must_use]
pub fn new(token_estimator: TokenEstimator, model_limit: u32) -> Self {
Self {
token_estimator,
model_limit,
trigger_threshold: 0.8,
consecutive_failures: AtomicUsize::new(0),
}
}
#[must_use]
pub fn with_threshold(mut self, threshold: f32) -> Self {
self.trigger_threshold = threshold.clamp(0.0, 1.0);
self
}
pub fn should_compact(&self, messages: &[Message]) -> bool {
let estimated = self.token_estimator.estimate(messages);
let threshold_tokens = (self.model_limit as f32 * self.trigger_threshold) as u32;
estimated > threshold_tokens
}
pub async fn compact(
&mut self,
messages: Vec<Message>,
provider: &dyn LanguageModel,
) -> CompactionResult<Vec<Message>> {
if self.consecutive_failures.load(Ordering::SeqCst) >= CIRCUIT_BREAKER_THRESHOLD {
return Err(CompactionError::CircuitBreakerTripped);
}
let mut current = messages;
current = self.apply_budget(current);
if self.fits(¤t) {
self.consecutive_failures.store(0, Ordering::SeqCst);
return Ok(current);
}
current = self.apply_trim(current);
if self.fits(¤t) {
self.consecutive_failures.store(0, Ordering::SeqCst);
return Ok(current);
}
current = self.apply_microcompact(current);
if self.fits(¤t) {
self.consecutive_failures.store(0, Ordering::SeqCst);
return Ok(current);
}
current = match self.apply_collapse(current, provider).await {
Ok(msgs) => msgs,
Err(e) => {
self.record_failure();
return Err(e);
}
};
if self.fits(¤t) {
self.consecutive_failures.store(0, Ordering::SeqCst);
return Ok(current);
}
current = match self.apply_autocompact(current, provider).await {
Ok(msgs) => msgs,
Err(e) => {
self.record_failure();
return Err(e);
}
};
if self.fits(¤t) {
self.consecutive_failures.store(0, Ordering::SeqCst);
Ok(current)
} else {
self.record_failure();
Err(CompactionError::CompactionFailed(
"all compaction layers applied but context still exceeds limit".into(),
))
}
}
#[must_use]
pub fn apply_budget(&self, messages: Vec<Message>) -> Vec<Message> {
messages
.into_iter()
.map(|msg| match msg {
Message::Tool { result } => {
if result.content.chars().count() > MAX_TOOL_RESULT_CHARS {
let mut truncated: String =
result.content.chars().take(MAX_TOOL_RESULT_CHARS).collect();
truncated.push_str(TRUNCATION_SUFFIX);
Message::Tool {
result: MessageToolResult {
content: truncated,
..result
},
}
} else {
Message::Tool { result }
}
}
other => other,
})
.collect()
}
#[must_use]
pub fn apply_trim(&self, messages: Vec<Message>) -> Vec<Message> {
let total_turns = count_turns(&messages);
if total_turns <= TRIM_TURN_THRESHOLD {
return messages;
}
let turns_to_trim = total_turns - TRIM_TURN_THRESHOLD;
let mut current_turn: usize = 0;
let mut in_trimmed_turn = false;
messages
.into_iter()
.map(|msg| {
if matches!(&msg, Message::User { .. }) {
if current_turn > 0 {
in_trimmed_turn = false;
}
current_turn += 1;
if current_turn <= turns_to_trim {
in_trimmed_turn = true;
}
}
if in_trimmed_turn && matches!(&msg, Message::Tool { .. }) {
if let Message::Tool { result } = msg {
Message::Tool {
result: MessageToolResult {
content: String::new(),
..result
},
}
} else {
msg
}
} else {
msg
}
})
.collect()
}
#[must_use]
pub fn apply_microcompact(&self, messages: Vec<Message>) -> Vec<Message> {
let mut last_occurrence: std::collections::HashMap<String, usize> =
std::collections::HashMap::new();
for (i, msg) in messages.iter().enumerate() {
if let Message::Tool { result } = msg {
last_occurrence.insert(result.tool_use_id.clone(), i);
}
}
messages
.into_iter()
.enumerate()
.map(|(i, msg)| {
if let Message::Tool { result } = msg {
if last_occurrence.get(&result.tool_use_id) == Some(&i) {
Message::Tool { result }
} else {
Message::Tool {
result: MessageToolResult {
content: String::new(),
..result
},
}
}
} else {
msg
}
})
.collect()
}
pub async fn apply_collapse(
&self,
messages: Vec<Message>,
provider: &dyn LanguageModel,
) -> CompactionResult<Vec<Message>> {
let total_turns = count_turns(&messages);
if total_turns <= COLLAPSE_TURN_THRESHOLD {
return Ok(messages);
}
let (old_messages, recent_messages) = split_at_turn(&messages, COLLAPSE_TURN_THRESHOLD);
if old_messages.is_empty() {
return Ok(messages);
}
let summary = self.summarize_with_llm(&old_messages, provider).await?;
let mut result = Vec::with_capacity(1 + recent_messages.len());
result.push(Message::User {
content: format!(
"[Conversation summary of {} earlier turns]\n{}",
total_turns - COLLAPSE_TURN_THRESHOLD,
summary
),
});
result.extend(recent_messages);
Ok(result)
}
pub async fn apply_autocompact(
&self,
messages: Vec<Message>,
provider: &dyn LanguageModel,
) -> CompactionResult<Vec<Message>> {
let total_turns = count_turns(&messages);
if total_turns <= PRESERVED_TURNS {
return Ok(messages);
}
let (old_messages, recent_messages) = split_at_turn(&messages, PRESERVED_TURNS);
if old_messages.is_empty() {
return Ok(messages);
}
let summary = self.summarize_with_llm(&old_messages, provider).await?;
let mut result = Vec::with_capacity(1 + recent_messages.len());
result.push(Message::User {
content: format!("[Full conversation summary]\n{}", summary),
});
result.extend(recent_messages);
Ok(result)
}
async fn summarize_with_llm(
&self,
messages: &[Message],
provider: &dyn LanguageModel,
) -> CompactionResult<String> {
let conversation_text = messages_to_text(messages);
let prompt_messages = vec![Message::User {
content: format!(
"Summarize the following conversation concisely. \
Preserve key decisions, tool call outcomes, and important context. \
Keep the summary under 500 words.\n\n\
Conversation:\n{conversation_text}"
),
}];
let mut rx = provider
.stream(&prompt_messages)
.await
.map_err(|e| CompactionError::ProviderError(e.to_string()))?;
let mut summary = String::new();
while let Some(event) = rx.recv().await {
if let AgentEvent::TextDelta { delta } = event {
summary.push_str(&delta);
}
}
if summary.is_empty() {
summary = "[No summary generated]".into();
}
Ok(summary)
}
fn fits(&self, messages: &[Message]) -> bool {
let estimated = self.token_estimator.estimate(messages);
estimated <= self.model_limit
}
#[must_use]
pub fn compact_deterministic(
&self,
messages: Vec<Message>,
) -> (Vec<Message>, CompactionStatus) {
let tokens_before = self.token_estimator.estimate(&messages);
let mut current = messages;
let mut layers = Vec::new();
current = self.apply_budget(current);
layers.push("budget");
if self.fits(¤t) {
let tokens_after = self.token_estimator.estimate(¤t);
return (
current,
CompactionStatus::Applied {
layers_applied: layers,
tokens_before,
tokens_after,
},
);
}
current = self.apply_trim(current);
layers.push("trim");
if self.fits(¤t) {
let tokens_after = self.token_estimator.estimate(¤t);
return (
current,
CompactionStatus::Applied {
layers_applied: layers,
tokens_before,
tokens_after,
},
);
}
current = self.apply_microcompact(current);
layers.push("microcompact");
let tokens_after = self.token_estimator.estimate(¤t);
if self.fits(¤t) {
(
current,
CompactionStatus::Applied {
layers_applied: layers,
tokens_before,
tokens_after,
},
)
} else {
(
current,
CompactionStatus::Skipped {
reason: "deterministic layers insufficient; LLM layers required",
tokens_current: tokens_after,
},
)
}
}
pub async fn manual_compact(
&mut self,
messages: Vec<Message>,
provider: &dyn LanguageModel,
) -> (Vec<Message>, CompactionStatus) {
if !self.should_compact(&messages) {
let tokens = self.token_estimator.estimate(&messages);
return (
messages,
CompactionStatus::Skipped {
reason: "below trigger threshold",
tokens_current: tokens,
},
);
}
let tokens_before = self.token_estimator.estimate(&messages);
match self.compact(messages.clone(), provider).await {
Ok(compacted) => {
let tokens_after = self.token_estimator.estimate(&compacted);
(
compacted,
CompactionStatus::Applied {
layers_applied: vec!["manual"],
tokens_before,
tokens_after,
},
)
}
Err(e) => (
messages,
CompactionStatus::Failed {
error: e.to_string(),
},
),
}
}
fn record_failure(&self) {
self.consecutive_failures.fetch_add(1, Ordering::SeqCst);
}
#[cfg(test)]
pub(super) fn failure_count(&self) -> usize {
self.consecutive_failures.load(Ordering::SeqCst)
}
}
fn count_turns(messages: &[Message]) -> usize {
messages
.iter()
.filter(|m| matches!(m, Message::User { .. }))
.count()
}
fn split_at_turn(messages: &[Message], turns_from_end: usize) -> (Vec<Message>, Vec<Message>) {
let total_turns = count_turns(messages);
if total_turns <= turns_from_end {
return (Vec::new(), messages.to_vec());
}
let turns_to_keep = turns_from_end;
let turns_to_skip = total_turns - turns_to_keep;
let mut current_turn: usize = 0;
let mut split_idx = 0;
for (i, msg) in messages.iter().enumerate() {
if matches!(msg, Message::User { .. }) {
current_turn += 1;
if current_turn > turns_to_skip {
split_idx = i;
break;
}
}
}
let old = messages[..split_idx].to_vec();
let recent = messages[split_idx..].to_vec();
(old, recent)
}
fn messages_to_text(messages: &[Message]) -> String {
messages
.iter()
.map(|msg| match msg {
Message::User { content } => format!("User: {content}"),
Message::System { content, .. } => format!("System: {content}"),
Message::Context { content } => format!("Context: {content}"),
Message::Assistant {
content,
tool_calls,
..
} => {
let mut text = format!("Assistant: {content}");
for tc in tool_calls {
text.push_str(&format!("\n [Tool call: {}({})]", tc.name, tc.input));
}
text
}
Message::Tool { result } => {
format!("Tool result ({}): {}", result.tool_use_id, result.content)
}
Message::Multimodal { parts } => {
let text: String = parts
.iter()
.filter_map(|p| match p {
talos_core::message::ContentPart::Text { text } => Some(text.as_str()),
_ => None,
})
.collect::<Vec<_>>()
.join("\n");
format!("User: {text}")
}
})
.collect::<Vec<_>>()
.join("\n\n")
}