Skip to main content

rskit_agent/memory/
sliding_window.rs

1use std::sync::Arc;
2
3use async_trait::async_trait;
4use rskit_errors::{AppError, ErrorCode};
5use rskit_llm::types::Message;
6
7use super::Memory;
8
9// ── SlidingWindowMemory ─────────────────────────────────────────────────────
10
11/// Wraps any [`Memory`]
12/// and limits the stored history to the last `max_messages` non-system entries (always preserving a leading system message in addition to that limit).
13pub struct SlidingWindowMemory {
14    inner: Arc<dyn Memory>,
15    max_messages: usize,
16}
17
18impl SlidingWindowMemory {
19    /// Create a sliding-window wrapper.
20    ///
21    /// `max_messages` is the number of non-system messages to retain and must be at least 1.
22    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    /// Trim a message list:
36    /// keep the system prompt (if present) plus the last `max_messages` non-system messages.
37    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        // Re-load, trim, and persist.
72        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// ── Tests ───────────────────────────────────────────────────────────────────