Skip to main content

tokio_prompt_orchestrator/
context_compression.rs

1//! Context window compression and summarization utilities.
2//!
3//! Provides strategies for keeping LLM conversation context within token budget
4//! limits by dropping, summarizing, or windowing older messages.
5
6/// Strategy used when compressing a context window.
7#[derive(Debug, Clone)]
8pub enum CompressionStrategy {
9    /// Keep only the last `window_size` non-system messages (plus all system messages).
10    SlidingWindow { window_size: usize },
11    /// Replace dropped messages with a synthetic summary up to `max_summary_tokens`.
12    Summarize { max_summary_tokens: usize },
13    /// Keep system messages plus the last `keep_last` messages.
14    DropOldest { keep_last: usize },
15    /// Drop oldest until within window, then summarize if still over budget.
16    Hybrid { window_size: usize, summary_tokens: usize },
17}
18
19/// A single message in a conversation context.
20#[derive(Debug, Clone, PartialEq)]
21pub struct Message {
22    /// Role: "system", "user", or "assistant".
23    pub role: String,
24    /// Text content of the message.
25    pub content: String,
26    /// Pre-computed token count for this message.
27    pub token_count: usize,
28    /// Importance score in [0.0, 1.0]; higher = more important to retain.
29    pub importance: f32,
30}
31
32/// Summary of what happened during a compression pass.
33#[derive(Debug, Clone)]
34pub struct CompressionResult {
35    /// Number of messages before compression.
36    pub original_count: usize,
37    /// Number of messages after compression.
38    pub compressed_count: usize,
39    /// How many messages were dropped entirely.
40    pub dropped_messages: usize,
41    /// Optional synthetic summary text produced by the Summarize / Hybrid strategy.
42    pub summary: Option<String>,
43    /// Approximate tokens saved.
44    pub tokens_saved: usize,
45}
46
47/// Applies a [`CompressionStrategy`] to a slice of messages.
48#[derive(Debug, Clone)]
49pub struct ContextCompressor {
50    /// The strategy to apply.
51    pub strategy: CompressionStrategy,
52    /// Hard upper bound on total tokens across all kept messages.
53    pub max_context_tokens: usize,
54}
55
56impl ContextCompressor {
57    /// Create a new compressor with the given strategy and token budget.
58    pub fn new(strategy: CompressionStrategy, max_context_tokens: usize) -> Self {
59        Self { strategy, max_context_tokens }
60    }
61
62    /// Apply the compression strategy to `messages` and return the retained
63    /// messages together with a [`CompressionResult`] summary.
64    pub fn compress(&self, messages: &[Message]) -> (Vec<Message>, CompressionResult) {
65        let original_count = messages.len();
66        let original_tokens: usize = messages.iter().map(|m| m.token_count).sum();
67
68        let (kept, summary) = match &self.strategy {
69            CompressionStrategy::SlidingWindow { window_size } => {
70                let (sys, non_sys): (Vec<_>, Vec<_>) =
71                    messages.iter().partition(|m| m.role == "system");
72                let keep_non_sys = non_sys
73                    .iter()
74                    .rev()
75                    .take(*window_size)
76                    .rev()
77                    .cloned()
78                    .cloned()
79                    .collect::<Vec<_>>();
80                let mut result: Vec<Message> =
81                    sys.into_iter().cloned().collect();
82                result.extend(keep_non_sys);
83                (result, None)
84            }
85
86            CompressionStrategy::DropOldest { keep_last } => {
87                let (sys, non_sys): (Vec<_>, Vec<_>) =
88                    messages.iter().partition(|m| m.role == "system");
89                let keep_non_sys = non_sys
90                    .iter()
91                    .rev()
92                    .take(*keep_last)
93                    .rev()
94                    .cloned()
95                    .cloned()
96                    .collect::<Vec<_>>();
97                let mut result: Vec<Message> =
98                    sys.into_iter().cloned().collect();
99                result.extend(keep_non_sys);
100                (result, None)
101            }
102
103            CompressionStrategy::Summarize { max_summary_tokens: _ } => {
104                let (sys, non_sys): (Vec<_>, Vec<_>) =
105                    messages.iter().partition(|m| m.role == "system");
106
107                // Determine which non-system messages fit within budget.
108                let sys_tokens: usize = sys.iter().map(|m| m.token_count).sum();
109                let budget = self.max_context_tokens.saturating_sub(sys_tokens);
110
111                let mut kept_non_sys: Vec<&Message> = Vec::new();
112                let mut used = 0usize;
113                for msg in non_sys.iter().rev() {
114                    if used + msg.token_count <= budget {
115                        kept_non_sys.push(msg);
116                        used += msg.token_count;
117                    } else {
118                        break;
119                    }
120                }
121                kept_non_sys.reverse();
122
123                // Build summary from dropped messages.
124                let dropped: Vec<&Message> = non_sys
125                    .iter()
126                    .filter(|m| !kept_non_sys.contains(m))
127                    .cloned()
128                    .collect();
129
130                let summary_text = if dropped.is_empty() {
131                    None
132                } else {
133                    let key_points: Vec<String> = dropped
134                        .iter()
135                        .map(|m| first_sentence(&m.content))
136                        .collect();
137                    Some(format!(
138                        "Previous context summary: {}",
139                        key_points.join(" ")
140                    ))
141                };
142
143                let mut result: Vec<Message> = sys.into_iter().cloned().collect();
144                if let Some(ref s) = summary_text {
145                    result.push(Message {
146                        role: "system".to_string(),
147                        content: s.clone(),
148                        token_count: estimate_tokens(s),
149                        importance: 1.0,
150                    });
151                }
152                result.extend(kept_non_sys.into_iter().cloned());
153                (result, summary_text)
154            }
155
156            CompressionStrategy::Hybrid { window_size, summary_tokens: _ } => {
157                // Phase 1: drop oldest non-system messages until within window.
158                let (sys, non_sys): (Vec<_>, Vec<_>) =
159                    messages.iter().partition(|m| m.role == "system");
160
161                let windowed_non_sys: Vec<&Message> = non_sys
162                    .iter()
163                    .rev()
164                    .take(*window_size)
165                    .rev()
166                    .cloned()
167                    .collect();
168
169                let dropped_first: Vec<&Message> = non_sys
170                    .iter()
171                    .filter(|m| !windowed_non_sys.contains(m))
172                    .cloned()
173                    .collect();
174
175                // Phase 2: if still over budget, summarize.
176                let sys_tokens: usize = sys.iter().map(|m| m.token_count).sum();
177                let windowed_tokens: usize =
178                    windowed_non_sys.iter().map(|m| m.token_count).sum();
179                let total = sys_tokens + windowed_tokens;
180
181                let (final_non_sys, summary_text) = if total > self.max_context_tokens {
182                    // Summarize the windowed set further.
183                    let budget =
184                        self.max_context_tokens.saturating_sub(sys_tokens);
185                    let mut kept2: Vec<&Message> = Vec::new();
186                    let mut used = 0usize;
187                    for msg in windowed_non_sys.iter().rev() {
188                        if used + msg.token_count <= budget {
189                            kept2.push(msg);
190                            used += msg.token_count;
191                        } else {
192                            break;
193                        }
194                    }
195                    kept2.reverse();
196
197                    let all_dropped: Vec<&Message> = dropped_first
198                        .iter()
199                        .chain(
200                            windowed_non_sys
201                                .iter()
202                                .filter(|m| !kept2.contains(m)),
203                        )
204                        .cloned()
205                        .collect();
206
207                    let summary_text = if all_dropped.is_empty() {
208                        None
209                    } else {
210                        let key_points: Vec<String> = all_dropped
211                            .iter()
212                            .map(|m| first_sentence(&m.content))
213                            .collect();
214                        Some(format!(
215                            "Previous context summary: {}",
216                            key_points.join(" ")
217                        ))
218                    };
219                    (kept2, summary_text)
220                } else {
221                    // Window alone is sufficient; produce summary of phase-1 drops.
222                    let summary_text = if dropped_first.is_empty() {
223                        None
224                    } else {
225                        let key_points: Vec<String> = dropped_first
226                            .iter()
227                            .map(|m| first_sentence(&m.content))
228                            .collect();
229                        Some(format!(
230                            "Previous context summary: {}",
231                            key_points.join(" ")
232                        ))
233                    };
234                    (windowed_non_sys, summary_text)
235                };
236
237                let mut result: Vec<Message> = sys.into_iter().cloned().collect();
238                if let Some(ref s) = summary_text {
239                    result.push(Message {
240                        role: "system".to_string(),
241                        content: s.clone(),
242                        token_count: estimate_tokens(s),
243                        importance: 1.0,
244                    });
245                }
246                result.extend(final_non_sys.into_iter().cloned());
247                (result, summary_text)
248            }
249        };
250
251        let new_tokens: usize = kept.iter().map(|m| m.token_count).sum();
252        let tokens_saved = original_tokens.saturating_sub(new_tokens);
253        let compressed_count = kept.len();
254        let dropped_messages = original_count.saturating_sub(compressed_count);
255
256        (
257            kept,
258            CompressionResult {
259                original_count,
260                compressed_count,
261                dropped_messages,
262                summary,
263                tokens_saved,
264            },
265        )
266    }
267
268    /// Filter `messages` down to `target_count` by dropping the lowest-importance
269    /// messages first (system messages are always retained).
270    pub fn filter_by_importance(
271        &self,
272        messages: &[Message],
273        target_count: usize,
274    ) -> Vec<Message> {
275        if messages.len() <= target_count {
276            return messages.to_vec();
277        }
278        let total = messages.len();
279        let mut scored: Vec<(usize, f32, &Message)> = messages
280            .iter()
281            .enumerate()
282            .map(|(i, m)| (i, importance_score(m, i, total), m))
283            .collect();
284
285        // Always keep system messages.
286        let sys_count = messages.iter().filter(|m| m.role == "system").count();
287        let non_sys_target = target_count.saturating_sub(sys_count);
288
289        scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
290
291        let mut kept_indices: std::collections::HashSet<usize> = scored
292            .iter()
293            .filter(|(_, _, m)| m.role == "system")
294            .map(|(i, _, _)| *i)
295            .collect();
296
297        let mut non_sys_kept = 0usize;
298        for (i, _, m) in &scored {
299            if m.role != "system" && non_sys_kept < non_sys_target {
300                kept_indices.insert(*i);
301                non_sys_kept += 1;
302            }
303        }
304
305        messages
306            .iter()
307            .enumerate()
308            .filter(|(i, _)| kept_indices.contains(i))
309            .map(|(_, m)| m.clone())
310            .collect()
311    }
312}
313
314/// Rough token estimate: word count × 1.3 + punctuation count.
315pub fn estimate_tokens(text: &str) -> usize {
316    let word_count = text.split_whitespace().count();
317    let punct_count = text
318        .chars()
319        .filter(|c| c.is_ascii_punctuation())
320        .count();
321    ((word_count as f64 * 1.3) as usize) + punct_count
322}
323
324/// Compute an importance score for a message based on recency and role.
325///
326/// - system  → base weight 1.0
327/// - user    → base weight 0.8
328/// - assistant → base weight 0.7
329///
330/// Recency bonus: messages closer to the end of the conversation get a bonus up
331/// to 0.2 (most recent = +0.2, oldest = +0.0).
332pub fn importance_score(msg: &Message, position: usize, total: usize) -> f32 {
333    let role_weight: f32 = match msg.role.as_str() {
334        "system" => 1.0,
335        "user" => 0.8,
336        _ => 0.7, // assistant
337    };
338    let recency = if total <= 1 {
339        0.2_f32
340    } else {
341        0.2 * (position as f32 / (total - 1) as f32)
342    };
343    role_weight + recency
344}
345
346// ── helpers ──────────────────────────────────────────────────────────────────
347
348fn first_sentence(text: &str) -> String {
349    text.split(['.', '!', '?'])
350        .next()
351        .unwrap_or(text)
352        .trim()
353        .to_string()
354}
355
356// ── Context budget ────────────────────────────────────────────────────────────
357
358/// Tracks token usage against a fixed budget.
359#[derive(Debug, Clone)]
360pub struct ContextBudget {
361    /// Total token capacity.
362    pub total_tokens: usize,
363    /// Tokens already consumed.
364    pub used: usize,
365    /// Tokens reserved for the model's response (not available for context).
366    pub reserved_for_response: usize,
367}
368
369impl ContextBudget {
370    /// Tokens available for additional context.
371    pub fn available(&self) -> usize {
372        self.total_tokens
373            .saturating_sub(self.used)
374            .saturating_sub(self.reserved_for_response)
375    }
376
377    /// Returns `true` if `tokens` can be accommodated within the remaining budget.
378    pub fn can_fit(&self, tokens: usize) -> bool {
379        tokens <= self.available()
380    }
381
382    /// Attempt to consume `tokens` from the budget.  Returns `true` on success,
383    /// `false` if there is insufficient capacity.
384    pub fn consume(&mut self, tokens: usize) -> bool {
385        if self.can_fit(tokens) {
386            self.used += tokens;
387            true
388        } else {
389            false
390        }
391    }
392}
393
394// ── Tests ─────────────────────────────────────────────────────────────────────
395
396#[cfg(test)]
397mod tests {
398    use super::*;
399
400    fn msg(role: &str, content: &str) -> Message {
401        let tc = estimate_tokens(content);
402        Message {
403            role: role.to_string(),
404            content: content.to_string(),
405            token_count: tc,
406            importance: 0.5,
407        }
408    }
409
410    #[test]
411    fn sliding_window_keeps_system_msgs() {
412        let messages = vec![
413            msg("system", "You are a helpful assistant."),
414            msg("user", "Message 1"),
415            msg("assistant", "Reply 1"),
416            msg("user", "Message 2"),
417            msg("assistant", "Reply 2"),
418            msg("user", "Message 3"),
419        ];
420        let compressor =
421            ContextCompressor::new(CompressionStrategy::SlidingWindow { window_size: 2 }, 9999);
422        let (kept, result) = compressor.compress(&messages);
423
424        // System message must always be retained.
425        assert!(kept.iter().any(|m| m.role == "system"));
426        // Only 2 non-system messages should remain.
427        let non_sys: Vec<_> = kept.iter().filter(|m| m.role != "system").collect();
428        assert_eq!(non_sys.len(), 2);
429        assert_eq!(result.dropped_messages, 3);
430    }
431
432    #[test]
433    fn drop_oldest_count() {
434        let messages: Vec<Message> = (0..6)
435            .map(|i| msg(if i == 0 { "system" } else { "user" }, &format!("msg {i}")))
436            .collect();
437        let compressor =
438            ContextCompressor::new(CompressionStrategy::DropOldest { keep_last: 3 }, 9999);
439        let (kept, result) = compressor.compress(&messages);
440        // 1 system + 3 non-system
441        assert_eq!(kept.len(), 4);
442        assert_eq!(result.dropped_messages, 2);
443    }
444
445    #[test]
446    fn summarize_produces_summary_msg() {
447        let messages = vec![
448            msg("system", "System prompt."),
449            msg("user", "First user turn. Extra words here."),
450            msg("assistant", "First assistant reply. More words."),
451            msg("user", "Second user turn."),
452        ];
453        // Make the budget tiny so some messages are dropped.
454        let compressor = ContextCompressor::new(
455            CompressionStrategy::Summarize { max_summary_tokens: 50 },
456            20, // very small budget
457        );
458        let (kept, result) = compressor.compress(&messages);
459        // A summary should have been generated.
460        assert!(result.summary.is_some());
461        let summary_content = result.summary.unwrap();
462        assert!(summary_content.starts_with("Previous context summary:"));
463        // The summary message should appear in kept messages.
464        assert!(kept
465            .iter()
466            .any(|m| m.content.starts_with("Previous context summary:")));
467    }
468
469    #[test]
470    fn hybrid_strategy() {
471        let messages: Vec<Message> = (0..8)
472            .map(|i| msg(if i == 0 { "system" } else { "user" }, &format!("message number {i}")))
473            .collect();
474        let compressor = ContextCompressor::new(
475            CompressionStrategy::Hybrid { window_size: 4, summary_tokens: 50 },
476            9999,
477        );
478        let (kept, result) = compressor.compress(&messages);
479        // System + at most 4 non-sys + optional summary.
480        let non_sys: Vec<_> = kept
481            .iter()
482            .filter(|m| m.role != "system" && !m.content.starts_with("Previous context summary:"))
483            .collect();
484        assert!(non_sys.len() <= 4, "got {} non-sys messages", non_sys.len());
485        assert_eq!(result.original_count, 8);
486    }
487
488    #[test]
489    fn token_estimation() {
490        let t = estimate_tokens("Hello, world!");
491        // "Hello," + "world!" → 2 words → floor(2 * 1.3) = 2, punct = 2 → total 4
492        assert!(t >= 3);
493    }
494
495    #[test]
496    fn importance_scoring() {
497        let m_sys = msg("system", "sys");
498        let m_user = msg("user", "usr");
499        let m_asst = msg("assistant", "asst");
500
501        let s_sys = importance_score(&m_sys, 0, 3);
502        let s_user = importance_score(&m_user, 1, 3);
503        let s_asst = importance_score(&m_asst, 2, 3);
504
505        // System should always score highest for a given position.
506        assert!(s_sys >= 1.0);
507        // Recency bonus means last message gets a boost.
508        assert!(s_asst > importance_score(&m_asst, 0, 3));
509        // user outranks assistant at same position.
510        assert!(
511            importance_score(&m_user, 0, 3) > importance_score(&m_asst, 0, 3)
512        );
513        let _ = s_user; // suppress unused warning
514    }
515
516    #[test]
517    fn context_budget_consume() {
518        let mut budget = ContextBudget {
519            total_tokens: 100,
520            used: 0,
521            reserved_for_response: 20,
522        };
523        assert_eq!(budget.available(), 80);
524        assert!(budget.can_fit(50));
525        assert!(budget.consume(50));
526        assert_eq!(budget.available(), 30);
527        assert!(!budget.consume(40));
528        assert_eq!(budget.used, 50);
529    }
530}