use crate::compact::types::{CompactionContext, CompactionOutcome};
use crate::compact::{ContextCompactor, ContextManager};
use crate::message::{Message, MessagePart, Role};
use std::collections::HashSet;
use std::future::Future;
use std::pin::Pin;
#[derive(Debug, Clone)]
pub struct TruncatingCompactor {
preserve_recent: usize,
min_messages: usize,
}
impl TruncatingCompactor {
#[must_use]
pub fn new() -> Self {
Self {
preserve_recent: 4,
min_messages: 6,
}
}
#[must_use]
pub fn with_preserve_recent(mut self, count: usize) -> Self {
self.preserve_recent = count.max(1);
self
}
#[must_use]
pub fn with_min_messages(mut self, count: usize) -> Self {
self.min_messages = count.max(2);
self
}
#[must_use]
pub fn preserve_recent(&self) -> usize {
self.preserve_recent
}
#[must_use]
pub fn min_messages(&self) -> usize {
self.min_messages
}
}
impl Default for TruncatingCompactor {
fn default() -> Self {
Self::new()
}
}
impl ContextCompactor for TruncatingCompactor {
fn compact(
&self,
messages: Vec<Message>,
_target_tokens: u64,
context: CompactionContext,
) -> Pin<Box<dyn Future<Output = CompactionOutcome> + Send + '_>> {
Box::pin(async move {
let total = messages.len();
if total <= self.min_messages {
return CompactionOutcome::no_change(messages);
}
let initial_split = total.saturating_sub(self.preserve_recent);
let split = Self::adjust_for_tool_pairs(&messages, initial_split);
let recent: Vec<Message> = messages.get(split..).unwrap_or_default().to_vec();
let preserved = if split > 0 {
if let Some(first) = messages.first() {
let mut v = vec![first.clone()];
v.extend(recent);
v
} else {
recent
}
} else {
recent
};
let tokens_after = CompactionOutcome::estimate_tokens(&preserved);
CompactionOutcome {
messages: preserved,
tokens_after,
tokens_saved: context.tokens_before.saturating_sub(tokens_after),
success: true,
error: None,
}
})
}
}
impl TruncatingCompactor {
fn adjust_for_tool_pairs(messages: &[Message], split: usize) -> usize {
if split == 0 {
return 0;
}
let recent = messages.get(split..).unwrap_or_default();
let recent_call_ids: HashSet<&String> = recent
.iter()
.flat_map(|msg| msg.parts.iter())
.filter_map(|part| match part {
MessagePart::ToolCall { id, .. } => Some(id),
_ => None,
})
.collect();
let orphaned_ids: Vec<&String> = recent
.iter()
.flat_map(|msg| msg.parts.iter())
.filter_map(|part| match part {
MessagePart::ToolResult { call_id, .. } => {
if recent_call_ids.contains(call_id) {
None
} else {
Some(call_id)
}
}
_ => None,
})
.collect();
if orphaned_ids.is_empty() {
return split;
}
let mut new_split = split;
for i in (0..split).rev() {
let Some(msg) = messages.get(i) else {
continue;
};
let has_orphaned_call = msg.parts.iter().any(|part| match part {
MessagePart::ToolCall { id, .. } => orphaned_ids.contains(&id),
_ => false,
});
if has_orphaned_call {
new_split = i;
}
}
new_split
}
}
#[derive(Debug, Clone)]
pub struct TokenSplitter {
preserve_recent: usize,
min_messages: usize,
}
#[derive(Debug, Clone)]
pub struct SplitResult {
pub to_compact: Vec<Message>,
pub preserved: Vec<Message>,
pub compact_tokens: u64,
pub preserved_tokens: u64,
pub split_index: usize,
}
impl TokenSplitter {
#[must_use]
pub fn new() -> Self {
Self {
preserve_recent: 4,
min_messages: 6,
}
}
#[must_use]
pub fn with_preserve_recent(mut self, count: usize) -> Self {
self.preserve_recent = count.max(1);
self
}
#[must_use]
pub fn with_min_messages(mut self, count: usize) -> Self {
self.min_messages = count.max(2);
self
}
#[must_use]
pub fn split(&self, messages: &[Message]) -> SplitResult {
if messages.len() <= self.min_messages {
return SplitResult {
to_compact: vec![],
preserved: messages.to_vec(),
compact_tokens: 0,
preserved_tokens: ContextManager::estimate_tokens(messages),
split_index: 0,
};
}
let target_split = messages.len().saturating_sub(self.preserve_recent);
let split_index = Self::find_turn_boundary(messages, target_split);
let (to_compact, preserved) = messages.split_at(split_index);
SplitResult {
to_compact: to_compact.to_vec(),
preserved: preserved.to_vec(),
compact_tokens: ContextManager::estimate_tokens(to_compact),
preserved_tokens: ContextManager::estimate_tokens(preserved),
split_index,
}
}
fn find_turn_boundary(messages: &[Message], target: usize) -> usize {
if target == 0 {
return 0;
}
for i in (1..=target).rev() {
if i < messages.len() {
let Some(prev) = messages.get(i.saturating_sub(1)) else {
continue;
};
let Some(curr) = messages.get(i) else {
continue;
};
let prev_is_assistant = prev.role == Role::Assistant;
let curr_is_user = curr.role == Role::User;
if prev_is_assistant && curr_is_user {
return i;
}
}
}
target
}
}
impl Default for TokenSplitter {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::compact::ContextCompactor;
use crate::compact::types::{CompactReason, CompactionContext};
use crate::message::{Message, MessagePart, Role, ToolContent};
use serde_json::json;
fn tool_text(s: &str) -> ToolContent {
ToolContent::from_string(s)
}
fn make_context(msgs: &[Message]) -> CompactionContext {
CompactionContext {
tokens_before: CompactionOutcome::estimate_tokens(msgs),
reason: CompactReason::ThresholdExceeded,
context_window: 1_000,
turn: 5,
}
}
fn convo_with_straddling_tool_pair() -> Vec<Message> {
vec![
Message::user("msg0"),
Message::assistant("reply0"),
Message::user("msg1"),
Message::assistant("reply1"),
Message::user("msg2"),
Message::new(
Role::Assistant,
vec![MessagePart::tool_call(
"call_a",
"search",
json!({"q": "rust"}),
)],
),
Message::new(
Role::User,
vec![MessagePart::tool_result(
"call_a",
tool_text("result data"),
false,
)],
),
Message::assistant("final reply"),
]
}
fn has_tool_call(msgs: &[Message], id: &str) -> bool {
msgs.iter()
.flat_map(|m| m.parts.iter())
.any(|p| matches!(p, MessagePart::ToolCall { id: tool_id, .. } if tool_id == id))
}
fn has_tool_result(msgs: &[Message], call_id: &str) -> bool {
msgs.iter()
.flat_map(|m| m.parts.iter())
.any(|p| matches!(p, MessagePart::ToolResult { call_id: cid, .. } if cid == call_id))
}
#[tokio::test]
async fn compact_preserves_tool_call_when_result_is_in_recent() {
let messages = convo_with_straddling_tool_pair();
let compactor = TruncatingCompactor::new()
.with_preserve_recent(2)
.with_min_messages(4);
let context = make_context(&messages);
let outcome = compactor.compact(messages, 500, context).await;
assert!(
has_tool_call(&outcome.messages, "call_a"),
"tool-call 'call_a' must be preserved"
);
assert!(
has_tool_result(&outcome.messages, "call_a"),
"tool-result for 'call_a' must be preserved"
);
}
#[tokio::test]
async fn compact_does_not_orphan_when_pairs_are_together_in_recent() {
let messages = vec![
Message::user("msg0"),
Message::assistant("reply0"),
Message::user("msg1"),
Message::assistant("reply1"),
Message::user("msg2"),
Message::assistant("reply2"),
Message::new(
Role::Assistant,
vec![MessagePart::tool_call("call_b", "calc", json!({}))],
),
Message::new(
Role::User,
vec![MessagePart::tool_result("call_b", tool_text("42"), false)],
),
];
let compactor = TruncatingCompactor::new()
.with_preserve_recent(2)
.with_min_messages(4);
let context = make_context(&messages);
let outcome = compactor.compact(messages, 500, context).await;
assert!(
has_tool_call(&outcome.messages, "call_b"),
"tool-call 'call_b' must be preserved"
);
assert!(
has_tool_result(&outcome.messages, "call_b"),
"tool-result for 'call_b' must be preserved"
);
}
#[tokio::test]
async fn compact_drops_both_call_and_result_when_in_old_portion() {
let messages = vec![
Message::user("msg0"),
Message::new(
Role::Assistant,
vec![MessagePart::tool_call("call_c", "tool", json!({}))],
),
Message::new(
Role::User,
vec![MessagePart::tool_result("call_c", tool_text("done"), false)],
),
Message::assistant("reply1"),
Message::user("msg2"),
Message::assistant("reply2"),
Message::user("msg3"),
Message::assistant("reply3"),
];
let compactor = TruncatingCompactor::new()
.with_preserve_recent(4)
.with_min_messages(4);
let context = make_context(&messages);
let outcome = compactor.compact(messages, 500, context).await;
assert!(
!has_tool_call(&outcome.messages, "call_c"),
"tool-call 'call_c' should be dropped"
);
assert!(
!has_tool_result(&outcome.messages, "call_c"),
"tool-result for 'call_c' should be dropped"
);
}
#[test]
fn adjust_for_tool_pairs_returns_zero_when_split_is_zero() {
let messages = convo_with_straddling_tool_pair();
assert_eq!(TruncatingCompactor::adjust_for_tool_pairs(&messages, 0), 0);
}
#[test]
fn adjust_for_tool_pairs_no_orphans_returns_original_split() {
let messages = vec![
Message::user("a"),
Message::assistant("b"),
Message::user("c"),
Message::assistant("d"),
Message::user("e"),
Message::assistant("f"),
];
assert_eq!(TruncatingCompactor::adjust_for_tool_pairs(&messages, 4), 4);
}
#[test]
fn adjust_for_tool_pairs_moves_split_back_for_orphaned_result() {
let messages = convo_with_straddling_tool_pair();
assert_eq!(TruncatingCompactor::adjust_for_tool_pairs(&messages, 6), 5);
}
#[tokio::test]
async fn compact_short_conversation_returns_unchanged() {
let messages = vec![
Message::user("hello"),
Message::new(
Role::Assistant,
vec![MessagePart::tool_call("call_d", "tool", json!({}))],
),
Message::new(
Role::User,
vec![MessagePart::tool_result("call_d", tool_text("ok"), false)],
),
];
let compactor = TruncatingCompactor::new()
.with_preserve_recent(2)
.with_min_messages(6);
let context = make_context(&messages);
let outcome = compactor.compact(messages.clone(), 500, context).await;
assert_eq!(outcome.messages.len(), messages.len());
}
}