Skip to main content

meerkat_memory/
simple.rs

1//! SimpleMemoryStore — basic in-memory keyword-matching memory store.
2//!
3//! This is a simple implementation that uses substring matching for search.
4//! A production implementation would use vector embeddings (e.g., HNSW).
5
6use 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/// Entry in the simple memory store.
14#[derive(Debug, Clone)]
15struct MemoryEntry {
16    scope: MemoryIndexScope,
17    content: String,
18    metadata: MemoryMetadata,
19}
20
21/// Simple in-memory store using substring matching.
22///
23/// This is a **test-only** implementation. Production use cases should use a
24/// vector-embedding-based store (e.g. HNSW). The substring matching here is
25/// not suitable for semantic search.
26pub struct SimpleMemoryStore {
27    entries: RwLock<Vec<MemoryEntry>>,
28}
29
30impl SimpleMemoryStore {
31    /// Create a new empty memory store.
32    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        // Sort by score descending
103        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}