1use std::collections::HashSet;
33
34#[derive(Debug, Clone)]
38pub struct CompressionResult {
39 pub text: String,
41 pub ratio: f64,
43 pub chars_removed: usize,
45 pub strategy: String,
47}
48
49pub trait Compressor: Send + Sync {
51 fn compress(&self, input: &str) -> CompressionResult;
53
54 fn name(&self) -> &'static str;
56}
57
58pub struct WhitespaceCompressor;
66
67impl Compressor for WhitespaceCompressor {
68 fn compress(&self, input: &str) -> CompressionResult {
69 let mut result = String::with_capacity(input.len());
71 let mut prev_space = false;
72 let mut newline_run = 0u8;
73
74 for ch in input.chars() {
75 match ch {
76 '\n' => {
77 newline_run += 1;
78 prev_space = false;
79 if newline_run <= 2 {
80 result.push('\n');
81 }
82 }
83 ' ' | '\t' => {
84 newline_run = 0;
85 if !prev_space {
86 result.push(' ');
87 prev_space = true;
88 }
89 }
90 _ => {
91 newline_run = 0;
92 prev_space = false;
93 result.push(ch);
94 }
95 }
96 }
97
98 let chars_removed = input.len().saturating_sub(result.len());
99 let ratio = if input.is_empty() {
100 1.0
101 } else {
102 result.len() as f64 / input.len() as f64
103 };
104
105 CompressionResult {
106 text: result,
107 ratio,
108 chars_removed,
109 strategy: self.name().to_string(),
110 }
111 }
112
113 fn name(&self) -> &'static str {
114 "whitespace"
115 }
116}
117
118pub struct RepetitionRemover;
126
127impl Compressor for RepetitionRemover {
128 fn compress(&self, input: &str) -> CompressionResult {
129 let mut seen: HashSet<&str> = HashSet::new();
131 let mut output_parts: Vec<&str> = Vec::new();
132
133 for sentence in input.split_inclusive(". ") {
134 let trimmed = sentence.trim();
135 if trimmed.is_empty() {
136 continue;
137 }
138 if seen.insert(trimmed) {
139 output_parts.push(sentence);
140 }
141 }
142
143 let text = output_parts.join("");
144 let chars_removed = input.len().saturating_sub(text.len());
145 let ratio = if input.is_empty() {
146 1.0
147 } else {
148 text.len() as f64 / input.len() as f64
149 };
150
151 CompressionResult {
152 text,
153 ratio,
154 chars_removed,
155 strategy: self.name().to_string(),
156 }
157 }
158
159 fn name(&self) -> &'static str {
160 "repetition_remover"
161 }
162}
163
164pub struct StopWordFilter {
173 stop_words: HashSet<&'static str>,
174}
175
176impl Default for StopWordFilter {
177 fn default() -> Self {
178 Self::new()
179 }
180}
181
182impl StopWordFilter {
183 pub fn new() -> Self {
185 let words = [
186 "a", "an", "the", "is", "are", "was", "were", "be", "been", "being",
187 "have", "has", "had", "do", "does", "did", "will", "would", "could",
188 "should", "may", "might", "shall", "can", "need", "dare", "ought",
189 "used", "to", "of", "in", "on", "at", "by", "for", "with", "about",
190 "against", "between", "into", "through", "during", "before", "after",
191 "above", "below", "from", "up", "down", "out", "off", "over", "under",
192 "again", "further", "then", "once", "and", "but", "or", "nor", "so",
193 "yet", "both", "either", "neither", "not", "only", "own", "same",
194 "than", "too", "very", "s", "t", "just", "don", "now", "i", "me",
195 "my", "myself", "we", "our", "you", "your", "he", "she", "it", "they",
196 "them", "this", "that", "these", "those", "what", "which", "who",
197 ];
198 Self {
199 stop_words: words.iter().copied().collect(),
200 }
201 }
202}
203
204impl Compressor for StopWordFilter {
205 fn compress(&self, input: &str) -> CompressionResult {
206 let words: Vec<&str> = input.split_whitespace().collect();
207 let filtered: Vec<&str> = words
208 .iter()
209 .copied()
210 .filter(|w| {
211 let lower = w.to_lowercase();
212 let bare = lower.trim_matches(|c: char| !c.is_alphabetic());
213 !self.stop_words.contains(bare)
214 })
215 .collect();
216
217 let text = filtered.join(" ");
218 let chars_removed = input.len().saturating_sub(text.len());
219 let ratio = if input.is_empty() {
220 1.0
221 } else {
222 text.len() as f64 / input.len() as f64
223 };
224
225 CompressionResult {
226 text,
227 ratio,
228 chars_removed,
229 strategy: self.name().to_string(),
230 }
231 }
232
233 fn name(&self) -> &'static str {
234 "stop_word_filter"
235 }
236}
237
238pub struct SentenceRanker {
249 keep_fraction: f64,
250}
251
252impl SentenceRanker {
253 pub fn new(keep_fraction: f64) -> Self {
259 Self {
260 keep_fraction: keep_fraction.clamp(0.05, 1.0),
261 }
262 }
263
264 fn score_sentences(&self, sentences: &[&str]) -> Vec<(usize, f64)> {
266 use std::collections::HashMap;
267
268 if sentences.is_empty() {
269 return Vec::new();
270 }
271
272 let term_freqs: Vec<HashMap<String, f64>> = sentences
274 .iter()
275 .map(|s| {
276 let mut freq: HashMap<String, f64> = HashMap::new();
277 for word in s.split_whitespace() {
278 let w = word.to_lowercase();
279 let w = w.trim_matches(|c: char| !c.is_alphanumeric()).to_string();
280 if !w.is_empty() {
281 *freq.entry(w).or_insert(0.0) += 1.0;
282 }
283 }
284 let total: f64 = freq.values().sum();
286 if total > 0.0 {
287 freq.values_mut().for_each(|v| *v /= total);
288 }
289 freq
290 })
291 .collect();
292
293 let n = sentences.len() as f64;
295 let mut doc_freq: HashMap<String, f64> = HashMap::new();
296 for tf in &term_freqs {
297 for term in tf.keys() {
298 *doc_freq.entry(term.clone()).or_insert(0.0) += 1.0;
299 }
300 }
301
302 let mut scores: Vec<(usize, f64)> = term_freqs
304 .iter()
305 .enumerate()
306 .map(|(i, tf)| {
307 let score: f64 = tf
308 .iter()
309 .map(|(term, &tf_val)| {
310 let df = doc_freq.get(term).copied().unwrap_or(1.0);
311 let idf = (n / df).ln();
312 tf_val * idf
313 })
314 .sum();
315 (i, score)
316 })
317 .collect();
318
319 scores.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
320 scores
321 }
322}
323
324impl Compressor for SentenceRanker {
325 fn compress(&self, input: &str) -> CompressionResult {
326 let mut sentences: Vec<&str> = Vec::new();
328 for part in input.split(". ") {
329 sentences.push(part);
330 }
331
332 if sentences.len() <= 1 {
333 return CompressionResult {
334 text: input.to_string(),
335 ratio: 1.0,
336 chars_removed: 0,
337 strategy: self.name().to_string(),
338 };
339 }
340
341 let keep = ((sentences.len() as f64 * self.keep_fraction).ceil() as usize)
342 .max(1)
343 .min(sentences.len());
344
345 let scored = self.score_sentences(&sentences);
346
347 let mut keep_indices: Vec<usize> = scored.iter().take(keep).map(|(i, _)| *i).collect();
349 keep_indices.sort_unstable();
350
351 let text = keep_indices
352 .iter()
353 .map(|&i| sentences[i])
354 .collect::<Vec<_>>()
355 .join(". ");
356
357 let chars_removed = input.len().saturating_sub(text.len());
358 let ratio = if input.is_empty() {
359 1.0
360 } else {
361 text.len() as f64 / input.len() as f64
362 };
363
364 CompressionResult {
365 text,
366 ratio,
367 chars_removed,
368 strategy: self.name().to_string(),
369 }
370 }
371
372 fn name(&self) -> &'static str {
373 "sentence_ranker"
374 }
375}
376
377pub struct TruncationStrategy {
383 max_chars: usize,
384}
385
386impl TruncationStrategy {
387 pub fn new(max_chars: usize) -> Self {
389 Self { max_chars }
390 }
391}
392
393impl Compressor for TruncationStrategy {
394 fn compress(&self, input: &str) -> CompressionResult {
395 if input.len() <= self.max_chars {
396 return CompressionResult {
397 text: input.to_string(),
398 ratio: 1.0,
399 chars_removed: 0,
400 strategy: self.name().to_string(),
401 };
402 }
403
404 let candidate = &input[..self.max_chars];
406 let cut = candidate
407 .rfind(". ")
408 .or_else(|| candidate.rfind('\n'))
409 .map(|i| i + 1)
410 .unwrap_or(self.max_chars);
411
412 let text = input[..cut].trim().to_string();
413 let chars_removed = input.len() - text.len();
414 let ratio = text.len() as f64 / input.len() as f64;
415
416 CompressionResult {
417 text,
418 ratio,
419 chars_removed,
420 strategy: self.name().to_string(),
421 }
422 }
423
424 fn name(&self) -> &'static str {
425 "truncation"
426 }
427}
428
429pub struct CompressionPipeline {
438 stages: Vec<Box<dyn Compressor>>,
439}
440
441impl Default for CompressionPipeline {
442 fn default() -> Self {
443 Self::new()
444 }
445}
446
447impl CompressionPipeline {
448 pub fn new() -> Self {
450 Self { stages: Vec::new() }
451 }
452
453 pub fn with(mut self, stage: Box<dyn Compressor>) -> Self {
455 self.stages.push(stage);
456 self
457 }
458
459 pub fn default_for_rag() -> Self {
463 Self::new()
464 .with(Box::new(WhitespaceCompressor))
465 .with(Box::new(RepetitionRemover))
466 .with(Box::new(SentenceRanker::new(0.7)))
467 }
468
469 pub fn for_conversation_history() -> Self {
473 Self::new()
474 .with(Box::new(WhitespaceCompressor))
475 .with(Box::new(RepetitionRemover))
476 .with(Box::new(SentenceRanker::new(0.5)))
477 }
478
479 pub fn compress(&self, input: &str) -> (String, f64) {
483 if self.stages.is_empty() {
484 return (input.to_string(), 1.0);
485 }
486
487 let original_len = input.len().max(1);
488 let mut current = input.to_string();
489
490 for stage in &self.stages {
491 let result = stage.compress(¤t);
492 current = result.text;
493 }
494
495 let ratio = current.len() as f64 / original_len as f64;
496 (current, ratio)
497 }
498
499 pub fn stage_count(&self) -> usize {
501 self.stages.len()
502 }
503}
504
505#[cfg(test)]
506mod tests {
507 use super::*;
508
509 #[test]
510 fn test_whitespace_compressor_collapses_spaces() {
511 let c = WhitespaceCompressor;
512 let r = c.compress("hello world\n\n\n\nfoo");
513 assert_eq!(r.text, "hello world\n\nfoo");
514 assert!(r.ratio < 1.0);
515 }
516
517 #[test]
518 fn test_whitespace_compressor_empty() {
519 let c = WhitespaceCompressor;
520 let r = c.compress("");
521 assert_eq!(r.ratio, 1.0);
522 assert_eq!(r.text, "");
523 }
524
525 #[test]
526 fn test_repetition_remover_deduplicates() {
527 let c = RepetitionRemover;
528 let r = c.compress("Hello world. Hello world. Goodbye.");
529 assert!(!r.text.contains("Hello world. Hello world."), "text={}", r.text);
530 }
531
532 #[test]
533 fn test_sentence_ranker_respects_fraction() {
534 let ranker = SentenceRanker::new(0.5);
535 let input = "The quick brown fox. A lazy dog sat down. Rust is fast. Memory safety matters. Tokio is async. Channels provide backpressure.";
536 let result = ranker.compress(input);
537 assert!(result.ratio < 0.9, "ratio={}", result.ratio);
539 }
540
541 #[test]
542 fn test_truncation_respects_sentence_boundary() {
543 let t = TruncationStrategy::new(30);
544 let input = "Short sentence. Another sentence follows.";
545 let result = t.compress(input);
546 assert!(result.text.len() <= 30 || result.text.ends_with('.') || result.text.ends_with("sentence"));
547 }
548
549 #[test]
550 fn test_truncation_no_op_when_short() {
551 let t = TruncationStrategy::new(1000);
552 let input = "Short.";
553 let result = t.compress(input);
554 assert_eq!(result.ratio, 1.0);
555 assert_eq!(result.text, input);
556 }
557
558 #[test]
559 fn test_pipeline_chains_strategies() {
560 let pipeline = CompressionPipeline::new()
561 .with(Box::new(WhitespaceCompressor))
562 .with(Box::new(RepetitionRemover));
563 let input = "Hello world. Hello world.";
564 let (text, ratio) = pipeline.compress(input);
565 assert!(!text.contains(" "), "should have no double spaces");
566 assert!(ratio < 1.0 || text.len() <= input.len());
567 }
568
569 #[test]
570 fn test_default_rag_pipeline() {
571 let pipeline = CompressionPipeline::default_for_rag();
572 assert_eq!(pipeline.stage_count(), 3);
573 let long_input = "The system encountered an error. The system encountered an error. \
574 Rust provides memory safety. The Tokio runtime handles async I/O. \
575 Circuit breakers prevent cascading failures. Backpressure ensures stability. \
576 Deduplication reduces redundant work. Rate limiting protects downstream services.";
577 let (text, ratio) = pipeline.compress(long_input);
578 assert!(ratio < 1.0, "should compress something, ratio={ratio}");
579 assert!(!text.is_empty());
580 }
581
582 #[test]
583 fn test_empty_pipeline_is_noop() {
584 let pipeline = CompressionPipeline::new();
585 let input = "hello world";
586 let (text, ratio) = pipeline.compress(input);
587 assert_eq!(text, input);
588 assert_eq!(ratio, 1.0);
589 }
590}