use crate::types::*;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
pub fn estimate_tokens(text: &str) -> usize {
text.len().div_ceil(4)
}
pub fn message_tokens(msg: &AgentMessage) -> usize {
match msg {
AgentMessage::Llm(m) => match m {
Message::User { content, .. } => content_tokens(content) + 4,
Message::Assistant { content, .. } => content_tokens(content) + 4,
Message::ToolResult {
content, tool_name, ..
} => content_tokens(content) + estimate_tokens(tool_name) + 8,
},
AgentMessage::Extension(ext) => estimate_tokens(&ext.data.to_string()) + 4,
}
}
fn content_tokens(content: &[Content]) -> usize {
content
.iter()
.map(|c| match c {
Content::Text { text } => estimate_tokens(text),
Content::Image { data, .. } => {
let raw_bytes = data.len() * 3 / 4;
(raw_bytes / 750).clamp(85, 16_000)
}
Content::Thinking { thinking, .. } => estimate_tokens(thinking),
Content::ToolCall {
name, arguments, ..
} => estimate_tokens(name) + estimate_tokens(&arguments.to_string()) + 8,
})
.sum()
}
pub fn total_tokens(messages: &[AgentMessage]) -> usize {
messages.iter().map(message_tokens).sum()
}
pub struct ContextTracker {
last_usage_tokens: Option<usize>,
last_usage_index: Option<usize>,
}
impl ContextTracker {
pub fn new() -> Self {
Self {
last_usage_tokens: None,
last_usage_index: None,
}
}
pub fn record_usage(&mut self, usage: &Usage, message_index: usize) {
let total = usage.input + usage.output + usage.cache_read + usage.cache_write;
if total > 0 {
self.last_usage_tokens = Some(total as usize);
self.last_usage_index = Some(message_index);
}
}
pub fn estimate_context_tokens(&self, messages: &[AgentMessage]) -> usize {
match (self.last_usage_tokens, self.last_usage_index) {
(Some(usage_tokens), Some(idx)) if idx < messages.len() => {
let trailing: usize = messages[idx + 1..].iter().map(message_tokens).sum();
usage_tokens + trailing
}
_ => total_tokens(messages),
}
}
pub fn reset(&mut self) {
self.last_usage_tokens = None;
self.last_usage_index = None;
}
}
impl Default for ContextTracker {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ContextConfig {
pub max_context_tokens: usize,
pub system_prompt_tokens: usize,
pub keep_recent: usize,
pub keep_first: usize,
pub tool_output_max_lines: usize,
#[serde(default)]
pub tool_output_max_lines_overrides: HashMap<String, usize>,
pub compact_target_ratio: f32,
#[serde(default = "default_headroom_turns")]
pub compact_headroom_turns: Option<usize>,
pub truncate_tool_output_on_append: bool,
}
impl Default for ContextConfig {
fn default() -> Self {
Self {
max_context_tokens: 100_000,
system_prompt_tokens: 4_000,
keep_recent: 10,
keep_first: 2,
tool_output_max_lines: 200,
tool_output_max_lines_overrides: default_tool_output_overrides(),
compact_target_ratio: 0.7,
compact_headroom_turns: default_headroom_turns(),
truncate_tool_output_on_append: true,
}
}
}
pub const MIN_HEADROOM_RATIO: f32 = 0.15;
fn default_headroom_turns() -> Option<usize> {
Some(30)
}
fn default_tool_output_overrides() -> HashMap<String, usize> {
HashMap::from([
("read_file".to_string(), usize::MAX),
("shared_state".to_string(), usize::MAX),
])
}
impl ContextConfig {
pub fn from_context_window(context_window: u32) -> Self {
let max_context_tokens = (context_window as usize) * 80 / 100;
Self {
max_context_tokens,
..Default::default()
}
}
pub fn effective_target_ratio(&self, growth_tokens_per_turn: f64) -> f32 {
let Some(turns) = self.compact_headroom_turns else {
return self.compact_target_ratio;
};
if turns == 0 || !growth_tokens_per_turn.is_finite() || growth_tokens_per_turn <= 0.0 {
return self.compact_target_ratio;
}
let budget = self
.max_context_tokens
.saturating_sub(self.system_prompt_tokens) as f64;
if budget <= 0.0 {
return self.compact_target_ratio;
}
let target = budget - (turns as f64) * growth_tokens_per_turn;
let derived = (target / budget) as f32;
derived
.min(self.compact_target_ratio)
.max(MIN_HEADROOM_RATIO)
}
pub fn max_lines_for(&self, tool_name: &str) -> usize {
self.tool_output_max_lines_overrides
.get(tool_name)
.copied()
.unwrap_or(self.tool_output_max_lines)
}
fn compaction_target(&self, budget: usize) -> usize {
let ratio = if self.compact_target_ratio.is_finite() {
self.compact_target_ratio.clamp(0.05, 1.0)
} else {
1.0
};
((budget as f64) * (ratio as f64)) as usize
}
}
pub trait CompactionStrategy: Send + Sync {
fn compact(&self, messages: Vec<AgentMessage>, config: &ContextConfig) -> Vec<AgentMessage>;
}
pub struct DefaultCompaction;
impl CompactionStrategy for DefaultCompaction {
fn compact(&self, messages: Vec<AgentMessage>, config: &ContextConfig) -> Vec<AgentMessage> {
compact_messages(messages, config)
}
}
pub fn compact_messages(messages: Vec<AgentMessage>, config: &ContextConfig) -> Vec<AgentMessage> {
let budget = config
.max_context_tokens
.saturating_sub(config.system_prompt_tokens);
if total_tokens(&messages) <= budget {
return messages;
}
let target = config.compaction_target(budget);
let before = total_tokens(&messages);
let compacted = level1_truncate_tool_outputs(&messages, config);
let after_l1 = total_tokens(&compacted);
if after_l1 < before {
tracing::debug!(
"compaction level 1: tool outputs truncated, {} -> {} tokens",
before,
after_l1
);
}
if after_l1 <= budget {
return compacted;
}
let before_l2 = compacted.len();
let compacted = level2_summarize_old_turns(&compacted, config.keep_recent);
if compacted.len() != before_l2 {
tracing::debug!(
"compaction level 2: old turns summarized, {} -> {} messages ({} tokens)",
before_l2,
compacted.len(),
total_tokens(&compacted)
);
}
if total_tokens(&compacted) <= target {
return compacted;
}
level3_drop_middle(&compacted, config, target)
}
fn level1_truncate_tool_outputs(
messages: &[AgentMessage],
config: &ContextConfig,
) -> Vec<AgentMessage> {
messages
.iter()
.map(|msg| truncate_tool_output(msg.clone(), config))
.collect()
}
pub fn message_text(msg: &AgentMessage) -> String {
block_texts(msg)
.into_iter()
.map(|(_, s)| s)
.collect::<Vec<_>>()
.join("\n")
}
pub const TOOL_OUTPUT_KEY_PREFIX: &str = "tool-out-";
pub fn tool_output_key(tool_call_id: &str, full_output: &str) -> String {
format!(
"{TOOL_OUTPUT_KEY_PREFIX}{tool_call_id}-{:016x}",
fnv1a(full_output.as_bytes())
)
}
fn fnv1a(bytes: &[u8]) -> u64 {
let mut hash: u64 = 0xcbf2_9ce4_8422_2325;
for byte in bytes {
hash ^= u64::from(*byte);
hash = hash.wrapping_mul(0x0000_0100_0000_01b3);
}
hash
}
pub fn block_key(base: &str, block: usize) -> String {
format!("{base}-b{block}")
}
pub fn block_texts(msg: &AgentMessage) -> Vec<(usize, String)> {
match msg {
AgentMessage::Llm(Message::ToolResult { content, .. }) => content
.iter()
.enumerate()
.filter_map(|(i, c)| match c {
Content::Text { text } => Some((i, text.clone())),
_ => None,
})
.collect(),
_ => Vec::new(),
}
}
pub fn truncate_tool_output_keyed(
msg: AgentMessage,
config: &ContextConfig,
base_key: Option<&str>,
) -> (AgentMessage, Vec<usize>) {
match msg {
AgentMessage::Llm(Message::ToolResult {
tool_call_id,
tool_name,
content,
is_error,
timestamp,
}) => {
let max_lines = config.max_lines_for(&tool_name);
let mut marked = Vec::new();
let truncated_content: Vec<Content> = content
.into_iter()
.enumerate()
.map(|(i, c)| match c {
Content::Text { text } => {
let key = base_key.map(|b| block_key(b, i));
let (out, emitted) =
truncate_text_head_tail(&text, max_lines, key.as_deref());
if emitted {
marked.push(i);
}
Content::Text { text: out }
}
other => other,
})
.collect();
(
AgentMessage::Llm(Message::ToolResult {
tool_call_id,
tool_name,
content: truncated_content,
is_error,
timestamp,
}),
marked,
)
}
other => (other, Vec::new()),
}
}
pub fn truncate_tool_output(msg: AgentMessage, config: &ContextConfig) -> AgentMessage {
truncate_tool_output_keyed(msg, config, None).0
}
const TRUNCATION_MARKER_LINES: usize = 3;
fn truncate_text_head_tail(text: &str, max_lines: usize, key: Option<&str>) -> (String, bool) {
let lines: Vec<&str> = text.lines().collect();
if lines.len() <= max_lines {
return (text.to_string(), false);
}
if max_lines <= TRUNCATION_MARKER_LINES + 1 {
return (lines[..max_lines].join("\n"), false);
}
let keep = max_lines - TRUNCATION_MARKER_LINES;
let head = keep / 2;
let tail = keep - head;
let omitted = lines.len() - head - tail;
let marker = match key {
Some(k) => {
format!("[... {omitted} lines truncated — full output: shared_state get \"{k}\" ...]")
}
None => format!("[... {omitted} lines truncated ...]"),
};
let mut result = lines[..head].join("\n");
result.push_str(&format!("\n\n{marker}\n\n"));
result.push_str(&lines[lines.len() - tail..].join("\n"));
(result, true)
}
fn level2_summarize_old_turns(messages: &[AgentMessage], keep_recent: usize) -> Vec<AgentMessage> {
let len = messages.len();
if len <= keep_recent {
return messages.to_vec();
}
let boundary = safe_turn_start(messages, len - keep_recent);
if boundary == 0 {
return messages.to_vec();
}
let mut result = Vec::new();
let mut i = 0;
while i < boundary {
let msg = &messages[i];
match msg {
AgentMessage::Llm(Message::Assistant {
content, timestamp, ..
}) => {
let text_parts: Vec<&str> = content
.iter()
.filter_map(|c| match c {
Content::Text { text } => {
if text.len() > 200 {
None } else {
Some(text.as_str())
}
}
_ => None,
})
.collect();
let tool_count = content
.iter()
.filter(|c| matches!(c, Content::ToolCall { .. }))
.count();
let summary = if !text_parts.is_empty() {
text_parts.join(" ")
} else if tool_count > 0 {
format!("[Assistant used {} tool(s)]", tool_count)
} else {
"[Assistant response]".into()
};
result.push(AgentMessage::Llm(Message::User {
content: vec![Content::Text {
text: format!("[Summary] {}", summary),
}],
timestamp: *timestamp,
}));
i += 1;
while i < boundary {
if let AgentMessage::Llm(Message::ToolResult { .. }) = &messages[i] {
i += 1;
} else {
break;
}
}
continue;
}
AgentMessage::Llm(Message::ToolResult { .. }) => {
i += 1;
continue;
}
other => {
result.push(other.clone());
}
}
i += 1;
}
result.extend_from_slice(&messages[boundary..]);
result
}
pub(crate) const COMPACTION_MARKER: &str =
"[Context compacted: earlier messages removed to fit the context window]";
fn compaction_marker(timestamp: u64) -> AgentMessage {
AgentMessage::Llm(Message::User {
content: vec![Content::Text {
text: COMPACTION_MARKER.into(),
}],
timestamp,
})
}
pub(crate) fn message_timestamp(msg: &AgentMessage) -> u64 {
match msg {
AgentMessage::Llm(
Message::User { timestamp, .. }
| Message::Assistant { timestamp, .. }
| Message::ToolResult { timestamp, .. },
) => *timestamp,
AgentMessage::Extension(_) => 0,
}
}
fn opens_tool_calls(msg: &AgentMessage) -> bool {
matches!(
msg,
AgentMessage::Llm(Message::Assistant { content, .. })
if content.iter().any(|c| matches!(c, Content::ToolCall { .. }))
)
}
fn is_tool_result(msg: &AgentMessage) -> bool {
matches!(msg, AgentMessage::Llm(Message::ToolResult { .. }))
}
pub(crate) fn safe_head_end(messages: &[AgentMessage], mut end: usize) -> usize {
while end > 0 && opens_tool_calls(&messages[end - 1]) {
end -= 1;
}
end
}
fn safe_tail_start(messages: &[AgentMessage], mut start: usize) -> usize {
while start < messages.len() && is_tool_result(&messages[start]) {
start += 1;
}
start
}
pub(crate) fn safe_turn_start(messages: &[AgentMessage], mut start: usize) -> usize {
while start > 0 && start < messages.len() && is_tool_result(&messages[start]) {
start -= 1;
}
start
}
fn level3_drop_middle(
messages: &[AgentMessage],
config: &ContextConfig,
target: usize,
) -> Vec<AgentMessage> {
let len = messages.len();
let head_end = safe_head_end(messages, config.keep_first.min(len));
let max_cut = len.saturating_sub(config.keep_recent);
if head_end >= max_cut {
return keep_within_budget(messages, target);
}
let mut suffix = vec![0usize; len + 1];
for i in (0..len).rev() {
suffix[i] = suffix[i + 1] + message_tokens(&messages[i]);
}
let head_tokens: usize = messages[..head_end].iter().map(message_tokens).sum();
let marker_tokens = message_tokens(&compaction_marker(0));
let fixed = head_tokens + marker_tokens;
let cut = suffix[head_end..=max_cut]
.iter()
.position(|&tokens| fixed + tokens <= target)
.map_or(max_cut, |offset| head_end + offset);
let cut = safe_tail_start(messages, cut).max(head_end);
let removed = cut - head_end;
if removed == 0 {
return messages.to_vec();
}
tracing::debug!(
"compaction level 3: dropping {} of {} messages (target {} tokens)",
removed,
len,
target
);
let mut result = messages[..head_end].to_vec();
result.push(compaction_marker(message_timestamp(&messages[head_end])));
result.extend_from_slice(&messages[cut..]);
if total_tokens(&result) > target {
return keep_within_budget(&result, target);
}
result
}
fn keep_within_budget(messages: &[AgentMessage], budget: usize) -> Vec<AgentMessage> {
if messages.is_empty() {
return Vec::new();
}
let mut kept = 0usize;
let mut remaining = budget;
for msg in messages.iter().rev() {
let tokens = message_tokens(msg);
if tokens > remaining {
break;
}
remaining -= tokens;
kept += 1;
}
let start = safe_tail_start(messages, messages.len() - kept);
let mut result = messages[start..].to_vec();
if start > 0 {
tracing::debug!(
"compaction: keeping {} of {} messages within {} tokens",
result.len(),
messages.len(),
budget
);
result.insert(0, compaction_marker(message_timestamp(&messages[0])));
}
result
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[non_exhaustive]
pub struct ExecutionLimits {
pub max_turns: usize,
pub max_total_tokens: usize,
pub max_duration: std::time::Duration,
pub max_consecutive_identical_tool_calls: Option<usize>,
}
impl Default for ExecutionLimits {
fn default() -> Self {
Self {
max_turns: 50,
max_total_tokens: 1_000_000,
max_duration: std::time::Duration::from_secs(600),
max_consecutive_identical_tool_calls: Some(3),
}
}
}
impl ExecutionLimits {
pub fn with_max_turns(mut self, turns: usize) -> Self {
self.max_turns = turns;
self
}
pub fn with_max_total_tokens(mut self, tokens: usize) -> Self {
self.max_total_tokens = tokens;
self
}
pub fn with_max_duration(mut self, duration: std::time::Duration) -> Self {
self.max_duration = duration;
self
}
pub fn with_max_consecutive_identical_tool_calls(mut self, calls: Option<usize>) -> Self {
self.max_consecutive_identical_tool_calls = calls;
self
}
}
const MAX_STEERED_SIGNATURES: usize = 32;
fn signature_hash(name: &str, args: &serde_json::Value) -> u64 {
let mut h = fnv1a(name.as_bytes());
h ^= fnv1a(args.to_string().as_bytes());
h
}
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord)]
#[must_use = "a loop verdict that is not acted on silently disables loop detection, \
while still advancing the tracker's state"]
pub enum LoopVerdict {
Continue,
Steer {
tool_name: String,
repetitions: usize,
},
Abort {
tool_name: String,
repetitions: usize,
},
}
pub struct ExecutionTracker {
pub limits: ExecutionLimits,
pub turns: usize,
pub tokens_used: usize,
pub started_at: std::time::Instant,
last_signature: Option<(String, serde_json::Value)>,
consecutive: usize,
steered: Vec<u64>,
}
impl ExecutionTracker {
pub fn new(limits: ExecutionLimits) -> Self {
Self {
limits,
turns: 0,
tokens_used: 0,
started_at: std::time::Instant::now(),
last_signature: None,
consecutive: 0,
steered: Vec::new(),
}
}
pub fn record_tool_calls(&mut self, calls: &[(String, serde_json::Value)]) -> LoopVerdict {
let Some(threshold) = self.limits.max_consecutive_identical_tool_calls else {
return LoopVerdict::Continue;
};
if threshold == 0 {
return LoopVerdict::Continue;
}
let mut verdict = LoopVerdict::Continue;
for (name, args) in calls {
let sig = (name.clone(), args.clone());
match &self.last_signature {
Some(prev) if *prev == sig => self.consecutive += 1,
_ => {
self.last_signature = Some(sig.clone());
self.consecutive = 1;
}
}
if self.consecutive >= threshold {
let repetitions = self.consecutive;
let sig_hash = signature_hash(name, args);
let this_call = if self.steered.contains(&sig_hash) {
LoopVerdict::Abort {
tool_name: name.clone(),
repetitions,
}
} else {
if self.steered.len() >= MAX_STEERED_SIGNATURES {
self.steered.remove(0);
}
self.steered.push(sig_hash);
self.consecutive = 0;
LoopVerdict::Steer {
tool_name: name.clone(),
repetitions,
}
};
if this_call > verdict {
verdict = this_call;
}
}
}
verdict
}
pub fn record_turn(&mut self, tokens: usize) {
self.turns += 1;
self.tokens_used += tokens;
}
pub fn check_limits(&self) -> Option<String> {
if self.turns >= self.limits.max_turns {
return Some(format!(
"Max turns reached ({}/{})",
self.turns, self.limits.max_turns
));
}
if self.tokens_used >= self.limits.max_total_tokens {
return Some(format!(
"Max tokens reached ({}/{})",
self.tokens_used, self.limits.max_total_tokens
));
}
let elapsed = self.started_at.elapsed();
if elapsed >= self.limits.max_duration {
return Some(format!(
"Max duration reached ({:.0}s/{:.0}s)",
elapsed.as_secs_f64(),
self.limits.max_duration.as_secs_f64()
));
}
None
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_estimate_tokens() {
assert!(estimate_tokens("hello world") > 0);
assert!(estimate_tokens("hello world") < 10);
assert_eq!(estimate_tokens(""), 0);
}
#[test]
fn test_context_config_from_context_window() {
let config = ContextConfig::from_context_window(200_000);
assert_eq!(config.max_context_tokens, 160_000); assert_eq!(config.system_prompt_tokens, 4_000); assert_eq!(config.keep_recent, 10);
let config = ContextConfig::from_context_window(1_000_000);
assert_eq!(config.max_context_tokens, 800_000);
let config = ContextConfig::from_context_window(128_000);
assert_eq!(config.max_context_tokens, 102_400); }
#[test]
fn test_truncate_head_tail() {
let text = (1..=100)
.map(|i| format!("line {}", i))
.collect::<Vec<_>>()
.join("\n");
let result = truncate_text_head_tail(&text, 10, None).0;
assert!(result.contains("line 1")); assert!(result.contains("line 100")); assert!(result.contains("truncated"));
assert!(!result.contains("line 50"));
assert_eq!(result.lines().count(), 10);
assert_eq!(truncate_text_head_tail(&result, 10, None).0, result);
}
#[test]
fn test_truncate_is_idempotent_and_keeps_an_honest_count() {
let text = (1..=1000)
.map(|i| format!("line {}", i))
.collect::<Vec<_>>()
.join("\n");
let once = truncate_text_head_tail(&text, 50, None).0;
assert!(once.contains("[... 953 lines truncated ...]"));
assert_eq!(once.lines().count(), 50);
let twice = truncate_text_head_tail(&once, 50, None).0;
assert_eq!(twice, once);
assert_eq!(truncate_text_head_tail(&twice, 50, None).0, once);
}
#[test]
fn test_truncate_degenerate_line_budget() {
let text = (1..=100)
.map(|i| format!("line {}", i))
.collect::<Vec<_>>()
.join("\n");
for max_lines in [1usize, 2, 3, 4] {
let result = truncate_text_head_tail(&text, max_lines, None).0;
assert_eq!(result.lines().count(), max_lines);
assert_eq!(truncate_text_head_tail(&result, max_lines, None).0, result);
}
}
fn tool_turn(i: usize, output_lines: usize) -> Vec<AgentMessage> {
let output = (0..output_lines)
.map(|n| format!("turn {} line {} {}", i, n, "x".repeat(40)))
.collect::<Vec<_>>()
.join("\n");
vec![
AgentMessage::Llm(Message::Assistant {
content: vec![Content::tool_call(
format!("tc-{}", i),
"bash",
serde_json::json!({ "command": format!("ls {}", i) }),
)],
stop_reason: StopReason::ToolUse,
model: "test".into(),
provider: "test".into(),
usage: Usage::default(),
timestamp: i as u64,
error_message: None,
}),
AgentMessage::Llm(Message::ToolResult {
tool_call_id: format!("tc-{}", i),
tool_name: "bash".into(),
content: vec![Content::Text { text: output }],
is_error: false,
timestamp: i as u64,
}),
]
}
fn tool_session(turns: usize, output_lines: usize) -> Vec<AgentMessage> {
let mut messages = vec![AgentMessage::Llm(Message::user("start"))];
for i in 0..turns {
messages.extend(tool_turn(i, output_lines));
}
messages
}
#[test]
fn test_compact_target_ratio_leaves_headroom() {
let messages = tool_session(40, 120);
let config = ContextConfig {
max_context_tokens: 8_000,
system_prompt_tokens: 0,
compact_target_ratio: 0.5,
..Default::default()
};
assert!(total_tokens(&messages) > config.max_context_tokens);
let result = compact_messages(messages, &config);
assert!(
total_tokens(&result) <= 4_000,
"compacted to {} tokens, expected <= 4000 (50% of budget)",
total_tokens(&result)
);
}
#[test]
fn test_headroom_policy_derives_the_ratio_from_growth() {
let config = ContextConfig {
max_context_tokens: 100_000,
system_prompt_tokens: 0,
compact_target_ratio: 0.7,
compact_headroom_turns: Some(30),
..Default::default()
};
assert!((config.effective_target_ratio(1000.0) - 0.7).abs() < 1e-6);
assert!((config.effective_target_ratio(2000.0) - 0.4).abs() < 1e-6);
assert!((config.effective_target_ratio(100.0) - 0.7).abs() < 1e-6);
assert_eq!(
config.effective_target_ratio(1_000_000.0),
MIN_HEADROOM_RATIO
);
}
#[test]
fn test_headroom_policy_is_inert_without_a_growth_estimate() {
let config = ContextConfig {
max_context_tokens: 100_000,
system_prompt_tokens: 0,
compact_target_ratio: 0.6,
compact_headroom_turns: Some(30),
..Default::default()
};
assert_eq!(config.effective_target_ratio(0.0), 0.6);
assert_eq!(config.effective_target_ratio(-5.0), 0.6);
assert_eq!(config.effective_target_ratio(f64::NAN), 0.6);
assert_eq!(
ContextConfig {
compact_headroom_turns: Some(0),
..config.clone()
}
.effective_target_ratio(1000.0),
0.6
);
assert_eq!(
ContextConfig {
compact_headroom_turns: None,
..config
}
.effective_target_ratio(9999.0),
0.6
);
}
#[test]
fn test_ratio_of_one_restores_compact_to_just_fit() {
let config = ContextConfig {
max_context_tokens: 10_000,
system_prompt_tokens: 0,
compact_target_ratio: 1.0,
..Default::default()
};
assert_eq!(config.compaction_target(10_000), 10_000);
}
#[test]
fn test_compaction_target_rejects_nonsense_ratios() {
let mut config = ContextConfig::default();
for (ratio, expected) in [
(0.0f32, 500usize), (-1.0, 500), (2.0, 10_000), (f32::NAN, 10_000), (f32::INFINITY, 10_000), ] {
config.compact_target_ratio = ratio;
assert_eq!(
config.compaction_target(10_000),
expected,
"ratio {}",
ratio
);
}
}
#[test]
fn test_level3_drops_only_what_the_target_requires() {
let messages = tool_session(40, 20);
let config = ContextConfig {
max_context_tokens: total_tokens(&messages) * 3 / 4,
system_prompt_tokens: 0,
compact_target_ratio: 0.9,
..Default::default()
};
let result = compact_messages(messages.clone(), &config);
assert!(total_tokens(&result) <= config.max_context_tokens);
assert!(
result.len() > 13,
"kept only {} messages; Level 3 is still collapsing history it did not need to",
result.len()
);
}
#[test]
fn test_compaction_never_orphans_tool_calls() {
for (max_context_tokens, keep_recent, keep_first) in [
(400usize, 10usize, 2usize),
(1_000, 10, 2),
(4_000, 10, 2),
(20_000, 10, 2),
(4_000, 9, 1),
(4_000, 11, 3),
(8_000, 7, 4),
(8_000, 12, 5),
(12_000, 13, 0),
] {
let messages = tool_session(30, 60);
let config = ContextConfig {
max_context_tokens,
system_prompt_tokens: 0,
keep_recent,
keep_first,
..Default::default()
};
let result = compact_messages(messages, &config);
let mut open: Vec<String> = Vec::new();
for msg in &result {
match msg {
AgentMessage::Llm(Message::Assistant { content, .. }) => {
open = content
.iter()
.filter_map(|c| match c {
Content::ToolCall { id, .. } => Some(id.clone()),
_ => None,
})
.collect();
}
AgentMessage::Llm(Message::ToolResult { tool_call_id, .. }) => {
let answered = open.iter().position(|id| id == tool_call_id);
assert!(
answered.is_some(),
"orphaned tool result {} at budget {}",
tool_call_id,
max_context_tokens
);
open.remove(answered.unwrap());
}
_ => {
assert!(
open.is_empty(),
"unanswered tool call {:?} at budget {}",
open,
max_context_tokens
);
}
}
}
assert!(
open.is_empty(),
"history ends on an unanswered tool call at budget {}",
max_context_tokens
);
}
}
#[test]
fn test_compaction_marker_carries_no_drifting_count() {
let messages = tool_session(40, 120);
let config = ContextConfig {
max_context_tokens: 2_000,
system_prompt_tokens: 0,
..Default::default()
};
let result = compact_messages(messages, &config);
let marker = result
.iter()
.find_map(|m| match m {
AgentMessage::Llm(Message::User { content, .. }) => match content.first() {
Some(Content::Text { text }) if text.starts_with("[Context compacted") => {
Some(text.clone())
}
_ => None,
},
_ => None,
})
.expect("expected a compaction marker");
assert_eq!(marker, COMPACTION_MARKER);
assert!(!marker.chars().any(|c| c.is_ascii_digit()));
}
#[test]
fn test_truncate_tool_output_helper() {
let big = (0..500)
.map(|i| format!("l{}", i))
.collect::<Vec<_>>()
.join("\n");
let msg = AgentMessage::Llm(Message::ToolResult {
tool_call_id: "tc-1".into(),
tool_name: "bash".into(),
content: vec![Content::Text { text: big }],
is_error: false,
timestamp: 7,
});
let config = ContextConfig {
tool_output_max_lines: 50,
..Default::default()
};
let once = truncate_tool_output(msg, &config);
let twice = truncate_tool_output(once.clone(), &config);
assert_eq!(once, twice, "on-append truncation must be idempotent");
match &once {
AgentMessage::Llm(Message::ToolResult {
content, timestamp, ..
}) => {
assert_eq!(*timestamp, 7, "timestamp must survive truncation");
let Content::Text { text } = &content[0] else {
panic!("expected text")
};
assert_eq!(text.lines().count(), 50);
}
_ => panic!("expected a tool result"),
}
let user = AgentMessage::Llm(Message::user("hello"));
assert_eq!(truncate_tool_output(user.clone(), &config), user);
}
#[test]
fn test_level1_truncation() {
let big_output = (1..=200)
.map(|i| format!("output line {}", i))
.collect::<Vec<_>>()
.join("\n");
let messages = vec![
AgentMessage::Llm(Message::user("do something")),
AgentMessage::Llm(Message::ToolResult {
tool_call_id: "tc-1".into(),
tool_name: "bash".into(),
content: vec![Content::Text { text: big_output }],
is_error: false,
timestamp: 0,
}),
];
let compacted = level1_truncate_tool_outputs(
&messages,
&ContextConfig {
tool_output_max_lines: 20,
..Default::default()
},
);
let tool_msg = &compacted[1];
if let AgentMessage::Llm(Message::ToolResult { content, .. }) = tool_msg {
if let Content::Text { text } = &content[0] {
assert!(text.contains("truncated"));
assert!(text.contains("output line 1")); assert!(text.contains("output line 200")); assert!(text.lines().count() < 50);
} else {
panic!("expected text content");
}
} else {
panic!("expected tool result");
}
}
#[test]
fn test_compact_within_budget() {
let messages = vec![
AgentMessage::Llm(Message::user("Hello")),
AgentMessage::Llm(Message::user("World")),
];
let config = ContextConfig::default();
let result = compact_messages(messages.clone(), &config);
assert_eq!(result.len(), 2);
}
#[test]
fn test_compact_drops_middle_when_needed() {
let mut messages = Vec::new();
for i in 0..100 {
messages.push(AgentMessage::Llm(Message::user(format!(
"Message {} {}",
i,
"x".repeat(200)
))));
}
let config = ContextConfig {
max_context_tokens: 500,
system_prompt_tokens: 100,
keep_recent: 5,
keep_first: 2,
tool_output_max_lines: 20,
..Default::default()
};
let result = compact_messages(messages, &config);
assert!(result.len() < 100);
assert!(result.len() >= 2);
}
#[test]
fn test_context_tracker_no_usage() {
let tracker = ContextTracker::new();
let messages = vec![
AgentMessage::Llm(Message::user("Hello")),
AgentMessage::Llm(Message::user("World")),
];
let tokens = tracker.estimate_context_tokens(&messages);
assert!(tokens > 0);
assert_eq!(tokens, total_tokens(&messages));
}
#[test]
fn test_context_tracker_with_usage() {
let mut tracker = ContextTracker::new();
let messages = vec![
AgentMessage::Llm(Message::user("Hello")),
AgentMessage::Llm(Message::Assistant {
content: vec![Content::Text {
text: "Hi there!".into(),
}],
stop_reason: StopReason::Stop,
model: "test".into(),
provider: "test".into(),
usage: Usage {
input: 100,
output: 50,
..Default::default()
},
timestamp: 0,
error_message: None,
}),
AgentMessage::Llm(Message::user("Follow up question here")),
];
tracker.record_usage(
&Usage {
input: 100,
output: 50,
..Default::default()
},
1,
);
let tokens = tracker.estimate_context_tokens(&messages);
let trailing_estimate = message_tokens(&messages[2]);
assert_eq!(tokens, 150 + trailing_estimate);
}
#[test]
fn test_context_tracker_reset() {
let mut tracker = ContextTracker::new();
tracker.record_usage(
&Usage {
input: 1000,
output: 500,
..Default::default()
},
5,
);
tracker.reset();
let messages = vec![AgentMessage::Llm(Message::user("test"))];
assert_eq!(
tracker.estimate_context_tokens(&messages),
total_tokens(&messages)
);
}
#[test]
fn test_execution_limits() {
let limits = ExecutionLimits {
max_turns: 3,
max_total_tokens: 1000,
max_duration: std::time::Duration::from_secs(60),
..Default::default()
};
let mut tracker = ExecutionTracker::new(limits);
assert!(tracker.check_limits().is_none());
tracker.record_turn(100);
tracker.record_turn(100);
assert!(tracker.check_limits().is_none());
tracker.record_turn(100);
assert!(tracker.check_limits().is_some());
}
}
#[cfg(test)]
mod compaction_retention {
use super::*;
use crate::types::{Content, StopReason, Usage};
fn assistant_call(i: usize) -> AgentMessage {
AgentMessage::Llm(Message::Assistant {
content: vec![Content::ToolCall {
id: format!("tc-{i}"),
name: "fetch_record".into(),
arguments: serde_json::json!({"n": i}),
provider_metadata: None,
}],
stop_reason: StopReason::ToolUse,
model: "m".into(),
provider: "p".into(),
usage: Usage::default(),
timestamp: i as u64,
error_message: None,
})
}
fn tool_result(i: usize) -> AgentMessage {
AgentMessage::Llm(Message::ToolResult {
tool_call_id: format!("tc-{i}"),
tool_name: "fetch_record".into(),
content: vec![Content::Text {
text: "record line, nominal, no action required ".repeat(400),
}],
is_error: false,
timestamp: i as u64,
})
}
#[test]
fn compaction_keeps_a_usable_well_formed_transcript() {
let cfg = ContextConfig {
max_context_tokens: 30_000,
keep_recent: 6,
keep_first: 2,
..Default::default()
};
for pairs in [6usize, 11, 12, 20, 40] {
let mut msgs = vec![AgentMessage::Llm(Message::user("go"))];
for i in 0..pairs {
msgs.push(assistant_call(i));
msgs.push(tool_result(i));
}
let before = msgs.len();
let out = compact_messages(msgs, &cfg);
assert!(
!out.is_empty(),
"compaction of {before} messages emptied the transcript"
);
let mut pending: Vec<String> = Vec::new();
for m in &out {
let AgentMessage::Llm(msg) = m else { continue };
match msg {
Message::Assistant { content, .. } => {
pending = content
.iter()
.filter_map(|c| match c {
Content::ToolCall { id, .. } => Some(id.clone()),
_ => None,
})
.collect();
}
Message::ToolResult { tool_call_id, .. } => {
if let Some(at) = pending.iter().position(|p| p == tool_call_id) {
pending.remove(at);
}
}
Message::User { .. } => {}
}
}
assert!(
pending.is_empty(),
"compaction of {before} messages orphaned tool calls {pending:?}"
);
}
}
}