tokio_prompt_orchestrator/
context_compression.rs1#[derive(Debug, Clone)]
8pub enum CompressionStrategy {
9 SlidingWindow { window_size: usize },
11 Summarize { max_summary_tokens: usize },
13 DropOldest { keep_last: usize },
15 Hybrid { window_size: usize, summary_tokens: usize },
17}
18
19#[derive(Debug, Clone, PartialEq)]
21pub struct Message {
22 pub role: String,
24 pub content: String,
26 pub token_count: usize,
28 pub importance: f32,
30}
31
32#[derive(Debug, Clone)]
34pub struct CompressionResult {
35 pub original_count: usize,
37 pub compressed_count: usize,
39 pub dropped_messages: usize,
41 pub summary: Option<String>,
43 pub tokens_saved: usize,
45}
46
47#[derive(Debug, Clone)]
49pub struct ContextCompressor {
50 pub strategy: CompressionStrategy,
52 pub max_context_tokens: usize,
54}
55
56impl ContextCompressor {
57 pub fn new(strategy: CompressionStrategy, max_context_tokens: usize) -> Self {
59 Self { strategy, max_context_tokens }
60 }
61
62 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 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 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 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 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 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 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 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 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
314pub 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
324pub 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, };
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
346fn first_sentence(text: &str) -> String {
349 text.split(['.', '!', '?'])
350 .next()
351 .unwrap_or(text)
352 .trim()
353 .to_string()
354}
355
356#[derive(Debug, Clone)]
360pub struct ContextBudget {
361 pub total_tokens: usize,
363 pub used: usize,
365 pub reserved_for_response: usize,
367}
368
369impl ContextBudget {
370 pub fn available(&self) -> usize {
372 self.total_tokens
373 .saturating_sub(self.used)
374 .saturating_sub(self.reserved_for_response)
375 }
376
377 pub fn can_fit(&self, tokens: usize) -> bool {
379 tokens <= self.available()
380 }
381
382 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#[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 assert!(kept.iter().any(|m| m.role == "system"));
426 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 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 let compressor = ContextCompressor::new(
455 CompressionStrategy::Summarize { max_summary_tokens: 50 },
456 20, );
458 let (kept, result) = compressor.compress(&messages);
459 assert!(result.summary.is_some());
461 let summary_content = result.summary.unwrap();
462 assert!(summary_content.starts_with("Previous context summary:"));
463 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 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 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 assert!(s_sys >= 1.0);
507 assert!(s_asst > importance_score(&m_asst, 0, 3));
509 assert!(
511 importance_score(&m_user, 0, 3) > importance_score(&m_asst, 0, 3)
512 );
513 let _ = s_user; }
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}