1use std::sync::Arc;
4
5use async_trait::async_trait;
6use parking_lot::RwLock;
7use serde::{Deserialize, Serialize};
8use tokio::sync::Mutex as AsyncMutex;
9
10use ai_agents_core::{ChatMessage, MemorySnapshot, Result};
11
12use super::Memory;
13use super::context::{CompressResult, ConversationContext, estimate_tokens};
14use super::summarizer::Summarizer;
15
16fn prefix_at_char_boundary(text: &str, max_chars: usize) -> &str {
17 if max_chars == 0 {
18 return "";
19 }
20
21 match text.char_indices().nth(max_chars) {
22 Some((idx, _)) => &text[..idx],
23 None => text,
24 }
25}
26
27pub struct CompactingMemory {
28 operation_lock: AsyncMutex<()>,
29 summary: RwLock<Option<String>>,
30 messages: RwLock<Vec<ChatMessage>>,
31 summarized_count: RwLock<usize>,
32 config: CompactingMemoryConfig,
33 summarizer: Arc<dyn Summarizer>,
34 compression_history: RwLock<Vec<CompressionEvent>>,
35}
36
37#[derive(Debug, Clone, Serialize, Deserialize)]
38#[serde(deny_unknown_fields)]
39pub struct CompactingMemoryConfig {
40 #[serde(default = "default_max_recent_messages")]
42 pub max_recent_messages: usize,
43
44 #[serde(default = "default_compress_threshold")]
45 pub compress_threshold: usize,
46
47 #[serde(default = "default_summarize_batch_size")]
48 pub summarize_batch_size: usize,
49
50 #[serde(default = "default_max_summary_length")]
52 pub max_summary_length: usize,
53}
54
55#[derive(Debug, Clone, Serialize, Deserialize)]
56pub struct CompressionEvent {
57 pub timestamp: chrono::DateTime<chrono::Utc>,
58 pub messages_compressed: usize,
59 pub summary_length_before: usize,
60 pub summary_length_after: usize,
61}
62
63fn default_max_recent_messages() -> usize {
64 50
65}
66
67fn default_compress_threshold() -> usize {
68 30
69}
70
71fn default_summarize_batch_size() -> usize {
72 10
73}
74
75fn default_max_summary_length() -> usize {
76 2000
77}
78
79fn protected_recent_count(config: &CompactingMemoryConfig, message_count: usize) -> usize {
80 if config.max_recent_messages < config.compress_threshold {
81 return config.max_recent_messages.min(message_count);
82 }
83
84 let batch_at_threshold = config
85 .summarize_batch_size
86 .max(1)
87 .min(config.compress_threshold);
88 let retention_cap = config.compress_threshold.saturating_sub(batch_at_threshold);
89 config
90 .max_recent_messages
91 .min(retention_cap)
92 .min(message_count)
93}
94
95impl Default for CompactingMemoryConfig {
96 fn default() -> Self {
97 Self {
98 max_recent_messages: default_max_recent_messages(),
99 compress_threshold: default_compress_threshold(),
100 summarize_batch_size: default_summarize_batch_size(),
101 max_summary_length: default_max_summary_length(),
102 }
103 }
104}
105
106impl CompactingMemory {
107 pub fn new(summarizer: Arc<dyn Summarizer>, config: CompactingMemoryConfig) -> Self {
108 Self {
109 operation_lock: AsyncMutex::new(()),
110 summary: RwLock::new(None),
111 messages: RwLock::new(Vec::new()),
112 summarized_count: RwLock::new(0),
113 config,
114 summarizer,
115 compression_history: RwLock::new(Vec::new()),
116 }
117 }
118
119 pub fn with_default_config(summarizer: Arc<dyn Summarizer>) -> Self {
120 Self::new(summarizer, CompactingMemoryConfig::default())
121 }
122
123 pub fn config(&self) -> &CompactingMemoryConfig {
124 &self.config
125 }
126
127 pub fn summary(&self) -> Option<String> {
128 self.summary.read().clone()
129 }
130
131 pub fn summarized_count(&self) -> usize {
132 *self.summarized_count.read()
133 }
134
135 pub fn compression_history(&self) -> Vec<CompressionEvent> {
136 self.compression_history.read().clone()
137 }
138
139 fn record_compression(&self, messages_compressed: usize, before: usize, after: usize) {
140 let event = CompressionEvent {
141 timestamp: chrono::Utc::now(),
142 messages_compressed,
143 summary_length_before: before,
144 summary_length_after: after,
145 };
146 self.compression_history.write().push(event);
147 }
148}
149
150#[async_trait]
151impl ai_agents_core::Memory for CompactingMemory {
152 async fn add_message(&self, message: ChatMessage) -> Result<()> {
153 let _operation = self.operation_lock.lock().await;
154 self.messages.write().push(message);
155 Ok(())
156 }
157
158 async fn get_messages(&self, limit: Option<usize>) -> Result<Vec<ChatMessage>> {
159 let messages = self.messages.read();
160 match limit {
161 Some(n) => {
162 let start = messages.len().saturating_sub(n);
163 Ok(messages[start..].to_vec())
164 }
165 None => Ok(messages.clone()),
166 }
167 }
168
169 async fn clear(&self) -> Result<()> {
170 let _operation = self.operation_lock.lock().await;
171 *self.summary.write() = None;
172 self.messages.write().clear();
173 *self.summarized_count.write() = 0;
174 self.compression_history.write().clear();
175 Ok(())
176 }
177
178 fn len(&self) -> usize {
179 self.messages.read().len()
180 }
181
182 async fn snapshot(&self) -> Result<MemorySnapshot> {
183 let _operation = self.operation_lock.lock().await;
184 let messages = self.messages.read().clone();
185 let summary = self.summary.read().clone();
186 let summarized_count = *self.summarized_count.read();
187
188 let mut snapshot = MemorySnapshot::new(messages).with_summarized_count(summarized_count);
189 if let Some(s) = summary {
190 snapshot = snapshot.with_summary(s);
191 }
192 Ok(snapshot)
193 }
194
195 async fn restore(&self, snapshot: MemorySnapshot) -> Result<()> {
196 let _operation = self.operation_lock.lock().await;
197 *self.messages.write() = snapshot.messages;
198 *self.summary.write() = snapshot.summary;
199 *self.summarized_count.write() = snapshot.summarized_count;
200 self.compression_history.write().clear();
201 Ok(())
202 }
203
204 async fn evict_oldest(&self, count: usize) -> Result<Vec<ChatMessage>> {
205 let _operation = self.operation_lock.lock().await;
206 let mut messages = self.messages.write();
207 let evict_count = count.min(messages.len());
208 let evicted: Vec<ChatMessage> = messages.drain(..evict_count).collect();
209 Ok(evicted)
210 }
211}
212
213#[async_trait]
214impl Memory for CompactingMemory {
215 async fn get_context(&self) -> Result<ConversationContext> {
216 let _operation = self.operation_lock.lock().await;
217 let messages = self.messages.read().clone();
218 let summary = self.summary.read().clone();
219 let summarized_count = *self.summarized_count.read();
220 let total_messages = messages.len() + summarized_count;
221
222 let mut ctx = ConversationContext::with_messages(messages);
223 ctx.total_messages = total_messages;
224
225 if let Some(s) = summary {
226 ctx = ctx.with_summary(s, summarized_count);
227 }
228
229 Ok(ctx)
230 }
231
232 async fn compress(&self, summarizer: Option<&dyn Summarizer>) -> Result<CompressResult> {
233 let _operation = self.operation_lock.lock().await;
234 let message_count = self.messages.read().len();
235
236 if message_count == 0 || message_count < self.config.compress_threshold {
237 return Ok(CompressResult::NotNeeded);
238 }
239
240 let summarizer = summarizer.unwrap_or(self.summarizer.as_ref());
241 let protected_count = protected_recent_count(&self.config, message_count);
242 let compressible_count = message_count - protected_count;
243 let batch_size = self
244 .config
245 .summarize_batch_size
246 .max(1)
247 .min(compressible_count);
248
249 let messages_to_summarize: Vec<ChatMessage> = {
250 let messages = self.messages.read();
251 messages[..batch_size].to_vec()
252 };
253
254 let new_summary = summarizer.summarize(&messages_to_summarize).await?;
255
256 let summary_before_len = self.summary.read().as_ref().map(|s| s.len()).unwrap_or(0);
257
258 let existing_summary = self.summary.read().clone();
259 let existing_summary_tokens = existing_summary
260 .as_deref()
261 .map(estimate_tokens)
262 .unwrap_or(0);
263 let combined_summary = match existing_summary {
264 Some(existing) => summarizer.merge_summaries(&[existing, new_summary]).await?,
265 None => new_summary,
266 };
267
268 let truncated = prefix_at_char_boundary(&combined_summary, self.config.max_summary_length);
269 let final_summary = if truncated.len() < combined_summary.len() {
270 truncated.to_string()
271 } else {
272 combined_summary
273 };
274
275 let summary_after_len = final_summary.len();
276
277 {
278 let mut messages = self.messages.write();
279 messages.drain(..batch_size);
280 }
281
282 *self.summary.write() = Some(final_summary.clone());
283 *self.summarized_count.write() += batch_size;
284
285 self.record_compression(batch_size, summary_before_len, summary_after_len);
286
287 let tokens_before: u32 = existing_summary_tokens.saturating_add(
288 messages_to_summarize
289 .iter()
290 .map(|m| estimate_tokens(&m.content))
291 .sum(),
292 );
293 let tokens_after = estimate_tokens(&final_summary);
294 let tokens_saved = tokens_before.saturating_sub(tokens_after);
295
296 Ok(CompressResult::Compressed {
297 messages_summarized: batch_size,
298 new_summary_length: summary_after_len,
299 tokens_saved,
300 })
301 }
302
303 fn needs_compression(&self) -> bool {
304 let message_count = self.messages.read().len();
305 message_count > 0 && message_count >= self.config.compress_threshold
306 }
307}
308
309#[cfg(test)]
310mod tests {
311 use super::*;
312
313 use crate::summarizer::NoopSummarizer;
314 use ai_agents_core::{AgentError, Memory as CoreMemory, Role};
315 use tokio::time::{Duration, timeout};
316
317 fn make_message(content: &str) -> ChatMessage {
318 ChatMessage {
319 role: Role::User,
320 content: content.to_string(),
321 name: None,
322 timestamp: None,
323 }
324 }
325
326 fn message_contents(messages: &[ChatMessage]) -> Vec<&str> {
327 messages
328 .iter()
329 .map(|message| message.content.as_str())
330 .collect()
331 }
332
333 struct BlockingSummarizer {
334 started: tokio::sync::Notify,
335 release: tokio::sync::Notify,
336 batches: RwLock<Vec<Vec<String>>>,
337 }
338
339 impl BlockingSummarizer {
340 fn new() -> Self {
341 Self {
342 started: tokio::sync::Notify::new(),
343 release: tokio::sync::Notify::new(),
344 batches: RwLock::new(Vec::new()),
345 }
346 }
347
348 async fn wait_until_started(&self) {
349 self.started.notified().await;
350 }
351
352 fn release(&self) {
353 self.release.notify_one();
354 }
355
356 fn batches(&self) -> Vec<Vec<String>> {
357 self.batches.read().clone()
358 }
359 }
360
361 #[async_trait]
362 impl Summarizer for BlockingSummarizer {
363 async fn summarize(&self, messages: &[ChatMessage]) -> Result<String> {
364 let contents: Vec<_> = messages
365 .iter()
366 .map(|message| message.content.clone())
367 .collect();
368 self.batches.write().push(contents.clone());
369 self.started.notify_one();
370 self.release.notified().await;
371 Ok(contents.join(" | "))
372 }
373 }
374
375 struct FailingSummarizer;
376
377 #[async_trait]
378 impl Summarizer for FailingSummarizer {
379 async fn summarize(&self, _messages: &[ChatMessage]) -> Result<String> {
380 Err(AgentError::MemoryError("summary failed".to_string()))
381 }
382 }
383
384 struct FailingMergeSummarizer;
385
386 #[async_trait]
387 impl Summarizer for FailingMergeSummarizer {
388 async fn summarize(&self, _messages: &[ChatMessage]) -> Result<String> {
389 Ok("new summary".to_string())
390 }
391
392 async fn merge_summaries(&self, _summaries: &[String]) -> Result<String> {
393 Err(AgentError::MemoryError("merge failed".to_string()))
394 }
395 }
396
397 fn create_test_memory() -> CompactingMemory {
398 let summarizer = Arc::new(NoopSummarizer);
399 let config = CompactingMemoryConfig {
400 max_recent_messages: 3,
401 compress_threshold: 5,
402 summarize_batch_size: 3,
403 max_summary_length: 500,
404 };
405 CompactingMemory::new(summarizer, config)
406 }
407
408 #[tokio::test]
409 async fn test_basic_add_and_get() {
410 let memory = create_test_memory();
411
412 memory.add_message(make_message("Hello")).await.unwrap();
413 memory.add_message(make_message("World")).await.unwrap();
414
415 let messages = memory.get_messages(None).await.unwrap();
416 assert_eq!(messages.len(), 2);
417 assert_eq!(messages[0].content, "Hello");
418 assert_eq!(messages[1].content, "World");
419 }
420
421 #[tokio::test]
422 async fn test_get_messages_with_limit() {
423 let memory = create_test_memory();
424
425 for i in 0..5 {
426 memory
427 .add_message(make_message(&format!("msg{}", i)))
428 .await
429 .unwrap();
430 }
431
432 let messages = memory.get_messages(Some(2)).await.unwrap();
433 assert_eq!(messages.len(), 2);
434 assert_eq!(messages[0].content, "msg3");
435 assert_eq!(messages[1].content, "msg4");
436 }
437
438 #[tokio::test]
439 async fn test_clear() {
440 let memory = create_test_memory();
441
442 memory.add_message(make_message("test")).await.unwrap();
443 assert!(!memory.is_empty());
444
445 memory.clear().await.unwrap();
446 assert!(memory.is_empty());
447 assert!(memory.summary().is_none());
448 }
449
450 #[tokio::test]
451 async fn test_needs_compression() {
452 let memory = create_test_memory();
453
454 for i in 0..4 {
455 memory
456 .add_message(make_message(&format!("msg{}", i)))
457 .await
458 .unwrap();
459 }
460 assert!(!memory.needs_compression());
461
462 memory.add_message(make_message("msg4")).await.unwrap();
463 assert!(memory.needs_compression());
464 }
465
466 #[tokio::test]
467 async fn test_compress_not_needed() {
468 let memory = create_test_memory();
469
470 memory.add_message(make_message("msg1")).await.unwrap();
471 memory.add_message(make_message("msg2")).await.unwrap();
472
473 let result = memory.compress(None).await.unwrap();
474 assert!(matches!(result, CompressResult::NotNeeded));
475 }
476
477 #[tokio::test]
478 async fn test_compress_when_needed() {
479 let memory = create_test_memory();
480
481 for i in 0..6 {
482 memory
483 .add_message(make_message(&format!("message number {}", i)))
484 .await
485 .unwrap();
486 }
487
488 assert!(memory.needs_compression());
489
490 let result = memory.compress(None).await.unwrap();
491
492 if let CompressResult::Compressed {
493 messages_summarized,
494 ..
495 } = result
496 {
497 assert_eq!(messages_summarized, 3);
498 } else {
499 panic!("Expected Compressed result");
500 }
501
502 assert_eq!(memory.len(), 3);
503 assert!(memory.summary().is_some());
504 assert_eq!(memory.summarized_count(), 3);
505 }
506
507 #[tokio::test]
508 async fn test_compress_preserves_configured_recent_tail() {
509 let config = CompactingMemoryConfig {
510 max_recent_messages: 3,
511 compress_threshold: 5,
512 summarize_batch_size: 10,
513 max_summary_length: 500,
514 };
515 let memory = CompactingMemory::new(Arc::new(NoopSummarizer), config);
516
517 for i in 0..7 {
518 memory
519 .add_message(make_message(&format!("msg{}", i)))
520 .await
521 .unwrap();
522 }
523
524 let result = memory.compress(None).await.unwrap();
525 assert!(matches!(
526 result,
527 CompressResult::Compressed {
528 messages_summarized: 4,
529 ..
530 }
531 ));
532 let remaining = memory.get_messages(None).await.unwrap();
533 let contents: Vec<_> = remaining
534 .iter()
535 .map(|message| message.content.as_str())
536 .collect();
537 assert_eq!(contents, vec!["msg4", "msg5", "msg6"]);
538 }
539
540 #[tokio::test]
541 async fn test_multilingual_recent_tail_remains_verbatim() {
542 let config = CompactingMemoryConfig {
543 max_recent_messages: 4,
544 compress_threshold: 6,
545 summarize_batch_size: 3,
546 max_summary_length: 100_000,
547 };
548 let memory = CompactingMemory::new(Arc::new(NoopSummarizer), config);
549 let large_mixed_message = "protected 한국어 日本語 العربية English emoji 🧭".to_string();
550 let source = vec![
551 "Old English context".to_string(),
552 "이전 한국어 문맥".to_string(),
553 "古い日本語の文脈".to_string(),
554 "سياق عربي قديم".to_string(),
555 "पुराना हिंदी संदर्भ".to_string(),
556 "Keep 한국어 and English together".to_string(),
557 "保留する日本語メッセージ".to_string(),
558 "احتفظ بهذه الرسالة العربية".to_string(),
559 large_mixed_message,
560 ];
561 let expected_tail = source[source.len() - 4..].to_vec();
562
563 for content in &source {
564 memory.add_message(make_message(content)).await.unwrap();
565 }
566
567 memory.compress(None).await.unwrap();
568 memory.compress(None).await.unwrap();
569 let remaining = memory.get_messages(None).await.unwrap();
570
571 assert!(!memory.needs_compression());
572 assert_eq!(
573 remaining
574 .iter()
575 .map(|message| message.content.clone())
576 .collect::<Vec<_>>(),
577 expected_tail
578 );
579 assert_eq!(memory.summarized_count() + memory.len(), source.len());
580 }
581
582 #[tokio::test]
583 async fn test_compress_clamps_recent_tail_for_configured_batch() {
584 let config = CompactingMemoryConfig {
585 max_recent_messages: 100,
586 compress_threshold: 5,
587 summarize_batch_size: 10,
588 max_summary_length: 500,
589 };
590 let memory = CompactingMemory::new(Arc::new(NoopSummarizer), config);
591
592 for i in 0..5 {
593 memory
594 .add_message(make_message(&format!("msg{}", i)))
595 .await
596 .unwrap();
597 }
598
599 let result = memory.compress(None).await.unwrap();
600 assert!(matches!(
601 result,
602 CompressResult::Compressed {
603 messages_summarized: 5,
604 ..
605 }
606 ));
607 assert!(memory.get_messages(None).await.unwrap().is_empty());
608 assert!(!memory.needs_compression());
609 }
610
611 #[tokio::test]
612 async fn test_default_config_compresses_full_batches_at_steady_state() {
613 let memory =
614 CompactingMemory::new(Arc::new(NoopSummarizer), CompactingMemoryConfig::default());
615
616 for i in 0..30 {
617 memory
618 .add_message(make_message(&format!("msg{}", i)))
619 .await
620 .unwrap();
621 }
622
623 for round in 0..4 {
624 let result = memory.compress(None).await.unwrap();
625 assert!(matches!(
626 result,
627 CompressResult::Compressed {
628 messages_summarized: 10,
629 ..
630 }
631 ));
632 assert_eq!(memory.len(), 20);
633 assert!(!memory.needs_compression());
634
635 if round < 3 {
636 let start = 30 + round * 10;
637 for i in start..start + 10 {
638 memory
639 .add_message(make_message(&format!("msg{}", i)))
640 .await
641 .unwrap();
642 }
643 assert_eq!(memory.len(), 30);
644 assert!(memory.needs_compression());
645 }
646 }
647
648 assert_eq!(memory.summarized_count(), 40);
649 let remaining = memory.get_messages(None).await.unwrap();
650 assert_eq!(remaining.first().unwrap().content, "msg40");
651 assert_eq!(remaining.last().unwrap().content, "msg59");
652 }
653
654 #[test]
655 fn test_protected_recent_count_edge_cases() {
656 let non_conflicting = CompactingMemoryConfig {
657 max_recent_messages: 25,
658 compress_threshold: 30,
659 summarize_batch_size: 10,
660 max_summary_length: 500,
661 };
662 assert_eq!(protected_recent_count(&non_conflicting, 30), 25);
663
664 let conflicting = CompactingMemoryConfig {
665 max_recent_messages: 50,
666 compress_threshold: 30,
667 summarize_batch_size: 10,
668 max_summary_length: 500,
669 };
670 assert_eq!(protected_recent_count(&conflicting, 30), 20);
671
672 let oversized_batch = CompactingMemoryConfig {
673 max_recent_messages: 5,
674 compress_threshold: 5,
675 summarize_batch_size: 10,
676 max_summary_length: 500,
677 };
678 assert_eq!(protected_recent_count(&oversized_batch, 5), 0);
679
680 let zero_batch = CompactingMemoryConfig {
681 max_recent_messages: 5,
682 compress_threshold: 5,
683 summarize_batch_size: 0,
684 max_summary_length: 500,
685 };
686 assert_eq!(protected_recent_count(&zero_batch, 5), 4);
687
688 let zero_threshold = CompactingMemoryConfig {
689 max_recent_messages: 5,
690 compress_threshold: 0,
691 summarize_batch_size: 10,
692 max_summary_length: 500,
693 };
694 assert_eq!(protected_recent_count(&zero_threshold, 5), 0);
695 }
696
697 #[tokio::test]
698 async fn test_initial_summary_failure_rolls_back_all_accounting() {
699 let config = CompactingMemoryConfig {
700 max_recent_messages: 3,
701 compress_threshold: 5,
702 summarize_batch_size: 2,
703 max_summary_length: 100_000,
704 };
705 let memory = CompactingMemory::new(Arc::new(FailingSummarizer), config);
706 let source = vec![
707 "English before failure".to_string(),
708 "실패 전 한국어".to_string(),
709 "失敗前の日本語".to_string(),
710 "قبل الفشل".to_string(),
711 "large mixed message 한界🙂abc".to_string(),
712 ];
713 for content in &source {
714 memory.add_message(make_message(content)).await.unwrap();
715 }
716
717 let snapshot_before = memory.snapshot().await.unwrap();
718 let context_before = memory.get_context().await.unwrap();
719 let error = memory.compress(None).await.unwrap_err();
720 let snapshot_after = memory.snapshot().await.unwrap();
721 let context_after = memory.get_context().await.unwrap();
722
723 assert!(error.to_string().contains("summary failed"));
724 assert_eq!(
725 message_contents(&snapshot_after.messages),
726 message_contents(&snapshot_before.messages)
727 );
728 assert_eq!(snapshot_after.summary, snapshot_before.summary);
729 assert_eq!(
730 snapshot_after.summarized_count,
731 snapshot_before.summarized_count
732 );
733 assert_eq!(context_after.total_messages, context_before.total_messages);
734 assert_eq!(
735 context_after.summarized_count,
736 context_before.summarized_count
737 );
738 assert_eq!(
739 context_after.estimated_tokens(),
740 context_before.estimated_tokens()
741 );
742 assert!(memory.compression_history().is_empty());
743 assert!(memory.needs_compression());
744
745 let retry = memory.compress(Some(&NoopSummarizer)).await.unwrap();
746 assert!(matches!(
747 retry,
748 CompressResult::Compressed {
749 messages_summarized: 2,
750 ..
751 }
752 ));
753 assert_eq!(memory.summarized_count() + memory.len(), source.len());
754 }
755
756 #[tokio::test]
757 async fn test_compression_failure_is_non_destructive() {
758 let config = CompactingMemoryConfig {
759 max_recent_messages: 2,
760 compress_threshold: 5,
761 summarize_batch_size: 3,
762 max_summary_length: 500,
763 };
764 let memory = CompactingMemory::new(Arc::new(NoopSummarizer), config);
765
766 for i in 0..5 {
767 memory
768 .add_message(make_message(&format!("msg{}", i)))
769 .await
770 .unwrap();
771 }
772 memory.compress(None).await.unwrap();
773 for i in 5..8 {
774 memory
775 .add_message(make_message(&format!("msg{}", i)))
776 .await
777 .unwrap();
778 }
779
780 let messages_before = memory.get_messages(None).await.unwrap();
781 let summary_before = memory.summary();
782 let summarized_count_before = memory.summarized_count();
783 let history_len_before = memory.compression_history().len();
784
785 let error = memory.compress(Some(&FailingMergeSummarizer)).await;
786 assert!(error.is_err());
787 let messages_after = memory.get_messages(None).await.unwrap();
788 let before_contents: Vec<_> = messages_before
789 .iter()
790 .map(|message| message.content.as_str())
791 .collect();
792 let after_contents: Vec<_> = messages_after
793 .iter()
794 .map(|message| message.content.as_str())
795 .collect();
796 assert_eq!(after_contents, before_contents);
797 assert_eq!(memory.summary(), summary_before);
798 assert_eq!(memory.summarized_count(), summarized_count_before);
799 assert_eq!(memory.compression_history().len(), history_len_before);
800 }
801
802 #[tokio::test]
803 async fn test_concurrent_compressions_are_serialized() {
804 let summarizer = Arc::new(BlockingSummarizer::new());
805 let config = CompactingMemoryConfig {
806 max_recent_messages: 2,
807 compress_threshold: 5,
808 summarize_batch_size: 3,
809 max_summary_length: 500,
810 };
811 let memory = Arc::new(CompactingMemory::new(summarizer.clone(), config));
812 for i in 0..5 {
813 memory
814 .add_message(make_message(&format!("msg{}", i)))
815 .await
816 .unwrap();
817 }
818
819 let first_memory = memory.clone();
820 let first = tokio::spawn(async move { first_memory.compress(None).await });
821 summarizer.wait_until_started().await;
822
823 let second_memory = memory.clone();
824 let mut second = tokio::spawn(async move { second_memory.compress(None).await });
825 assert!(
826 timeout(Duration::from_millis(50), &mut second)
827 .await
828 .is_err()
829 );
830
831 summarizer.release();
832 assert!(matches!(
833 first.await.unwrap().unwrap(),
834 CompressResult::Compressed {
835 messages_summarized: 3,
836 ..
837 }
838 ));
839 assert!(matches!(
840 second.await.unwrap().unwrap(),
841 CompressResult::NotNeeded
842 ));
843 assert_eq!(summarizer.batches(), vec![vec!["msg0", "msg1", "msg2"]]);
844 let remaining = memory.get_messages(None).await.unwrap();
845 let contents: Vec<_> = remaining
846 .iter()
847 .map(|message| message.content.as_str())
848 .collect();
849 assert_eq!(contents, vec!["msg3", "msg4"]);
850 }
851
852 #[tokio::test]
853 async fn test_compression_serializes_add_message() {
854 let summarizer = Arc::new(BlockingSummarizer::new());
855 let config = CompactingMemoryConfig {
856 max_recent_messages: 2,
857 compress_threshold: 5,
858 summarize_batch_size: 3,
859 max_summary_length: 500,
860 };
861 let memory = Arc::new(CompactingMemory::new(summarizer.clone(), config));
862 for i in 0..5 {
863 memory
864 .add_message(make_message(&format!("msg{}", i)))
865 .await
866 .unwrap();
867 }
868
869 let compress_memory = memory.clone();
870 let compress = tokio::spawn(async move { compress_memory.compress(None).await });
871 summarizer.wait_until_started().await;
872
873 let add_memory = memory.clone();
874 let mut add =
875 tokio::spawn(async move { add_memory.add_message(make_message("msg5")).await });
876 assert!(timeout(Duration::from_millis(50), &mut add).await.is_err());
877
878 summarizer.release();
879 compress.await.unwrap().unwrap();
880 add.await.unwrap().unwrap();
881 assert_eq!(summarizer.batches(), vec![vec!["msg0", "msg1", "msg2"]]);
882 let remaining = memory.get_messages(None).await.unwrap();
883 let contents: Vec<_> = remaining
884 .iter()
885 .map(|message| message.content.as_str())
886 .collect();
887 assert_eq!(contents, vec!["msg3", "msg4", "msg5"]);
888 }
889
890 #[tokio::test]
891 async fn test_compression_serializes_eviction() {
892 let summarizer = Arc::new(BlockingSummarizer::new());
893 let config = CompactingMemoryConfig {
894 max_recent_messages: 2,
895 compress_threshold: 5,
896 summarize_batch_size: 3,
897 max_summary_length: 500,
898 };
899 let memory = Arc::new(CompactingMemory::new(summarizer.clone(), config));
900 for i in 0..5 {
901 memory
902 .add_message(make_message(&format!("msg{}", i)))
903 .await
904 .unwrap();
905 }
906
907 let compress_memory = memory.clone();
908 let compress = tokio::spawn(async move { compress_memory.compress(None).await });
909 summarizer.wait_until_started().await;
910
911 let evict_memory = memory.clone();
912 let mut evict = tokio::spawn(async move { evict_memory.evict_oldest(1).await });
913 assert!(
914 timeout(Duration::from_millis(50), &mut evict)
915 .await
916 .is_err()
917 );
918
919 summarizer.release();
920 compress.await.unwrap().unwrap();
921 let evicted = evict.await.unwrap().unwrap();
922 assert_eq!(evicted.len(), 1);
923 assert_eq!(evicted[0].content, "msg3");
924 assert_eq!(summarizer.batches(), vec![vec!["msg0", "msg1", "msg2"]]);
925 let remaining = memory.get_messages(None).await.unwrap();
926 assert_eq!(remaining.len(), 1);
927 assert_eq!(remaining[0].content, "msg4");
928 }
929
930 #[tokio::test]
931 async fn test_get_context() {
932 let memory = create_test_memory();
933
934 for i in 0..6 {
935 memory
936 .add_message(make_message(&format!("msg{}", i)))
937 .await
938 .unwrap();
939 }
940
941 memory.compress(None).await.unwrap();
942
943 let ctx = memory.get_context().await.unwrap();
944 assert!(ctx.summary.is_some());
945 assert_eq!(ctx.messages.len(), 3);
946 assert_eq!(ctx.summarized_count, 3);
947 }
948
949 #[tokio::test]
950 async fn test_snapshot_restore() {
951 let memory = create_test_memory();
952
953 memory.add_message(make_message("msg1")).await.unwrap();
954 memory.add_message(make_message("msg2")).await.unwrap();
955
956 let snapshot = memory.snapshot().await.unwrap();
957 assert_eq!(snapshot.messages.len(), 2);
958
959 memory.clear().await.unwrap();
960 assert!(memory.is_empty());
961
962 memory.restore(snapshot).await.unwrap();
963 let messages = memory.get_messages(None).await.unwrap();
964 assert_eq!(messages.len(), 2);
965 }
966
967 #[tokio::test]
968 async fn test_snapshot_restore_preserves_recent_tail() {
969 let config = CompactingMemoryConfig {
970 max_recent_messages: 3,
971 compress_threshold: 5,
972 summarize_batch_size: 10,
973 max_summary_length: 500,
974 };
975 let memory = CompactingMemory::new(Arc::new(NoopSummarizer), config);
976
977 for i in 0..7 {
978 memory
979 .add_message(make_message(&format!("msg{}", i)))
980 .await
981 .unwrap();
982 }
983 memory.compress(None).await.unwrap();
984 let snapshot = memory.snapshot().await.unwrap();
985 assert_eq!(snapshot.summarized_count, 4);
986 let serialized = serde_json::to_string(&snapshot).unwrap();
987 let persisted: MemorySnapshot = serde_json::from_str(&serialized).unwrap();
988
989 memory.clear().await.unwrap();
990 memory.restore(persisted).await.unwrap();
991 assert!(memory.summary().is_some());
992 assert_eq!(memory.summarized_count(), 4);
993 for i in 7..9 {
994 memory
995 .add_message(make_message(&format!("msg{}", i)))
996 .await
997 .unwrap();
998 }
999 memory.compress(None).await.unwrap();
1000
1001 let remaining = memory.get_messages(None).await.unwrap();
1002 let contents: Vec<_> = remaining
1003 .iter()
1004 .map(|message| message.content.as_str())
1005 .collect();
1006 assert_eq!(contents, vec!["msg6", "msg7", "msg8"]);
1007 let context = memory.get_context().await.unwrap();
1008 assert_eq!(context.summarized_count, 6);
1009 assert_eq!(context.total_messages, 9);
1010 }
1011
1012 #[tokio::test]
1013 async fn test_compression_history() {
1014 let memory = create_test_memory();
1015
1016 for i in 0..6 {
1017 memory
1018 .add_message(make_message(&format!("msg{}", i)))
1019 .await
1020 .unwrap();
1021 }
1022
1023 memory.compress(None).await.unwrap();
1024
1025 let history = memory.compression_history();
1026 assert_eq!(history.len(), 1);
1027 assert_eq!(history[0].messages_compressed, 3);
1028 }
1029
1030 #[test]
1031 fn test_config_default() {
1032 let config = CompactingMemoryConfig::default();
1033 assert_eq!(config.max_recent_messages, 50);
1034 assert_eq!(config.compress_threshold, 30);
1035 assert_eq!(config.summarize_batch_size, 10);
1036 assert_eq!(config.max_summary_length, 2000);
1037 }
1038
1039 #[test]
1040 fn test_config_rejects_unknown_fields() {
1041 let yaml = r#"
1042max_recent_messages: 5
1043compress_thresold: 10
1044"#;
1045 let error = serde_yaml::from_str::<CompactingMemoryConfig>(yaml).unwrap_err();
1046 assert!(
1047 error
1048 .to_string()
1049 .contains("unknown field `compress_thresold`")
1050 );
1051 }
1052
1053 #[tokio::test]
1054 async fn test_evict_oldest() {
1055 let memory = create_test_memory();
1056 for i in 0..5 {
1057 memory
1058 .add_message(make_message(&format!("msg{}", i)))
1059 .await
1060 .unwrap();
1061 }
1062
1063 let evicted = memory.evict_oldest(2).await.unwrap();
1064 assert_eq!(evicted.len(), 2);
1065 assert_eq!(evicted[0].content, "msg0");
1066 assert_eq!(evicted[1].content, "msg1");
1067
1068 let remaining = memory.get_messages(None).await.unwrap();
1069 assert_eq!(remaining.len(), 3);
1070 assert_eq!(remaining[0].content, "msg2");
1071 }
1072
1073 #[test]
1074 fn test_prefix_at_char_boundary_handles_unicode() {
1075 let text = "계약서 내용을 확인하고 싶어서";
1076 let prefix = prefix_at_char_boundary(text, 5);
1077 assert_eq!(prefix.chars().count(), 5);
1078 assert!(text.starts_with(prefix));
1079 }
1080}