1use async_trait::async_trait;
7use meerkat_core::memory::{
8 MemoryEnumerationPage, MemoryEnumerationRequest, MemoryIndexBatch, MemoryIndexReceipt,
9 MemoryIndexScope, MemoryMetadata, MemoryOwner, MemoryRecord, MemoryResult,
10 MemoryScopeDropReceipt, MemorySearchScope, MemoryStore, MemoryStoreError,
11};
12use tokio::sync::RwLock;
13
14#[derive(Debug, Clone)]
16struct MemoryEntry {
17 scope: MemoryIndexScope,
18 content: String,
19 metadata: MemoryMetadata,
20}
21
22pub struct SimpleMemoryStore {
28 entries: RwLock<Vec<MemoryEntry>>,
29}
30
31impl SimpleMemoryStore {
32 pub fn new() -> Self {
34 Self {
35 entries: RwLock::new(Vec::new()),
36 }
37 }
38}
39
40impl Default for SimpleMemoryStore {
41 fn default() -> Self {
42 Self::new()
43 }
44}
45
46#[async_trait]
47impl MemoryStore for SimpleMemoryStore {
48 async fn index_scoped_batch(
49 &self,
50 batch: MemoryIndexBatch,
51 ) -> Result<MemoryIndexReceipt, MemoryStoreError> {
52 let (receipt_scope, requests) = batch.into_parts();
53 let mut entries = self.entries.write().await;
54 let mut indexed_entries = 0usize;
55 for request in requests {
56 let (scope, content, metadata) = request.into_parts();
57 if !content.is_indexable() {
60 continue;
61 }
62 entries.push(MemoryEntry {
63 scope,
64 content: content.into_indexable_text(),
65 metadata,
66 });
67 indexed_entries += 1;
68 }
69 Ok(MemoryIndexReceipt {
70 scope: receipt_scope,
71 indexed_entries,
72 })
73 }
74
75 async fn search(
76 &self,
77 scope: &MemorySearchScope,
78 query: &str,
79 limit: usize,
80 ) -> Result<Vec<MemoryResult>, MemoryStoreError> {
81 let entries = self.entries.read().await;
82
83 let query_lower = query.to_lowercase();
84 let query_words: Vec<&str> = query_lower.split_whitespace().collect();
85
86 let mut results: Vec<MemoryResult> = entries
87 .iter()
88 .filter(|entry| entry.scope.owner == scope.owner && scope.includes(&entry.metadata))
89 .filter_map(|entry| {
90 let content_lower = entry.content.to_lowercase();
91 let matching_words = query_words
92 .iter()
93 .filter(|w| content_lower.contains(**w))
94 .count();
95
96 if matching_words == 0 {
97 return None;
98 }
99
100 let score = matching_words as f32 / query_words.len().max(1) as f32;
101 Some(MemoryResult {
102 content: entry.content.clone(),
103 metadata: entry.metadata.clone(),
104 score,
105 })
106 })
107 .collect();
108
109 results.sort_by(|a, b| {
111 b.score
112 .partial_cmp(&a.score)
113 .unwrap_or(std::cmp::Ordering::Equal)
114 });
115 results.truncate(limit);
116
117 Ok(results)
118 }
119
120 async fn drop_scope(
121 &self,
122 owner: &MemoryOwner,
123 ) -> Result<MemoryScopeDropReceipt, MemoryStoreError> {
124 let mut entries = self.entries.write().await;
125 let before = entries.len();
126 entries.retain(|entry| entry.scope.owner != *owner);
127 Ok(MemoryScopeDropReceipt {
128 owner: owner.clone(),
129 dropped_entries: before - entries.len(),
130 })
131 }
132
133 async fn enumerate_scoped(
134 &self,
135 scope: &MemorySearchScope,
136 request: MemoryEnumerationRequest,
137 ) -> Result<MemoryEnumerationPage, MemoryStoreError> {
138 if request.limit == 0 {
139 return Err(MemoryStoreError::EnumerationLimitZero);
140 }
141 let entries = self.entries.read().await;
142
143 let scoped: Vec<&MemoryEntry> = entries
147 .iter()
148 .filter(|entry| entry.scope.owner == scope.owner && scope.includes(&entry.metadata))
149 .collect();
150 let window_start = request.offset.min(scoped.len());
151 let window_end = request
152 .offset
153 .saturating_add(request.limit)
154 .min(scoped.len());
155 let rows_scanned = window_end - window_start;
156
157 let records = scoped[window_start..window_end]
158 .iter()
159 .filter(|entry| request.admits(&entry.metadata))
160 .map(|entry| MemoryRecord {
161 content: entry.content.clone(),
162 metadata: entry.metadata.clone(),
163 })
164 .collect();
165 let next_offset =
166 (window_end < scoped.len()).then(|| request.offset.saturating_add(rows_scanned));
167
168 Ok(MemoryEnumerationPage {
169 records,
170 next_offset,
171 })
172 }
173}
174
175#[cfg(test)]
176#[allow(clippy::unwrap_used, clippy::expect_used)]
177mod tests {
178 use super::*;
179 use meerkat_core::memory::{MemoryIndexRequest, MemorySource, MessageRange};
180 use meerkat_core::types::SessionId;
181 use std::time::{Duration, SystemTime, UNIX_EPOCH};
182
183 fn meta(session_id: &SessionId) -> MemoryMetadata {
184 MemoryMetadata {
185 session_id: session_id.clone(),
186 source: MemorySource::Compaction {
187 source_range: MessageRange::single(0),
188 },
189 indexed_at: SystemTime::now(),
190 }
191 }
192
193 fn request(content: impl Into<String>, session_id: &SessionId) -> MemoryIndexRequest {
194 MemoryIndexRequest::new(
195 MemoryIndexScope::for_session(session_id.clone()),
196 meerkat_core::MemoryIndexableContent::Indexable(content.into()),
197 meta(session_id),
198 )
199 .unwrap()
200 }
201
202 #[tokio::test]
203 async fn test_index_and_search() {
204 let store = SimpleMemoryStore::new();
205 let session_id = SessionId::new();
206 let scope = MemorySearchScope::for_session(session_id.clone());
207 let other_session_id = SessionId::new();
208
209 store
210 .index_scoped(request(
211 "The user wants to implement a REST API",
212 &session_id,
213 ))
214 .await
215 .unwrap();
216 store
217 .index_scoped(request("Configuration uses TOML format", &session_id))
218 .await
219 .unwrap();
220 store
221 .index_scoped(request("Authentication uses JWT tokens", &other_session_id))
222 .await
223 .unwrap();
224
225 {
226 let entries = store.entries.read().await;
227 assert!(
228 entries
229 .iter()
230 .all(|entry| entry.scope.includes(&entry.metadata))
231 );
232 assert_eq!(entries[0].scope.session_id(), &session_id);
233 }
234
235 let results = store.search(&scope, "REST API", 10).await.unwrap();
236 assert!(!results.is_empty());
237 assert!(results[0].content.contains("REST API"));
238 assert!(
239 results
240 .iter()
241 .all(|result| scope.includes(&result.metadata))
242 );
243 }
244
245 #[tokio::test]
246 async fn test_search_empty_store() {
247 let store = SimpleMemoryStore::new();
248 let scope = MemorySearchScope::for_session(SessionId::new());
249 let results = store.search(&scope, "anything", 10).await.unwrap();
250 assert!(results.is_empty());
251 }
252
253 #[tokio::test]
254 async fn test_search_limit() {
255 let store = SimpleMemoryStore::new();
256 let session_id = SessionId::new();
257 let scope = MemorySearchScope::for_session(session_id.clone());
258
259 for i in 0..10 {
260 store
261 .index_scoped(request(format!("Item {i} with keyword test"), &session_id))
262 .await
263 .unwrap();
264 }
265
266 let results = store.search(&scope, "test", 3).await.unwrap();
267 assert_eq!(results.len(), 3);
268 }
269
270 #[tokio::test]
271 async fn test_search_no_match() {
272 let store = SimpleMemoryStore::new();
273 let session_id = SessionId::new();
274 let scope = MemorySearchScope::for_session(session_id.clone());
275 store
276 .index_scoped(request("Hello world", &session_id))
277 .await
278 .unwrap();
279
280 let results = store.search(&scope, "quantum computing", 10).await.unwrap();
281 assert!(results.is_empty());
282 }
283
284 #[test]
285 fn test_index_request_rejects_metadata_outside_scope() {
286 let session_id = SessionId::new();
287 let other_session_id = SessionId::new();
288 let error = MemoryIndexRequest::new(
289 MemoryIndexScope::for_session(session_id),
290 meerkat_core::MemoryIndexableContent::Indexable("outside scope".to_string()),
291 meta(&other_session_id),
292 )
293 .unwrap_err();
294
295 assert!(matches!(error, MemoryStoreError::Scope(_)));
296 }
297
298 fn request_with(
299 content: impl Into<String>,
300 session_id: &SessionId,
301 source_range: MessageRange,
302 indexed_at: SystemTime,
303 ) -> MemoryIndexRequest {
304 MemoryIndexRequest::new(
305 MemoryIndexScope::for_session(session_id.clone()),
306 meerkat_core::MemoryIndexableContent::Indexable(content.into()),
307 MemoryMetadata {
308 session_id: session_id.clone(),
309 source: MemorySource::Compaction { source_range },
310 indexed_at,
311 },
312 )
313 .unwrap()
314 }
315
316 fn enumeration(limit: usize, offset: usize) -> MemoryEnumerationRequest {
317 MemoryEnumerationRequest {
318 limit,
319 offset,
320 source_overlap: None,
321 indexed_after: None,
322 }
323 }
324
325 #[tokio::test]
328 async fn test_drop_scope_removes_only_owner_entries() {
329 let store = SimpleMemoryStore::new();
330 let session_a = SessionId::new();
331 let session_b = SessionId::new();
332 let scope_a = MemorySearchScope::for_session(session_a.clone());
333 let scope_b = MemorySearchScope::for_session(session_b.clone());
334
335 store
336 .index_scoped(request("doomed alpha entry", &session_a))
337 .await
338 .unwrap();
339 store
340 .index_scoped(request("doomed beta entry", &session_a))
341 .await
342 .unwrap();
343 store
344 .index_scoped(request("surviving gamma entry", &session_b))
345 .await
346 .unwrap();
347
348 let receipt = store
349 .drop_scope(&MemoryOwner::canonical_session(session_a.clone()))
350 .await
351 .unwrap();
352 assert_eq!(receipt.dropped_entries, 2);
353 assert_eq!(receipt.owner.session_id(), &session_a);
354
355 let dropped = store.search(&scope_a, "doomed", 10).await.unwrap();
356 assert!(dropped.is_empty());
357 let surviving = store.search(&scope_b, "surviving gamma", 10).await.unwrap();
358 assert_eq!(surviving.len(), 1);
359
360 let repeat = store
362 .drop_scope(&MemoryOwner::canonical_session(session_a))
363 .await
364 .unwrap();
365 assert_eq!(repeat.dropped_entries, 0);
366 }
367
368 #[tokio::test]
371 async fn test_enumerate_scoped_pages_in_insertion_order() {
372 let store = SimpleMemoryStore::new();
373 let session_id = SessionId::new();
374 let other_session = SessionId::new();
375 let scope = MemorySearchScope::for_session(session_id.clone());
376
377 let texts = ["entry zero", "entry one", "entry two", "entry three"];
378 for (i, text) in texts.iter().enumerate() {
379 store
380 .index_scoped(request(*text, &session_id))
381 .await
382 .unwrap();
383 store
384 .index_scoped(request(format!("interloper {i}"), &other_session))
385 .await
386 .unwrap();
387 }
388
389 let first = store
390 .enumerate_scoped(&scope, enumeration(3, 0))
391 .await
392 .unwrap();
393 assert_eq!(first.records.len(), 3);
394 assert_eq!(first.records[0].content, "entry zero");
395 assert_eq!(first.records[1].content, "entry one");
396 assert_eq!(first.records[2].content, "entry two");
397 assert_eq!(first.next_offset, Some(3));
398
399 let last = store
400 .enumerate_scoped(&scope, enumeration(3, 3))
401 .await
402 .unwrap();
403 assert_eq!(last.records.len(), 1);
404 assert_eq!(last.records[0].content, "entry three");
405 assert_eq!(last.next_offset, None);
406
407 let beyond = store
408 .enumerate_scoped(&scope, enumeration(3, 9))
409 .await
410 .unwrap();
411 assert!(beyond.records.is_empty());
412 assert_eq!(beyond.next_offset, None);
413 }
414
415 #[tokio::test]
419 async fn test_enumerate_scoped_filters_apply_after_raw_paging() {
420 let store = SimpleMemoryStore::new();
421 let session_id = SessionId::new();
422 let scope = MemorySearchScope::for_session(session_id.clone());
423 let early = UNIX_EPOCH + Duration::from_secs(1_000);
424 let late = UNIX_EPOCH + Duration::from_secs(2_000);
425
426 store
427 .index_scoped(request_with(
428 "covers zero to five early",
429 &session_id,
430 MessageRange::new(0, 5).unwrap(),
431 early,
432 ))
433 .await
434 .unwrap();
435 store
436 .index_scoped(request_with(
437 "covers five to ten late",
438 &session_id,
439 MessageRange::new(5, 10).unwrap(),
440 late,
441 ))
442 .await
443 .unwrap();
444 store
445 .index_scoped(request_with(
446 "covers ten to fifteen late",
447 &session_id,
448 MessageRange::new(10, 15).unwrap(),
449 late,
450 ))
451 .await
452 .unwrap();
453
454 let overlap = store
457 .enumerate_scoped(
458 &scope,
459 MemoryEnumerationRequest {
460 limit: 10,
461 offset: 0,
462 source_overlap: Some(MessageRange::new(6, 8).unwrap()),
463 indexed_after: None,
464 },
465 )
466 .await
467 .unwrap();
468 assert_eq!(overlap.records.len(), 1);
469 assert_eq!(overlap.records[0].content, "covers five to ten late");
470 assert_eq!(overlap.next_offset, None);
471
472 let after = store
474 .enumerate_scoped(
475 &scope,
476 MemoryEnumerationRequest {
477 limit: 10,
478 offset: 0,
479 source_overlap: None,
480 indexed_after: Some(early),
481 },
482 )
483 .await
484 .unwrap();
485 assert_eq!(after.records.len(), 2);
486 assert!(after.records.iter().all(|r| r.content.contains("late")));
487
488 let error = store
492 .enumerate_scoped(&scope, enumeration(0, 1))
493 .await
494 .expect_err("limit zero must be rejected");
495 assert!(matches!(
496 error,
497 meerkat_core::memory::MemoryStoreError::EnumerationLimitZero
498 ));
499 }
500}