use async_openai::types::chat::{
ChatCompletionRequestMessage, ChatCompletionRequestUserMessage,
ChatCompletionRequestUserMessageContent,
};
use robit_ai::config::ContextConfig;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum TruncationAction {
NewSegment,
MergeSegments {
summaries: Vec<String>,
start_position: usize,
count: usize,
},
TruncateOnly,
}
#[derive(Debug)]
pub struct TruncationResult {
pub rounds_removed: usize,
pub messages_removed: usize,
pub removed_messages: Vec<ChatCompletionRequestMessage>,
pub insert_position: usize,
pub needs_compression: bool,
pub action: TruncationAction,
}
pub fn truncate_output(content: &str, max_lines: usize, max_bytes: usize) -> String {
let lines: Vec<&str> = content.lines().collect();
let total_lines = lines.len();
let total_bytes = content.len();
let line_truncated = total_lines > max_lines;
let byte_truncated = total_bytes > max_bytes;
if !line_truncated && !byte_truncated {
return content.to_string();
}
let mut output = String::new();
let mut byte_count = 0;
let mut displayed_lines = 0;
for (i, line) in lines.iter().enumerate() {
if i >= max_lines {
break;
}
let line_with_newline = if i < total_lines - 1 {
format!("{}\n", line)
} else {
line.to_string()
};
if byte_count + line_with_newline.len() > max_bytes {
break;
}
output.push_str(&line_with_newline);
byte_count += line_with_newline.len();
displayed_lines += 1;
}
if line_truncated {
output.push_str(&format!(
"\n... (Output truncated, {} lines total, showing first {}. Use offset/limit to read more)",
total_lines, displayed_lines
));
} else if byte_truncated {
output.push_str(&format!(
"\n... (Output truncated, {} bytes total, showing first {} bytes)",
total_bytes, byte_count
));
}
output
}
pub fn estimate_tokens(text: &str) -> usize {
if text.is_empty() {
return 0;
}
let mut ascii_alnum = 0usize;
let mut cjk = 0usize;
let mut whitespace = 0usize;
let mut other = 0usize;
for ch in text.chars() {
if ch.is_whitespace() {
whitespace += 1;
} else if ch.is_ascii_alphanumeric() {
ascii_alnum += 1;
} else {
let cp = ch as u32;
if (0x4E00..=0x9FFF).contains(&cp)
|| (0x3400..=0x4DBF).contains(&cp)
|| (0xF900..=0xFAFF).contains(&cp)
|| (0xFF00..=0xFFEF).contains(&cp)
|| (0x3000..=0x303F).contains(&cp)
|| (0x3040..=0x309F).contains(&cp)
|| (0x30A0..=0x30FF).contains(&cp)
|| (0xAC00..=0xD7AF).contains(&cp)
{
cjk += 1;
} else {
other += 1;
}
}
}
let ascii_tokens = (ascii_alnum as f64 / 3.5).ceil() as usize;
let cjk_tokens = (cjk as f64 / 1.5).ceil() as usize;
let whitespace_tokens = (whitespace as f64 / 10.0).ceil() as usize;
let other_tokens = other;
ascii_tokens + cjk_tokens + whitespace_tokens + other_tokens
}
pub fn estimate_messages_tokens(messages: &[ChatCompletionRequestMessage]) -> usize {
let mut total = 0;
for msg in messages {
total += 4;
total += estimate_message_content_tokens(msg);
}
total
}
pub fn estimate_messages_tokens_with_margin(
messages: &[ChatCompletionRequestMessage],
safety_margin: f32,
) -> usize {
let raw = estimate_messages_tokens(messages);
(raw as f32 * safety_margin).ceil() as usize
}
fn estimate_message_content_tokens(msg: &ChatCompletionRequestMessage) -> usize {
use async_openai::types::chat::ChatCompletionRequestUserMessageContentPart;
if let ChatCompletionRequestMessage::User(user_msg) = msg {
if let ChatCompletionRequestUserMessageContent::Array(parts) = &user_msg.content {
let mut total = 0;
for part in parts {
match part {
ChatCompletionRequestUserMessageContentPart::Text(t) => {
total += estimate_tokens(&t.text);
}
ChatCompletionRequestUserMessageContentPart::ImageUrl(_) => {
total += 2000;
}
_ => {
}
}
}
return total;
}
}
match serde_json::to_string(msg) {
Ok(json) => estimate_tokens(&json),
Err(_) => 0,
}
}
pub struct ContextManager {
pub max_tokens: usize,
pub reserve_ratio: f32,
pub truncation_ratio: f32,
pub min_keep_rounds: usize,
pub token_safety_margin: f32,
pub max_output_lines: usize,
pub max_output_bytes: usize,
pub compression_token_threshold: usize,
pub compression_enabled: bool,
pub max_tool_calls_per_turn: usize,
pub progressive_compression: bool,
pub rounds_per_summary: usize,
pub max_summary_segments: usize,
pub merge_count: usize,
pub max_merges_per_segment: usize,
}
impl ContextManager {
pub fn new(context_window: Option<u64>, config: Option<&ContextConfig>) -> Self {
let max_tokens = context_window.unwrap_or(65536) as usize;
let (
max_output_lines,
max_output_bytes,
reserve_ratio,
truncation_ratio,
min_keep_rounds,
token_safety_margin,
compression_token_threshold,
compression_enabled,
max_tool_calls_per_turn,
progressive_compression,
rounds_per_summary,
max_summary_segments,
merge_count,
max_merges_per_segment,
) = match config {
Some(c) => (
c.max_output_lines.unwrap_or(500),
c.max_output_bytes.unwrap_or(51200),
c.reserve_ratio.unwrap_or(0.2),
c.truncation_ratio.unwrap_or(0.7),
c.min_keep_rounds.unwrap_or(3),
c.token_safety_margin.unwrap_or(1.3),
c.compression_token_threshold.unwrap_or(5000),
c.compression_enabled.unwrap_or(true),
c.max_tool_calls_per_turn.unwrap_or(30),
c.progressive_compression.unwrap_or(true),
c.rounds_per_summary.unwrap_or(3),
c.max_summary_segments.unwrap_or(5),
c.merge_count.unwrap_or(2),
c.max_merges_per_segment.unwrap_or(2),
),
None => (500, 51200, 0.2, 0.7, 3, 1.3, 5000, true, 30, true, 3, 5, 2, 2),
};
Self {
max_tokens,
reserve_ratio,
truncation_ratio,
min_keep_rounds,
token_safety_margin,
max_output_lines,
max_output_bytes,
compression_token_threshold,
compression_enabled,
max_tool_calls_per_turn,
progressive_compression,
rounds_per_summary,
max_summary_segments,
merge_count,
max_merges_per_segment,
}
}
pub fn truncation_threshold(&self) -> usize {
(self.max_tokens as f32 * self.truncation_ratio) as usize
}
pub fn available_tokens(&self) -> usize {
(self.max_tokens as f32 * (1.0 - self.reserve_ratio)) as usize
}
pub fn truncate_tool_output(&self, content: &str) -> String {
truncate_output(content, self.max_output_lines, self.max_output_bytes)
}
pub fn maybe_truncate(
&self,
messages: &mut Vec<ChatCompletionRequestMessage>,
) -> TruncationResult {
let estimated = estimate_messages_tokens_with_margin(messages, self.token_safety_margin);
let threshold = self.truncation_threshold();
tracing::debug!("maybe_truncate: initial message count = {}, estimated tokens = {}, threshold = {}",
messages.len(), estimated, threshold);
if estimated <= threshold {
tracing::debug!("maybe_truncate: No truncation needed ({} <= {})", estimated, threshold);
return TruncationResult {
rounds_removed: 0,
messages_removed: 0,
removed_messages: Vec::new(),
insert_position: 0,
needs_compression: false,
action: TruncationAction::TruncateOnly,
};
}
tracing::info!(
"=== Context truncation triggered ==="
);
tracing::info!(
"Estimated: {} tokens (with {:.1}x margin), Threshold: {} tokens, Max: {} tokens",
estimated,
self.token_safety_margin,
threshold,
self.max_tokens
);
tracing::info!("Compression enabled: {}, Progressive: {}",
self.compression_enabled, self.progressive_compression);
if !self.progressive_compression {
tracing::info!("Progressive compression disabled, using legacy single-shot truncation");
return self.legacy_truncate(messages, threshold);
}
if let Some(result) = self.try_new_segment(messages) {
tracing::info!("Progressive action: NewSegment ({} rounds)", result.rounds_removed);
return result;
}
if let Some(result) = self.try_merge_segments(messages) {
tracing::info!("Progressive action: MergeSegments");
return result;
}
if let Some(result) = self.try_discard_oldest_segment(messages) {
tracing::info!("Progressive action: Discard oldest segment");
return result;
}
tracing::warn!("All progressive actions exhausted, falling back to legacy truncation");
self.legacy_truncate(messages, threshold)
}
fn try_new_segment(
&self,
messages: &mut Vec<ChatCompletionRequestMessage>,
) -> Option<TruncationResult> {
if !self.compression_enabled {
return None;
}
let round_starts: Vec<usize> = messages
.iter()
.enumerate()
.filter(|(_, m)| is_user_message(m) && !is_summary_segment(m) && !is_discard_notice(m))
.map(|(i, _)| i)
.collect();
if round_starts.len() <= self.min_keep_rounds + self.rounds_per_summary {
tracing::debug!("Not enough full rounds to compress (have {}, need min_keep + rounds_per_summary = {})",
round_starts.len(), self.min_keep_rounds + self.rounds_per_summary);
return None;
}
let take_rounds = self.rounds_per_summary.min(round_starts.len() - self.min_keep_rounds);
if take_rounds == 0 {
return None;
}
let start_idx = round_starts[0];
let end_idx = if take_rounds < round_starts.len() {
round_starts[take_rounds]
} else {
messages.len()
};
let removed_messages: Vec<ChatCompletionRequestMessage> =
messages[start_idx..end_idx].to_vec();
let messages_removed = removed_messages.len();
let removed_tokens = estimate_messages_tokens(&removed_messages);
if removed_tokens < self.compression_token_threshold {
tracing::debug!("Removed tokens ({}) below compression threshold ({})",
removed_tokens, self.compression_token_threshold);
return None;
}
messages.drain(start_idx..end_idx);
let system_msg_count = messages.iter().take_while(|m| is_system_message(m)).count();
let has_discard = find_discard_notice_pos(messages).is_some();
let insert_pos = system_msg_count + if has_discard { 1 } else { 0 };
let placeholder = make_summary_placeholder(take_rounds);
messages.insert(insert_pos, placeholder);
tracing::info!(
"New summary segment: removed {} rounds ({} messages, ~{} tokens), insert at {}",
take_rounds, messages_removed, removed_tokens, insert_pos
);
Some(TruncationResult {
rounds_removed: take_rounds,
messages_removed,
removed_messages,
insert_position: insert_pos,
needs_compression: true,
action: TruncationAction::NewSegment,
})
}
fn try_merge_segments(
&self,
messages: &mut Vec<ChatCompletionRequestMessage>,
) -> Option<TruncationResult> {
if !self.compression_enabled {
return None;
}
let segments = find_summary_segments(messages);
if segments.len() <= self.max_summary_segments {
tracing::debug!("Segment count ({}) within limit ({})", segments.len(), self.max_summary_segments);
return None;
}
let oldest = segments.first()?;
if oldest.merge_level >= self.max_merges_per_segment {
tracing::debug!("Oldest segment at merge level {} >= max {}, will discard instead",
oldest.merge_level, self.max_merges_per_segment);
return None;
}
let take_count = self.merge_count.min(segments.len());
if take_count < 2 {
return None;
}
let start_pos = segments[0].index;
let end_pos = segments[take_count - 1].index + 1;
let summaries: Vec<String> = segments[..take_count]
.iter()
.map(|s| s.content.clone())
.collect();
let max_level = segments[..take_count]
.iter()
.map(|s| s.merge_level)
.max()
.unwrap_or(0);
let new_level = max_level + 1;
messages.drain(start_pos..end_pos);
let placeholder = make_merge_placeholder(new_level, take_count);
messages.insert(start_pos, placeholder);
tracing::info!(
"Merging {} summary segments into one (level {}), start pos {}",
take_count, new_level, start_pos
);
Some(TruncationResult {
rounds_removed: 0,
messages_removed: take_count,
removed_messages: Vec::new(),
insert_position: start_pos,
needs_compression: true,
action: TruncationAction::MergeSegments {
summaries,
start_position: start_pos,
count: take_count,
},
})
}
fn try_discard_oldest_segment(
&self,
messages: &mut Vec<ChatCompletionRequestMessage>,
) -> Option<TruncationResult> {
let segments = find_summary_segments(messages);
if segments.len() <= self.max_summary_segments {
return None;
}
let oldest = segments.first()?;
if oldest.merge_level < self.max_merges_per_segment {
return None;
}
let pos = oldest.index;
messages.remove(pos);
tracing::info!("Discarded oldest summary segment at position {}", pos);
let system_msg_count = messages.iter().take_while(|m| is_system_message(m)).count();
if find_discard_notice_pos(messages).is_none() {
let notice = make_discard_notice();
messages.insert(system_msg_count, notice);
tracing::debug!("Added discard notice at position {}", system_msg_count);
}
Some(TruncationResult {
rounds_removed: 0,
messages_removed: 1,
removed_messages: Vec::new(),
insert_position: pos,
needs_compression: false,
action: TruncationAction::TruncateOnly,
})
}
fn legacy_truncate(
&self,
messages: &mut Vec<ChatCompletionRequestMessage>,
threshold: usize,
) -> TruncationResult {
let mut round_starts: Vec<usize> = Vec::new();
for (i, msg) in messages.iter().enumerate() {
if is_user_message(msg) {
round_starts.push(i);
}
}
tracing::debug!("Found {} user message round boundaries", round_starts.len());
if round_starts.is_empty() {
tracing::debug!("No user messages found, no truncation performed");
return TruncationResult {
rounds_removed: 0,
messages_removed: 0,
removed_messages: Vec::new(),
insert_position: 0,
needs_compression: false,
action: TruncationAction::TruncateOnly,
};
}
let total_rounds = round_starts.len();
let must_keep = self.min_keep_rounds.min(total_rounds);
tracing::debug!("Total rounds: {}, Must keep at least: {} rounds", total_rounds, must_keep);
let mut removed_messages: Vec<ChatCompletionRequestMessage> = Vec::new();
let mut rounds_removed = 0;
let mut messages_removed = 0;
while round_starts.len() > must_keep
&& estimate_messages_tokens_with_margin(messages, self.token_safety_margin) > threshold
{
let start_idx = round_starts[0];
let end_idx = if round_starts.len() > 1 {
round_starts[1]
} else {
messages.len()
};
if self.compression_enabled {
removed_messages.extend(messages[start_idx..end_idx].to_vec());
}
let count = end_idx - start_idx;
messages.drain(start_idx..end_idx);
round_starts.remove(0);
for idx in round_starts.iter_mut() {
*idx = idx.saturating_sub(count);
}
rounds_removed += 1;
messages_removed += count;
}
if rounds_removed == 0 {
tracing::debug!("No rounds removed after checks");
return TruncationResult {
rounds_removed: 0,
messages_removed: 0,
removed_messages: Vec::new(),
insert_position: 0,
needs_compression: false,
action: TruncationAction::TruncateOnly,
};
}
let removed_tokens = estimate_messages_tokens(&removed_messages);
let needs_compression =
self.compression_enabled && removed_tokens >= self.compression_token_threshold;
let system_msg_count = messages
.iter()
.take_while(|m| is_system_message(m))
.count();
let notice = if needs_compression {
format!(
"[Context compressed: {} earlier rounds ({} messages, ~{} tokens) have been summarized. {} most recent rounds preserved.]",
rounds_removed, messages_removed, removed_tokens, round_starts.len()
)
} else {
format!(
"[Context truncated: {} earlier rounds ({} messages) removed to stay within token limit. {} most recent rounds preserved.]",
rounds_removed, messages_removed, round_starts.len()
)
};
let notice_msg = ChatCompletionRequestMessage::User(
async_openai::types::chat::ChatCompletionRequestUserMessage {
content: notice.into(),
name: Some("system_notice".to_string()),
}
.into(),
);
messages.insert(system_msg_count, notice_msg);
tracing::info!(
"Legacy truncation: removed {} rounds ({} messages), kept {} rounds",
rounds_removed, messages_removed, round_starts.len()
);
TruncationResult {
rounds_removed,
messages_removed,
removed_messages,
insert_position: system_msg_count,
needs_compression,
action: if needs_compression {
TruncationAction::NewSegment
} else {
TruncationAction::TruncateOnly
},
}
}
}
fn is_user_message(msg: &ChatCompletionRequestMessage) -> bool {
matches!(msg, ChatCompletionRequestMessage::User(_))
}
fn is_system_message(msg: &ChatCompletionRequestMessage) -> bool {
matches!(msg, ChatCompletionRequestMessage::System(_))
}
const SUMMARY_SEGMENT_PREFIX: &str = "summary_segment";
const DISCARD_NOTICE_NAME: &str = "discard_notice";
const LEGACY_NOTICE_NAME: &str = "system_notice";
fn user_message_name(msg: &ChatCompletionRequestMessage) -> Option<&str> {
match msg {
ChatCompletionRequestMessage::User(u) => u.name.as_deref(),
_ => None,
}
}
fn user_message_text(msg: &ChatCompletionRequestMessage) -> String {
match msg {
ChatCompletionRequestMessage::User(u) => match &u.content {
ChatCompletionRequestUserMessageContent::Text(t) => t.clone(),
ChatCompletionRequestUserMessageContent::Array(parts) => parts
.iter()
.filter_map(|p| match p {
async_openai::types::chat::ChatCompletionRequestUserMessageContentPart::Text(t) => Some(t.text.as_str()),
_ => None,
})
.collect::<Vec<_>>()
.join(" "),
},
_ => String::new(),
}
}
pub fn is_summary_segment(msg: &ChatCompletionRequestMessage) -> bool {
match user_message_name(msg) {
Some(name) => {
name.starts_with(SUMMARY_SEGMENT_PREFIX) || name == LEGACY_NOTICE_NAME
}
None => false,
}
}
pub fn get_merge_level(msg: &ChatCompletionRequestMessage) -> usize {
match user_message_name(msg) {
Some(name) => {
if name == LEGACY_NOTICE_NAME {
0
} else if let Some(suffix) = name.strip_prefix(SUMMARY_SEGMENT_PREFIX) {
if suffix.is_empty() {
0
} else if let Some(num_str) = suffix.strip_prefix("_m") {
num_str.parse::<usize>().unwrap_or(0)
} else {
0
}
} else {
0
}
}
None => 0,
}
}
fn is_discard_notice(msg: &ChatCompletionRequestMessage) -> bool {
matches!(user_message_name(msg), Some(name) if name == DISCARD_NOTICE_NAME)
}
fn summary_segment_name(merge_level: usize) -> String {
if merge_level == 0 {
SUMMARY_SEGMENT_PREFIX.to_string()
} else {
format!("{}_m{}", SUMMARY_SEGMENT_PREFIX, merge_level)
}
}
fn make_summary_placeholder(rounds_removed: usize) -> ChatCompletionRequestMessage {
let text = format!(
"[Compressing {} earlier conversation rounds into a summary...]",
rounds_removed
);
ChatCompletionRequestMessage::User(
ChatCompletionRequestUserMessage {
content: text.into(),
name: Some(summary_segment_name(0)),
}
.into(),
)
}
fn make_merge_placeholder(merge_level: usize, count: usize) -> ChatCompletionRequestMessage {
let text = format!(
"[Merging {} earlier summary segments...]",
count
);
ChatCompletionRequestMessage::User(
ChatCompletionRequestUserMessage {
content: text.into(),
name: Some(summary_segment_name(merge_level)),
}
.into(),
)
}
fn make_discard_notice() -> ChatCompletionRequestMessage {
let text = "[Note: Earlier conversation history beyond the earliest summary has been discarded to save context space.]";
ChatCompletionRequestMessage::User(
ChatCompletionRequestUserMessage {
content: text.into(),
name: Some(DISCARD_NOTICE_NAME.to_string()),
}
.into(),
)
}
#[derive(Debug, Clone)]
struct SummarySegmentInfo {
index: usize,
merge_level: usize,
content: String,
}
fn find_summary_segments(messages: &[ChatCompletionRequestMessage]) -> Vec<SummarySegmentInfo> {
let mut segments = Vec::new();
for (i, msg) in messages.iter().enumerate() {
if is_summary_segment(msg) {
segments.push(SummarySegmentInfo {
index: i,
merge_level: get_merge_level(msg),
content: user_message_text(msg),
});
}
}
segments
}
fn find_discard_notice_pos(messages: &[ChatCompletionRequestMessage]) -> Option<usize> {
messages.iter().position(is_discard_notice)
}
pub fn format_removed_messages_as_transcript(
messages: &[ChatCompletionRequestMessage],
) -> String {
let mut transcript = String::new();
for msg in messages {
match msg {
ChatCompletionRequestMessage::User(user_msg) => {
let text = match &user_msg.content {
async_openai::types::chat::ChatCompletionRequestUserMessageContent::Text(t) => {
t.clone()
}
async_openai::types::chat::ChatCompletionRequestUserMessageContent::Array(parts) => {
parts.iter()
.filter_map(|p| match p {
async_openai::types::chat::ChatCompletionRequestUserMessageContentPart::Text(t) => Some(t.text.as_str()),
_ => None,
})
.collect::<Vec<_>>()
.join(" ")
}
};
let truncated = truncate_str(&text, 200);
transcript.push_str(&format!("User: {}\n", truncated));
}
ChatCompletionRequestMessage::Assistant(assistant_msg) => {
let content_str = assistant_msg.content.as_ref().map(|c| {
match serde_json::to_string(c) {
Ok(json) => json.trim_matches('"').to_string(),
Err(_) => format!("{:?}", c),
}
}).unwrap_or_default();
let truncated = truncate_str(&content_str, 300);
transcript.push_str(&format!("Assistant: {}\n", truncated));
if let Some(tool_calls) = &assistant_msg.tool_calls {
for tc in tool_calls {
if let async_openai::types::chat::ChatCompletionMessageToolCalls::Function(f) = tc {
transcript.push_str(&format!(
" [Tool: {}({})]\n",
f.function.name,
truncate_str(&f.function.arguments, 100)
));
}
}
}
}
ChatCompletionRequestMessage::Tool(tool_msg) => {
let content_str = match serde_json::to_string(&tool_msg.content) {
Ok(json) => json.trim_matches('"').to_string(),
Err(_) => format!("{:?}", tool_msg.content),
};
let truncated = truncate_str(&content_str, 150);
transcript.push_str(&format!(" [Result: {}]\n", truncated));
}
_ => {}
}
}
if transcript.is_empty() {
transcript.push_str("(no conversation content)");
}
transcript
}
fn truncate_str(s: &str, max_chars: usize) -> String {
if s.len() <= max_chars {
s.to_string()
} else {
let mut end = max_chars;
while end > 0 && !s.is_char_boundary(end) {
end -= 1;
}
format!("{}...", &s[..end])
}
}
#[cfg(test)]
mod tests {
use super::*;
use async_openai::types::chat::ChatCompletionRequestUserMessage;
fn make_user_message(content: &str) -> ChatCompletionRequestMessage {
ChatCompletionRequestMessage::User(
ChatCompletionRequestUserMessage {
content: content.into(),
name: None,
}
.into(),
)
}
fn make_system_message(content: &str) -> ChatCompletionRequestMessage {
ChatCompletionRequestMessage::System(
async_openai::types::chat::ChatCompletionRequestSystemMessage {
content: content.into(),
name: None,
}
.into(),
)
}
fn make_test_config() -> ContextConfig {
ContextConfig {
max_output_lines: Some(500),
max_output_bytes: Some(51200),
reserve_ratio: Some(0.2),
truncation_ratio: Some(0.7),
min_keep_rounds: Some(3),
token_safety_margin: Some(1.3),
compression_token_threshold: Some(5000),
compression_enabled: Some(true),
max_tool_calls_per_turn: Some(30),
progressive_compression: Some(true),
rounds_per_summary: Some(3),
max_summary_segments: Some(5),
merge_count: Some(2),
max_merges_per_segment: Some(2),
}
}
fn make_user_message_named(content: &str, name: &str) -> ChatCompletionRequestMessage {
ChatCompletionRequestMessage::User(
ChatCompletionRequestUserMessage {
content: content.into(),
name: Some(name.to_string()),
}
.into(),
)
}
fn make_summary_segment(content: &str, merge_level: usize) -> ChatCompletionRequestMessage {
let name = if merge_level == 0 {
"summary_segment".to_string()
} else {
format!("summary_segment_m{}", merge_level)
};
make_user_message_named(content, &name)
}
fn make_legacy_notice(content: &str) -> ChatCompletionRequestMessage {
make_user_message_named(content, "system_notice")
}
#[test]
fn test_estimate_tokens_english() {
let text = "Hello world, this is a test of the token estimation system.";
let tokens = estimate_tokens(text);
assert!(tokens >= 10, "Expected at least 10 tokens, got {}", tokens);
assert!(tokens <= 30, "Expected at most 30 tokens, got {}", tokens);
}
#[test]
fn test_estimate_tokens_chinese() {
let chinese = "你好世界,这是一个测试。";
let tokens = estimate_tokens(chinese);
assert!(tokens >= 5, "Expected at least 5 tokens, got {}", tokens);
assert!(tokens <= 15, "Expected at most 15 tokens, got {}", tokens);
}
#[test]
fn test_estimate_tokens_code() {
let code = "fn main() {\n println!(\"Hello\");\n}";
let tokens = estimate_tokens(code);
assert!(tokens >= 10, "Expected at least 10 tokens, got {}", tokens);
assert!(tokens <= 40, "Expected at most 40 tokens, got {}", tokens);
}
#[test]
fn test_estimate_tokens_empty() {
assert_eq!(estimate_tokens(""), 0);
}
#[test]
fn test_estimate_tokens_mixed() {
let mixed = "Hello 你好 world 世界!fn test() {}";
let tokens = estimate_tokens(mixed);
assert!(tokens > 0);
assert!(tokens <= 40, "Expected at most 40 tokens, got {}", tokens);
}
#[test]
fn test_estimate_messages_tokens_with_margin() {
let messages = vec![
make_system_message("You are a helpful assistant"),
make_user_message("Hello world"),
];
let raw = estimate_messages_tokens(&messages);
let with_margin = estimate_messages_tokens_with_margin(&messages, 1.3);
assert!(with_margin > raw);
let expected = (raw as f32 * 1.3).ceil() as usize;
assert_eq!(with_margin, expected);
}
#[test]
fn test_truncation_threshold() {
let config = make_test_config();
let manager = ContextManager::new(Some(65536), Some(&config));
assert_eq!(manager.truncation_threshold(), 45875);
}
#[test]
fn test_truncation_result_no_truncation() {
let mut messages = vec![
make_system_message("You are a helpful assistant"),
make_user_message("Hello"),
];
let config = make_test_config();
let manager = ContextManager::new(Some(65536), Some(&config));
let result = manager.maybe_truncate(&mut messages);
assert_eq!(result.rounds_removed, 0);
assert!(!result.needs_compression);
}
#[test]
fn test_truncation_respects_min_keep_rounds() {
let mut messages = vec![
make_system_message("You are a helpful assistant"),
];
for i in 0..10 {
let content = format!("User message {}: {}", i, "x".repeat(2000));
messages.push(make_user_message(&content));
}
let mut config = make_test_config();
config.min_keep_rounds = Some(3);
let manager = ContextManager::new(Some(5000), Some(&config));
let result = manager.maybe_truncate(&mut messages);
assert!(
result.rounds_removed > 0,
"Should have removed some rounds"
);
let user_count = messages
.iter()
.filter(|m| matches!(m, ChatCompletionRequestMessage::User(_)))
.count();
assert!(
user_count >= 4,
"Should have at least 3 user rounds + notice, got {}",
user_count
);
}
#[test]
fn test_truncation_early_trigger() {
let mut messages = vec![
make_system_message("You are a helpful assistant"),
];
for i in 0..8 {
let content = format!("User message {}: {}", i, "x".repeat(2000));
messages.push(make_user_message(&content));
}
let mut config = make_test_config();
config.truncation_ratio = Some(0.7);
config.min_keep_rounds = Some(2);
config.token_safety_margin = Some(1.3);
let manager = ContextManager::new(Some(65536), Some(&config));
let result = manager.maybe_truncate(&mut messages);
assert_eq!(
result.rounds_removed, 0,
"Should not truncate small messages in large window"
);
let manager2 = ContextManager::new(Some(8000), Some(&config));
let mut messages2 = messages.clone();
let result2 = manager2.maybe_truncate(&mut messages2);
assert!(
result2.rounds_removed > 0,
"Should truncate when exceeding small window"
);
}
#[test]
fn test_token_safety_margin_effect() {
let mut messages = vec![
make_system_message("You are a helpful assistant"),
];
for i in 0..10 {
let content = format!("User message {}: {}", i, "x".repeat(500));
messages.push(make_user_message(&content));
}
let mut config_low = make_test_config();
config_low.token_safety_margin = Some(1.0);
config_low.truncation_ratio = Some(0.7);
config_low.min_keep_rounds = Some(1);
let mut msgs_low = messages.clone();
let manager_low = ContextManager::new(Some(8000), Some(&config_low));
let result_low = manager_low.maybe_truncate(&mut msgs_low);
let mut config_high = make_test_config();
config_high.token_safety_margin = Some(2.0);
config_high.truncation_ratio = Some(0.7);
config_high.min_keep_rounds = Some(1);
let mut msgs_high = messages.clone();
let manager_high = ContextManager::new(Some(8000), Some(&config_high));
let result_high = manager_high.maybe_truncate(&mut msgs_high);
assert!(
result_high.rounds_removed >= result_low.rounds_removed,
"Higher safety margin should trigger at least as much truncation: high={}, low={}",
result_high.rounds_removed,
result_low.rounds_removed
);
}
#[test]
fn test_compression_flag_in_old_tests() {
let mut messages = vec![
make_system_message("You are a helpful assistant"),
];
for i in 0..20 {
let content = format!("User message {}: {}", i, "x".repeat(2000));
messages.push(make_user_message(&content));
}
let mut config = make_test_config();
config.compression_enabled = Some(false);
config.compression_token_threshold = Some(1000);
config.min_keep_rounds = Some(1);
let manager = ContextManager::new(Some(5000), Some(&config));
let result = manager.maybe_truncate(&mut messages);
assert!(result.rounds_removed > 0);
assert!(
!result.needs_compression,
"Should be false when compression disabled"
);
}
#[test]
fn test_truncate_output() {
let content = "line1\nline2\nline3\nline4\nline5";
let truncated = truncate_output(content, 3, 100);
assert!(truncated.contains("line1"));
assert!(truncated.contains("line2"));
assert!(truncated.contains("line3"));
assert!(!truncated.contains("line4"));
assert!(truncated.contains("Output truncated"));
}
#[test]
fn test_truncate_str_no_truncation() {
let result = truncate_str("hello", 10);
assert_eq!(result, "hello");
}
#[test]
fn test_truncate_str_with_truncation() {
let result = truncate_str("hello world this is long", 10);
assert_eq!(result, "hello worl...");
}
#[test]
fn test_format_transcript_user_and_assistant() {
let messages = vec![
make_user_message("Fix the bug in auth.rs"),
make_system_message("System message should be skipped"),
];
let transcript = format_removed_messages_as_transcript(&messages);
assert!(transcript.contains("User: Fix the bug in auth.rs"));
assert!(!transcript.contains("System message"), "System messages should be skipped");
}
#[test]
fn test_format_transcript_empty() {
let messages: Vec<ChatCompletionRequestMessage> = vec![];
let transcript = format_removed_messages_as_transcript(&messages);
assert!(transcript.contains("no conversation content"));
}
#[test]
fn test_format_transcript_truncates_long_messages() {
let long_text = "x".repeat(500);
let messages = vec![
make_user_message(&long_text),
];
let transcript = format_removed_messages_as_transcript(&messages);
assert!(transcript.contains("..."));
assert!(transcript.len() < long_text.len() + 50);
}
#[test]
fn test_is_summary_segment_recognizes_all_levels() {
let m0 = make_summary_segment("summary 0", 0);
let m1 = make_summary_segment("summary 1", 1);
let m2 = make_summary_segment("summary 2", 2);
let legacy = make_legacy_notice("old notice");
let normal = make_user_message("hello");
assert!(is_summary_segment(&m0));
assert!(is_summary_segment(&m1));
assert!(is_summary_segment(&m2));
assert!(is_summary_segment(&legacy));
assert!(!is_summary_segment(&normal));
}
#[test]
fn test_get_merge_level() {
assert_eq!(get_merge_level(&make_summary_segment("a", 0)), 0);
assert_eq!(get_merge_level(&make_summary_segment("b", 1)), 1);
assert_eq!(get_merge_level(&make_summary_segment("c", 2)), 2);
assert_eq!(get_merge_level(&make_legacy_notice("d")), 0);
assert_eq!(get_merge_level(&make_user_message("e")), 0);
}
#[test]
fn test_progressive_no_truncation_needed() {
let mut messages = vec![
make_system_message("sys"),
make_user_message("hi"),
];
let config = make_test_config();
let manager = ContextManager::new(Some(65536), Some(&config));
let result = manager.maybe_truncate(&mut messages);
assert_eq!(result.rounds_removed, 0);
assert!(!result.needs_compression);
assert_eq!(result.action, TruncationAction::TruncateOnly);
}
#[test]
fn test_progressive_new_segment() {
let mut messages = vec![make_system_message("sys")];
for i in 0..10 {
let content = format!("Round {}: {}", i, "x".repeat(2000));
messages.push(make_user_message(&content));
}
let mut config = make_test_config();
config.min_keep_rounds = Some(3);
config.rounds_per_summary = Some(3);
config.compression_token_threshold = Some(100);
let manager = ContextManager::new(Some(8000), Some(&config));
let result = manager.maybe_truncate(&mut messages);
assert_eq!(result.action, TruncationAction::NewSegment);
assert!(result.needs_compression);
assert_eq!(result.rounds_removed, 3);
assert!(result.removed_messages.len() > 0);
let has_seg = messages.iter().any(|m| is_summary_segment(m));
assert!(has_seg, "Should have a summary segment placeholder");
}
#[test]
fn test_progressive_disabled_falls_back_to_legacy() {
let mut messages = vec![make_system_message("sys")];
for i in 0..20 {
let content = format!("Round {}: {}", i, "x".repeat(2000));
messages.push(make_user_message(&content));
}
let mut config = make_test_config();
config.progressive_compression = Some(false);
config.min_keep_rounds = Some(3);
config.compression_token_threshold = Some(100);
let manager = ContextManager::new(Some(8000), Some(&config));
let result = manager.maybe_truncate(&mut messages);
assert!(result.rounds_removed > 3,
"Legacy should remove more than rounds_per_summary (3) rounds, removed {}",
result.rounds_removed);
assert!(matches!(result.action, TruncationAction::NewSegment));
}
#[test]
fn test_progressive_merge_segments() {
let mut messages = vec![make_system_message("sys")];
for i in 0..6 {
let content = format!("Summary {}: {}", i, "x".repeat(500));
messages.push(make_summary_segment(&content, 0));
}
for i in 0..2 {
let content = format!("User {}: {}", i, "x".repeat(300));
messages.push(make_user_message(&content));
}
let mut config = make_test_config();
config.max_summary_segments = Some(5);
config.merge_count = Some(2);
config.min_keep_rounds = Some(3);
config.rounds_per_summary = Some(3);
config.compression_token_threshold = Some(10);
let manager = ContextManager::new(Some(2000), Some(&config));
let estimated = estimate_messages_tokens_with_margin(&messages, 1.3);
assert!(estimated > manager.truncation_threshold(),
"Test setup error: should be over threshold, est={}, threshold={}",
estimated, manager.truncation_threshold());
let result = manager.maybe_truncate(&mut messages);
match &result.action {
TruncationAction::MergeSegments { summaries, start_position, count } => {
assert_eq!(*count, 2, "Should merge 2 segments");
assert_eq!(summaries.len(), 2);
assert!(*start_position >= 1, "Start after system message");
}
other => {
panic!("Expected MergeSegments, got {:?}", other);
}
}
let seg_count = messages.iter().filter(|m| is_summary_segment(m)).count();
assert_eq!(seg_count, 5, "Should have 5 segments after merge");
}
#[test]
fn test_progressive_discard_after_merge_limit() {
let mut messages = vec![make_system_message("sys")];
for i in 0..6 {
messages.push(make_summary_segment(&format!("Old summary {}", i), 2));
}
for i in 0..5 {
let content = format!("User {}: {}", i, "x".repeat(2000));
messages.push(make_user_message(&content));
}
let mut config = make_test_config();
config.max_summary_segments = Some(5);
config.max_merges_per_segment = Some(2);
config.merge_count = Some(2);
let manager = ContextManager::new(Some(8000), Some(&config));
let mut did_discard = false;
for _ in 0..5 {
let result = manager.maybe_truncate(&mut messages);
if result.rounds_removed == 0 && result.messages_removed > 0 && !result.needs_compression {
if find_discard_notice_pos(&messages).is_some() {
did_discard = true;
break;
}
}
}
let seg_count = messages.iter().filter(|m| is_summary_segment(m)).count();
assert!(
seg_count <= 6,
"Segment count should decrease or stay same, got {}",
seg_count
);
let _ = did_discard;
}
#[test]
fn test_legacy_notice_recognized_as_summary_segment() {
let mut messages = vec![
make_system_message("sys"),
make_legacy_notice("[Old compressed context notice]"),
];
for i in 0..8 {
let content = format!("User {}: {}", i, "x".repeat(2000));
messages.push(make_user_message(&content));
}
let mut config = make_test_config();
config.min_keep_rounds = Some(3);
config.max_summary_segments = Some(5);
config.compression_token_threshold = Some(100);
let manager = ContextManager::new(Some(8000), Some(&config));
let result = manager.maybe_truncate(&mut messages);
assert!(result.messages_removed > 0 || result.rounds_removed > 0);
}
#[test]
fn test_discard_notice_inserted_once() {
let mut messages = vec![
make_system_message("sys"),
make_summary_segment("old seg 1", 2),
make_summary_segment("old seg 2", 2),
make_summary_segment("old seg 3", 2),
make_summary_segment("old seg 4", 2),
make_summary_segment("old seg 5", 2),
make_summary_segment("old seg 6", 2),
];
for i in 0..4 {
let content = format!("User {}: {}", i, "x".repeat(2000));
messages.push(make_user_message(&content));
}
let mut config = make_test_config();
config.max_summary_segments = Some(5);
config.max_merges_per_segment = Some(2);
let manager = ContextManager::new(Some(6000), Some(&config));
for _ in 0..3 {
let _ = manager.maybe_truncate(&mut messages);
}
let discard_count = messages.iter().filter(|m| is_discard_notice(m)).count();
assert!(
discard_count <= 1,
"Should have at most 1 discard notice, found {}",
discard_count
);
}
#[test]
fn test_multiple_progressive_rounds_gradual() {
let mut messages = vec![make_system_message("sys")];
for i in 0..20 {
let content = format!("Round {} user message: {}", i, "x".repeat(1500));
messages.push(make_user_message(&content));
}
let mut config = make_test_config();
config.min_keep_rounds = Some(3);
config.rounds_per_summary = Some(3);
config.max_summary_segments = Some(4);
config.compression_token_threshold = Some(100);
let manager = ContextManager::new(Some(10000), Some(&config));
let mut seg_count_before = 0;
let mut did_new_segment = false;
let mut did_merge = false;
for round in 0..10 {
let result = manager.maybe_truncate(&mut messages);
let seg_count = messages.iter().filter(|m| is_summary_segment(m)).count();
match &result.action {
TruncationAction::NewSegment => {
did_new_segment = true;
assert_eq!(result.rounds_removed, 3);
assert!(seg_count > seg_count_before || seg_count_before == 0);
}
TruncationAction::MergeSegments { .. } => {
did_merge = true;
assert!(seg_count <= seg_count_before);
}
TruncationAction::TruncateOnly => {
}
}
seg_count_before = seg_count;
let estimated = estimate_messages_tokens_with_margin(&messages, 1.3);
if estimated <= manager.truncation_threshold() {
break;
}
tracing::debug!("Round {}: segments={}, estimated={}", round, seg_count, estimated);
}
assert!(did_new_segment, "Should have created at least one new summary segment");
let _ = did_merge;
}
}