1use async_openai::types::chat::{
4 ChatCompletionRequestMessage, ChatCompletionRequestUserMessage,
5 ChatCompletionRequestUserMessageContent,
6};
7use robit_ai::config::ContextConfig;
8
9#[derive(Debug, Clone, PartialEq, Eq)]
15pub enum TruncationAction {
16 NewSegment,
19 MergeSegments {
24 summaries: Vec<String>,
25 start_position: usize,
26 count: usize,
27 },
28 TruncateOnly,
31}
32
33#[derive(Debug)]
35pub struct TruncationResult {
36 pub rounds_removed: usize,
38 pub messages_removed: usize,
40 pub removed_messages: Vec<ChatCompletionRequestMessage>,
42 pub insert_position: usize,
44 pub needs_compression: bool,
46 pub action: TruncationAction,
48}
49
50pub fn truncate_output(content: &str, max_lines: usize, max_bytes: usize) -> String {
56 let lines: Vec<&str> = content.lines().collect();
57 let total_lines = lines.len();
58 let total_bytes = content.len();
59
60 let line_truncated = total_lines > max_lines;
62 let byte_truncated = total_bytes > max_bytes;
63
64 if !line_truncated && !byte_truncated {
65 return content.to_string();
66 }
67
68 let mut output = String::new();
69 let mut byte_count = 0;
70 let mut displayed_lines = 0;
71
72 for (i, line) in lines.iter().enumerate() {
73 if i >= max_lines {
74 break;
75 }
76 let line_with_newline = if i < total_lines - 1 {
77 format!("{}\n", line)
78 } else {
79 line.to_string()
80 };
81
82 if byte_count + line_with_newline.len() > max_bytes {
83 break;
84 }
85
86 output.push_str(&line_with_newline);
87 byte_count += line_with_newline.len();
88 displayed_lines += 1;
89 }
90
91 if line_truncated {
92 output.push_str(&format!(
93 "\n... (Output truncated, {} lines total, showing first {}. Use offset/limit to read more)",
94 total_lines, displayed_lines
95 ));
96 } else if byte_truncated {
97 output.push_str(&format!(
98 "\n... (Output truncated, {} bytes total, showing first {} bytes)",
99 total_bytes, byte_count
100 ));
101 }
102
103 output
104}
105
106pub fn estimate_tokens(text: &str) -> usize {
121 if text.is_empty() {
122 return 0;
123 }
124
125 let mut ascii_alnum = 0usize;
126 let mut cjk = 0usize;
127 let mut whitespace = 0usize;
128 let mut other = 0usize; for ch in text.chars() {
131 if ch.is_whitespace() {
132 whitespace += 1;
133 } else if ch.is_ascii_alphanumeric() {
134 ascii_alnum += 1;
135 } else {
136 let cp = ch as u32;
137 if (0x4E00..=0x9FFF).contains(&cp)
140 || (0x3400..=0x4DBF).contains(&cp)
141 || (0xF900..=0xFAFF).contains(&cp)
142 || (0xFF00..=0xFFEF).contains(&cp)
143 || (0x3000..=0x303F).contains(&cp)
144 || (0x3040..=0x309F).contains(&cp)
145 || (0x30A0..=0x30FF).contains(&cp)
146 || (0xAC00..=0xD7AF).contains(&cp)
147 {
148 cjk += 1;
149 } else {
150 other += 1;
151 }
152 }
153 }
154
155 let ascii_tokens = (ascii_alnum as f64 / 3.5).ceil() as usize;
161 let cjk_tokens = (cjk as f64 / 1.5).ceil() as usize;
162 let whitespace_tokens = (whitespace as f64 / 10.0).ceil() as usize;
163 let other_tokens = other; ascii_tokens + cjk_tokens + whitespace_tokens + other_tokens
166}
167
168pub fn estimate_messages_tokens(messages: &[ChatCompletionRequestMessage]) -> usize {
170 let mut total = 0;
171 for msg in messages {
172 total += 4;
174 total += estimate_message_content_tokens(msg);
175 }
176 total
177}
178
179pub fn estimate_messages_tokens_with_margin(
181 messages: &[ChatCompletionRequestMessage],
182 safety_margin: f32,
183) -> usize {
184 let raw = estimate_messages_tokens(messages);
185 (raw as f32 * safety_margin).ceil() as usize
186}
187
188fn estimate_message_content_tokens(msg: &ChatCompletionRequestMessage) -> usize {
190 use async_openai::types::chat::ChatCompletionRequestUserMessageContentPart;
191
192 if let ChatCompletionRequestMessage::User(user_msg) = msg {
198 if let ChatCompletionRequestUserMessageContent::Array(parts) = &user_msg.content {
199 let mut total = 0;
200 for part in parts {
201 match part {
202 ChatCompletionRequestUserMessageContentPart::Text(t) => {
203 total += estimate_tokens(&t.text);
204 }
205 ChatCompletionRequestUserMessageContentPart::ImageUrl(_) => {
206 total += 2000;
210 }
211 _ => {
212 }
214 }
215 }
216 return total;
217 }
218 }
219
220 match serde_json::to_string(msg) {
222 Ok(json) => estimate_tokens(&json),
223 Err(_) => 0,
224 }
225}
226
227pub struct ContextManager {
233 pub max_tokens: usize,
235 pub reserve_ratio: f32,
237 pub truncation_ratio: f32,
239 pub min_keep_rounds: usize,
241 pub token_safety_margin: f32,
243 pub max_output_lines: usize,
245 pub max_output_bytes: usize,
247 pub compression_token_threshold: usize,
249 pub compression_enabled: bool,
251 pub max_tool_calls_per_turn: usize,
253 pub progressive_compression: bool,
255 pub rounds_per_summary: usize,
257 pub max_summary_segments: usize,
259 pub merge_count: usize,
261 pub max_merges_per_segment: usize,
263 pub max_image_dimension: u32,
267}
268
269impl ContextManager {
270 pub fn new(context_window: Option<u64>, config: Option<&ContextConfig>) -> Self {
271 let max_tokens = context_window.unwrap_or(65536) as usize;
272
273 let (
274 max_output_lines,
275 max_output_bytes,
276 reserve_ratio,
277 truncation_ratio,
278 min_keep_rounds,
279 token_safety_margin,
280 compression_token_threshold,
281 compression_enabled,
282 max_tool_calls_per_turn,
283 progressive_compression,
284 rounds_per_summary,
285 max_summary_segments,
286 merge_count,
287 max_merges_per_segment,
288 max_image_dimension,
289 ) = match config {
290 Some(c) => (
291 c.max_output_lines.unwrap_or(500),
292 c.max_output_bytes.unwrap_or(51200),
293 c.reserve_ratio.unwrap_or(0.2),
294 c.truncation_ratio.unwrap_or(0.7),
295 c.min_keep_rounds.unwrap_or(3),
296 c.token_safety_margin.unwrap_or(1.3),
297 c.compression_token_threshold.unwrap_or(5000),
298 c.compression_enabled.unwrap_or(true),
299 c.max_tool_calls_per_turn.unwrap_or(30),
300 c.progressive_compression.unwrap_or(true),
301 c.rounds_per_summary.unwrap_or(3),
302 c.max_summary_segments.unwrap_or(5),
303 c.merge_count.unwrap_or(2),
304 c.max_merges_per_segment.unwrap_or(2),
305 c.max_image_dimension.unwrap_or(1024),
306 ),
307 None => (500, 51200, 0.2, 0.7, 3, 1.3, 5000, true, 30, true, 3, 5, 2, 2, 1024),
308 };
309
310 Self {
311 max_tokens,
312 reserve_ratio,
313 truncation_ratio,
314 min_keep_rounds,
315 token_safety_margin,
316 max_output_lines,
317 max_output_bytes,
318 compression_token_threshold,
319 compression_enabled,
320 max_tool_calls_per_turn,
321 progressive_compression,
322 rounds_per_summary,
323 max_summary_segments,
324 merge_count,
325 max_merges_per_segment,
326 max_image_dimension,
327 }
328 }
329
330 pub fn truncation_threshold(&self) -> usize {
334 (self.max_tokens as f32 * self.truncation_ratio) as usize
335 }
336
337 pub fn available_tokens(&self) -> usize {
341 (self.max_tokens as f32 * (1.0 - self.reserve_ratio)) as usize
342 }
343
344 pub fn truncate_tool_output(&self, content: &str) -> String {
346 truncate_output(content, self.max_output_lines, self.max_output_bytes)
347 }
348
349 pub fn estimate_context_tokens(
362 &self,
363 messages: &[ChatCompletionRequestMessage],
364 last_known_prompt_tokens: Option<u32>,
365 snapshot_message_count: usize,
366 ) -> usize {
367 match last_known_prompt_tokens {
368 Some(known) if messages.len() >= snapshot_message_count => {
369 let delta = estimate_messages_tokens(&messages[snapshot_message_count..]);
373 let total = known as usize + delta;
374 tracing::trace!(
375 "estimate_context_tokens: calibrated (baseline={}, snapshot={}, new={}, delta={}, total={})",
376 known, snapshot_message_count, messages.len() - snapshot_message_count, delta, total
377 );
378 total
379 }
380 _ => {
381 let estimated = estimate_messages_tokens_with_margin(messages, self.token_safety_margin);
383 tracing::trace!(
384 "estimate_context_tokens: fallback heuristic ({} messages, {} tokens with {:.1}x margin)",
385 messages.len(), estimated, self.token_safety_margin
386 );
387 estimated
388 }
389 }
390 }
391
392 pub fn maybe_truncate(
404 &self,
405 messages: &mut Vec<ChatCompletionRequestMessage>,
406 last_known_prompt_tokens: Option<u32>,
407 snapshot_message_count: usize,
408 ) -> TruncationResult {
409 let estimated = self.estimate_context_tokens(messages, last_known_prompt_tokens, snapshot_message_count);
410 let threshold = self.truncation_threshold();
411
412 if estimated <= threshold {
413 tracing::trace!("maybe_truncate: no truncation needed ({} messages, {} <= {} tokens)",
416 messages.len(), estimated, threshold);
417 return TruncationResult {
418 rounds_removed: 0,
419 messages_removed: 0,
420 removed_messages: Vec::new(),
421 insert_position: 0,
422 needs_compression: false,
423 action: TruncationAction::TruncateOnly,
424 };
425 }
426
427 tracing::info!(
428 "=== Context truncation triggered ==="
429 );
430 tracing::info!(
431 "Estimated: {} tokens (with {:.1}x margin), Threshold: {} tokens, Max: {} tokens",
432 estimated,
433 self.token_safety_margin,
434 threshold,
435 self.max_tokens
436 );
437 tracing::info!("Compression enabled: {}, Progressive: {}",
438 self.compression_enabled, self.progressive_compression);
439
440 if !self.progressive_compression {
442 tracing::info!("Progressive compression disabled, using legacy single-shot truncation");
443 return self.legacy_truncate(messages, threshold);
444 }
445
446 if let Some(result) = self.try_new_segment(messages) {
448 tracing::info!("Progressive action: NewSegment ({} rounds)", result.rounds_removed);
449 return result;
450 }
451
452 if let Some(result) = self.try_merge_segments(messages) {
453 tracing::info!("Progressive action: MergeSegments");
454 return result;
455 }
456
457 if let Some(result) = self.try_discard_oldest_segment(messages) {
458 tracing::info!("Progressive action: Discard oldest segment");
459 return result;
460 }
461
462 tracing::warn!("All progressive actions exhausted, falling back to legacy truncation");
464 self.legacy_truncate(messages, threshold)
465 }
466
467 fn try_new_segment(
472 &self,
473 messages: &mut Vec<ChatCompletionRequestMessage>,
474 ) -> Option<TruncationResult> {
475 if !self.compression_enabled {
476 return None;
477 }
478
479 let round_starts: Vec<usize> = messages
482 .iter()
483 .enumerate()
484 .filter(|(_, m)| is_user_message(m) && !is_summary_segment(m) && !is_discard_notice(m))
485 .map(|(i, _)| i)
486 .collect();
487
488 if round_starts.len() <= self.min_keep_rounds + self.rounds_per_summary {
489 tracing::debug!("Not enough full rounds to compress (have {}, need min_keep + rounds_per_summary = {})",
490 round_starts.len(), self.min_keep_rounds + self.rounds_per_summary);
491 return None;
492 }
493
494 let take_rounds = self.rounds_per_summary.min(round_starts.len() - self.min_keep_rounds);
496 if take_rounds == 0 {
497 return None;
498 }
499
500 let start_idx = round_starts[0];
501 let end_idx = if take_rounds < round_starts.len() {
502 round_starts[take_rounds]
503 } else {
504 messages.len()
505 };
506
507 let removed_messages: Vec<ChatCompletionRequestMessage> =
508 messages[start_idx..end_idx].to_vec();
509 let messages_removed = removed_messages.len();
510 let removed_tokens = estimate_messages_tokens(&removed_messages);
511
512 if removed_tokens < self.compression_token_threshold {
514 tracing::debug!("Removed tokens ({}) below compression threshold ({})",
515 removed_tokens, self.compression_token_threshold);
516 return None;
517 }
518
519 messages.drain(start_idx..end_idx);
521
522 let system_msg_count = messages.iter().take_while(|m| is_system_message(m)).count();
525 let has_discard = find_discard_notice_pos(messages).is_some();
526 let insert_pos = system_msg_count + if has_discard { 1 } else { 0 };
527
528 let placeholder = make_summary_placeholder(take_rounds);
529 messages.insert(insert_pos, placeholder);
530
531 tracing::info!(
532 "New summary segment: removed {} rounds ({} messages, ~{} tokens), insert at {}",
533 take_rounds, messages_removed, removed_tokens, insert_pos
534 );
535
536 Some(TruncationResult {
537 rounds_removed: take_rounds,
538 messages_removed,
539 removed_messages,
540 insert_position: insert_pos,
541 needs_compression: true,
542 action: TruncationAction::NewSegment,
543 })
544 }
545
546 fn try_merge_segments(
551 &self,
552 messages: &mut Vec<ChatCompletionRequestMessage>,
553 ) -> Option<TruncationResult> {
554 if !self.compression_enabled {
555 return None;
556 }
557
558 let segments = find_summary_segments(messages);
559 if segments.len() <= self.max_summary_segments {
560 tracing::debug!("Segment count ({}) within limit ({})", segments.len(), self.max_summary_segments);
561 return None;
562 }
563
564 let oldest = segments.first()?;
566 if oldest.merge_level >= self.max_merges_per_segment {
567 tracing::debug!("Oldest segment at merge level {} >= max {}, will discard instead",
568 oldest.merge_level, self.max_merges_per_segment);
569 return None;
570 }
571
572 let take_count = self.merge_count.min(segments.len());
574 if take_count < 2 {
575 return None;
576 }
577
578 let start_pos = segments[0].index;
579 let end_pos = segments[take_count - 1].index + 1;
580
581 let summaries: Vec<String> = segments[..take_count]
582 .iter()
583 .map(|s| s.content.clone())
584 .collect();
585
586 let max_level = segments[..take_count]
587 .iter()
588 .map(|s| s.merge_level)
589 .max()
590 .unwrap_or(0);
591 let new_level = max_level + 1;
592
593 messages.drain(start_pos..end_pos);
595
596 let placeholder = make_merge_placeholder(new_level, take_count);
598 messages.insert(start_pos, placeholder);
599
600 tracing::info!(
601 "Merging {} summary segments into one (level {}), start pos {}",
602 take_count, new_level, start_pos
603 );
604
605 Some(TruncationResult {
606 rounds_removed: 0,
607 messages_removed: take_count,
608 removed_messages: Vec::new(),
609 insert_position: start_pos,
610 needs_compression: true,
611 action: TruncationAction::MergeSegments {
612 summaries,
613 start_position: start_pos,
614 count: take_count,
615 },
616 })
617 }
618
619 fn try_discard_oldest_segment(
624 &self,
625 messages: &mut Vec<ChatCompletionRequestMessage>,
626 ) -> Option<TruncationResult> {
627 let segments = find_summary_segments(messages);
628 if segments.len() <= self.max_summary_segments {
629 return None;
630 }
631
632 let oldest = segments.first()?;
633 if oldest.merge_level < self.max_merges_per_segment {
634 return None;
636 }
637
638 let pos = oldest.index;
640 messages.remove(pos);
641
642 tracing::info!("Discarded oldest summary segment at position {}", pos);
643
644 let system_msg_count = messages.iter().take_while(|m| is_system_message(m)).count();
646 if find_discard_notice_pos(messages).is_none() {
647 let notice = make_discard_notice();
648 messages.insert(system_msg_count, notice);
649 tracing::debug!("Added discard notice at position {}", system_msg_count);
650 }
651
652 Some(TruncationResult {
653 rounds_removed: 0,
654 messages_removed: 1,
655 removed_messages: Vec::new(),
656 insert_position: pos,
657 needs_compression: false,
658 action: TruncationAction::TruncateOnly,
659 })
660 }
661
662 fn legacy_truncate(
667 &self,
668 messages: &mut Vec<ChatCompletionRequestMessage>,
669 threshold: usize,
670 ) -> TruncationResult {
671 let mut round_starts: Vec<usize> = Vec::new();
673 for (i, msg) in messages.iter().enumerate() {
674 if is_user_message(msg) {
675 round_starts.push(i);
676 }
677 }
678 tracing::debug!("Found {} user message round boundaries", round_starts.len());
679
680 if round_starts.is_empty() {
681 tracing::debug!("No user messages found, no truncation performed");
682 return TruncationResult {
683 rounds_removed: 0,
684 messages_removed: 0,
685 removed_messages: Vec::new(),
686 insert_position: 0,
687 needs_compression: false,
688 action: TruncationAction::TruncateOnly,
689 };
690 }
691
692 let total_rounds = round_starts.len();
693 let must_keep = self.min_keep_rounds.min(total_rounds);
694 tracing::debug!("Total rounds: {}, Must keep at least: {} rounds", total_rounds, must_keep);
695
696 let mut removed_messages: Vec<ChatCompletionRequestMessage> = Vec::new();
697 let mut rounds_removed = 0;
698 let mut messages_removed = 0;
699
700 while round_starts.len() > must_keep
701 && estimate_messages_tokens_with_margin(messages, self.token_safety_margin) > threshold
702 {
703 let start_idx = round_starts[0];
704 let end_idx = if round_starts.len() > 1 {
705 round_starts[1]
706 } else {
707 messages.len()
708 };
709
710 if self.compression_enabled {
711 removed_messages.extend(messages[start_idx..end_idx].to_vec());
712 }
713
714 let count = end_idx - start_idx;
715 messages.drain(start_idx..end_idx);
716
717 round_starts.remove(0);
718 for idx in round_starts.iter_mut() {
719 *idx = idx.saturating_sub(count);
720 }
721
722 rounds_removed += 1;
723 messages_removed += count;
724 }
725
726 if rounds_removed == 0 {
727 tracing::debug!("No rounds removed after checks");
728 return TruncationResult {
729 rounds_removed: 0,
730 messages_removed: 0,
731 removed_messages: Vec::new(),
732 insert_position: 0,
733 needs_compression: false,
734 action: TruncationAction::TruncateOnly,
735 };
736 }
737
738 let removed_tokens = estimate_messages_tokens(&removed_messages);
739 let needs_compression =
740 self.compression_enabled && removed_tokens >= self.compression_token_threshold;
741
742 let system_msg_count = messages
743 .iter()
744 .take_while(|m| is_system_message(m))
745 .count();
746
747 let notice = if needs_compression {
748 format!(
749 "[Context compressed: {} earlier rounds ({} messages, ~{} tokens) have been summarized. {} most recent rounds preserved.]",
750 rounds_removed, messages_removed, removed_tokens, round_starts.len()
751 )
752 } else {
753 format!(
754 "[Context truncated: {} earlier rounds ({} messages) removed to stay within token limit. {} most recent rounds preserved.]",
755 rounds_removed, messages_removed, round_starts.len()
756 )
757 };
758
759 let notice_msg = ChatCompletionRequestMessage::User(
760 async_openai::types::chat::ChatCompletionRequestUserMessage {
761 content: notice.into(),
762 name: Some("system_notice".to_string()),
763 }
764 .into(),
765 );
766
767 messages.insert(system_msg_count, notice_msg);
768
769 tracing::info!(
770 "Legacy truncation: removed {} rounds ({} messages), kept {} rounds",
771 rounds_removed, messages_removed, round_starts.len()
772 );
773
774 TruncationResult {
775 rounds_removed,
776 messages_removed,
777 removed_messages,
778 insert_position: system_msg_count,
779 needs_compression,
780 action: if needs_compression {
781 TruncationAction::NewSegment
784 } else {
785 TruncationAction::TruncateOnly
786 },
787 }
788 }
789}
790
791fn is_user_message(msg: &ChatCompletionRequestMessage) -> bool {
792 matches!(msg, ChatCompletionRequestMessage::User(_))
793}
794
795fn is_system_message(msg: &ChatCompletionRequestMessage) -> bool {
796 matches!(msg, ChatCompletionRequestMessage::System(_))
797}
798
799const SUMMARY_SEGMENT_PREFIX: &str = "summary_segment";
804const DISCARD_NOTICE_NAME: &str = "discard_notice";
805const LEGACY_NOTICE_NAME: &str = "system_notice";
806
807fn user_message_name(msg: &ChatCompletionRequestMessage) -> Option<&str> {
809 match msg {
810 ChatCompletionRequestMessage::User(u) => u.name.as_deref(),
811 _ => None,
812 }
813}
814
815fn user_message_text(msg: &ChatCompletionRequestMessage) -> String {
817 match msg {
818 ChatCompletionRequestMessage::User(u) => match &u.content {
819 ChatCompletionRequestUserMessageContent::Text(t) => t.clone(),
820 ChatCompletionRequestUserMessageContent::Array(parts) => parts
821 .iter()
822 .filter_map(|p| match p {
823 async_openai::types::chat::ChatCompletionRequestUserMessageContentPart::Text(t) => Some(t.text.as_str()),
824 _ => None,
825 })
826 .collect::<Vec<_>>()
827 .join(" "),
828 },
829 _ => String::new(),
830 }
831}
832
833pub fn is_summary_segment(msg: &ChatCompletionRequestMessage) -> bool {
835 match user_message_name(msg) {
836 Some(name) => {
837 name.starts_with(SUMMARY_SEGMENT_PREFIX) || name == LEGACY_NOTICE_NAME
838 }
839 None => false,
840 }
841}
842
843pub fn get_merge_level(msg: &ChatCompletionRequestMessage) -> usize {
846 match user_message_name(msg) {
847 Some(name) => {
848 if name == LEGACY_NOTICE_NAME {
849 0
850 } else if let Some(suffix) = name.strip_prefix(SUMMARY_SEGMENT_PREFIX) {
851 if suffix.is_empty() {
852 0
853 } else if let Some(num_str) = suffix.strip_prefix("_m") {
854 num_str.parse::<usize>().unwrap_or(0)
855 } else {
856 0
857 }
858 } else {
859 0
860 }
861 }
862 None => 0,
863 }
864}
865
866fn is_discard_notice(msg: &ChatCompletionRequestMessage) -> bool {
868 matches!(user_message_name(msg), Some(name) if name == DISCARD_NOTICE_NAME)
869}
870
871fn summary_segment_name(merge_level: usize) -> String {
873 if merge_level == 0 {
874 SUMMARY_SEGMENT_PREFIX.to_string()
875 } else {
876 format!("{}_m{}", SUMMARY_SEGMENT_PREFIX, merge_level)
877 }
878}
879
880fn make_summary_placeholder(rounds_removed: usize) -> ChatCompletionRequestMessage {
882 let text = format!(
883 "[Compressing {} earlier conversation rounds into a summary...]",
884 rounds_removed
885 );
886 ChatCompletionRequestMessage::User(
887 ChatCompletionRequestUserMessage {
888 content: text.into(),
889 name: Some(summary_segment_name(0)),
890 }
891 .into(),
892 )
893}
894
895fn make_merge_placeholder(merge_level: usize, count: usize) -> ChatCompletionRequestMessage {
897 let text = format!(
898 "[Merging {} earlier summary segments...]",
899 count
900 );
901 ChatCompletionRequestMessage::User(
902 ChatCompletionRequestUserMessage {
903 content: text.into(),
904 name: Some(summary_segment_name(merge_level)),
905 }
906 .into(),
907 )
908}
909
910fn make_discard_notice() -> ChatCompletionRequestMessage {
912 let text = "[Note: Earlier conversation history beyond the earliest summary has been discarded to save context space.]";
913 ChatCompletionRequestMessage::User(
914 ChatCompletionRequestUserMessage {
915 content: text.into(),
916 name: Some(DISCARD_NOTICE_NAME.to_string()),
917 }
918 .into(),
919 )
920}
921
922#[derive(Debug, Clone)]
924struct SummarySegmentInfo {
925 index: usize,
926 merge_level: usize,
927 content: String,
928}
929
930fn find_summary_segments(messages: &[ChatCompletionRequestMessage]) -> Vec<SummarySegmentInfo> {
932 let mut segments = Vec::new();
933 for (i, msg) in messages.iter().enumerate() {
934 if is_summary_segment(msg) {
935 segments.push(SummarySegmentInfo {
936 index: i,
937 merge_level: get_merge_level(msg),
938 content: user_message_text(msg),
939 });
940 }
941 }
942 segments
943}
944
945fn find_discard_notice_pos(messages: &[ChatCompletionRequestMessage]) -> Option<usize> {
947 messages.iter().position(is_discard_notice)
948}
949
950pub fn format_removed_messages_as_transcript(
958 messages: &[ChatCompletionRequestMessage],
959) -> String {
960 let mut transcript = String::new();
961
962 for msg in messages {
963 match msg {
964 ChatCompletionRequestMessage::User(user_msg) => {
965 let text = match &user_msg.content {
966 async_openai::types::chat::ChatCompletionRequestUserMessageContent::Text(t) => {
967 t.clone()
968 }
969 async_openai::types::chat::ChatCompletionRequestUserMessageContent::Array(parts) => {
970 parts.iter()
971 .filter_map(|p| match p {
972 async_openai::types::chat::ChatCompletionRequestUserMessageContentPart::Text(t) => Some(t.text.as_str()),
973 _ => None,
974 })
975 .collect::<Vec<_>>()
976 .join(" ")
977 }
978 };
979 let truncated = truncate_str(&text, 200);
980 transcript.push_str(&format!("User: {}\n", truncated));
981 }
982 ChatCompletionRequestMessage::Assistant(assistant_msg) => {
983 let content_str = assistant_msg.content.as_ref().map(|c| {
984 match serde_json::to_string(c) {
985 Ok(json) => json.trim_matches('"').to_string(),
986 Err(_) => format!("{:?}", c),
987 }
988 }).unwrap_or_default();
989 let truncated = truncate_str(&content_str, 300);
990 transcript.push_str(&format!("Assistant: {}\n", truncated));
991 if let Some(tool_calls) = &assistant_msg.tool_calls {
992 for tc in tool_calls {
993 if let async_openai::types::chat::ChatCompletionMessageToolCalls::Function(f) = tc {
994 transcript.push_str(&format!(
995 " [Tool: {}({})]\n",
996 f.function.name,
997 truncate_str(&f.function.arguments, 100)
998 ));
999 }
1000 }
1001 }
1002 }
1003 ChatCompletionRequestMessage::Tool(tool_msg) => {
1004 let content_str = match serde_json::to_string(&tool_msg.content) {
1005 Ok(json) => json.trim_matches('"').to_string(),
1006 Err(_) => format!("{:?}", tool_msg.content),
1007 };
1008 let truncated = truncate_str(&content_str, 150);
1009 transcript.push_str(&format!(" [Result: {}]\n", truncated));
1010 }
1011 _ => {}
1012 }
1013 }
1014
1015 if transcript.is_empty() {
1016 transcript.push_str("(no conversation content)");
1017 }
1018
1019 transcript
1020}
1021
1022fn truncate_str(s: &str, max_chars: usize) -> String {
1025 if s.len() <= max_chars {
1026 s.to_string()
1027 } else {
1028 let mut end = max_chars;
1029 while end > 0 && !s.is_char_boundary(end) {
1030 end -= 1;
1031 }
1032 format!("{}...", &s[..end])
1033 }
1034}
1035
1036#[cfg(test)]
1041mod tests {
1042 use super::*;
1043 use async_openai::types::chat::ChatCompletionRequestUserMessage;
1044
1045 fn make_user_message(content: &str) -> ChatCompletionRequestMessage {
1046 ChatCompletionRequestMessage::User(
1047 ChatCompletionRequestUserMessage {
1048 content: content.into(),
1049 name: None,
1050 }
1051 .into(),
1052 )
1053 }
1054
1055 fn make_system_message(content: &str) -> ChatCompletionRequestMessage {
1056 ChatCompletionRequestMessage::System(
1057 async_openai::types::chat::ChatCompletionRequestSystemMessage {
1058 content: content.into(),
1059 name: None,
1060 }
1061 .into(),
1062 )
1063 }
1064
1065 fn make_test_config() -> ContextConfig {
1066 ContextConfig {
1067 max_output_lines: Some(500),
1068 max_output_bytes: Some(51200),
1069 reserve_ratio: Some(0.2),
1070 truncation_ratio: Some(0.7),
1071 min_keep_rounds: Some(3),
1072 token_safety_margin: Some(1.3),
1073 compression_token_threshold: Some(5000),
1074 compression_enabled: Some(true),
1075 max_tool_calls_per_turn: Some(30),
1076 progressive_compression: Some(true),
1077 rounds_per_summary: Some(3),
1078 max_summary_segments: Some(5),
1079 merge_count: Some(2),
1080 max_merges_per_segment: Some(2),
1081 max_image_dimension: Some(1024),
1082 }
1083 }
1084
1085 fn make_user_message_named(content: &str, name: &str) -> ChatCompletionRequestMessage {
1086 ChatCompletionRequestMessage::User(
1087 ChatCompletionRequestUserMessage {
1088 content: content.into(),
1089 name: Some(name.to_string()),
1090 }
1091 .into(),
1092 )
1093 }
1094
1095 fn make_summary_segment(content: &str, merge_level: usize) -> ChatCompletionRequestMessage {
1096 let name = if merge_level == 0 {
1097 "summary_segment".to_string()
1098 } else {
1099 format!("summary_segment_m{}", merge_level)
1100 };
1101 make_user_message_named(content, &name)
1102 }
1103
1104 fn make_legacy_notice(content: &str) -> ChatCompletionRequestMessage {
1105 make_user_message_named(content, "system_notice")
1106 }
1107
1108 #[test]
1120 fn test_estimate_tokens_english() {
1121 let text = "Hello world, this is a test of the token estimation system.";
1122 let tokens = estimate_tokens(text);
1123 assert!(tokens >= 10, "Expected at least 10 tokens, got {}", tokens);
1124 assert!(tokens <= 30, "Expected at most 30 tokens, got {}", tokens);
1125 }
1126
1127 #[test]
1128 fn test_estimate_tokens_chinese() {
1129 let chinese = "你好世界,这是一个测试。";
1130 let tokens = estimate_tokens(chinese);
1131 assert!(tokens >= 5, "Expected at least 5 tokens, got {}", tokens);
1132 assert!(tokens <= 15, "Expected at most 15 tokens, got {}", tokens);
1133 }
1134
1135 #[test]
1136 fn test_estimate_tokens_code() {
1137 let code = "fn main() {\n println!(\"Hello\");\n}";
1138 let tokens = estimate_tokens(code);
1139 assert!(tokens >= 10, "Expected at least 10 tokens, got {}", tokens);
1140 assert!(tokens <= 40, "Expected at most 40 tokens, got {}", tokens);
1141 }
1142
1143 #[test]
1144 fn test_estimate_tokens_empty() {
1145 assert_eq!(estimate_tokens(""), 0);
1146 }
1147
1148 #[test]
1149 fn test_estimate_tokens_mixed() {
1150 let mixed = "Hello 你好 world 世界!fn test() {}";
1151 let tokens = estimate_tokens(mixed);
1152 assert!(tokens > 0);
1153 assert!(tokens <= 40, "Expected at most 40 tokens, got {}", tokens);
1154 }
1155
1156 #[test]
1157 fn test_estimate_messages_tokens_with_margin() {
1158 let messages = vec![
1159 make_system_message("You are a helpful assistant"),
1160 make_user_message("Hello world"),
1161 ];
1162 let raw = estimate_messages_tokens(&messages);
1163 let with_margin = estimate_messages_tokens_with_margin(&messages, 1.3);
1164 assert!(with_margin > raw);
1165 let expected = (raw as f32 * 1.3).ceil() as usize;
1167 assert_eq!(with_margin, expected);
1168 }
1169
1170 #[test]
1175 fn test_truncation_threshold() {
1176 let config = make_test_config();
1177 let manager = ContextManager::new(Some(65536), Some(&config));
1178 assert_eq!(manager.truncation_threshold(), 45875);
1180 }
1181
1182 #[test]
1183 fn test_truncation_result_no_truncation() {
1184 let mut messages = vec![
1185 make_system_message("You are a helpful assistant"),
1186 make_user_message("Hello"),
1187 ];
1188
1189 let config = make_test_config();
1190 let manager = ContextManager::new(Some(65536), Some(&config));
1191 let result = manager.maybe_truncate(&mut messages, None, 0);
1192
1193 assert_eq!(result.rounds_removed, 0);
1194 assert!(!result.needs_compression);
1195 }
1196
1197 #[test]
1198 fn test_truncation_respects_min_keep_rounds() {
1199 let mut messages = vec![
1200 make_system_message("You are a helpful assistant"),
1201 ];
1202
1203 for i in 0..10 {
1205 let content = format!("User message {}: {}", i, "x".repeat(2000));
1206 messages.push(make_user_message(&content));
1207 }
1208
1209 let mut config = make_test_config();
1210 config.min_keep_rounds = Some(3); let manager = ContextManager::new(Some(5000), Some(&config));
1214 let result = manager.maybe_truncate(&mut messages, None, 0);
1215
1216 assert!(
1218 result.rounds_removed > 0,
1219 "Should have removed some rounds"
1220 );
1221 let user_count = messages
1223 .iter()
1224 .filter(|m| matches!(m, ChatCompletionRequestMessage::User(_)))
1225 .count();
1226 assert!(
1227 user_count >= 4,
1228 "Should have at least 3 user rounds + notice, got {}",
1229 user_count
1230 );
1231 }
1232
1233 #[test]
1234 fn test_truncation_early_trigger() {
1235 let mut messages = vec![
1236 make_system_message("You are a helpful assistant"),
1237 ];
1238
1239 for i in 0..8 {
1241 let content = format!("User message {}: {}", i, "x".repeat(2000));
1242 messages.push(make_user_message(&content));
1243 }
1244
1245 let mut config = make_test_config();
1246 config.truncation_ratio = Some(0.7);
1247 config.min_keep_rounds = Some(2);
1248 config.token_safety_margin = Some(1.3);
1249
1250 let manager = ContextManager::new(Some(65536), Some(&config));
1253 let result = manager.maybe_truncate(&mut messages, None, 0);
1254 assert_eq!(
1255 result.rounds_removed, 0,
1256 "Should not truncate small messages in large window"
1257 );
1258
1259 let manager2 = ContextManager::new(Some(8000), Some(&config));
1261 let mut messages2 = messages.clone();
1262 let result2 = manager2.maybe_truncate(&mut messages2, None, 0);
1263 assert!(
1264 result2.rounds_removed > 0,
1265 "Should truncate when exceeding small window"
1266 );
1267 }
1268
1269 #[test]
1270 fn test_token_safety_margin_effect() {
1271 let mut messages = vec![
1272 make_system_message("You are a helpful assistant"),
1273 ];
1274
1275 for i in 0..10 {
1276 let content = format!("User message {}: {}", i, "x".repeat(500));
1277 messages.push(make_user_message(&content));
1278 }
1279
1280 let mut config_low = make_test_config();
1282 config_low.token_safety_margin = Some(1.0);
1283 config_low.truncation_ratio = Some(0.7);
1284 config_low.min_keep_rounds = Some(1);
1285
1286 let mut msgs_low = messages.clone();
1287 let manager_low = ContextManager::new(Some(8000), Some(&config_low));
1288 let result_low = manager_low.maybe_truncate(&mut msgs_low, None, 0);
1289
1290 let mut config_high = make_test_config();
1292 config_high.token_safety_margin = Some(2.0);
1293 config_high.truncation_ratio = Some(0.7);
1294 config_high.min_keep_rounds = Some(1);
1295
1296 let mut msgs_high = messages.clone();
1297 let manager_high = ContextManager::new(Some(8000), Some(&config_high));
1298 let result_high = manager_high.maybe_truncate(&mut msgs_high, None, 0);
1299
1300 assert!(
1302 result_high.rounds_removed >= result_low.rounds_removed,
1303 "Higher safety margin should trigger at least as much truncation: high={}, low={}",
1304 result_high.rounds_removed,
1305 result_low.rounds_removed
1306 );
1307 }
1308
1309 #[test]
1310 fn test_compression_flag_in_old_tests() {
1311 let mut messages = vec![
1312 make_system_message("You are a helpful assistant"),
1313 ];
1314
1315 for i in 0..20 {
1317 let content = format!("User message {}: {}", i, "x".repeat(2000));
1318 messages.push(make_user_message(&content));
1319 }
1320
1321 let mut config = make_test_config();
1322 config.compression_enabled = Some(false);
1323 config.compression_token_threshold = Some(1000);
1324 config.min_keep_rounds = Some(1);
1325
1326 let manager = ContextManager::new(Some(5000), Some(&config));
1327 let result = manager.maybe_truncate(&mut messages, None, 0);
1328
1329 assert!(result.rounds_removed > 0);
1330 assert!(
1331 !result.needs_compression,
1332 "Should be false when compression disabled"
1333 );
1334 }
1335
1336 #[test]
1337 fn test_truncate_output() {
1338 let content = "line1\nline2\nline3\nline4\nline5";
1339 let truncated = truncate_output(content, 3, 100);
1340 assert!(truncated.contains("line1"));
1341 assert!(truncated.contains("line2"));
1342 assert!(truncated.contains("line3"));
1343 assert!(!truncated.contains("line4"));
1344 assert!(truncated.contains("Output truncated"));
1345 }
1346
1347 #[test]
1352 fn test_truncate_str_no_truncation() {
1353 let result = truncate_str("hello", 10);
1354 assert_eq!(result, "hello");
1355 }
1356
1357 #[test]
1358 fn test_truncate_str_with_truncation() {
1359 let result = truncate_str("hello world this is long", 10);
1360 assert_eq!(result, "hello worl...");
1361 }
1362
1363 #[test]
1364 fn test_format_transcript_user_and_assistant() {
1365 let messages = vec![
1366 make_user_message("Fix the bug in auth.rs"),
1367 make_system_message("System message should be skipped"),
1368 ];
1369
1370 let transcript = format_removed_messages_as_transcript(&messages);
1371 assert!(transcript.contains("User: Fix the bug in auth.rs"));
1372 assert!(!transcript.contains("System message"), "System messages should be skipped");
1373 }
1374
1375 #[test]
1376 fn test_format_transcript_empty() {
1377 let messages: Vec<ChatCompletionRequestMessage> = vec![];
1378 let transcript = format_removed_messages_as_transcript(&messages);
1379 assert!(transcript.contains("no conversation content"));
1380 }
1381
1382 #[test]
1383 fn test_format_transcript_truncates_long_messages() {
1384 let long_text = "x".repeat(500);
1385 let messages = vec![
1386 make_user_message(&long_text),
1387 ];
1388
1389 let transcript = format_removed_messages_as_transcript(&messages);
1390 assert!(transcript.contains("..."));
1391 assert!(transcript.len() < long_text.len() + 50);
1393 }
1394
1395 #[test]
1400 fn test_is_summary_segment_recognizes_all_levels() {
1401 let m0 = make_summary_segment("summary 0", 0);
1402 let m1 = make_summary_segment("summary 1", 1);
1403 let m2 = make_summary_segment("summary 2", 2);
1404 let legacy = make_legacy_notice("old notice");
1405 let normal = make_user_message("hello");
1406
1407 assert!(is_summary_segment(&m0));
1408 assert!(is_summary_segment(&m1));
1409 assert!(is_summary_segment(&m2));
1410 assert!(is_summary_segment(&legacy));
1411 assert!(!is_summary_segment(&normal));
1412 }
1413
1414 #[test]
1415 fn test_get_merge_level() {
1416 assert_eq!(get_merge_level(&make_summary_segment("a", 0)), 0);
1417 assert_eq!(get_merge_level(&make_summary_segment("b", 1)), 1);
1418 assert_eq!(get_merge_level(&make_summary_segment("c", 2)), 2);
1419 assert_eq!(get_merge_level(&make_legacy_notice("d")), 0);
1420 assert_eq!(get_merge_level(&make_user_message("e")), 0);
1421 }
1422
1423 #[test]
1424 fn test_progressive_no_truncation_needed() {
1425 let mut messages = vec![
1426 make_system_message("sys"),
1427 make_user_message("hi"),
1428 ];
1429 let config = make_test_config();
1430 let manager = ContextManager::new(Some(65536), Some(&config));
1431 let result = manager.maybe_truncate(&mut messages, None, 0);
1432
1433 assert_eq!(result.rounds_removed, 0);
1434 assert!(!result.needs_compression);
1435 assert_eq!(result.action, TruncationAction::TruncateOnly);
1436 }
1437
1438 #[test]
1439 fn test_progressive_new_segment() {
1440 let mut messages = vec![make_system_message("sys")];
1441 for i in 0..10 {
1443 let content = format!("Round {}: {}", i, "x".repeat(2000));
1444 messages.push(make_user_message(&content));
1445 }
1446
1447 let mut config = make_test_config();
1448 config.min_keep_rounds = Some(3);
1449 config.rounds_per_summary = Some(3);
1450 config.compression_token_threshold = Some(100); let manager = ContextManager::new(Some(8000), Some(&config));
1453 let result = manager.maybe_truncate(&mut messages, None, 0);
1454
1455 assert_eq!(result.action, TruncationAction::NewSegment);
1456 assert!(result.needs_compression);
1457 assert_eq!(result.rounds_removed, 3);
1458 assert!(result.removed_messages.len() > 0);
1459
1460 let has_seg = messages.iter().any(|m| is_summary_segment(m));
1462 assert!(has_seg, "Should have a summary segment placeholder");
1463 }
1464
1465 #[test]
1466 fn test_progressive_disabled_falls_back_to_legacy() {
1467 let mut messages = vec![make_system_message("sys")];
1468 for i in 0..20 {
1470 let content = format!("Round {}: {}", i, "x".repeat(2000));
1471 messages.push(make_user_message(&content));
1472 }
1473
1474 let mut config = make_test_config();
1475 config.progressive_compression = Some(false);
1476 config.min_keep_rounds = Some(3);
1477 config.compression_token_threshold = Some(100);
1478
1479 let manager = ContextManager::new(Some(8000), Some(&config));
1480 let result = manager.maybe_truncate(&mut messages, None, 0);
1481
1482 assert!(result.rounds_removed > 3,
1485 "Legacy should remove more than rounds_per_summary (3) rounds, removed {}",
1486 result.rounds_removed);
1487 assert!(matches!(result.action, TruncationAction::NewSegment));
1489 }
1490
1491 #[test]
1492 fn test_progressive_merge_segments() {
1493 let mut messages = vec![make_system_message("sys")];
1496 for i in 0..6 {
1498 let content = format!("Summary {}: {}", i, "x".repeat(500));
1499 messages.push(make_summary_segment(&content, 0));
1500 }
1501 for i in 0..2 {
1503 let content = format!("User {}: {}", i, "x".repeat(300));
1504 messages.push(make_user_message(&content));
1505 }
1506
1507 let mut config = make_test_config();
1508 config.max_summary_segments = Some(5);
1509 config.merge_count = Some(2);
1510 config.min_keep_rounds = Some(3);
1511 config.rounds_per_summary = Some(3);
1512 config.compression_token_threshold = Some(10);
1513
1514 let manager = ContextManager::new(Some(2000), Some(&config));
1516
1517 let estimated = estimate_messages_tokens_with_margin(&messages, 1.3);
1519 assert!(estimated > manager.truncation_threshold(),
1520 "Test setup error: should be over threshold, est={}, threshold={}",
1521 estimated, manager.truncation_threshold());
1522
1523 let result = manager.maybe_truncate(&mut messages, None, 0);
1524
1525 match &result.action {
1528 TruncationAction::MergeSegments { summaries, start_position, count } => {
1529 assert_eq!(*count, 2, "Should merge 2 segments");
1530 assert_eq!(summaries.len(), 2);
1531 assert!(*start_position >= 1, "Start after system message");
1532 }
1533 other => {
1534 panic!("Expected MergeSegments, got {:?}", other);
1535 }
1536 }
1537
1538 let seg_count = messages.iter().filter(|m| is_summary_segment(m)).count();
1540 assert_eq!(seg_count, 5, "Should have 5 segments after merge");
1541 }
1542
1543 #[test]
1544 fn test_progressive_discard_after_merge_limit() {
1545 let mut messages = vec![make_system_message("sys")];
1547 for i in 0..6 {
1548 messages.push(make_summary_segment(&format!("Old summary {}", i), 2));
1549 }
1550 for i in 0..5 {
1551 let content = format!("User {}: {}", i, "x".repeat(2000));
1552 messages.push(make_user_message(&content));
1553 }
1554
1555 let mut config = make_test_config();
1556 config.max_summary_segments = Some(5);
1557 config.max_merges_per_segment = Some(2);
1558 config.merge_count = Some(2);
1559
1560 let manager = ContextManager::new(Some(8000), Some(&config));
1561
1562 let mut did_discard = false;
1564 for _ in 0..5 {
1565 let result = manager.maybe_truncate(&mut messages, None, 0);
1566 if result.rounds_removed == 0 && result.messages_removed > 0 && !result.needs_compression {
1567 if find_discard_notice_pos(&messages).is_some() {
1569 did_discard = true;
1570 break;
1571 }
1572 }
1573 }
1574
1575 let seg_count = messages.iter().filter(|m| is_summary_segment(m)).count();
1577 assert!(
1578 seg_count <= 6,
1579 "Segment count should decrease or stay same, got {}",
1580 seg_count
1581 );
1582
1583 let _ = did_discard;
1585 }
1586
1587 #[test]
1588 fn test_legacy_notice_recognized_as_summary_segment() {
1589 let mut messages = vec![
1590 make_system_message("sys"),
1591 make_legacy_notice("[Old compressed context notice]"),
1592 ];
1593 for i in 0..8 {
1595 let content = format!("User {}: {}", i, "x".repeat(2000));
1596 messages.push(make_user_message(&content));
1597 }
1598
1599 let mut config = make_test_config();
1600 config.min_keep_rounds = Some(3);
1601 config.max_summary_segments = Some(5);
1602 config.compression_token_threshold = Some(100);
1603
1604 let manager = ContextManager::new(Some(8000), Some(&config));
1605 let result = manager.maybe_truncate(&mut messages, None, 0);
1606
1607 assert!(result.messages_removed > 0 || result.rounds_removed > 0);
1609 }
1610
1611 #[test]
1612 fn test_discard_notice_inserted_once() {
1613 let mut messages = vec![
1614 make_system_message("sys"),
1615 make_summary_segment("old seg 1", 2),
1616 make_summary_segment("old seg 2", 2),
1617 make_summary_segment("old seg 3", 2),
1618 make_summary_segment("old seg 4", 2),
1619 make_summary_segment("old seg 5", 2),
1620 make_summary_segment("old seg 6", 2),
1621 ];
1622 for i in 0..4 {
1623 let content = format!("User {}: {}", i, "x".repeat(2000));
1624 messages.push(make_user_message(&content));
1625 }
1626
1627 let mut config = make_test_config();
1628 config.max_summary_segments = Some(5);
1629 config.max_merges_per_segment = Some(2);
1630
1631 let manager = ContextManager::new(Some(6000), Some(&config));
1632
1633 for _ in 0..3 {
1635 let _ = manager.maybe_truncate(&mut messages, None, 0);
1636 }
1637
1638 let discard_count = messages.iter().filter(|m| is_discard_notice(m)).count();
1640 assert!(
1641 discard_count <= 1,
1642 "Should have at most 1 discard notice, found {}",
1643 discard_count
1644 );
1645 }
1646
1647 #[test]
1648 fn test_multiple_progressive_rounds_gradual() {
1649 let mut messages = vec![make_system_message("sys")];
1651 for i in 0..20 {
1652 let content = format!("Round {} user message: {}", i, "x".repeat(1500));
1653 messages.push(make_user_message(&content));
1654 }
1655
1656 let mut config = make_test_config();
1657 config.min_keep_rounds = Some(3);
1658 config.rounds_per_summary = Some(3);
1659 config.max_summary_segments = Some(4);
1660 config.compression_token_threshold = Some(100);
1661
1662 let manager = ContextManager::new(Some(10000), Some(&config));
1663
1664 let mut seg_count_before = 0;
1665 let mut did_new_segment = false;
1666 let mut did_merge = false;
1667
1668 for round in 0..10 {
1669 let result = manager.maybe_truncate(&mut messages, None, 0);
1670
1671 let seg_count = messages.iter().filter(|m| is_summary_segment(m)).count();
1672
1673 match &result.action {
1674 TruncationAction::NewSegment => {
1675 did_new_segment = true;
1676 assert_eq!(result.rounds_removed, 3);
1677 assert!(seg_count > seg_count_before || seg_count_before == 0);
1678 }
1679 TruncationAction::MergeSegments { .. } => {
1680 did_merge = true;
1681 assert!(seg_count <= seg_count_before);
1682 }
1683 TruncationAction::TruncateOnly => {
1684 }
1686 }
1687
1688 seg_count_before = seg_count;
1689
1690 let estimated = estimate_messages_tokens_with_margin(&messages, 1.3);
1691 if estimated <= manager.truncation_threshold() {
1692 break;
1693 }
1694
1695 tracing::debug!("Round {}: segments={}, estimated={}", round, seg_count, estimated);
1696 }
1697
1698 assert!(did_new_segment, "Should have created at least one new summary segment");
1700 let _ = did_merge;
1701 }
1702
1703 #[test]
1708 fn test_estimate_context_tokens_fallback_no_calibration() {
1709 let messages = vec![
1710 make_system_message("You are a helpful assistant"),
1711 make_user_message("Hello world"),
1712 ];
1713 let config = make_test_config();
1714 let manager = ContextManager::new(Some(65536), Some(&config));
1715
1716 let result = manager.estimate_context_tokens(&messages, None, 0);
1718 let expected = estimate_messages_tokens_with_margin(&messages, 1.3);
1719 assert_eq!(result, expected);
1720 }
1721
1722 #[test]
1723 fn test_estimate_context_tokens_calibrated_exact_match() {
1724 let messages = vec![
1725 make_system_message("sys"),
1726 make_user_message("hello"),
1727 ];
1728 let config = make_test_config();
1729 let manager = ContextManager::new(Some(65536), Some(&config));
1730
1731 let result = manager.estimate_context_tokens(&messages, Some(100), 2);
1733 assert_eq!(result, 100);
1735 }
1736
1737 #[test]
1738 fn test_estimate_context_tokens_calibrated_with_new_messages() {
1739 let mut messages = vec![
1740 make_system_message("sys"),
1741 make_user_message("hello"),
1742 ];
1743 let config = make_test_config();
1744 let manager = ContextManager::new(Some(65536), Some(&config));
1745
1746 messages.push(make_user_message("world"));
1749 let result = manager.estimate_context_tokens(&messages, Some(100), 2);
1750
1751 let new_msg_tokens = estimate_messages_tokens(&messages[2..]);
1753 assert_eq!(result, 100 + new_msg_tokens);
1754 let heuristic = estimate_messages_tokens_with_margin(&messages, 1.3);
1757 assert!(new_msg_tokens > 0, "New message should have non-zero token estimate");
1760 let _ = heuristic; }
1762
1763 #[test]
1764 fn test_estimate_context_tokens_stale_snapshot_falls_back() {
1765 let messages = vec![
1766 make_system_message("sys"),
1767 make_user_message("hello"),
1768 ];
1769 let config = make_test_config();
1770 let manager = ContextManager::new(Some(65536), Some(&config));
1771
1772 let result = manager.estimate_context_tokens(&messages, Some(100), 5);
1774 let expected = estimate_messages_tokens_with_margin(&messages, 1.3);
1776 assert_eq!(result, expected);
1777 }
1778
1779 #[test]
1780 fn test_maybe_truncate_with_calibration_avoids_premature_truncation() {
1781 let mut messages = vec![make_system_message("sys")];
1783 for i in 0..5 {
1784 let content = format!("User {}: {}", i, "x".repeat(500));
1785 messages.push(make_user_message(&content));
1786 }
1787
1788 let mut config = make_test_config();
1789 config.min_keep_rounds = Some(3);
1790
1791 let manager = ContextManager::new(Some(3000), Some(&config));
1793
1794 let mut msgs_no_cal = messages.clone();
1796 let result_no_cal = manager.maybe_truncate(&mut msgs_no_cal, None, 0);
1797
1798 let mut msgs_cal = messages.clone();
1800 let result_cal = manager.maybe_truncate(&mut msgs_cal, Some(500), messages.len());
1801
1802 assert_eq!(result_cal.rounds_removed, 0,
1804 "Calibrated estimation should not trigger premature truncation");
1805
1806 tracing::debug!(
1808 "no_cal: rounds_removed={}, cal: rounds_removed={}",
1809 result_no_cal.rounds_removed, result_cal.rounds_removed
1810 );
1811 }
1812}