Skip to main content

tokio_prompt_orchestrator/
compression.rs

1//! Prompt compression — reduce token count before sending to the model.
2//!
3//! Compressing prompts before inference reduces cost and latency, especially
4//! for long conversation histories or large context retrievals. This module
5//! provides several complementary strategies that can be chained.
6//!
7//! ## Strategies
8//!
9//! | Strategy | What it does | Tokens saved (typical) |
10//! |----------|-------------|------------------------|
11//! | [`WhitespaceCompressor`] | Collapses redundant whitespace/newlines | 2–5% |
12//! | [`RepetitionRemover`] | Removes duplicate paragraphs/sentences | 5–15% |
13//! | [`StopWordFilter`] | Removes low-information filler words | 5–20% |
14//! | [`SentenceRanker`] | Keeps only the top-K most relevant sentences | 20–60% |
15//! | [`TruncationStrategy`] | Hard truncation with smart boundary detection | variable |
16//! | [`CompressionPipeline`] | Chains multiple strategies | additive |
17//!
18//! ## Usage
19//!
20//! ```rust,no_run
21//! use tokio_prompt_orchestrator::compression::{CompressionPipeline, SentenceRanker, WhitespaceCompressor};
22//!
23//! let pipeline = CompressionPipeline::new()
24//!     .with(Box::new(WhitespaceCompressor))
25//!     .with(Box::new(SentenceRanker::new(0.5))); // keep top 50% sentences
26//!
27//! let original = "Long prompt with lots of redundancy...";
28//! let (compressed, ratio) = pipeline.compress(original);
29//! println!("Reduced by {:.1}%", (1.0 - ratio) * 100.0);
30//! ```
31
32use std::collections::HashSet;
33
34/// The output of a compression pass: compressed text + compression ratio.
35///
36/// `ratio` is in `[0.0, 1.0]` where `1.0` means no compression.
37#[derive(Debug, Clone)]
38pub struct CompressionResult {
39    /// The compressed text.
40    pub text: String,
41    /// Ratio of output length to input length (1.0 = unchanged, 0.5 = halved).
42    pub ratio: f64,
43    /// Number of characters removed.
44    pub chars_removed: usize,
45    /// Strategy that produced this result.
46    pub strategy: String,
47}
48
49/// Trait for a single compression strategy.
50pub trait Compressor: Send + Sync {
51    /// Apply this compression strategy to the input text.
52    fn compress(&self, input: &str) -> CompressionResult;
53
54    /// Human-readable name for metrics/logging.
55    fn name(&self) -> &'static str;
56}
57
58// ---------------------------------------------------------------------------
59// WhitespaceCompressor
60// ---------------------------------------------------------------------------
61
62/// Collapses redundant whitespace: multiple spaces → one, 3+ newlines → two.
63///
64/// This is the cheapest compression to apply and should always be first.
65pub struct WhitespaceCompressor;
66
67impl Compressor for WhitespaceCompressor {
68    fn compress(&self, input: &str) -> CompressionResult {
69        // Collapse multiple spaces within lines
70        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
118// ---------------------------------------------------------------------------
119// RepetitionRemover
120// ---------------------------------------------------------------------------
121
122/// Removes duplicate sentences or paragraphs.
123///
124/// Uses exact-match deduplication on sentences split by `. ` or `\n\n`.
125pub struct RepetitionRemover;
126
127impl Compressor for RepetitionRemover {
128    fn compress(&self, input: &str) -> CompressionResult {
129        // Split into sentences on ". " or paragraph breaks
130        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
164// ---------------------------------------------------------------------------
165// StopWordFilter
166// ---------------------------------------------------------------------------
167
168/// Removes common English stop words from the text.
169///
170/// This is aggressive — only use when the downstream task tolerates it
171/// (e.g. topic classification, keyword extraction, not generation).
172pub 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    /// Create a stop word filter with a built-in English stop word list.
184    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
238// ---------------------------------------------------------------------------
239// SentenceRanker (TF-IDF based)
240// ---------------------------------------------------------------------------
241
242/// Keeps only the top-K most relevant sentences using TF-IDF scoring.
243///
244/// `keep_fraction` is in `(0, 1]`: `0.5` keeps the top 50% of sentences.
245/// Sentences are re-assembled in their original order after ranking.
246///
247/// This is the highest-impact compressor for long RAG-retrieved contexts.
248pub struct SentenceRanker {
249    keep_fraction: f64,
250}
251
252impl SentenceRanker {
253    /// Create a sentence ranker.
254    ///
255    /// # Arguments
256    /// * `keep_fraction` — fraction of sentences to retain (0.0–1.0).
257    ///   Clamped to [0.05, 1.0] to avoid empty output.
258    pub fn new(keep_fraction: f64) -> Self {
259        Self {
260            keep_fraction: keep_fraction.clamp(0.05, 1.0),
261        }
262    }
263
264    /// Score sentences by TF-IDF relevance against the entire document.
265    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        // Build term frequency per sentence
273        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                // Normalise by sentence length
285                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        // IDF: log(N / df) for each term
294        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        // Score each sentence as sum of TF-IDF of its terms
303        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        // Split on ". " or "\n\n"
327        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        // Take top-K indices, sorted back to original order
348        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
377// ---------------------------------------------------------------------------
378// TruncationStrategy
379// ---------------------------------------------------------------------------
380
381/// Hard truncation at a character limit, respecting sentence boundaries.
382pub struct TruncationStrategy {
383    max_chars: usize,
384}
385
386impl TruncationStrategy {
387    /// Create a truncator with the given character limit.
388    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        // Find last sentence boundary before limit
405        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
429// ---------------------------------------------------------------------------
430// CompressionPipeline
431// ---------------------------------------------------------------------------
432
433/// Chains multiple compression strategies sequentially.
434///
435/// Each strategy receives the output of the previous one.
436/// Tracks overall statistics across the full pipeline.
437pub 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    /// Create an empty pipeline.
449    pub fn new() -> Self {
450        Self { stages: Vec::new() }
451    }
452
453    /// Add a compression stage to the end of the pipeline.
454    pub fn with(mut self, stage: Box<dyn Compressor>) -> Self {
455        self.stages.push(stage);
456        self
457    }
458
459    /// Build a sensible default pipeline for typical RAG prompts.
460    ///
461    /// Applies: whitespace → repetition → sentence ranking (keep 70%)
462    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    /// Build a pipeline for long conversation history compression.
470    ///
471    /// Applies: whitespace → repetition → sentence ranking (keep 50%)
472    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    /// Compress the input through all stages, returning the final result.
480    ///
481    /// The returned `ratio` reflects the overall compression across all stages.
482    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(&current);
492            current = result.text;
493        }
494
495        let ratio = current.len() as f64 / original_len as f64;
496        (current, ratio)
497    }
498
499    /// Return the number of stages in this pipeline.
500    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        // Should be roughly half the sentences
538        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}