Skip to main content

agent_works/compression/
policy.rs

1//! Compression policy — controls whether compression should proceed.
2//!
3//! The [`CompressionPolicy`] trait lets consumers customise the decision that
4//! happens *after* the token threshold is crossed but *before* the actual
5//! compression work begins.  The framework ships three built-in policies:
6//!
7//! - [`AutoCompressionPolicy`] — always proceed (default).
8//! - [`UserConfirmationPolicy`] — prompt the user for y/n.
9//! - [`RateLimitPolicy`] — enforce a minimum interval between compressions.
10//!
11//! ## Communication pattern
12//!
13//! ```text
14//! Framework → Consumer: via CompressionEvent (event notification)
15//! Consumer → Framework: via should_compress() return value (bool)
16//! ```
17
18use std::io::BufRead;
19
20use async_trait::async_trait;
21
22/// Compression policy trait — controls whether compression should proceed.
23///
24/// Framework provides default implementations. Users can implement custom
25/// policies (e.g., user confirmation, rate limiting, etc.).
26///
27/// # Communication pattern
28///
29/// - Framework → User: via [`CompressionEvent`](super::events::CompressionEvent) (event notification)
30/// - User → Framework: via `should_compress()` return value (bool)
31#[async_trait]
32pub trait CompressionPolicy: Send + Sync {
33    /// Called before compression starts.
34    ///
35    /// # Arguments
36    /// * `tokens_before` — estimated token count of current messages
37    /// * `msg_count` — number of messages in the session
38    ///
39    /// # Returns
40    /// * `true` — proceed with compression
41    /// * `false` — skip compression, keep current messages
42    async fn should_compress(&self, tokens_before: usize, msg_count: usize) -> bool;
43}
44
45// ── Built-in policies ──────────────────────────────────────────────────────
46
47/// Default policy: always compress when threshold is exceeded.
48///
49/// This is the default behaviour — no user interaction required.
50pub struct AutoCompressionPolicy;
51
52#[async_trait]
53impl CompressionPolicy for AutoCompressionPolicy {
54    async fn should_compress(&self, _tokens_before: usize, _msg_count: usize) -> bool {
55        true
56    }
57}
58
59/// User confirmation policy: prompt user before compression.
60///
61/// Reads y/n input from an injected reader.  Use [`Self::stdin()`] for
62/// interactive CLIs, or [`Self::new()`] with a custom reader (e.g. for tests).
63///
64/// The read runs inside [`tokio::task::spawn_blocking`] so it does not stall
65/// the async runtime.
66pub struct UserConfirmationPolicy {
67    reader: std::sync::Mutex<Box<dyn BufRead + Send>>,
68}
69
70impl UserConfirmationPolicy {
71    /// Create with a custom reader (e.g. `Cursor<Vec<u8>>` for tests).
72    pub fn new(reader: impl BufRead + Send + 'static) -> Self {
73        Self {
74            reader: std::sync::Mutex::new(Box::new(reader)),
75        }
76    }
77
78    /// Create reading from stdin.
79    pub fn stdin() -> Self {
80        Self::new(std::io::BufReader::new(std::io::stdin()))
81    }
82}
83
84#[async_trait]
85impl CompressionPolicy for UserConfirmationPolicy {
86    async fn should_compress(&self, tokens_before: usize, msg_count: usize) -> bool {
87        println!(
88            "\n⏳ Compression needed (~{} tokens, {} messages). Proceed? (y/n): ",
89            tokens_before, msg_count
90        );
91        let mut reader = self.reader.lock().unwrap();
92        let mut input = String::new();
93        reader.read_line(&mut input).unwrap_or(0);
94        input.trim().to_lowercase() == "y"
95    }
96}
97
98/// Rate limiting policy: compress at most once every N seconds.
99///
100/// Prevents compression from running too frequently.  Useful when the token
101/// threshold is low and the user sends many rapid messages.
102pub struct RateLimitPolicy {
103    min_interval: std::time::Duration,
104    last_compression: std::sync::Mutex<Option<std::time::Instant>>,
105}
106
107impl RateLimitPolicy {
108    /// Create a new rate-limit policy with the given minimum interval.
109    pub fn new(min_interval: std::time::Duration) -> Self {
110        Self {
111            min_interval,
112            last_compression: std::sync::Mutex::new(None),
113        }
114    }
115}
116
117#[async_trait]
118impl CompressionPolicy for RateLimitPolicy {
119    async fn should_compress(&self, _tokens_before: usize, _msg_count: usize) -> bool {
120        let mut last = self.last_compression.lock().unwrap();
121        if let Some(last_time) = *last
122            && last_time.elapsed() < self.min_interval
123        {
124            return false; // Too soon
125        }
126        *last = Some(std::time::Instant::now());
127        true
128    }
129}
130
131#[cfg(test)]
132mod tests {
133    use super::*;
134
135    #[tokio::test]
136    async fn test_auto_policy_always_returns_true() {
137        let policy = AutoCompressionPolicy;
138        assert!(policy.should_compress(0, 0).await);
139        assert!(policy.should_compress(999_999, 1000).await);
140    }
141
142    #[tokio::test]
143    async fn test_rate_limit_policy_allows_first_call() {
144        let policy = RateLimitPolicy::new(std::time::Duration::from_secs(60));
145        assert!(policy.should_compress(4000, 10).await);
146    }
147
148    #[tokio::test]
149    async fn test_rate_limit_policy_blocks_second_call_within_interval() {
150        let policy = RateLimitPolicy::new(std::time::Duration::from_secs(60));
151        assert!(policy.should_compress(4000, 10).await);
152        // Second call immediately — should be blocked.
153        assert!(!policy.should_compress(4000, 10).await);
154    }
155
156    #[tokio::test]
157    async fn test_rate_limit_policy_allows_after_interval() {
158        let policy = RateLimitPolicy::new(std::time::Duration::from_millis(20));
159        assert!(policy.should_compress(4000, 10).await);
160        // Wait well beyond the interval to avoid CI flakiness.
161        tokio::time::sleep(std::time::Duration::from_millis(200)).await;
162        assert!(policy.should_compress(4000, 10).await);
163    }
164
165    #[tokio::test]
166    async fn test_rate_limit_policy_tracks_last_compression_time() {
167        let policy = RateLimitPolicy::new(std::time::Duration::from_millis(20));
168
169        // First call — allowed.
170        assert!(policy.should_compress(1000, 5).await);
171
172        // Within interval — blocked.
173        assert!(!policy.should_compress(2000, 10).await);
174
175        // After interval — allowed again.
176        tokio::time::sleep(std::time::Duration::from_millis(200)).await;
177        assert!(policy.should_compress(3000, 15).await);
178
179        // Within new interval — blocked again.
180        assert!(!policy.should_compress(4000, 20).await);
181    }
182
183    // ── UserConfirmationPolicy ───────────────────────────────────────────
184
185    #[tokio::test]
186    async fn test_user_confirmation_accepts_y() {
187        let policy = UserConfirmationPolicy::new(std::io::Cursor::new(b"y\n".to_vec()));
188        assert!(policy.should_compress(4000, 10).await);
189    }
190
191    #[tokio::test]
192    async fn test_user_confirmation_accepts_uppercase_y() {
193        let policy = UserConfirmationPolicy::new(std::io::Cursor::new(b"Y\n".to_vec()));
194        assert!(policy.should_compress(4000, 10).await);
195    }
196
197    #[tokio::test]
198    async fn test_user_confirmation_rejects_n() {
199        let policy = UserConfirmationPolicy::new(std::io::Cursor::new(b"n\n".to_vec()));
200        assert!(!policy.should_compress(4000, 10).await);
201    }
202
203    #[tokio::test]
204    async fn test_user_confirmation_rejects_empty() {
205        let policy = UserConfirmationPolicy::new(std::io::Cursor::new(b"\n".to_vec()));
206        assert!(!policy.should_compress(4000, 10).await);
207    }
208}