1use async_trait::async_trait;
7use meerkat_core::memory::{
8 MemoryIndexBatch, MemoryIndexReceipt, MemoryIndexScope, MemoryMetadata, MemoryResult,
9 MemorySearchScope, MemoryStore, MemoryStoreError,
10};
11use tokio::sync::RwLock;
12
13#[derive(Debug, Clone)]
15struct MemoryEntry {
16 scope: MemoryIndexScope,
17 content: String,
18 metadata: MemoryMetadata,
19}
20
21pub struct SimpleMemoryStore {
27 entries: RwLock<Vec<MemoryEntry>>,
28}
29
30impl SimpleMemoryStore {
31 pub fn new() -> Self {
33 Self {
34 entries: RwLock::new(Vec::new()),
35 }
36 }
37}
38
39impl Default for SimpleMemoryStore {
40 fn default() -> Self {
41 Self::new()
42 }
43}
44
45#[async_trait]
46impl MemoryStore for SimpleMemoryStore {
47 async fn index_scoped_batch(
48 &self,
49 batch: MemoryIndexBatch,
50 ) -> Result<MemoryIndexReceipt, MemoryStoreError> {
51 let (receipt_scope, requests) = batch.into_parts();
52 let indexed_entries = requests.len();
53 let mut entries = self.entries.write().await;
54 for request in requests {
55 let (scope, content, metadata) = request.into_parts();
56 entries.push(MemoryEntry {
57 scope,
58 content,
59 metadata,
60 });
61 }
62 Ok(MemoryIndexReceipt {
63 scope: receipt_scope,
64 indexed_entries,
65 })
66 }
67
68 async fn search(
69 &self,
70 scope: &MemorySearchScope,
71 query: &str,
72 limit: usize,
73 ) -> Result<Vec<MemoryResult>, MemoryStoreError> {
74 let entries = self.entries.read().await;
75
76 let query_lower = query.to_lowercase();
77 let query_words: Vec<&str> = query_lower.split_whitespace().collect();
78
79 let mut results: Vec<MemoryResult> = entries
80 .iter()
81 .filter(|entry| entry.scope.owner == scope.owner && scope.includes(&entry.metadata))
82 .filter_map(|entry| {
83 let content_lower = entry.content.to_lowercase();
84 let matching_words = query_words
85 .iter()
86 .filter(|w| content_lower.contains(**w))
87 .count();
88
89 if matching_words == 0 {
90 return None;
91 }
92
93 let score = matching_words as f32 / query_words.len().max(1) as f32;
94 Some(MemoryResult {
95 content: entry.content.clone(),
96 metadata: entry.metadata.clone(),
97 score,
98 })
99 })
100 .collect();
101
102 results.sort_by(|a, b| {
104 b.score
105 .partial_cmp(&a.score)
106 .unwrap_or(std::cmp::Ordering::Equal)
107 });
108 results.truncate(limit);
109
110 Ok(results)
111 }
112}
113
114#[cfg(test)]
115#[allow(clippy::unwrap_used, clippy::expect_used)]
116mod tests {
117 use super::*;
118 use meerkat_core::memory::MemoryIndexRequest;
119 use meerkat_core::types::SessionId;
120 use std::time::SystemTime;
121
122 fn meta(session_id: &SessionId) -> MemoryMetadata {
123 MemoryMetadata {
124 session_id: session_id.clone(),
125 turn: Some(1),
126 indexed_at: SystemTime::now(),
127 }
128 }
129
130 fn request(content: impl Into<String>, session_id: &SessionId) -> MemoryIndexRequest {
131 MemoryIndexRequest::new(
132 MemoryIndexScope::for_session(session_id.clone()),
133 content.into(),
134 meta(session_id),
135 )
136 .unwrap()
137 }
138
139 #[tokio::test]
140 async fn test_index_and_search() {
141 let store = SimpleMemoryStore::new();
142 let session_id = SessionId::new();
143 let scope = MemorySearchScope::for_session(session_id.clone());
144 let other_session_id = SessionId::new();
145
146 store
147 .index_scoped(request(
148 "The user wants to implement a REST API",
149 &session_id,
150 ))
151 .await
152 .unwrap();
153 store
154 .index_scoped(request("Configuration uses TOML format", &session_id))
155 .await
156 .unwrap();
157 store
158 .index_scoped(request("Authentication uses JWT tokens", &other_session_id))
159 .await
160 .unwrap();
161
162 {
163 let entries = store.entries.read().await;
164 assert!(
165 entries
166 .iter()
167 .all(|entry| entry.scope.includes(&entry.metadata))
168 );
169 assert_eq!(entries[0].scope.session_id(), &session_id);
170 }
171
172 let results = store.search(&scope, "REST API", 10).await.unwrap();
173 assert!(!results.is_empty());
174 assert!(results[0].content.contains("REST API"));
175 assert!(
176 results
177 .iter()
178 .all(|result| scope.includes(&result.metadata))
179 );
180 }
181
182 #[tokio::test]
183 async fn test_search_empty_store() {
184 let store = SimpleMemoryStore::new();
185 let scope = MemorySearchScope::for_session(SessionId::new());
186 let results = store.search(&scope, "anything", 10).await.unwrap();
187 assert!(results.is_empty());
188 }
189
190 #[tokio::test]
191 async fn test_search_limit() {
192 let store = SimpleMemoryStore::new();
193 let session_id = SessionId::new();
194 let scope = MemorySearchScope::for_session(session_id.clone());
195
196 for i in 0..10 {
197 store
198 .index_scoped(request(format!("Item {i} with keyword test"), &session_id))
199 .await
200 .unwrap();
201 }
202
203 let results = store.search(&scope, "test", 3).await.unwrap();
204 assert_eq!(results.len(), 3);
205 }
206
207 #[tokio::test]
208 async fn test_search_no_match() {
209 let store = SimpleMemoryStore::new();
210 let session_id = SessionId::new();
211 let scope = MemorySearchScope::for_session(session_id.clone());
212 store
213 .index_scoped(request("Hello world", &session_id))
214 .await
215 .unwrap();
216
217 let results = store.search(&scope, "quantum computing", 10).await.unwrap();
218 assert!(results.is_empty());
219 }
220
221 #[test]
222 fn test_index_request_rejects_metadata_outside_scope() {
223 let session_id = SessionId::new();
224 let other_session_id = SessionId::new();
225 let error = MemoryIndexRequest::new(
226 MemoryIndexScope::for_session(session_id),
227 "outside scope".to_string(),
228 meta(&other_session_id),
229 )
230 .unwrap_err();
231
232 assert!(matches!(error, MemoryStoreError::Scope(_)));
233 }
234}