use std::borrow::Cow;
use std::cmp::min;
use crate::tokenizer::Tokenizer;
use crate::types::{ContentBlock, Message, Role};
const COMPRESSED_TOOL_RESULT_STUB: &str = "[tool result compressed]";
const HEAD_TRIM_PROBE_MSGS: usize = 64;
fn contains_rule_sentinel(text: &str) -> bool {
text.contains(crate::rules::RULE_SENTINEL)
}
fn message_has_rule_injection(m: &Message) -> bool {
m.content.iter().any(|b| match b {
ContentBlock::Text { text } => contains_rule_sentinel(text),
ContentBlock::ToolResult { content, .. } => contains_rule_sentinel(content),
_ => false,
})
}
pub trait HistoryCompressor: Send + Sync {
fn compress(
&self,
messages: &mut Vec<Message>,
target_tokens: u32,
tokenizer: &dyn Tokenizer,
) -> u32;
}
#[derive(Debug, Default, Clone, Copy)]
pub struct DefaultHistoryCompressor;
impl HistoryCompressor for DefaultHistoryCompressor {
fn compress(
&self,
messages: &mut Vec<Message>,
target_tokens: u32,
tokenizer: &dyn Tokenizer,
) -> u32 {
let mut total = count_total(messages, tokenizer);
if total <= target_tokens {
return total;
}
for i in 0..messages.len() {
for block in messages[i].content.iter_mut() {
if let ContentBlock::ToolResult { content, .. } = block
&& content.as_str() != COMPRESSED_TOOL_RESULT_STUB
&& !contains_rule_sentinel(content)
{
*content = COMPRESSED_TOOL_RESULT_STUB.to_string();
}
}
total = count_total(messages, tokenizer);
if total <= target_tokens {
return total;
}
}
let preserved_first = !messages.is_empty() && is_user_anchor(&messages[0]);
let anchor_floor = if preserved_first { 1 } else { 0 };
let mut floor = anchor_floor;
let mut attempts = 0;
while count_total(messages, tokenizer) > target_tokens
&& attempts < HEAD_TRIM_PROBE_MSGS
&& messages.len() > floor + 1
{
let Some(turn_start) = messages
.iter()
.enumerate()
.skip(floor)
.find(|(_, m)| is_user_anchor(m))
.map(|(i, _)| i)
else {
break;
};
let turn_end = messages
.iter()
.enumerate()
.skip(turn_start + 1)
.find(|(_, m)| is_user_anchor(m))
.map_or(messages.len(), |(i, _)| i);
if turn_end == messages.len() {
break;
}
if messages[turn_start..turn_end]
.iter()
.any(message_has_rule_injection)
{
floor = turn_end;
continue;
}
messages.drain(turn_start..min(turn_end, messages.len()));
attempts += 1;
}
count_total(messages, tokenizer)
}
}
fn is_user_anchor(m: &Message) -> bool {
if m.role != Role::User {
return false;
}
if m.content
.iter()
.any(|b| matches!(b, ContentBlock::ToolResult { .. }))
{
return false;
}
m.content
.iter()
.any(|b| matches!(b, ContentBlock::Text { .. }))
}
fn count_total(messages: &[Message], tokenizer: &dyn Tokenizer) -> u32 {
let mut total: u32 = 0;
for m in messages {
total = total.saturating_add(4); for b in &m.content {
total = total.saturating_add(tokenizer.count(&block_text(b)));
}
}
total
}
fn block_text(b: &ContentBlock) -> Cow<'_, str> {
match b {
ContentBlock::Text { text } => Cow::Borrowed(text),
ContentBlock::Thinking { thinking, .. } => Cow::Borrowed(thinking),
ContentBlock::ToolUse { input, .. } => Cow::Owned(input.to_string()),
ContentBlock::ToolResult { content, .. } => Cow::Borrowed(content),
}
}
#[cfg(test)]
#[allow(clippy::expect_used, clippy::unwrap_used)]
mod tests {
use super::*;
use crate::tokenizer::ApproxTokenizer;
use crate::types::{Message, Role};
fn tool_result(id: &str, content: &str) -> Message {
Message {
role: Role::User,
content: vec![ContentBlock::ToolResult {
tool_use_id: id.into(),
content: content.into(),
is_error: Some(false),
}],
}
}
#[test]
fn noop_when_already_under_budget() {
let mut msgs = vec![Message::user_text("hi")];
let t = ApproxTokenizer;
let c = DefaultHistoryCompressor;
let n = c.compress(&mut msgs, 1000, &t);
assert!(n < 100);
assert_eq!(msgs.len(), 1);
}
#[test]
fn stubs_tool_results_when_over_budget() {
let big = "x".repeat(4000); let mut msgs = vec![
Message::user_text("hi"),
tool_result("t1", &big),
Message::assistant_text("ok"),
];
let t = ApproxTokenizer;
let c = DefaultHistoryCompressor;
let _ = c.compress(&mut msgs, 50, &t);
match &msgs[1].content[0] {
ContentBlock::ToolResult { content, .. } => {
assert_eq!(content, COMPRESSED_TOOL_RESULT_STUB);
}
_ => panic!("expected tool result"),
}
}
#[test]
fn trims_head_preserving_first_message() {
let t = ApproxTokenizer;
let c = DefaultHistoryCompressor;
let mut msgs = vec![
Message::user_text("<brief>"),
Message::user_text("a".repeat(4000)),
Message::assistant_text("b".repeat(4000)),
Message::user_text("c".repeat(4000)),
Message::assistant_text("final"),
];
let before = msgs.len();
c.compress(&mut msgs, 20, &t);
assert!(msgs.len() < before);
match &msgs[0].content[0] {
ContentBlock::Text { text } => assert_eq!(text, "<brief>"),
_ => panic!("first block should be the brief"),
}
}
#[test]
fn stops_when_only_two_messages_remain() {
let t = ApproxTokenizer;
let c = DefaultHistoryCompressor;
let mut msgs = vec![
Message::user_text("a".repeat(1000)),
Message::assistant_text("b".repeat(1000)),
];
c.compress(&mut msgs, 1, &t);
assert!(!msgs.is_empty());
}
#[test]
fn trim_preserves_tool_use_tool_result_pairing() {
let t = ApproxTokenizer;
let c = DefaultHistoryCompressor;
fn tool_use_msg(id: &str, name: &str) -> Message {
Message {
role: Role::Assistant,
content: vec![ContentBlock::ToolUse {
id: id.into(),
name: name.into(),
input: serde_json::json!({"path": "x".repeat(2000)}),
}],
}
}
fn tool_result_msg(id: &str, body: &str) -> Message {
Message {
role: Role::User,
content: vec![ContentBlock::ToolResult {
tool_use_id: id.into(),
content: body.into(),
is_error: Some(false),
}],
}
}
let mut msgs = vec![
Message::user_text("<brief>"),
Message::user_text("turn1 user"),
tool_use_msg("t1", "file_read"),
tool_result_msg("t1", &"r".repeat(4000)),
Message::assistant_text("turn1 done"),
Message::user_text("turn2 user"),
tool_use_msg("t2", "file_write"),
tool_result_msg("t2", &"r".repeat(4000)),
Message::assistant_text("turn2 done"),
Message::user_text("turn3 user"),
];
c.compress(&mut msgs, 50, &t);
let mut tool_use_ids = std::collections::BTreeSet::new();
let mut tool_result_ids = std::collections::BTreeSet::new();
for m in &msgs {
for b in &m.content {
match b {
ContentBlock::ToolUse { id, .. } => {
tool_use_ids.insert(id.clone());
}
ContentBlock::ToolResult { tool_use_id, .. } => {
tool_result_ids.insert(tool_use_id.clone());
}
_ => {}
}
}
}
assert_eq!(
tool_use_ids, tool_result_ids,
"tool_use ids and tool_result ids must remain in sync after trim"
);
match &msgs.last().unwrap().content[0] {
ContentBlock::Text { text } => assert_eq!(text, "turn3 user"),
_ => panic!("expected the live tail to be preserved"),
}
}
#[test]
fn tool_result_with_rule_note_is_never_stubbed() {
let t = ApproxTokenizer;
let c = DefaultHistoryCompressor;
let big = "x".repeat(4000);
let mut msgs = vec![
Message::user_text("hi"),
tool_result("t1", &big),
tool_result(
"t2",
&format!("{big}\n\n[stream rule `r` reminder]\nnote body"),
),
Message::assistant_text("ok"),
];
let _ = c.compress(&mut msgs, 50, &t);
match &msgs[1].content[0] {
ContentBlock::ToolResult { content, .. } => {
assert_eq!(content, COMPRESSED_TOOL_RESULT_STUB)
}
_ => panic!("expected tool result"),
}
match &msgs[2].content[0] {
ContentBlock::ToolResult { content, .. } => {
assert!(
content.contains("[stream rule `r` reminder]"),
"sentinel-bearing tool result must not be stubbed"
)
}
_ => panic!("expected tool result"),
}
}
#[test]
fn turn_with_rule_injection_survives_head_trim() {
let t = ApproxTokenizer;
let c = DefaultHistoryCompressor;
let mut msgs = vec![
Message::user_text("<brief>"),
Message::user_text("do the thing\n\n[stream rule `r` fired]\n\nrule body"),
Message::assistant_text("a".repeat(4000)),
Message::user_text("b".repeat(4000)),
Message::assistant_text("c".repeat(4000)),
Message::user_text("d".repeat(4000)),
Message::assistant_text("e".repeat(4000)),
Message::user_text("latest question"),
];
c.compress(&mut msgs, 50, &t);
assert!(
msgs.iter()
.any(|m| m.text_content().contains("[stream rule `r` fired]")),
"the turn containing the rule injection must survive compression"
);
assert_eq!(msgs.last().unwrap().text_content(), "latest question");
}
#[test]
fn tool_result_with_appended_text_is_never_an_anchor() {
let t = ApproxTokenizer;
let c = DefaultHistoryCompressor;
fn tool_use_msg(id: &str, name: &str) -> Message {
Message {
role: Role::Assistant,
content: vec![ContentBlock::ToolUse {
id: id.into(),
name: name.into(),
input: serde_json::json!({"path": "x".repeat(2000)}),
}],
}
}
let mut msgs = vec![
Message::user_text("<brief>"),
Message::user_text("turn1 user"),
tool_use_msg("t1", "file_read"),
Message {
role: Role::User,
content: vec![
ContentBlock::ToolResult {
tool_use_id: "t1".into(),
content: "r".repeat(4000),
is_error: Some(false),
},
ContentBlock::Text {
text: "\n\n[stream rule `r` fired]\n\nbody".into(),
},
],
},
Message::assistant_text("turn1 done"),
Message::user_text("b".repeat(4000)),
Message::assistant_text("c".repeat(4000)),
Message::user_text("d".repeat(4000)),
Message::assistant_text("e".repeat(4000)),
Message::user_text("latest question"),
];
c.compress(&mut msgs, 50, &t);
let mut tool_use_ids = std::collections::BTreeSet::new();
let mut tool_result_ids = std::collections::BTreeSet::new();
for m in &msgs {
for b in &m.content {
match b {
ContentBlock::ToolUse { id, .. } => {
tool_use_ids.insert(id.clone());
}
ContentBlock::ToolResult { tool_use_id, .. } => {
tool_result_ids.insert(tool_use_id.clone());
}
_ => {}
}
}
}
assert_eq!(
tool_use_ids, tool_result_ids,
"tool_use ids and tool_result ids must remain in sync after trim"
);
assert!(
msgs.iter()
.any(|m| m.text_content().contains("[stream rule `r` fired]")),
"the rule injection must survive compression"
);
}
}