rskit_agent/memory/
sliding_window.rs1use std::sync::Arc;
2
3use async_trait::async_trait;
4use rskit_errors::{AppError, ErrorCode};
5use rskit_llm::types::Message;
6
7use super::Memory;
8
9pub struct SlidingWindowMemory {
14 inner: Arc<dyn Memory>,
15 max_messages: usize,
16}
17
18impl SlidingWindowMemory {
19 pub fn new(inner: Arc<dyn Memory>, max_messages: usize) -> Result<Self, AppError> {
23 if max_messages == 0 {
24 return Err(AppError::new(
25 ErrorCode::InvalidInput,
26 "max_messages must be at least 1",
27 ));
28 }
29 Ok(Self {
30 inner,
31 max_messages,
32 })
33 }
34
35 fn trim(&self, messages: &[Message]) -> Vec<Message> {
38 if messages.len() <= self.max_messages {
39 return messages.to_vec();
40 }
41
42 let mut result = Vec::with_capacity(self.max_messages + 1);
43
44 let has_system = matches!(messages.first(), Some(Message::System(_)));
45 if has_system {
46 result.push(messages[0].clone());
47 }
48
49 let non_system = if has_system { &messages[1..] } else { messages };
50 let start = non_system.len().saturating_sub(self.max_messages);
51 result.extend_from_slice(&non_system[start..]);
52
53 result
54 }
55}
56
57#[async_trait]
58impl Memory for SlidingWindowMemory {
59 async fn load(&self, session_id: &str) -> Result<Vec<Message>, AppError> {
60 let messages = self.inner.load(session_id).await?;
61 Ok(self.trim(&messages))
62 }
63
64 async fn save(&self, session_id: &str, messages: &[Message]) -> Result<(), AppError> {
65 let trimmed = self.trim(messages);
66 self.inner.save(session_id, &trimmed).await
67 }
68
69 async fn append(&self, session_id: &str, messages: &[Message]) -> Result<(), AppError> {
70 self.inner.append(session_id, messages).await?;
71 let all = self.inner.load(session_id).await?;
73 let trimmed = self.trim(&all);
74 self.inner.save(session_id, &trimmed).await
75 }
76
77 async fn clear(&self, session_id: &str) -> Result<(), AppError> {
78 self.inner.clear(session_id).await
79 }
80}
81
82