autoagents_core/agent/memory/
sliding_window.rs1use async_trait::async_trait;
6use autoagents_llm::{chat::ChatMessage, error::LLMError};
7use std::collections::VecDeque;
8
9use super::{MemoryProvider, MemoryType};
10
11#[derive(Debug, Clone)]
13pub enum TrimStrategy {
14 Drop,
16 Summarize,
22}
23
24#[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 pub fn new(window_size: usize) -> Self {
52 Self::with_strategy(window_size, TrimStrategy::Drop)
53 }
54
55 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 pub fn window_size(&self) -> usize {
80 self.window_size
81 }
82
83 pub fn messages(&self) -> Vec<ChatMessage> {
89 Vec::from(self.messages.clone())
90 }
91
92 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 pub fn needs_summary(&self) -> bool {
109 self.needs_summary
110 }
111
112 pub fn mark_for_summary(&mut self) {
114 self.needs_summary = true;
115 }
116
117 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 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 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 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 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 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); }
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 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 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 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 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}