Skip to main content

autoagents_core/agent/memory/
sliding_window.rs

1//! Simple sliding window memory implementation.
2//!
3//! This module provides a basic FIFO (First In, First Out) memory that maintains
4//! a fixed-size window of the most recent conversation messages.
5use async_trait::async_trait;
6use autoagents_llm::{chat::ChatMessage, error::LLMError};
7use std::collections::VecDeque;
8
9use super::{MemoryProvider, MemoryType};
10
11/// Strategy for handling memory when window size limit is reached
12#[derive(Debug, Clone)]
13pub enum TrimStrategy {
14    /// Drop oldest messages (FIFO behavior)
15    Drop,
16    /// Mark memory for manual summarization when the window overflows.
17    ///
18    /// The first overflow is retained so a caller can summarize the full
19    /// previous window plus the triggering message. Further writes fail until
20    /// [`SlidingWindowMemory::replace_with_summary`] clears the pending state.
21    Summarize,
22}
23
24/// Simple sliding window memory that keeps the N most recent messages.
25///
26/// This implementation uses a FIFO strategy where old messages are automatically
27/// removed when the window size limit is reached. It's suitable for:
28/// - Simple conversation contexts
29/// - Memory-constrained environments
30/// - Cases where only recent context matters
31///
32#[derive(Debug, Clone)]
33pub struct SlidingWindowMemory {
34    messages: VecDeque<ChatMessage>,
35    window_size: usize,
36    trim_strategy: TrimStrategy,
37    needs_summary: bool,
38}
39
40impl SlidingWindowMemory {
41    /// Create a new sliding window memory with the specified window size.
42    ///
43    /// # Arguments
44    ///
45    /// * `window_size` - Maximum number of messages to keep in memory
46    ///
47    /// # Panics
48    ///
49    /// Panics if `window_size` is 0
50    ///
51    pub fn new(window_size: usize) -> Self {
52        Self::with_strategy(window_size, TrimStrategy::Drop)
53    }
54
55    /// Create a new sliding window memory with specified trim strategy
56    ///
57    /// # Arguments
58    ///
59    /// * `window_size` - Maximum number of messages to keep in memory
60    /// * `strategy` - How to handle overflow when window is full
61    pub fn with_strategy(window_size: usize, strategy: TrimStrategy) -> Self {
62        if window_size == 0 {
63            panic!("Window size must be greater than 0");
64        }
65
66        Self {
67            messages: VecDeque::with_capacity(window_size),
68            window_size,
69            trim_strategy: strategy,
70            needs_summary: false,
71        }
72    }
73
74    /// Get the configured window size.
75    ///
76    /// # Returns
77    ///
78    /// The maximum number of messages this memory can hold
79    pub fn window_size(&self) -> usize {
80        self.window_size
81    }
82
83    /// Get all stored messages in chronological order.
84    ///
85    /// # Returns
86    ///
87    /// A vector containing all messages from oldest to newest
88    pub fn messages(&self) -> Vec<ChatMessage> {
89        Vec::from(self.messages.clone())
90    }
91
92    /// Get the most recent N messages.
93    ///
94    /// # Arguments
95    ///
96    /// * `limit` - Maximum number of recent messages to return
97    ///
98    /// # Returns
99    ///
100    /// A vector containing the most recent messages, up to `limit`
101    pub fn recent_messages(&self, limit: usize) -> Vec<ChatMessage> {
102        let len = self.messages.len();
103        let start = len.saturating_sub(limit);
104        self.messages.range(start..).cloned().collect()
105    }
106
107    /// Check if memory needs summarization
108    pub fn needs_summary(&self) -> bool {
109        self.needs_summary
110    }
111
112    /// Mark memory as needing summarization
113    pub fn mark_for_summary(&mut self) {
114        self.needs_summary = true;
115    }
116
117    /// Replace all messages with a summary
118    ///
119    /// # Arguments
120    ///
121    /// * `summary` - The summary text to replace all messages with
122    pub fn replace_with_summary(&mut self, summary: String) {
123        self.messages.clear();
124        self.messages
125            .push_back(ChatMessage::assistant().content(summary).build());
126        self.needs_summary = false;
127    }
128}
129
130#[async_trait]
131impl MemoryProvider for SlidingWindowMemory {
132    async fn remember(&mut self, message: &ChatMessage) -> Result<(), LLMError> {
133        if self.messages.len() >= self.window_size {
134            match self.trim_strategy {
135                TrimStrategy::Drop => {
136                    self.messages.pop_front();
137                }
138                TrimStrategy::Summarize => {
139                    if self.needs_summary {
140                        return Err(summary_required_error());
141                    }
142                    self.mark_for_summary();
143                }
144            }
145        }
146        self.messages.push_back(message.clone());
147        Ok(())
148    }
149
150    async fn remember_many(&mut self, messages: &[ChatMessage]) -> Result<(), LLMError> {
151        if messages.is_empty() {
152            return Ok(());
153        }
154
155        match self.trim_strategy {
156            TrimStrategy::Drop => {
157                for message in messages {
158                    self.remember(message).await?;
159                }
160                Ok(())
161            }
162            TrimStrategy::Summarize => {
163                if self.needs_summary {
164                    return Err(summary_required_error());
165                }
166
167                let projected_len = self.messages.len().saturating_add(messages.len());
168                if projected_len > self.window_size.saturating_add(1) {
169                    return Err(summary_required_error());
170                }
171
172                if projected_len > self.window_size {
173                    self.mark_for_summary();
174                }
175                self.messages.extend(messages.iter().cloned());
176                Ok(())
177            }
178        }
179    }
180
181    async fn recall(
182        &self,
183        _query: &str,
184        limit: Option<usize>,
185    ) -> Result<Vec<ChatMessage>, LLMError> {
186        let limit = limit.unwrap_or(self.messages.len());
187        Ok(self.recent_messages(limit))
188    }
189
190    async fn clear(&mut self) -> Result<(), LLMError> {
191        self.messages.clear();
192        self.needs_summary = false;
193        Ok(())
194    }
195
196    fn memory_type(&self) -> MemoryType {
197        MemoryType::SlidingWindow
198    }
199
200    fn size(&self) -> usize {
201        self.messages.len()
202    }
203
204    fn needs_summary(&self) -> bool {
205        self.needs_summary
206    }
207
208    fn mark_for_summary(&mut self) {
209        self.needs_summary = true;
210    }
211
212    fn replace_with_summary(&mut self, summary: String) {
213        self.messages.clear();
214        self.messages
215            .push_back(ChatMessage::assistant().content(summary).build());
216        self.needs_summary = false;
217    }
218
219    fn clone_box(&self) -> Box<dyn MemoryProvider> {
220        Box::new(self.clone())
221    }
222
223    fn preload(&mut self, data: Vec<ChatMessage>) -> bool {
224        self.messages.clear();
225        for msg in data {
226            self.messages.push_back(msg);
227        }
228        true
229    }
230
231    fn export(&self) -> Vec<ChatMessage> {
232        Vec::from(self.messages.clone())
233    }
234}
235
236fn summary_required_error() -> LLMError {
237    LLMError::ProviderError(
238        "SlidingWindowMemory requires summarization before accepting new messages".to_string(),
239    )
240}
241
242#[cfg(test)]
243mod tests {
244    use super::*;
245    use autoagents_llm::chat::{ChatMessage, ChatRole, MessageType};
246
247    #[test]
248    fn test_new_sliding_window_memory() {
249        let memory = SlidingWindowMemory::new(5);
250        assert_eq!(memory.window_size(), 5);
251        assert_eq!(memory.size(), 0);
252        assert!(memory.is_empty());
253        assert_eq!(memory.memory_type(), MemoryType::SlidingWindow);
254    }
255
256    #[test]
257    fn test_sliding_window_memory_with_strategy() {
258        let memory = SlidingWindowMemory::with_strategy(3, TrimStrategy::Summarize);
259        assert_eq!(memory.window_size(), 3);
260        assert_eq!(memory.size(), 0);
261        assert!(memory.is_empty());
262    }
263
264    #[test]
265    #[should_panic(expected = "Window size must be greater than 0")]
266    fn test_new_sliding_window_memory_zero_size() {
267        SlidingWindowMemory::new(0);
268    }
269
270    #[tokio::test]
271    async fn test_remember_single_message() {
272        let mut memory = SlidingWindowMemory::new(3);
273        let message = ChatMessage {
274            role: ChatRole::User,
275            message_type: MessageType::Text,
276            content: "Hello".to_string(),
277        };
278
279        memory.remember(&message).await.unwrap();
280        assert_eq!(memory.size(), 1);
281        assert!(!memory.is_empty());
282
283        let messages = memory.messages();
284        assert_eq!(messages.len(), 1);
285        assert_eq!(messages[0].content, "Hello");
286    }
287
288    #[tokio::test]
289    async fn test_remember_multiple_messages() {
290        let mut memory = SlidingWindowMemory::new(3);
291
292        for i in 1..=3 {
293            let message = ChatMessage {
294                role: ChatRole::User,
295                message_type: MessageType::Text,
296                content: format!("Message {i}"),
297            };
298            memory.remember(&message).await.unwrap();
299        }
300
301        assert_eq!(memory.size(), 3);
302        let messages = memory.messages();
303        assert_eq!(messages.len(), 3);
304        assert_eq!(messages[0].content, "Message 1");
305        assert_eq!(messages[2].content, "Message 3");
306    }
307
308    #[tokio::test]
309    async fn test_sliding_window_overflow_drop_strategy() {
310        let mut memory = SlidingWindowMemory::with_strategy(2, TrimStrategy::Drop);
311
312        // Add 3 messages to a window of size 2
313        for i in 1..=3 {
314            let message = ChatMessage {
315                role: ChatRole::User,
316                message_type: MessageType::Text,
317                content: format!("Message {i}"),
318            };
319            memory.remember(&message).await.unwrap();
320        }
321
322        // Should only keep the last 2 messages
323        assert_eq!(memory.size(), 2);
324        let messages = memory.messages();
325        assert_eq!(messages[0].content, "Message 2");
326        assert_eq!(messages[1].content, "Message 3");
327    }
328
329    #[tokio::test]
330    async fn test_sliding_window_overflow_summarize_strategy() {
331        let mut memory = SlidingWindowMemory::with_strategy(2, TrimStrategy::Summarize);
332
333        // Add first message
334        let message1 = ChatMessage {
335            role: ChatRole::User,
336            message_type: MessageType::Text,
337            content: "First message".to_string(),
338        };
339        memory.remember(&message1).await.unwrap();
340
341        // Add second message
342        let message2 = ChatMessage {
343            role: ChatRole::User,
344            message_type: MessageType::Text,
345            content: "Second message".to_string(),
346        };
347        memory.remember(&message2).await.unwrap();
348
349        // Add third message - should trigger summarize flag
350        let message3 = ChatMessage {
351            role: ChatRole::User,
352            message_type: MessageType::Text,
353            content: "Third message".to_string(),
354        };
355        memory.remember(&message3).await.unwrap();
356
357        assert!(memory.needs_summary());
358        assert_eq!(memory.size(), 3); // Keeps one overflow message for manual summarization.
359    }
360
361    #[tokio::test]
362    async fn test_summarize_rejects_writes_while_summary_pending() {
363        let mut memory = SlidingWindowMemory::with_strategy(2, TrimStrategy::Summarize);
364
365        for i in 1..=3 {
366            let message = ChatMessage {
367                role: ChatRole::User,
368                message_type: MessageType::Text,
369                content: format!("Message {i}"),
370            };
371            memory.remember(&message).await.unwrap();
372        }
373
374        let rejected = ChatMessage {
375            role: ChatRole::User,
376            message_type: MessageType::Text,
377            content: "Message 4".to_string(),
378        };
379        let result = memory.remember(&rejected).await;
380
381        assert!(matches!(result, Err(LLMError::ProviderError(_))));
382        assert!(memory.needs_summary());
383        assert_eq!(memory.size(), 3);
384        let messages = memory.messages();
385        assert_eq!(messages[0].content, "Message 1");
386        assert_eq!(messages[2].content, "Message 3");
387    }
388
389    #[tokio::test]
390    async fn test_summarize_accepts_writes_after_replace_with_summary() {
391        let mut memory = SlidingWindowMemory::with_strategy(2, TrimStrategy::Summarize);
392
393        for i in 1..=3 {
394            let message = ChatMessage {
395                role: ChatRole::User,
396                message_type: MessageType::Text,
397                content: format!("Message {i}"),
398            };
399            memory.remember(&message).await.unwrap();
400        }
401
402        memory.replace_with_summary("summary".to_string());
403
404        let next = ChatMessage {
405            role: ChatRole::User,
406            message_type: MessageType::Text,
407            content: "Message 4".to_string(),
408        };
409        memory.remember(&next).await.unwrap();
410
411        assert!(!memory.needs_summary());
412        assert_eq!(memory.size(), 2);
413        let messages = memory.messages();
414        assert_eq!(messages[0].content, "summary");
415        assert_eq!(messages[1].content, "Message 4");
416    }
417
418    #[tokio::test]
419    async fn test_clear_resets_pending_summary_state() {
420        let mut memory = SlidingWindowMemory::with_strategy(2, TrimStrategy::Summarize);
421
422        for i in 1..=3 {
423            let message = ChatMessage {
424                role: ChatRole::User,
425                message_type: MessageType::Text,
426                content: format!("Message {i}"),
427            };
428            memory.remember(&message).await.unwrap();
429        }
430        assert!(memory.needs_summary());
431
432        memory.clear().await.unwrap();
433
434        assert!(!memory.needs_summary());
435        assert_eq!(memory.size(), 0);
436
437        for i in 4..=6 {
438            let message = ChatMessage {
439                role: ChatRole::User,
440                message_type: MessageType::Text,
441                content: format!("Message {i}"),
442            };
443            memory.remember(&message).await.unwrap();
444        }
445
446        assert!(memory.needs_summary());
447        assert_eq!(memory.size(), 3);
448    }
449
450    #[tokio::test]
451    async fn test_summarize_batch_rejects_without_mutating_when_over_cap() {
452        let mut memory = SlidingWindowMemory::with_strategy(2, TrimStrategy::Summarize);
453
454        for i in 1..=2 {
455            let message = ChatMessage {
456                role: ChatRole::User,
457                message_type: MessageType::Text,
458                content: format!("Message {i}"),
459            };
460            memory.remember(&message).await.unwrap();
461        }
462
463        let batch = [
464            ChatMessage {
465                role: ChatRole::Assistant,
466                message_type: MessageType::Text,
467                content: "Message 3".to_string(),
468            },
469            ChatMessage {
470                role: ChatRole::Tool,
471                message_type: MessageType::Text,
472                content: "Message 4".to_string(),
473            },
474        ];
475        let result = memory.remember_many(&batch).await;
476
477        assert!(matches!(result, Err(LLMError::ProviderError(_))));
478        assert!(!memory.needs_summary());
479        assert_eq!(memory.size(), 2);
480        let messages = memory.messages();
481        assert_eq!(messages[0].content, "Message 1");
482        assert_eq!(messages[1].content, "Message 2");
483    }
484
485    #[tokio::test]
486    async fn test_recall_all_messages() {
487        let mut memory = SlidingWindowMemory::new(3);
488
489        for i in 1..=3 {
490            let message = ChatMessage {
491                role: ChatRole::User,
492                message_type: MessageType::Text,
493                content: format!("Message {i}"),
494            };
495            memory.remember(&message).await.unwrap();
496        }
497
498        let recalled = memory.recall("", None).await.unwrap();
499        assert_eq!(recalled.len(), 3);
500        assert_eq!(recalled[0].content, "Message 1");
501        assert_eq!(recalled[2].content, "Message 3");
502    }
503
504    #[tokio::test]
505    async fn test_recall_with_limit() {
506        let mut memory = SlidingWindowMemory::new(5);
507
508        for i in 1..=5 {
509            let message = ChatMessage {
510                role: ChatRole::User,
511                message_type: MessageType::Text,
512                content: format!("Message {i}"),
513            };
514            memory.remember(&message).await.unwrap();
515        }
516
517        let recalled = memory.recall("", Some(2)).await.unwrap();
518        assert_eq!(recalled.len(), 2);
519        assert_eq!(recalled[0].content, "Message 4");
520        assert_eq!(recalled[1].content, "Message 5");
521    }
522
523    #[tokio::test]
524    async fn test_clear_memory() {
525        let mut memory = SlidingWindowMemory::new(3);
526
527        let message = ChatMessage {
528            role: ChatRole::User,
529            message_type: MessageType::Text,
530            content: "Test message".to_string(),
531        };
532        memory.remember(&message).await.unwrap();
533
534        assert_eq!(memory.size(), 1);
535        memory.clear().await.unwrap();
536        assert_eq!(memory.size(), 0);
537        assert!(memory.is_empty());
538    }
539
540    #[test]
541    fn test_recent_messages() {
542        let mut memory = SlidingWindowMemory::new(5);
543
544        // Add messages directly to the internal deque for testing
545        for i in 1..=5 {
546            let message = ChatMessage {
547                role: ChatRole::User,
548                message_type: MessageType::Text,
549                content: format!("Message {i}"),
550            };
551            memory.messages.push_back(message);
552        }
553
554        let recent = memory.recent_messages(3);
555        assert_eq!(recent.len(), 3);
556        assert_eq!(recent[0].content, "Message 3");
557        assert_eq!(recent[2].content, "Message 5");
558    }
559
560    #[test]
561    fn test_recent_messages_limit_exceeds_size() {
562        let mut memory = SlidingWindowMemory::new(5);
563
564        // Add only 2 messages
565        for i in 1..=2 {
566            let message = ChatMessage {
567                role: ChatRole::User,
568                message_type: MessageType::Text,
569                content: format!("Message {i}"),
570            };
571            memory.messages.push_back(message);
572        }
573
574        let recent = memory.recent_messages(10);
575        assert_eq!(recent.len(), 2);
576        assert_eq!(recent[0].content, "Message 1");
577        assert_eq!(recent[1].content, "Message 2");
578    }
579
580    #[test]
581    fn test_mark_for_summary() {
582        let mut memory = SlidingWindowMemory::new(3);
583        assert!(!memory.needs_summary());
584
585        memory.mark_for_summary();
586        assert!(memory.needs_summary());
587    }
588
589    #[test]
590    fn test_replace_with_summary() {
591        let mut memory = SlidingWindowMemory::new(3);
592
593        // Add some messages
594        for i in 1..=3 {
595            let message = ChatMessage {
596                role: ChatRole::User,
597                message_type: MessageType::Text,
598                content: format!("Message {i}"),
599            };
600            memory.messages.push_back(message);
601        }
602
603        memory.mark_for_summary();
604        assert!(memory.needs_summary());
605        assert_eq!(memory.size(), 3);
606
607        memory.replace_with_summary("This is a summary".to_string());
608
609        assert!(!memory.needs_summary());
610        assert_eq!(memory.size(), 1);
611        let messages = memory.messages();
612        assert_eq!(messages[0].content, "This is a summary");
613        assert_eq!(messages[0].role, ChatRole::Assistant);
614    }
615
616    #[test]
617    fn test_memory_provider_trait_methods() {
618        let memory = SlidingWindowMemory::new(3);
619
620        // Test trait methods
621        assert_eq!(memory.memory_type(), MemoryType::SlidingWindow);
622        assert_eq!(memory.size(), 0);
623        assert!(memory.is_empty());
624        assert!(!memory.needs_summary());
625        assert!(memory.get_event_receiver().is_none());
626    }
627}