Skip to main content

agent_works/compression/
middleware.rs

1//! Middleware integration for context compression.
2//!
3//! [`CompressionMiddleware`] implements the [`agent_base::Middleware`] trait,
4//! compressing the per-LLM-call message list on every `on_pre_llm` invocation.
5//!
6//! The middleware delegates to [`ContextCompactor`] which handles the hybrid
7//! retention strategy (system prompt + summary + recent N messages) and the
8//! stable-prefix cache.  The session's stored history is **never** modified;
9//! only the mutable message copy in [`PreLlmCtx`] is replaced.
10//!
11//! A [`CompressionPolicy`] controls whether compression proceeds once the token
12//! threshold is crossed.  The default [`AutoCompressionPolicy`] always proceeds;
13//! custom policies can add user confirmation, rate limiting, etc.
14
15use std::sync::Arc;
16use std::sync::atomic::{AtomicUsize, Ordering};
17
18use agent_base::{AgentResult, ChatMessage, Middleware, PreLlmCtx};
19
20use crate::compression::compactor::{ContextCompactor, estimate_message_tokens};
21use std::sync::Once;
22
23use crate::compression::config::CompressionConfig;
24use crate::compression::events::{CompressionEvent, CompressionTrigger};
25use crate::compression::filter::is_summary_message;
26use crate::compression::policy::{AutoCompressionPolicy, CompressionPolicy};
27
28/// Write compression before/after log to `/tmp/phi-agent-compression/`.
29fn write_compression_log(
30    session_id: u64,
31    before: &[ChatMessage],
32    after: &[ChatMessage],
33    tokens_before: usize,
34    tokens_after: usize,
35    reduction_pct: i32,
36) {
37    use std::io::Write;
38
39    let dir = std::path::Path::new("/tmp/phi-agent-compression");
40    if let Err(e) = std::fs::create_dir_all(dir) {
41        tracing::warn!("Failed to create compression log dir: {e}");
42        return;
43    }
44
45    let ts = std::time::SystemTime::now()
46        .duration_since(std::time::UNIX_EPOCH)
47        .unwrap_or_default()
48        .as_secs();
49    let filename = dir.join(format!("session_{session_id}_{ts}.json"));
50
51    let log = serde_json::json!({
52        "session_id": session_id,
53        "timestamp": ts,
54        "evaluation": {
55            "tokens_before": tokens_before,
56            "tokens_after": tokens_after,
57            "reduction_pct": reduction_pct,
58            "msg_count_before": before.len(),
59            "msg_count_after": after.len(),
60        },
61        "before": before.iter().map(format_message).collect::<Vec<_>>(),
62        "after": after.iter().map(format_message).collect::<Vec<_>>(),
63    });
64
65    match std::fs::File::create(&filename) {
66        Ok(mut f) => {
67            if let Err(e) = f.write_all(serde_json::to_string_pretty(&log).unwrap().as_bytes()) {
68                tracing::warn!("Failed to write compression log: {e}");
69            } else {
70                tracing::info!("Compression log written to {}", filename.display());
71            }
72        }
73        Err(e) => tracing::warn!("Failed to create compression log file: {e}"),
74    }
75}
76
77/// Format a ChatMessage into a readable JSON value for logging.
78fn format_message(msg: &ChatMessage) -> serde_json::Value {
79    match msg {
80        ChatMessage::System { content, .. } => serde_json::json!({
81            "role": "system",
82            "content": content,
83        }),
84        ChatMessage::User { content, .. } => serde_json::json!({
85            "role": "user",
86            "content": content,
87        }),
88        ChatMessage::Assistant {
89            content,
90            tool_calls,
91            ..
92        } => {
93            let mut m = serde_json::json!({
94                "role": "assistant",
95                "content": content,
96            });
97            if let Some(tc) = tool_calls {
98                m["tool_calls"] = serde_json::to_value(tc).unwrap_or_default();
99            }
100            m
101        }
102        ChatMessage::Tool {
103            content,
104            tool_call_id,
105            ..
106        } => serde_json::json!({
107            "role": "tool",
108            "tool_call_id": tool_call_id,
109            "content": content,
110        }),
111        ChatMessage::Custom { role, data } => serde_json::json!({
112            "role": role,
113            "data": data,
114        }),
115    }
116}
117
118/// Middleware that compresses conversation history before each LLM call.
119///
120/// Wraps a [`ContextCompactor`] and forwards `on_pre_llm` to its [`compact`]
121/// method.  When the estimated token count is below [`CompressionConfig::trigger_tokens`],
122/// the middleware is a no-op and the original messages pass through unchanged.
123///
124/// A [`CompressionPolicy`] controls whether compression actually proceeds once
125/// the threshold is crossed.  Use [`with_policy`](Self::with_policy) to supply
126/// a custom policy (e.g. user confirmation, rate limiting).
127///
128/// # Usage
129///
130/// ```ignore
131/// use agent_works::compression::{CompressionConfig, CompressionMiddleware};
132///
133/// let mw = CompressionMiddleware::new(config, client);
134/// builder.middleware(mw);
135/// ```
136///
137/// [`compact`]: ContextCompactor::compact
138pub struct CompressionMiddleware {
139    compactor: ContextCompactor,
140    policy: Box<dyn CompressionPolicy>,
141    /// Message count after the last successful compression.
142    /// Used to calculate only the *new* old-block tokens for threshold checks,
143    /// preventing re-triggering caused by recent messages rolling into the old block.
144    last_compressed_msg_count: AtomicUsize,
145}
146
147#[allow(missing_docs)]
148impl CompressionMiddleware {
149    /// Create a new middleware with the given config and LLM client.
150    ///
151    /// Uses [`AutoCompressionPolicy`] (always compress when threshold is hit).
152    pub fn new(
153        config: CompressionConfig,
154        client: Arc<dyn agent_base::llm_trait::LlmProvider>,
155    ) -> Self {
156        Self {
157            compactor: ContextCompactor::new(client, config),
158            policy: Box::new(AutoCompressionPolicy),
159            last_compressed_msg_count: AtomicUsize::new(0),
160        }
161    }
162
163    /// Create a new middleware from an existing [`ContextCompactor`].
164    ///
165    /// Uses [`AutoCompressionPolicy`].
166    pub fn from_compactor(compactor: ContextCompactor) -> Self {
167        Self {
168            compactor,
169            policy: Box::new(AutoCompressionPolicy),
170            last_compressed_msg_count: AtomicUsize::new(0),
171        }
172    }
173
174    /// Create a new middleware with a custom [`CompressionPolicy`].
175    pub fn with_policy(
176        config: CompressionConfig,
177        client: Arc<dyn agent_base::llm_trait::LlmProvider>,
178        policy: Box<dyn CompressionPolicy>,
179    ) -> Self {
180        Self {
181            compactor: ContextCompactor::new(client, config),
182            policy,
183            last_compressed_msg_count: AtomicUsize::new(0),
184        }
185    }
186
187    /// Access the inner compactor (e.g. for `/compact` or cache clearing).
188    pub fn compactor(&self) -> &ContextCompactor {
189        &self.compactor
190    }
191
192    /// Create a cloned handle of the inner compactor.
193    ///
194    /// The clone shares the same cache — clearing it through either handle
195    /// affects both.  Useful for storing a separate handle outside the
196    /// middleware (e.g. in `PhiAgent` for `/compact` access).
197    pub fn clone_compactor(&self) -> ContextCompactor {
198        self.compactor.clone_handle()
199    }
200
201    /// Access the compression config.
202    pub fn config(&self) -> &CompressionConfig {
203        self.compactor.config()
204    }
205}
206
207#[async_trait::async_trait]
208impl Middleware for CompressionMiddleware {
209    async fn on_pre_llm(&self, ctx: &mut PreLlmCtx) -> AgentResult<()> {
210        let t0 = std::time::Instant::now();
211        let msg_count = ctx.messages.len();
212
213        // Quick check — below minimum message count, skip entirely.
214        // Need at least system + keep_recent + 1 to have anything to compress.
215        let keep = self.config().keep_recent_messages;
216        if msg_count <= keep + 1 {
217            return Ok(());
218        }
219
220        // Only count tokens added SINCE the last compression (new old-block
221        // content).  This prevents re-triggering caused by recent messages
222        // rolling into the old block after compression.
223        //
224        // Also skip system and existing summary messages.
225        let last_compressed = self.last_compressed_msg_count.load(Ordering::Relaxed);
226        let tokens_before: usize = ctx
227            .messages
228            .iter()
229            .skip(last_compressed)
230            .filter(|m| !matches!(m, ChatMessage::System { .. }))
231            .filter(|m| !is_summary_message(m))
232            .map(estimate_message_tokens)
233            .sum();
234        tracing::info!(
235            tokens_before,
236            msg_count,
237            trigger = self.config().trigger_tokens,
238            "[compression-timing] threshold check"
239        );
240
241        // Quick check — below threshold, skip entirely.
242        if tokens_before <= self.config().trigger_tokens {
243            return Ok(());
244        }
245
246        // Defensive: warn once if trigger_tokens >= 1M.
247        // The caller should ensure trigger_tokens < context_window.
248        static TRIGGER_WARN: Once = Once::new();
249        if self.config().trigger_tokens >= 1_000_000 {
250            TRIGGER_WARN.call_once(|| {
251                tracing::warn!(
252                    trigger_tokens = self.config().trigger_tokens,
253                    "trigger_tokens >= 1M — likely exceeds context_window; compression may never fire efficiently"
254                );
255            });
256        }
257
258        tracing::info!(
259            elapsed_ms = t0.elapsed().as_millis() as u64,
260            "[compression-timing] threshold passed"
261        );
262
263        // Policy check — ask policy whether to proceed.
264        if !self.policy.should_compress(tokens_before, msg_count).await {
265            return Ok(());
266        }
267
268        tracing::info!(
269            elapsed_ms = t0.elapsed().as_millis() as u64,
270            "[compression-timing] policy passed, entering compact"
271        );
272
273        // Determine trigger type (manual if /compact, auto otherwise).
274        // For now, middleware is always auto; /compact goes through a different path.
275        let trigger = CompressionTrigger::Auto;
276        let sid = ctx.session_id.id;
277
278        // Discard old summary before re-compressing.
279        // compact() handles preserving user messages and summarizing assistant/tool.
280        // We only need to remove the old summary so it doesn't get re-summarized.
281        let filtered: Vec<ChatMessage> = ctx
282            .messages
283            .iter()
284            .filter(|m| !is_summary_message(m))
285            .cloned()
286            .collect();
287
288        // After filtering, check if there's enough content to compress.
289        // Need more than keep_recent messages to have an old block at all.
290        let keep = self.config().keep_recent_messages;
291        if filtered.len() <= keep + 1 {
292            self.last_compressed_msg_count
293                .store(msg_count, Ordering::Relaxed);
294            return Ok(());
295        }
296
297        // Save messages before compression for logging.
298        let messages_before = ctx.messages.clone();
299        let t_compact = std::time::Instant::now();
300
301        match self
302            .compactor
303            .compact(
304                sid,
305                &filtered,
306                trigger.clone(),
307                Some(&|ev| ctx.emit(ev.into_user_event())),
308            )
309            .await?
310        {
311            Some(compressed) => {
312                tracing::info!(
313                    elapsed_ms = t_compact.elapsed().as_millis() as u64,
314                    "[compression-timing] compact() returned Some"
315                );
316                ctx.messages = compressed;
317                // Record the pre-compression message count so the next threshold
318                // check only counts new content added since this compression.
319                // Using pre-compression count (not compressed length) because
320                // skip() operates on the full message array structure.
321                self.last_compressed_msg_count
322                    .store(msg_count, Ordering::Relaxed);
323                // Count tokens after compression (same scope as tokens_before).
324                // tokens_before counts all non-system, non-summary messages from last_compressed.
325                // tokens_after should count the same scope after compression.
326                // Since compression replaces the old block with a summary, we count
327                // all non-system messages in the compressed result (including the new summary).
328                let tokens_after: usize = ctx
329                    .messages
330                    .iter()
331                    .filter(|m| !matches!(m, ChatMessage::System { .. }))
332                    .map(estimate_message_tokens)
333                    .sum();
334                let reduction_pct = if tokens_before > 0 {
335                    ((tokens_before as f64 - tokens_after as f64) / tokens_before as f64 * 100.0)
336                        .round() as i32
337                } else {
338                    0
339                };
340
341                // Send Completed event.
342                ctx.emit(
343                    CompressionEvent::Completed {
344                        session_id: sid,
345                        tokens_before,
346                        tokens_after,
347                        reduction_pct,
348                        msg_count_before: msg_count,
349                        msg_count_after: ctx.messages.len(),
350                        trigger,
351                    }
352                    .into_user_event(),
353                );
354
355                // Write before/after log to /tmp/phi-agent-compression/.
356                write_compression_log(
357                    ctx.session_id.id,
358                    &messages_before,
359                    &ctx.messages,
360                    tokens_before,
361                    tokens_after,
362                    reduction_pct,
363                );
364            }
365            None => {
366                // Compression skipped (threshold not reached, disabled, or too few messages).
367                // This is a normal no-op — do NOT send Failed event.
368                // compact() is pure and never modified ctx.messages, so no restore needed.
369                //
370                // Still update the checkpoint so the next threshold check only
371                // counts truly new content, preventing a re-trigger loop.
372                self.last_compressed_msg_count
373                    .store(msg_count, Ordering::Relaxed);
374            }
375        }
376        Ok(())
377    }
378}
379
380#[cfg(test)]
381mod tests {
382    use super::*;
383    use crate::compression::events::CompressionEvent;
384    use crate::compression::policy::{CompressionPolicy, RateLimitPolicy};
385    use agent_base::llm_trait::response::FinishReason;
386    use agent_base::llm_trait::types::UsageInfo;
387    use agent_base::llm_trait::{
388        Capabilities, ChatRequest, ChatResponse, ChatStream, LlmError, LlmProvider, ProviderInfo,
389    };
390    use agent_base::{ChatMessage, Middleware, PreLlmCtx, SessionId};
391
392    // ── Test helpers ──────────────────────────────────────────────────────
393
394    /// Minimal mock that returns a fixed string and counts calls.
395    struct MockClient {
396        response: &'static str,
397        calls: std::sync::Arc<std::sync::atomic::AtomicUsize>,
398    }
399
400    #[async_trait::async_trait]
401    impl LlmProvider for MockClient {
402        async fn stream(&self, _request: ChatRequest) -> Result<ChatStream, LlmError> {
403            self.calls.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
404            let response = self.response.to_string();
405            Ok(ChatStream::new(Box::pin(futures_util::stream::once(
406                async move { Ok(agent_base::StreamChunk::Text(response)) },
407            ))))
408        }
409
410        async fn chat(&self, _request: ChatRequest) -> Result<ChatResponse, LlmError> {
411            self.calls.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
412            Ok(ChatResponse {
413                content: self.response.to_string(),
414                tool_calls: vec![],
415                usage: UsageInfo::default(),
416                finish_reason: FinishReason::Stop,
417                raw: None,
418                reasoning_content: None,
419                thinking_signature: None,
420            })
421        }
422
423        fn capabilities(&self) -> Capabilities {
424            Capabilities::default()
425        }
426
427        fn info(&self) -> ProviderInfo {
428            ProviderInfo {
429                name: "stub".to_string(),
430                model: "stub-model".to_string(),
431                version: None,
432            }
433        }
434    }
435
436    /// Mock that always returns an error.
437    struct FailingClient;
438
439    #[async_trait::async_trait]
440    impl LlmProvider for FailingClient {
441        async fn stream(&self, _request: ChatRequest) -> Result<ChatStream, LlmError> {
442            Err(LlmError::llm("summarisation failed"))
443        }
444
445        async fn chat(&self, _request: ChatRequest) -> Result<ChatResponse, LlmError> {
446            Err(LlmError::llm("summarisation failed"))
447        }
448
449        fn capabilities(&self) -> Capabilities {
450            Capabilities::default()
451        }
452
453        fn info(&self) -> ProviderInfo {
454            ProviderInfo {
455                name: "stub".to_string(),
456                model: "stub-model".to_string(),
457                version: None,
458            }
459        }
460    }
461
462    /// Policy that records calls and returns a fixed value.
463    struct SpyPolicy {
464        /// `(tokens_before, msg_count)` for each call.
465        observed: std::sync::Arc<std::sync::Mutex<Vec<(usize, usize)>>>,
466        result: bool,
467    }
468
469    impl SpyPolicy {
470        #[allow(clippy::type_complexity)]
471        fn new(result: bool) -> (Self, std::sync::Arc<std::sync::Mutex<Vec<(usize, usize)>>>) {
472            let observed = std::sync::Arc::new(std::sync::Mutex::new(Vec::new()));
473            (
474                Self {
475                    observed: observed.clone(),
476                    result,
477                },
478                observed,
479            )
480        }
481    }
482
483    #[async_trait::async_trait]
484    impl CompressionPolicy for SpyPolicy {
485        async fn should_compress(&self, tokens_before: usize, msg_count: usize) -> bool {
486            self.observed
487                .lock()
488                .unwrap()
489                .push((tokens_before, msg_count));
490            self.result
491        }
492    }
493
494    fn make_ctx(messages: Vec<ChatMessage>) -> PreLlmCtx {
495        PreLlmCtx {
496            session_id: SessionId::new(1),
497            messages,
498            tools: vec![],
499            emit_fn: None,
500            turn_count: 1,
501            max_turns: 50,
502        }
503    }
504
505    fn make_ctx_with_events(
506        messages: Vec<ChatMessage>,
507        events: std::sync::Arc<std::sync::Mutex<Vec<CompressionEvent>>>,
508    ) -> PreLlmCtx {
509        let events_clone = events.clone();
510        PreLlmCtx {
511            session_id: SessionId::new(1),
512            messages,
513            tools: vec![],
514            emit_fn: Some(Box::new(move |event: agent_base::UserEvent| {
515                if let Some(ev) = CompressionEvent::from_user_event(&event) {
516                    events_clone.lock().unwrap().push(ev);
517                }
518            })),
519            turn_count: 1,
520            max_turns: 50,
521        }
522    }
523
524    fn make_messages(count: usize) -> Vec<ChatMessage> {
525        let mut msgs = vec![ChatMessage::system("You are a test agent.")];
526        for i in 0..count {
527            msgs.push(ChatMessage::user(format!("question {i}")));
528            msgs.push(ChatMessage::assistant(format!(
529                "answer {i} with some extra content to make the old block large enough"
530            )));
531        }
532        msgs
533    }
534
535    // ── CompressionMiddleware::on_pre_llm ─────────────────────────────────
536
537    #[tokio::test]
538    async fn test_middleware_noop_when_below_threshold() {
539        let calls = std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0));
540        let client = std::sync::Arc::new(MockClient {
541            response: "summary",
542            calls: calls.clone(),
543        });
544        let config = CompressionConfig::default().with_trigger_tokens(999_999);
545        let mw = CompressionMiddleware::new(config, client);
546
547        let msgs = make_messages(5);
548        let original_len = msgs.len();
549        let mut ctx = make_ctx(msgs);
550
551        mw.on_pre_llm(&mut ctx).await.unwrap();
552
553        assert_eq!(ctx.messages.len(), original_len);
554        assert_eq!(calls.load(std::sync::atomic::Ordering::SeqCst), 0);
555    }
556
557    #[tokio::test]
558    async fn test_middleware_compresses_when_above_threshold() {
559        let calls = std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0));
560        let client = std::sync::Arc::new(MockClient {
561            response: "compressed summary of earlier work",
562            calls: calls.clone(),
563        });
564        let config = CompressionConfig::default()
565            .with_trigger_tokens(1)
566            .with_keep_recent_messages(4);
567        let mw = CompressionMiddleware::new(config, client);
568
569        let msgs = make_messages(20);
570        let mut ctx = make_ctx(msgs);
571
572        mw.on_pre_llm(&mut ctx).await.unwrap();
573
574        assert!(
575            ctx.messages.len() < 41,
576            "expected compression, got {} messages",
577            ctx.messages.len()
578        );
579        assert!(matches!(&ctx.messages[0], ChatMessage::System { .. }));
580
581        use crate::compression::SUMMARY_PREFIX;
582        let has_summary = ctx.messages.iter().any(|m| match m {
583            ChatMessage::User { content, .. } => content.starts_with(SUMMARY_PREFIX),
584            _ => false,
585        });
586        assert!(has_summary, "expected summary in compressed output");
587        assert_eq!(calls.load(std::sync::atomic::Ordering::SeqCst), 1);
588    }
589
590    #[tokio::test]
591    async fn test_middleware_uses_cache_on_repeated_calls() {
592        let calls = std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0));
593        let client = std::sync::Arc::new(MockClient {
594            response: "cached summary",
595            calls: calls.clone(),
596        });
597        let config = CompressionConfig::default()
598            .with_trigger_tokens(1)
599            .with_keep_recent_messages(4);
600        let mw = CompressionMiddleware::new(config, client);
601
602        let msgs = make_messages(20);
603
604        let mut ctx1 = make_ctx(msgs.clone());
605        mw.on_pre_llm(&mut ctx1).await.unwrap();
606        assert_eq!(calls.load(std::sync::atomic::Ordering::SeqCst), 1);
607
608        let mut ctx2 = make_ctx(msgs);
609        mw.on_pre_llm(&mut ctx2).await.unwrap();
610        assert_eq!(calls.load(std::sync::atomic::Ordering::SeqCst), 1);
611    }
612
613    #[tokio::test]
614    async fn test_middleware_fallback_on_summarisation_failure() {
615        let client = std::sync::Arc::new(FailingClient);
616        let config = CompressionConfig::default()
617            .with_trigger_tokens(1)
618            .with_keep_recent_messages(4);
619        let mw = CompressionMiddleware::new(config, client);
620
621        let msgs = make_messages(20);
622        let mut ctx = make_ctx(msgs);
623
624        mw.on_pre_llm(&mut ctx).await.unwrap();
625
626        assert!(matches!(&ctx.messages[0], ChatMessage::System { .. }));
627
628        use crate::compression::SUMMARY_PREFIX;
629        let has_summary = ctx.messages.iter().any(|m| match m {
630            ChatMessage::User { content, .. } => content.starts_with(SUMMARY_PREFIX),
631            _ => false,
632        });
633        assert!(!has_summary, "should have no summary on failure");
634    }
635
636    #[tokio::test]
637    async fn test_middleware_noop_when_disabled() {
638        let calls = std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0));
639        let client = std::sync::Arc::new(MockClient {
640            response: "summary",
641            calls: calls.clone(),
642        });
643        let config = CompressionConfig::default().with_enabled(false);
644        let mw = CompressionMiddleware::new(config, client);
645
646        let msgs = make_messages(20);
647        let original_len = msgs.len();
648        let mut ctx = make_ctx(msgs);
649
650        mw.on_pre_llm(&mut ctx).await.unwrap();
651
652        assert_eq!(ctx.messages.len(), original_len);
653        assert_eq!(calls.load(std::sync::atomic::Ordering::SeqCst), 0);
654    }
655
656    #[tokio::test]
657    async fn test_middleware_preserves_system_prompt() {
658        let client = std::sync::Arc::new(MockClient {
659            response: "summary",
660            calls: std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0)),
661        });
662        let config = CompressionConfig::default()
663            .with_trigger_tokens(1)
664            .with_keep_recent_messages(4);
665        let mw = CompressionMiddleware::new(config, client);
666
667        let mut msgs = vec![
668            ChatMessage::system("You are helpful."),
669            ChatMessage::system("Extra instructions."),
670        ];
671        for i in 0..20 {
672            msgs.push(ChatMessage::user(format!("q{i}")));
673            msgs.push(ChatMessage::assistant(format!("a{i}")));
674        }
675
676        let mut ctx = make_ctx(msgs);
677        mw.on_pre_llm(&mut ctx).await.unwrap();
678
679        assert!(matches!(&ctx.messages[0], ChatMessage::System { .. }));
680        assert!(matches!(&ctx.messages[1], ChatMessage::System { .. }));
681        match &ctx.messages[0] {
682            ChatMessage::System { content, .. } => assert_eq!(content, "You are helpful."),
683            _ => unreachable!(),
684        }
685        match &ctx.messages[1] {
686            ChatMessage::System { content, .. } => assert_eq!(content, "Extra instructions."),
687            _ => unreachable!(),
688        }
689    }
690
691    #[tokio::test]
692    async fn test_from_compactor_and_accessor() {
693        let client = std::sync::Arc::new(MockClient {
694            response: "s",
695            calls: std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0)),
696        });
697        let config = CompressionConfig::default().with_trigger_tokens(42);
698        let compactor = ContextCompactor::new(client, config.clone());
699        let mw = CompressionMiddleware::from_compactor(compactor);
700
701        assert_eq!(mw.config().trigger_tokens, 42);
702    }
703
704    // ── Policy integration ────────────────────────────────────────────────
705
706    #[tokio::test]
707    async fn test_policy_called_with_correct_args() {
708        let client = std::sync::Arc::new(MockClient {
709            response: "summary",
710            calls: std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0)),
711        });
712        let config = CompressionConfig::default()
713            .with_trigger_tokens(1)
714            .with_keep_recent_messages(4);
715        let (policy, observed) = SpyPolicy::new(true);
716        let mw = CompressionMiddleware::with_policy(config, client, Box::new(policy));
717
718        let msgs = make_messages(10);
719        let msg_count = msgs.len();
720        let tokens_est: usize = msgs.iter().map(estimate_message_tokens).sum();
721        let mut ctx = make_ctx(msgs);
722        mw.on_pre_llm(&mut ctx).await.unwrap();
723
724        // Policy should have been called once.
725        let spy_calls = observed.lock().unwrap();
726        assert_eq!(spy_calls.len(), 1);
727        assert_eq!(spy_calls[0].1, msg_count);
728        // tokens_before should be > trigger_tokens (which is 1).
729        assert!(spy_calls[0].0 > 1);
730        // Approximate match — spy should receive roughly the same token estimate.
731        assert!(
732            (spy_calls[0].0 as i64 - tokens_est as i64).unsigned_abs() < 100,
733            "token estimate mismatch: spy={}, computed={}",
734            spy_calls[0].0,
735            tokens_est
736        );
737    }
738
739    #[tokio::test]
740    async fn test_policy_deny_skips_compression() {
741        let calls = std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0));
742        let client = std::sync::Arc::new(MockClient {
743            response: "summary",
744            calls: calls.clone(),
745        });
746        let config = CompressionConfig::default()
747            .with_trigger_tokens(1)
748            .with_keep_recent_messages(4);
749        let (policy, _observed) = SpyPolicy::new(false);
750        let mw = CompressionMiddleware::with_policy(config, client, Box::new(policy));
751
752        let msgs = make_messages(20);
753        let original_len = msgs.len();
754        let mut ctx = make_ctx(msgs);
755
756        mw.on_pre_llm(&mut ctx).await.unwrap();
757
758        // Messages should be unchanged — policy denied.
759        assert_eq!(ctx.messages.len(), original_len);
760        // No LLM call should have been made.
761        assert_eq!(calls.load(std::sync::atomic::Ordering::SeqCst), 0);
762    }
763
764    #[tokio::test]
765    async fn test_policy_not_called_when_below_threshold() {
766        let client = std::sync::Arc::new(MockClient {
767            response: "summary",
768            calls: std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0)),
769        });
770        let config = CompressionConfig::default().with_trigger_tokens(999_999);
771        let (policy, observed) = SpyPolicy::new(true);
772        let mw = CompressionMiddleware::with_policy(config, client, Box::new(policy));
773
774        let msgs = make_messages(5);
775        let mut ctx = make_ctx(msgs);
776
777        mw.on_pre_llm(&mut ctx).await.unwrap();
778
779        // Policy should NOT have been called — below threshold.
780        let spy_calls = observed.lock().unwrap();
781        assert_eq!(spy_calls.len(), 0);
782    }
783
784    #[tokio::test]
785    async fn test_rate_limit_policy_blocks_repeated_compression() {
786        let calls = std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0));
787        let client = std::sync::Arc::new(MockClient {
788            response: "summary",
789            calls: calls.clone(),
790        });
791        let config = CompressionConfig::default()
792            .with_trigger_tokens(1)
793            .with_keep_recent_messages(4);
794        // 60-second rate limit — second call should be blocked.
795        let policy = Box::new(RateLimitPolicy::new(std::time::Duration::from_secs(60)));
796        let mw = CompressionMiddleware::with_policy(config, client, policy);
797
798        let msgs = make_messages(20);
799
800        // First call — should compress.
801        let mut ctx1 = make_ctx(msgs.clone());
802        mw.on_pre_llm(&mut ctx1).await.unwrap();
803        assert_eq!(calls.load(std::sync::atomic::Ordering::SeqCst), 1);
804
805        // Second call immediately — rate limited, no LLM call.
806        let mut ctx2 = make_ctx(msgs);
807        mw.on_pre_llm(&mut ctx2).await.unwrap();
808        assert_eq!(calls.load(std::sync::atomic::Ordering::SeqCst), 1);
809    }
810
811    // ── Typed event emission ──────────────────────────────────────────────
812
813    #[tokio::test]
814    async fn test_emits_preparing_and_completed_events() {
815        let client = std::sync::Arc::new(MockClient {
816            response: "compressed summary text",
817            calls: std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0)),
818        });
819        let config = CompressionConfig::default()
820            .with_trigger_tokens(1)
821            .with_keep_recent_messages(4);
822        let mw = CompressionMiddleware::new(config, client);
823
824        let msgs = make_messages(20);
825        let events = std::sync::Arc::new(std::sync::Mutex::new(Vec::<CompressionEvent>::new()));
826        let mut ctx = make_ctx_with_events(msgs, events.clone());
827
828        mw.on_pre_llm(&mut ctx).await.unwrap();
829
830        let events = events.lock().unwrap();
831        // Should have: Preparing (from compact), Started (from compact), Progress, Completed.
832        assert!(
833            events.len() >= 3,
834            "expected at least 3 events, got {}",
835            events.len()
836        );
837
838        // First event should be Preparing.
839        assert!(
840            matches!(&events[0], CompressionEvent::Preparing { .. }),
841            "first event should be Preparing, got {:?}",
842            events[0]
843        );
844
845        // Second event should be Started.
846        assert!(
847            matches!(&events[1], CompressionEvent::Started { .. }),
848            "second event should be Started, got {:?}",
849            events[1]
850        );
851
852        // Last event should be Completed.
853        assert!(
854            matches!(events.last().unwrap(), CompressionEvent::Completed { .. }),
855            "last event should be Completed, got {:?}",
856            events.last().unwrap()
857        );
858
859        // Verify Preparing has trigger = Auto.
860        if let CompressionEvent::Preparing { trigger, .. } = &events[0] {
861            assert_eq!(*trigger, CompressionTrigger::Auto);
862        }
863    }
864
865    #[tokio::test]
866    async fn test_no_events_when_below_threshold() {
867        let client = std::sync::Arc::new(MockClient {
868            response: "summary",
869            calls: std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0)),
870        });
871        let config = CompressionConfig::default().with_trigger_tokens(999_999);
872        let mw = CompressionMiddleware::new(config, client);
873
874        let msgs = make_messages(5);
875        let events = std::sync::Arc::new(std::sync::Mutex::new(Vec::<CompressionEvent>::new()));
876        let mut ctx = make_ctx_with_events(msgs, events.clone());
877
878        mw.on_pre_llm(&mut ctx).await.unwrap();
879
880        let events = events.lock().unwrap();
881        assert!(events.is_empty(), "no events expected below threshold");
882    }
883
884    #[tokio::test]
885    async fn test_no_events_when_policy_denies() {
886        let client = std::sync::Arc::new(MockClient {
887            response: "summary",
888            calls: std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0)),
889        });
890        let config = CompressionConfig::default()
891            .with_trigger_tokens(1)
892            .with_keep_recent_messages(4);
893        let (policy, _observed) = SpyPolicy::new(false);
894        let mw = CompressionMiddleware::with_policy(config, client, Box::new(policy));
895
896        let msgs = make_messages(20);
897        let events = std::sync::Arc::new(std::sync::Mutex::new(Vec::<CompressionEvent>::new()));
898        let mut ctx = make_ctx_with_events(msgs, events.clone());
899
900        mw.on_pre_llm(&mut ctx).await.unwrap();
901
902        let events = events.lock().unwrap();
903        assert!(events.is_empty(), "no events expected when policy denies");
904    }
905
906    #[tokio::test]
907    async fn test_messages_unchanged_when_compact_returns_none() {
908        // compact() returns None when disabled or too few messages —
909        // compact() is pure so ctx.messages is never mutated.
910        let calls = std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0));
911        let client = std::sync::Arc::new(MockClient {
912            response: "summary",
913            calls: calls.clone(),
914        });
915        let config = CompressionConfig::default()
916            .with_trigger_tokens(1)
917            .with_keep_recent_messages(4)
918            .with_enabled(false); // disabled → compact() returns None
919        let (policy, _observed) = SpyPolicy::new(true);
920        let mw = CompressionMiddleware::with_policy(config, client, Box::new(policy));
921
922        let msgs = make_messages(20);
923        let original_len = msgs.len();
924        let mut ctx = make_ctx(msgs);
925
926        mw.on_pre_llm(&mut ctx).await.unwrap();
927
928        // Messages unchanged — compact() never modified them.
929        assert_eq!(ctx.messages.len(), original_len);
930    }
931
932    #[tokio::test]
933    async fn test_exact_event_sequence_on_cache_miss() {
934        let client = std::sync::Arc::new(MockClient {
935            response: "compressed summary text",
936            calls: std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0)),
937        });
938        let config = CompressionConfig::default()
939            .with_trigger_tokens(1)
940            .with_keep_recent_messages(4);
941        let mw = CompressionMiddleware::new(config, client);
942
943        let msgs = make_messages(20);
944        let events = std::sync::Arc::new(std::sync::Mutex::new(Vec::<CompressionEvent>::new()));
945        let mut ctx = make_ctx_with_events(msgs, events.clone());
946
947        mw.on_pre_llm(&mut ctx).await.unwrap();
948
949        let events = events.lock().unwrap();
950        // Exact sequence: Preparing → Started → Progress(0) → Progress(N) → Completed
951        // Preparing = compact emits Preparing first, then Started,
952        // Progress(0) = "Connecting to LLM", Progress(N) = streaming chars from summarizer.
953        assert_eq!(
954            events.len(),
955            5,
956            "expected 5 events, got {}: {:?}",
957            events.len(),
958            *events
959        );
960        assert!(matches!(&events[0], CompressionEvent::Preparing { .. }));
961        assert!(matches!(&events[1], CompressionEvent::Started { .. }));
962        assert!(matches!(
963            &events[2],
964            CompressionEvent::Progress { chars: 0, .. }
965        ));
966        // events[3] is Progress with actual char count from the mock response.
967        assert!(matches!(&events[3], CompressionEvent::Progress { chars, .. } if *chars > 0));
968        assert!(matches!(&events[4], CompressionEvent::Completed { .. }));
969    }
970
971    #[tokio::test]
972    async fn test_event_sequence_on_no_retrigger() {
973        // After compression, calling with the same messages should NOT re-trigger
974        // because last_compressed_msg_count was set to the pre-compression count.
975        let client = std::sync::Arc::new(MockClient {
976            response: "cached summary",
977            calls: std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0)),
978        });
979        let config = CompressionConfig::default()
980            .with_trigger_tokens(1)
981            .with_keep_recent_messages(4);
982        let mw = CompressionMiddleware::new(config, client);
983
984        let msgs = make_messages(20);
985        let events = std::sync::Arc::new(std::sync::Mutex::new(Vec::<CompressionEvent>::new()));
986
987        // First call — should compress and emit events.
988        let mut ctx1 = make_ctx_with_events(msgs.clone(), events.clone());
989        mw.on_pre_llm(&mut ctx1).await.unwrap();
990        assert!(
991            events.lock().unwrap().len() >= 3,
992            "first call should emit events"
993        );
994        events.lock().unwrap().clear();
995
996        // Second call with the same messages — should NOT re-trigger.
997        // last_compressed=20, skip(20) → 0 new messages → below threshold.
998        let mut ctx2 = make_ctx_with_events(msgs, events.clone());
999        mw.on_pre_llm(&mut ctx2).await.unwrap();
1000
1001        let events = events.lock().unwrap();
1002        assert_eq!(
1003            events.len(),
1004            0,
1005            "same messages: expected 0 events (no re-trigger), got {}",
1006            events.len()
1007        );
1008    }
1009
1010    #[tokio::test]
1011    async fn test_no_retrigger_after_compression() {
1012        // After successful compression, the next call with the compressed messages
1013        // should NOT re-trigger compression because last_compressed_msg_count
1014        // was set to the pre-compression message count, so skip() returns 0 new messages.
1015        let calls = std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0));
1016        let client = std::sync::Arc::new(MockClient {
1017            response: "short",
1018            calls: calls.clone(),
1019        });
1020        let config = CompressionConfig::default()
1021            .with_trigger_tokens(1)
1022            .with_keep_recent_messages(4);
1023        let mw = CompressionMiddleware::new(config, client);
1024
1025        let msgs = make_messages(20);
1026
1027        // First call — should compress (20 messages > trigger_tokens=1).
1028        let mut ctx1 = make_ctx(msgs.clone());
1029        mw.on_pre_llm(&mut ctx1).await.unwrap();
1030        assert_eq!(calls.load(std::sync::atomic::Ordering::SeqCst), 1);
1031        let compressed_len = ctx1.messages.len();
1032        assert!(compressed_len < msgs.len(), "should have compressed");
1033
1034        // Second call with the compressed messages (realistic: next turn reloads
1035        // compressed state from disk). last_compressed=20 (pre-compression count),
1036        // skip(20) on compressed_len messages → 0 messages → no re-trigger.
1037        let mut ctx2 = make_ctx(ctx1.messages.clone());
1038        mw.on_pre_llm(&mut ctx2).await.unwrap();
1039        assert_eq!(
1040            calls.load(std::sync::atomic::Ordering::SeqCst),
1041            1,
1042            "should NOT have re-triggered compression"
1043        );
1044    }
1045
1046    /// End-to-end simulation: 30 turns of long conversation.
1047    /// Verifies compression triggers, no re-trigger, and old block stability.
1048    #[tokio::test]
1049    async fn test_e2e_simulation_30_turns() {
1050        let calls = std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0));
1051        let client = std::sync::Arc::new(MockClient {
1052            response: "compressed summary of the conversation so far",
1053            calls: calls.clone(),
1054        });
1055        let config = CompressionConfig::default()
1056            .with_trigger_tokens(2000)
1057            .with_keep_recent_messages(4);
1058        let mw = CompressionMiddleware::new(config, client);
1059
1060        let mut messages = vec![ChatMessage::system("You are a helpful assistant.")];
1061        let events = std::sync::Arc::new(std::sync::Mutex::new(Vec::<CompressionEvent>::new()));
1062        let mut compression_count = 0;
1063        let mut compression_turns = Vec::new();
1064        let mut old_block_sizes = Vec::new();
1065
1066        for turn in 0..30 {
1067            messages.push(ChatMessage::user(format!(
1068                "请继续讲故事,这是第{}轮对话。我想听一个关于古代英雄的故事,要有曲折的情节和深刻的寓意,最好能让人有所启发和思考。",
1069                turn + 1
1070            )));
1071            // Longer assistant response (~500 chars) to accumulate tokens faster.
1072            messages.push(ChatMessage::assistant(format!(
1073                "好的,让我继续讲第{}轮的故事。从前有座山,山里有座庙,庙里有个老和尚在讲故事。\
1074                 这个故事讲的是从前有座山,山里有座庙,庙里有个老和尚在讲故事。\
1075                 故事的内容是关于一个勇敢的冒险者,他走遍了千山万水,经历了无数磨难。\
1076                 他遇到了各种各样的人,有善良的农夫,有狡猾的商人,有智慧的老者。\
1077                 每个人都给了他不同的启示,让他对人生有了更深的理解。\
1078                 他学会了坚韧不拔,学会了与人为善,学会了在困境中寻找希望。\
1079                 最终,他回到了家乡,成为了一个受人尊敬的长者,把自己的故事讲给后人听。\
1080                 这个故事告诉我们,人生就是一场旅行,重要的不是目的地,而是沿途的风景。",
1081                turn + 1
1082            )));
1083
1084            let mut ctx = make_ctx_with_events(messages.clone(), events.clone());
1085            mw.on_pre_llm(&mut ctx).await.unwrap();
1086            messages = ctx.messages.clone();
1087
1088            let new_compressions = events
1089                .lock()
1090                .unwrap()
1091                .iter()
1092                .filter(|e| matches!(e, CompressionEvent::Completed { .. }))
1093                .count();
1094            let turn_compressions = new_compressions - compression_count;
1095            compression_count = new_compressions;
1096
1097            if turn_compressions > 0 {
1098                compression_turns.push(turn + 1);
1099                // Record old block size from the compression event.
1100                if let CompressionEvent::Completed {
1101                    msg_count_before,
1102                    msg_count_after,
1103                    ..
1104                } = events.lock().unwrap().last().unwrap()
1105                {
1106                    old_block_sizes.push((*msg_count_before, *msg_count_after));
1107                }
1108                println!(
1109                    "[Turn {:2}] COMPRESSED | msgs={:2} → {:2} | llm_calls={}",
1110                    turn + 1,
1111                    old_block_sizes.last().unwrap().0,
1112                    old_block_sizes.last().unwrap().1,
1113                    calls.load(std::sync::atomic::Ordering::SeqCst)
1114                );
1115            } else {
1116                println!(
1117                    "[Turn {:2}] ok         | msgs={:2}",
1118                    turn + 1,
1119                    messages.len(),
1120                );
1121            }
1122        }
1123
1124        println!("\n=== Simulation Summary ===");
1125        println!("Compression turns: {:?}", compression_turns);
1126        println!("Total compressions: {}", compression_count);
1127        println!(
1128            "LLM calls: {}",
1129            calls.load(std::sync::atomic::Ordering::SeqCst)
1130        );
1131        println!("Final messages: {}", messages.len());
1132        println!("Old block before/after: {:?}", old_block_sizes);
1133
1134        // Verify: at least 2 compressions in 30 turns.
1135        assert!(
1136            compression_count >= 2,
1137            "expected at least 2 compressions, got {}",
1138            compression_count
1139        );
1140
1141        // Verify: LLM calls == compression count (no wasted calls).
1142        assert_eq!(
1143            calls.load(std::sync::atomic::Ordering::SeqCst),
1144            compression_count,
1145        );
1146
1147        // Verify: no two compressions on consecutive turns.
1148        for window in compression_turns.windows(2) {
1149            assert!(
1150                window[1] - window[0] >= 2,
1151                "consecutive compressions at turns {:?}",
1152                window
1153            );
1154        }
1155
1156        // Verify: final message count is reasonable.
1157        // With "preserve user messages" strategy, users accumulate, so the count
1158        // grows over time.  But it should be less than uncompressed (61 msgs).
1159        assert!(
1160            messages.len() < 61,
1161            "final messages should be < 61 (uncompressed), got {}",
1162            messages.len()
1163        );
1164    }
1165}