Skip to main content

talos_agent/
compression.rs

1//! Deterministic output compression for tool results entering model context.
2//!
3//! This module implements MEM-007 active context compression, starting with
4//! `bash` tool output. Compression applies only to the model-facing
5//! representation; raw output is preserved on the UI event/export surface.
6//!
7//! # Design Constraints
8//!
9//! - **Deterministic**: same input always produces same compressed bytes.
10//!   No timestamps, no random state, no session-dependent ordering.
11//! - **Default OFF**: compression must be explicitly enabled via
12//!   [`Agent::with_bash_compression`].
13//! - **Stable prefix safe**: compression applies only to tool results entering
14//!   the dynamic suffix, never to cached-prefix messages.
15//! - **bash only**: other tools are unaffected by this module.
16
17/// Default line threshold for bash output compression.
18///
19/// When bash output exceeds this many lines, it is compressed to the last N
20/// lines plus a truncation marker.
21pub const DEFAULT_BASH_LINE_THRESHOLD: usize = 30;
22
23/// Truncation marker prepended to compressed bash output.
24///
25/// The `{omitted}` placeholder is replaced with the count of omitted lines.
26const TRUNCATION_MARKER_TEMPLATE: &str =
27    "\n... (first {omitted} lines omitted, see /export for full output)\n";
28
29/// Result of compressing bash tool output.
30#[derive(Debug, Clone)]
31pub struct CompressedOutput {
32    /// The model-facing content (compressed or original).
33    pub content: String,
34    /// Number of characters in the original input.
35    pub original_size: usize,
36    /// Number of characters in the output content.
37    pub compressed_size: usize,
38    /// The compression strategy applied.
39    pub strategy: &'static str,
40}
41
42/// Deterministic compressor for bash tool output.
43///
44/// When enabled and bash output exceeds the configured line threshold, the
45/// model-facing content is compressed to the last N lines plus a truncation
46/// marker. Full output is preserved on the UI event/export surface.
47#[derive(Debug, Clone, Copy)]
48pub struct BashOutputCompressor {
49    /// Maximum number of lines to retain from the end of the output.
50    line_threshold: usize,
51}
52
53impl BashOutputCompressor {
54    /// Creates a new compressor with the default line threshold (30 lines).
55    #[must_use]
56    pub fn new() -> Self {
57        Self {
58            line_threshold: DEFAULT_BASH_LINE_THRESHOLD,
59        }
60    }
61
62    /// Creates a new compressor with a custom line threshold.
63    #[must_use]
64    pub fn with_threshold(line_threshold: usize) -> Self {
65        Self { line_threshold }
66    }
67
68    /// Compresses bash output if it exceeds the line threshold.
69    ///
70    /// When the number of lines is at or below the threshold, the content is
71    /// returned unchanged with `strategy = "none"`.
72    ///
73    /// When the number of lines exceeds the threshold, the output is compressed
74    /// to the last N lines with a truncation marker, and `strategy = "last_n_lines"`.
75    ///
76    /// # Determinism
77    ///
78    /// This method is fully deterministic: the same input string always produces
79    /// the same output bytes. No timestamps, random state, or external context
80    /// is used.
81    #[must_use]
82    pub fn compress(&self, content: &str) -> CompressedOutput {
83        let original_size = content.len();
84        let lines: Vec<&str> = content.lines().collect();
85        let line_count = lines.len();
86
87        if line_count <= self.line_threshold {
88            return CompressedOutput {
89                content: content.to_string(),
90                original_size,
91                compressed_size: original_size,
92                strategy: "none",
93            };
94        }
95
96        let omitted = line_count - self.line_threshold;
97        let retained = &lines[omitted..];
98
99        let marker = TRUNCATION_MARKER_TEMPLATE.replace("{omitted}", &omitted.to_string());
100        let mut compressed = String::with_capacity(
101            marker.len() + retained.iter().map(|l| l.len() + 1).sum::<usize>(),
102        );
103        compressed.push_str(&marker);
104        for (i, line) in retained.iter().enumerate() {
105            if i > 0 {
106                compressed.push('\n');
107            }
108            compressed.push_str(line);
109        }
110        // Preserve trailing newline if the original had one
111        if content.ends_with('\n') && !compressed.ends_with('\n') {
112            compressed.push('\n');
113        }
114
115        let compressed_size = compressed.len();
116        CompressedOutput {
117            content: compressed,
118            original_size,
119            compressed_size,
120            strategy: "last_n_lines",
121        }
122    }
123}
124
125impl Default for BashOutputCompressor {
126    fn default() -> Self {
127        Self::new()
128    }
129}
130
131#[derive(Debug, Clone, Default)]
132pub struct CompressionMetrics {
133    pub compression_events: u64,
134    pub bytes_before: u64,
135    pub bytes_after: u64,
136}
137
138impl CompressionMetrics {
139    pub fn record(&mut self, output: &CompressedOutput) {
140        if output.strategy == "none" {
141            return;
142        }
143        self.compression_events += 1;
144        self.bytes_before += output.original_size as u64;
145        self.bytes_after += output.compressed_size as u64;
146    }
147
148    pub fn bytes_saved(&self) -> u64 {
149        self.bytes_before.saturating_sub(self.bytes_after)
150    }
151
152    pub fn estimated_tokens_saved(&self) -> u64 {
153        self.bytes_saved() / 4
154    }
155}
156
157#[derive(Debug, Clone, Default)]
158pub struct RetrievalMetrics {
159    pub recall_calls: u64,
160    pub results_returned: u64,
161}
162
163impl RetrievalMetrics {
164    pub fn record_recall(&mut self, result_count: usize) {
165        self.recall_calls += 1;
166        self.results_returned += result_count as u64;
167    }
168
169    pub fn avg_results_per_call(&self) -> f64 {
170        if self.recall_calls == 0 {
171            0.0
172        } else {
173            self.results_returned as f64 / self.recall_calls as f64
174        }
175    }
176}
177
178#[cfg(test)]
179mod tests {
180    use super::*;
181
182    fn make_lines(n: usize) -> String {
183        (0..n)
184            .map(|i| format!("line {}", i))
185            .collect::<Vec<_>>()
186            .join("\n")
187    }
188
189    #[test]
190    fn short_output_no_compression() {
191        let compressor = BashOutputCompressor::new();
192        let content = make_lines(10);
193        let result = compressor.compress(&content);
194
195        assert_eq!(result.strategy, "none");
196        assert_eq!(result.content, content);
197        assert_eq!(result.original_size, result.compressed_size);
198    }
199
200    #[test]
201    fn exactly_threshold_no_compression() {
202        let compressor = BashOutputCompressor::new();
203        let content = make_lines(DEFAULT_BASH_LINE_THRESHOLD);
204        let result = compressor.compress(&content);
205
206        assert_eq!(result.strategy, "none");
207        assert_eq!(result.content, content);
208    }
209
210    #[test]
211    fn one_over_threshold_compressed() {
212        let compressor = BashOutputCompressor::new();
213        let content = make_lines(DEFAULT_BASH_LINE_THRESHOLD + 1);
214        let result = compressor.compress(&content);
215
216        assert_eq!(result.strategy, "last_n_lines");
217        assert!(result.content.contains("first 1 lines omitted"));
218        // Should contain last 30 lines
219        assert!(result.content.contains("line 1")); // line index 1 = second line (first retained)
220    }
221
222    #[test]
223    fn long_output_compressed_to_last_n() {
224        let compressor = BashOutputCompressor::new();
225        let content = make_lines(100);
226        let result = compressor.compress(&content);
227
228        assert_eq!(result.strategy, "last_n_lines");
229        assert!(result.content.contains("first 70 lines omitted"));
230        // Last retained line should be "line 99"
231        assert!(result.content.contains("line 99"));
232        // First retained line should be "line 70"
233        assert!(result.content.contains("line 70"));
234        // First omitted line "line 0" should NOT appear
235        assert!(!result.content.contains("line 0"));
236        assert!(!result.content.contains("line 69"));
237    }
238
239    #[test]
240    fn determinism() {
241        let compressor = BashOutputCompressor::new();
242        let content = make_lines(100);
243
244        let result1 = compressor.compress(&content);
245        let result2 = compressor.compress(&content);
246
247        assert_eq!(result1.content, result2.content);
248        assert_eq!(result1.original_size, result2.original_size);
249        assert_eq!(result1.compressed_size, result2.compressed_size);
250        assert_eq!(result1.strategy, result2.strategy);
251    }
252
253    #[test]
254    fn trailing_newline_preserved() {
255        let compressor = BashOutputCompressor::new();
256        let content = format!("{}\n", make_lines(100));
257        let result = compressor.compress(&content);
258
259        assert!(result.content.ends_with('\n'));
260    }
261
262    #[test]
263    fn no_trailing_newline_preserved() {
264        let compressor = BashOutputCompressor::new();
265        let content = make_lines(100);
266        let result = compressor.compress(&content);
267
268        assert!(!result.content.ends_with('\n'));
269    }
270
271    #[test]
272    fn empty_input_no_compression() {
273        let compressor = BashOutputCompressor::new();
274        let result = compressor.compress("");
275
276        assert_eq!(result.strategy, "none");
277        assert_eq!(result.content, "");
278    }
279
280    #[test]
281    fn custom_threshold() {
282        let compressor = BashOutputCompressor::with_threshold(5);
283        let content = make_lines(10);
284        let result = compressor.compress(&content);
285
286        assert_eq!(result.strategy, "last_n_lines");
287        assert!(result.content.contains("first 5 lines omitted"));
288        // Should retain last 5 lines: line 5 through line 9
289        assert!(result.content.contains("line 5"));
290        assert!(result.content.contains("line 9"));
291        assert!(!result.content.contains("line 4"));
292    }
293
294    #[test]
295    fn metadata_accuracy() {
296        let compressor = BashOutputCompressor::new();
297        let content = make_lines(100);
298        let result = compressor.compress(&content);
299
300        assert_eq!(result.original_size, content.len());
301        assert_eq!(result.compressed_size, result.content.len());
302    }
303
304    #[test]
305    fn metrics_accumulate_across_events() {
306        let compressor = BashOutputCompressor::new();
307        let mut metrics = CompressionMetrics::default();
308
309        let r1 = compressor.compress(&make_lines(50));
310        let r2 = compressor.compress(&make_lines(100));
311        let r3 = compressor.compress(&make_lines(10));
312
313        metrics.record(&r1);
314        metrics.record(&r2);
315        metrics.record(&r3);
316
317        assert_eq!(metrics.compression_events, 2);
318        assert!(metrics.bytes_saved() > 0);
319        assert!(metrics.bytes_before > metrics.bytes_after);
320    }
321
322    #[test]
323    fn metrics_skip_none_strategy() {
324        let compressor = BashOutputCompressor::new();
325        let mut metrics = CompressionMetrics::default();
326
327        let r = compressor.compress(&make_lines(10));
328        metrics.record(&r);
329
330        assert_eq!(metrics.compression_events, 0);
331        assert_eq!(metrics.bytes_saved(), 0);
332    }
333
334    #[test]
335    fn metrics_estimated_tokens_saved() {
336        let compressor = BashOutputCompressor::new();
337        let mut metrics = CompressionMetrics::default();
338
339        metrics.record(&compressor.compress(&make_lines(100)));
340
341        assert!(metrics.estimated_tokens_saved() > 0);
342    }
343
344    #[test]
345    fn retrieval_metrics_track_calls_and_results() {
346        let mut metrics = RetrievalMetrics::default();
347
348        metrics.record_recall(5);
349        metrics.record_recall(3);
350        metrics.record_recall(0);
351
352        assert_eq!(metrics.recall_calls, 3);
353        assert_eq!(metrics.results_returned, 8);
354        assert!((metrics.avg_results_per_call() - 2.666).abs() < 0.1);
355    }
356
357    #[test]
358    fn retrieval_metrics_empty_has_zero_avg() {
359        let metrics = RetrievalMetrics::default();
360        assert_eq!(metrics.avg_results_per_call(), 0.0);
361    }
362}