Skip to main content

xz_memory_engine/layered/
default.rs

1//! Default [`LayeredMemory`] implementation.
2//!
3//! Combines an [`EntryStore`] for append/query/evict operations with an
4//! [`IndexSearcher`] for future search capabilities.  Implements the
5//! [`ConversationMemory`] trait using a `conv:{session_id}` partition scheme.
6
7use std::sync::Arc;
8
9use async_trait::async_trait;
10use chrono::Utc;
11use uuid::Uuid;
12use xz_memory_core::{
13    Entry, EntryStore, IndexSearcher, QueryOptions, SortOrder, StoreError, TimeRange,
14};
15
16use super::traits::ConversationMemory;
17
18/// Default layered memory backed by an entry store and index searcher.
19pub struct LayeredMemory<S: EntryStore, I: IndexSearcher> {
20    pub(crate) store: Arc<S>,
21    pub(crate) searcher: Arc<I>,
22}
23
24impl<S: EntryStore, I: IndexSearcher> LayeredMemory<S, I> {
25    /// Create a new `LayeredMemory`.
26    pub fn new(store: Arc<S>, searcher: Arc<I>) -> Self {
27        Self { store, searcher }
28    }
29}
30
31#[async_trait]
32impl<S: EntryStore, I: IndexSearcher> ConversationMemory for LayeredMemory<S, I> {
33    async fn append(&self, session_id: &str, message: &str) -> Result<(), StoreError> {
34        let partition = format!("conv:{}", session_id);
35        let entry = Entry {
36            id: Uuid::new_v4().to_string(),
37            partition,
38            body: message.to_string(),
39            recorded_at: Utc::now().timestamp_millis() as u64,
40        };
41        self.store.append(entry).await
42    }
43
44    async fn recent(&self, session_id: &str, n: usize) -> Result<Vec<String>, StoreError> {
45        let partition = format!("conv:{}", session_id);
46        let limit = if n == 0 { usize::MAX } else { n };
47        let opts = QueryOptions { limit, sort: SortOrder::Descending };
48        let range = TimeRange { start: None, end: None };
49        let mut entries = self.store.query(&partition, &range, &opts).await?;
50        entries.reverse();
51        Ok(entries.into_iter().map(|e| e.body).collect())
52    }
53
54    async fn evict(&self, session_id: &str, keep: usize) -> Result<usize, StoreError> {
55        let partition = format!("conv:{}", session_id);
56        self.store.evict(&partition, keep).await
57    }
58}
59
60#[cfg(test)]
61mod tests {
62    use super::*;
63    use std::sync::Arc;
64    use tokio::time::{Duration, sleep};
65    use xz_memory_core::{ScoredEntry, SearchOptions};
66
67    use crate::backends::InMemoryEntryStore;
68
69    struct MockSearcher;
70
71    #[async_trait]
72    impl IndexSearcher for MockSearcher {
73        async fn search(
74            &self,
75            _partitions: &[String],
76            _query: &str,
77            _opts: &SearchOptions,
78        ) -> Result<Vec<ScoredEntry>, StoreError> {
79            Ok(vec![])
80        }
81    }
82
83    #[tokio::test]
84    async fn test_append_and_recent() {
85        let store = Arc::new(InMemoryEntryStore::new());
86        let searcher = Arc::new(MockSearcher);
87        let memory = LayeredMemory::new(store, searcher);
88
89        memory.append("sess-1", "Hello").await.unwrap();
90        sleep(Duration::from_millis(2)).await;
91        memory.append("sess-1", "World").await.unwrap();
92
93        let messages = memory.recent("sess-1", 10).await.unwrap();
94        assert_eq!(messages.len(), 2);
95        assert_eq!(messages[0], "Hello");
96        assert_eq!(messages[1], "World");
97    }
98
99    #[tokio::test]
100    async fn test_evict() {
101        let store = Arc::new(InMemoryEntryStore::new());
102        let searcher = Arc::new(MockSearcher);
103        let memory = LayeredMemory::new(store, searcher);
104
105        for msg in ["a", "b", "c", "d", "e"] {
106            memory.append("sess-1", msg).await.unwrap();
107            sleep(Duration::from_millis(2)).await;
108        }
109
110        let evicted = memory.evict("sess-1", 2).await.unwrap();
111        assert_eq!(evicted, 3);
112
113        let remaining = memory.recent("sess-1", 10).await.unwrap();
114        assert_eq!(remaining.len(), 2);
115        assert_eq!(remaining[0], "d");
116        assert_eq!(remaining[1], "e");
117    }
118
119    #[tokio::test]
120    async fn test_empty_recent() {
121        let store = Arc::new(InMemoryEntryStore::new());
122        let searcher = Arc::new(MockSearcher);
123        let memory = LayeredMemory::new(store, searcher);
124
125        let messages = memory.recent("nonexistent", 10).await.unwrap();
126        assert!(messages.is_empty());
127    }
128}