1use 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
28fn 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
77fn 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
118pub struct CompressionMiddleware {
139 compactor: ContextCompactor,
140 policy: Box<dyn CompressionPolicy>,
141 last_compressed_msg_count: AtomicUsize,
145}
146
147#[allow(missing_docs)]
148impl CompressionMiddleware {
149 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 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 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 pub fn compactor(&self) -> &ContextCompactor {
189 &self.compactor
190 }
191
192 pub fn clone_compactor(&self) -> ContextCompactor {
198 self.compactor.clone_handle()
199 }
200
201 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 let keep = self.config().keep_recent_messages;
216 if msg_count <= keep + 1 {
217 return Ok(());
218 }
219
220 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 if tokens_before <= self.config().trigger_tokens {
243 return Ok(());
244 }
245
246 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 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 let trigger = CompressionTrigger::Auto;
276 let sid = ctx.session_id.id;
277
278 let filtered: Vec<ChatMessage> = ctx
282 .messages
283 .iter()
284 .filter(|m| !is_summary_message(m))
285 .cloned()
286 .collect();
287
288 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 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 self.last_compressed_msg_count
322 .store(msg_count, Ordering::Relaxed);
323 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 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_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 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 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 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 struct SpyPolicy {
464 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 #[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 #[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 let spy_calls = observed.lock().unwrap();
726 assert_eq!(spy_calls.len(), 1);
727 assert_eq!(spy_calls[0].1, msg_count);
728 assert!(spy_calls[0].0 > 1);
730 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 assert_eq!(ctx.messages.len(), original_len);
760 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 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 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 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 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 #[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 assert!(
833 events.len() >= 3,
834 "expected at least 3 events, got {}",
835 events.len()
836 );
837
838 assert!(
840 matches!(&events[0], CompressionEvent::Preparing { .. }),
841 "first event should be Preparing, got {:?}",
842 events[0]
843 );
844
845 assert!(
847 matches!(&events[1], CompressionEvent::Started { .. }),
848 "second event should be Started, got {:?}",
849 events[1]
850 );
851
852 assert!(
854 matches!(events.last().unwrap(), CompressionEvent::Completed { .. }),
855 "last event should be Completed, got {:?}",
856 events.last().unwrap()
857 );
858
859 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 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); 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 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 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 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 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 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 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 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 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 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 #[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 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 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 assert!(
1136 compression_count >= 2,
1137 "expected at least 2 compressions, got {}",
1138 compression_count
1139 );
1140
1141 assert_eq!(
1143 calls.load(std::sync::atomic::Ordering::SeqCst),
1144 compression_count,
1145 );
1146
1147 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 assert!(
1160 messages.len() < 61,
1161 "final messages should be < 61 (uncompressed), got {}",
1162 messages.len()
1163 );
1164 }
1165}