xz_memory_engine/layered/
default.rs1use 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
18pub 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 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}