Skip to main content

robit_agent/
context.rs

1//! Context management — output truncation and history window management.
2
3use async_openai::types::chat::{
4    ChatCompletionRequestMessage, ChatCompletionRequestUserMessage,
5    ChatCompletionRequestUserMessageContent,
6};
7use robit_ai::config::ContextConfig;
8
9// ============================================================================
10// Truncation result
11// ============================================================================
12
13/// Type of truncation action, determines how the caller should handle the result.
14#[derive(Debug, Clone, PartialEq, Eq)]
15pub enum TruncationAction {
16    /// Generate a new summary segment from full conversation rounds.
17    /// The removed messages in `TruncationResult` are the full rounds to summarize.
18    NewSegment,
19    /// Merge multiple existing summary segments into one.
20    /// `summaries` contains the text of segments to merge (oldest first).
21    /// `start_position` is the index of the first segment in history.
22    /// `count` is how many consecutive segments to merge.
23    MergeSegments {
24        summaries: Vec<String>,
25        start_position: usize,
26        count: usize,
27    },
28    /// No compression needed — history was truncated but the removed content
29    /// is too small to justify an LLM summary call.
30    TruncateOnly,
31}
32
33/// Result of context truncation, used for async compression.
34#[derive(Debug)]
35pub struct TruncationResult {
36    /// Number of conversation rounds removed.
37    pub rounds_removed: usize,
38    /// Number of individual messages removed.
39    pub messages_removed: usize,
40    /// The removed messages (for generating summary — only for NewSegment action).
41    pub removed_messages: Vec<ChatCompletionRequestMessage>,
42    /// Position where summary should be inserted / replaced.
43    pub insert_position: usize,
44    /// Whether compression is needed (token count exceeds threshold).
45    pub needs_compression: bool,
46    /// The type of truncation action taken.
47    pub action: TruncationAction,
48}
49
50// ============================================================================
51// Tool output truncation (Layer 1)
52// ============================================================================
53
54/// Truncate tool output based on line count and byte limits.
55pub 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    // Check if truncation is needed
61    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
106// ============================================================================
107// Token estimation
108// ============================================================================
109
110/// Estimate token count for a string.
111///
112/// Uses a more nuanced heuristic based on character type:
113/// - ASCII letters/digits: ~3.5 chars/token (BPE tokenizer average)
114/// - CJK characters: ~1.5 chars/token (most CJK chars are 1-2 tokens each)
115/// - Whitespace: minimal token cost (usually merged with adjacent tokens)
116/// - Punctuation/symbols: ~1 token per char (often individual tokens)
117/// - Code (braces, operators): ~1 token per char
118///
119/// This is still an estimate; apply `token_safety_margin` at the message level.
120pub 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; // punctuation, symbols, code characters
129
130    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            // CJK Unified Ideographs + extensions + fullwidth forms
138            // + Hiragana, Katakana, Hangul, CJK punctuation
139            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    // BPE tokenizer averages:
156    // - ASCII alphanumeric: ~3.5 chars per token
157    // - CJK: ~1.5 chars per token (most are 1 token each, some pairs)
158    // - Whitespace: negligible (merged with adjacent tokens)
159    // - Other (punctuation/code): ~1 char per token
160    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; // ~1:1
164
165    ascii_tokens + cjk_tokens + whitespace_tokens + other_tokens
166}
167
168/// Estimate tokens for a list of messages.
169pub fn estimate_messages_tokens(messages: &[ChatCompletionRequestMessage]) -> usize {
170    let mut total = 0;
171    for msg in messages {
172        // Each message has ~4 tokens of overhead (role, delimiters)
173        total += 4;
174        total += estimate_message_content_tokens(msg);
175    }
176    total
177}
178
179/// Estimate tokens for messages, applying the configured safety margin.
180pub 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
188/// Estimate tokens for a single message's content.
189fn estimate_message_content_tokens(msg: &ChatCompletionRequestMessage) -> usize {
190    use async_openai::types::chat::ChatCompletionRequestUserMessageContentPart;
191
192    // For user messages with multimodal (array) content, estimate each part
193    // separately. Image base64 data URLs must NOT be counted by string length:
194    // a 2K image is ~10MB of base64 but only ~1-2k tokens to the vision API.
195    // Counting the raw base64 wildly overestimates tokens and triggers endless
196    // truncation loops.
197    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                        // Vision models count image tokens by resolution,
207                        // typically ~765-2000 tokens per image. Use a
208                        // conservative flat estimate.
209                        total += 2000;
210                    }
211                    _ => {
212                        // InputAudio, File, etc. - not used in this codebase.
213                    }
214                }
215            }
216            return total;
217        }
218    }
219
220    // Fallback: text-only messages - estimate from the JSON serialization.
221    match serde_json::to_string(msg) {
222        Ok(json) => estimate_tokens(&json),
223        Err(_) => 0,
224    }
225}
226
227// ============================================================================
228// Context manager (Layer 2: history truncation)
229// ============================================================================
230
231/// Manages the context window, truncating history when approaching token limits.
232pub struct ContextManager {
233    /// Model's context window size in tokens.
234    pub max_tokens: usize,
235    /// Ratio of context window to reserve for LLM response (default 0.2 = 20%).
236    pub reserve_ratio: f32,
237    /// Fraction of max_tokens at which truncation triggers (default 0.7).
238    pub truncation_ratio: f32,
239    /// Minimum conversation rounds to keep after truncation (default 3).
240    pub min_keep_rounds: usize,
241    /// Safety multiplier for token estimates (default 1.3).
242    pub token_safety_margin: f32,
243    /// Max output lines for tool results.
244    pub max_output_lines: usize,
245    /// Max output bytes for tool results.
246    pub max_output_bytes: usize,
247    /// Token threshold for triggering compression.
248    pub compression_token_threshold: usize,
249    /// Whether compression is enabled.
250    pub compression_enabled: bool,
251    /// Maximum tool calls per turn (default 30).
252    pub max_tool_calls_per_turn: usize,
253    /// Whether progressive segmented compression is enabled (default true).
254    pub progressive_compression: bool,
255    /// Number of full rounds per summary segment (default 3).
256    pub rounds_per_summary: usize,
257    /// Maximum number of summary segments to keep (default 5).
258    pub max_summary_segments: usize,
259    /// Number of segments to merge at a time (default 2).
260    pub merge_count: usize,
261    /// Maximum merges per segment before discarding (default 2).
262    pub max_merges_per_segment: usize,
263    /// Max dimension (longest side, px) for images encoded into the context
264    /// (default 1024). Larger images are downscaled and re-encoded as JPEG.
265    /// 0 disables compression.
266    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    /// Maximum tokens at which truncation is triggered.
331    /// Uses `truncation_ratio` (default 0.7) to trigger earlier than the
332    /// absolute limit, leaving headroom for estimation errors and LLM response.
333    pub fn truncation_threshold(&self) -> usize {
334        (self.max_tokens as f32 * self.truncation_ratio) as usize
335    }
336
337    /// Maximum tokens available for input (total - reserved for response).
338    /// Note: this is the absolute upper bound; truncation actually triggers
339    /// earlier via `truncation_threshold()`.
340    pub fn available_tokens(&self) -> usize {
341        (self.max_tokens as f32 * (1.0 - self.reserve_ratio)) as usize
342    }
343
344    /// Truncate tool output using the configured limits.
345    pub fn truncate_tool_output(&self, content: &str) -> String {
346        truncate_output(content, self.max_output_lines, self.max_output_bytes)
347    }
348
349    /// Hybrid token estimation: precise API anchor + incremental estimation.
350    ///
351    /// When `last_known_prompt_tokens` is `Some` and history hasn't been truncated
352    /// (i.e. `messages.len() >= snapshot_message_count`), uses the API-reported
353    /// `prompt_tokens` as an exact baseline and adds estimated tokens for only
354    /// the new messages appended since the snapshot. This is far more accurate
355    /// than full heuristic estimation because the baseline is from the model's
356    /// own tokenizer.
357    ///
358    /// Falls back to full heuristic estimation with `token_safety_margin` when:
359    /// - No calibration data yet (first call, before any API response)
360    /// - History was truncated/compressed (messages.len() < snapshot)
361    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                // Precise baseline + incremental estimate for new messages only.
370                // No safety_margin needed: the baseline is exact, and incremental
371                // error is small (only 1-3 new messages).
372                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                // Fallback: full heuristic estimation with safety margin
382                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    /// Check if history needs truncation and perform it if necessary.
393    /// Returns `TruncationResult` with removed messages for async compression.
394    ///
395    /// Strategy (progressive compression):
396    /// 1. Uses `truncation_threshold()` (default 70% of max_tokens) as trigger point.
397    /// 2. When over threshold, performs exactly one compression action per call:
398    ///    - Priority 1: Compress the oldest `rounds_per_summary` full rounds into a new summary segment.
399    ///    - Priority 2: Merge the oldest `merge_count` summary segments into one.
400    ///    - Priority 3: Discard the oldest summary segment (if merge limit reached).
401    ///    - Priority 4: Fall back to aggressive truncation (old behavior).
402    /// 3. If `progressive_compression` is disabled, falls back to single-shot truncation.
403    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            // The overwhelmingly common path — keep at trace to avoid log spam
414            // (this check runs twice per agent step).
415            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        // Fall back to legacy single-shot truncation if progressive is disabled
441        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        // Try progressive compression actions in priority order
447        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        // Fallback: aggressive single-shot truncation
463        tracing::warn!("All progressive actions exhausted, falling back to legacy truncation");
464        self.legacy_truncate(messages, threshold)
465    }
466
467    // ------------------------------------------------------------------------
468    // Progressive compression: Priority 1 — new summary segment
469    // ------------------------------------------------------------------------
470
471    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        // Find all full (non-summary, non-system) user rounds,
480        // skipping summary segments and the discard notice.
481        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        // Take the oldest `rounds_per_summary` rounds
495        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        // Need enough tokens to justify compression
513        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        // Remove the rounds
520        messages.drain(start_idx..end_idx);
521
522        // Insert placeholder after system messages + discard notice (if any)
523        // i.e. before any other summary segments
524        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    // ------------------------------------------------------------------------
547    // Progressive compression: Priority 2 — merge summary segments
548    // ------------------------------------------------------------------------
549
550    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        // Check if the oldest segment can still be merged
565        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        // Take the oldest `merge_count` segments
573        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        // Remove the old segments
594        messages.drain(start_pos..end_pos);
595
596        // Insert merged placeholder
597        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    // ------------------------------------------------------------------------
620    // Progressive compression: Priority 3 — discard oldest summary segment
621    // ------------------------------------------------------------------------
622
623    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            // Should have been handled by try_merge_segments
635            return None;
636        }
637
638        // Remove the oldest segment
639        let pos = oldest.index;
640        messages.remove(pos);
641
642        tracing::info!("Discarded oldest summary segment at position {}", pos);
643
644        // Ensure discard notice exists
645        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    // ------------------------------------------------------------------------
663    // Legacy single-shot truncation (fallback)
664    // ------------------------------------------------------------------------
665
666    fn legacy_truncate(
667        &self,
668        messages: &mut Vec<ChatCompletionRequestMessage>,
669        threshold: usize,
670    ) -> TruncationResult {
671        // Find round boundaries: a round starts with a User message
672        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                // For legacy mode, treat single-shot summary as NewSegment
782                // (one summary from many rounds)
783                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
799// ============================================================================
800// Summary segment helpers (progressive compression)
801// ============================================================================
802
803const SUMMARY_SEGMENT_PREFIX: &str = "summary_segment";
804const DISCARD_NOTICE_NAME: &str = "discard_notice";
805const LEGACY_NOTICE_NAME: &str = "system_notice";
806
807/// Returns the `name` field of a User message, if any.
808fn user_message_name(msg: &ChatCompletionRequestMessage) -> Option<&str> {
809    match msg {
810        ChatCompletionRequestMessage::User(u) => u.name.as_deref(),
811        _ => None,
812    }
813}
814
815/// Returns the text content of a User message, empty if not text.
816fn 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
833/// Check whether a message is a summary segment (any merge level, including legacy notice).
834pub 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
843/// Get the merge level (number of times this segment has been merged).
844/// 0 = fresh segment, 1 = merged once, etc.
845pub 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
866/// Check if a message is the discard notice.
867fn is_discard_notice(msg: &ChatCompletionRequestMessage) -> bool {
868    matches!(user_message_name(msg), Some(name) if name == DISCARD_NOTICE_NAME)
869}
870
871/// Build the `name` value for a summary segment at the given merge level.
872fn 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
880/// Build a placeholder summary segment message (pending LLM generation).
881fn 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
895/// Build a placeholder for merged segments (pending LLM generation).
896fn 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
910/// Build the discard notice message.
911fn 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/// Information about a summary segment found in history.
923#[derive(Debug, Clone)]
924struct SummarySegmentInfo {
925    index: usize,
926    merge_level: usize,
927    content: String,
928}
929
930/// Scan message history and collect all summary segments, ordered from oldest to newest.
931fn 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
945/// Find the position of the discard notice, if any.
946fn find_discard_notice_pos(messages: &[ChatCompletionRequestMessage]) -> Option<usize> {
947    messages.iter().position(is_discard_notice)
948}
949
950// ============================================================================
951// Transcript formatting for summary compression
952// ============================================================================
953
954/// Format removed messages into a compact transcript for summary generation.
955/// Extracts user messages, assistant text, and tool call names only.
956/// Truncates each message to keep the transcript concise.
957pub 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
1022/// Truncate a string to at most `max_chars` characters, adding "..." if truncated.
1023/// Respects UTF-8 character boundaries.
1024fn 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// ============================================================================
1037// Tests
1038// ============================================================================
1039
1040#[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    // fn make_discard_notice_msg() -> ChatCompletionRequestMessage {
1109    //     make_user_message_named(
1110    //         "[Note: Earlier conversation history beyond the earliest summary has been discarded.]",
1111    //         "discard_notice",
1112    //     )
1113    // }
1114
1115    // ==========================================================================
1116    // estimate_tokens tests
1117    // ==========================================================================
1118
1119    #[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        // 1.3x margin should be ~30% higher
1166        let expected = (raw as f32 * 1.3).ceil() as usize;
1167        assert_eq!(with_margin, expected);
1168    }
1169
1170    // ==========================================================================
1171    // ContextManager tests
1172    // ==========================================================================
1173
1174    #[test]
1175    fn test_truncation_threshold() {
1176        let config = make_test_config();
1177        let manager = ContextManager::new(Some(65536), Some(&config));
1178        // 65536 * 0.7 = 45875
1179        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        // Add 10 rounds of large messages
1204        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); // Must keep at least 3 rounds
1211
1212        // Use small context window to force aggressive truncation
1213        let manager = ContextManager::new(Some(5000), Some(&config));
1214        let result = manager.maybe_truncate(&mut messages, None, 0);
1215
1216        // Should have removed some rounds...
1217        assert!(
1218            result.rounds_removed > 0,
1219            "Should have removed some rounds"
1220        );
1221        // ...but should still have at least 3 user rounds + notice
1222        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        // Add 8 rounds of messages, each ~2000 chars
1240        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        // With 65536 context, truncation threshold = 45875
1251        // 8 rounds * ~2000 chars each ≈ much less than 45875, so no truncation
1252        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        // With 8000 context, truncation threshold = 5600
1260        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        // With margin 1.0 (no safety), truncation may not trigger
1281        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        // With margin 2.0 (very conservative), truncation more likely triggers
1291        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        // Higher margin should result in >= rounds removed
1301        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        // Add 20 rounds of large messages
1316        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    // ==========================================================================
1348    // Transcript formatting tests
1349    // ==========================================================================
1350
1351    #[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        // Should not contain the full 500 chars
1392        assert!(transcript.len() < long_text.len() + 50);
1393    }
1394
1395    // ==========================================================================
1396    // Progressive compression tests
1397    // ==========================================================================
1398
1399    #[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        // 10 rounds of large content
1442        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); // low threshold
1451
1452        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        // Verify the placeholder was inserted
1461        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        // 20 rounds of large content — definitely over threshold
1469        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        // Legacy mode removes as many rounds as needed to go below threshold,
1483        // which for 20 rounds of 2000 chars in an 8000-token window is more than 3.
1484        assert!(result.rounds_removed > 3,
1485            "Legacy should remove more than rounds_per_summary (3) rounds, removed {}",
1486            result.rounds_removed);
1487        // And the action type should be NewSegment (legacy single summary)
1488        assert!(matches!(result.action, TruncationAction::NewSegment));
1489    }
1490
1491    #[test]
1492    fn test_progressive_merge_segments() {
1493        // Build history where full rounds are within budget (fewer than min_keep + rounds_per_summary)
1494        // but we have too many summary segments, forcing a merge.
1495        let mut messages = vec![make_system_message("sys")];
1496        // 6 summary segments at level 0 — exceeds max of 5, each large enough to matter
1497        for i in 0..6 {
1498            let content = format!("Summary {}: {}", i, "x".repeat(500));
1499            messages.push(make_summary_segment(&content, 0));
1500        }
1501        // Only 2 full user messages — less than min_keep, so NewSegment won't trigger
1502        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        // Tiny context window to force over-threshold
1515        let manager = ContextManager::new(Some(2000), Some(&config));
1516
1517        // First check: confirm we're over threshold
1518        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        // With full rounds < min_keep + rounds_per_summary, and segments > max,
1526        // try_new_segment returns None, try_merge_segments should run
1527        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        // After merge: 6 - 2 + 1 = 5 segments (including the placeholder)
1539        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        // 6 summary segments at merge level 2 (at the max), plus some full rounds
1546        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        // First call may do NewSegment, so call a few times to reach discard
1563        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                // Likely a discard
1568                if find_discard_notice_pos(&messages).is_some() {
1569                    did_discard = true;
1570                    break;
1571                }
1572            }
1573        }
1574
1575        // Verify segment count went down or discard notice appeared
1576        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        // Just verify no panics and something happened
1584        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        // Add several full rounds to push over threshold
1594        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        // Should not crash; legacy notice is treated as a summary segment
1608        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        // Trigger a few discards
1634        for _ in 0..3 {
1635            let _ = manager.maybe_truncate(&mut messages, None, 0);
1636        }
1637
1638        // Count discard notices — should be at most 1
1639        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        // Build a long history and verify compression happens gradually
1650        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                    // Could be discard or nothing
1685                }
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        // With 20 rounds and a small window, we should see at least new segments
1699        assert!(did_new_segment, "Should have created at least one new summary segment");
1700        let _ = did_merge;
1701    }
1702
1703    // ==========================================================================
1704    // Token calibration (estimate_context_tokens) tests
1705    // ==========================================================================
1706
1707    #[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        // No calibration data → fallback to heuristic with safety margin
1717        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        // Calibration matches current history length exactly
1732        let result = manager.estimate_context_tokens(&messages, Some(100), 2);
1733        // No new messages → delta = 0, total = 100
1734        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        // Calibration snapshot was at 2 messages with 100 tokens.
1747        // Now we have 3 messages (1 new).
1748        messages.push(make_user_message("world"));
1749        let result = manager.estimate_context_tokens(&messages, Some(100), 2);
1750
1751        // result = 100 (baseline) + estimate(1 new message)
1752        let new_msg_tokens = estimate_messages_tokens(&messages[2..]);
1753        assert_eq!(result, 100 + new_msg_tokens);
1754        // Key property: calibrated uses the API baseline, NOT the heuristic.
1755        // The heuristic+margin may be larger or smaller depending on content.
1756        let heuristic = estimate_messages_tokens_with_margin(&messages, 1.3);
1757        // Just verify they're different (calibrated != heuristic for arbitrary baseline)
1758        // and that calibrated correctly adds the delta
1759        assert!(new_msg_tokens > 0, "New message should have non-zero token estimate");
1760        let _ = heuristic; // heuristic comparison is content-dependent, not asserted
1761    }
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        // Snapshot was at 5 messages, but we only have 2 → history was truncated
1773        let result = manager.estimate_context_tokens(&messages, Some(100), 5);
1774        // Should fallback to heuristic because messages.len() < snapshot
1775        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        // Build messages that are near threshold by heuristic but clearly under by calibration
1782        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        // Use small window so heuristic+margin would trigger truncation
1792        let manager = ContextManager::new(Some(3000), Some(&config));
1793
1794        // Without calibration: heuristic likely triggers truncation
1795        let mut msgs_no_cal = messages.clone();
1796        let result_no_cal = manager.maybe_truncate(&mut msgs_no_cal, None, 0);
1797
1798        // With calibration: API says we're at exactly 500 tokens (well under 3000*0.7=2100)
1799        let mut msgs_cal = messages.clone();
1800        let result_cal = manager.maybe_truncate(&mut msgs_cal, Some(500), messages.len());
1801
1802        // Calibrated should NOT truncate (500 < 2100 threshold)
1803        assert_eq!(result_cal.rounds_removed, 0,
1804            "Calibrated estimation should not trigger premature truncation");
1805
1806        // Log for visibility (the no-cal case may or may not truncate depending on exact heuristic)
1807        tracing::debug!(
1808            "no_cal: rounds_removed={}, cal: rounds_removed={}",
1809            result_no_cal.rounds_removed, result_cal.rounds_removed
1810        );
1811    }
1812}