Skip to main content

lc_vector_stores/
memory.rs

1// lc-vector-stores/src/memory.rs
2//! In-memory vector store
3//!
4//! Stores documents and vectors in memory, suitable for small-scale data and tests.
5
6use crate::{
7    cosine_similarity, Document, MetadataFilter, SearchResult, VectorDocument, VectorStore,
8    VectorStoreError,
9};
10use async_trait::async_trait;
11use std::collections::HashMap;
12use std::sync::Arc;
13use tokio::sync::RwLock;
14use uuid::Uuid;
15
16/// In-memory vector store
17pub struct InMemoryVectorStore {
18    /// Document storage
19    documents: Arc<RwLock<HashMap<String, VectorDocument>>>,
20}
21
22impl InMemoryVectorStore {
23    /// Creates a new in-memory vector store
24    pub fn new() -> Self {
25        Self {
26            documents: Arc::new(RwLock::new(HashMap::new())),
27        }
28    }
29}
30
31impl Default for InMemoryVectorStore {
32    fn default() -> Self {
33        Self::new()
34    }
35}
36
37#[async_trait]
38impl VectorStore for InMemoryVectorStore {
39    async fn add_documents(
40        &self,
41        documents: Vec<Document>,
42        embeddings: Vec<Vec<f32>>,
43    ) -> Result<Vec<String>, VectorStoreError> {
44        if documents.len() != embeddings.len() {
45            return Err(VectorStoreError::StorageError(
46                "document count and embedding count mismatch".to_string(),
47            ));
48        }
49
50        let mut store = self.documents.write().await;
51        let mut ids = Vec::new();
52
53        for (doc, embedding) in documents.into_iter().zip(embeddings.into_iter()) {
54            let id = doc.id.clone().unwrap_or_else(|| Uuid::new_v4().to_string());
55
56            let vector_doc = VectorDocument {
57                document: Document {
58                    id: Some(id.clone()),
59                    content: doc.content,
60                    metadata: doc.metadata,
61                },
62                embedding,
63            };
64
65            store.insert(id.clone(), vector_doc);
66            ids.push(id);
67        }
68
69        Ok(ids)
70    }
71
72    async fn similarity_search(
73        &self,
74        query_embedding: &[f32],
75        k: usize,
76    ) -> Result<Vec<SearchResult>, VectorStoreError> {
77        // Q2: no longer hard-filters score > 0 — under an all-negative corpus the top-k
78        // should still be returned; whether to set a threshold is the caller's explicit
79        // decision via similarity_search_with_min_score.
80        self.similarity_search_with_min_score(query_embedding, k, None)
81            .await
82    }
83
84    async fn similarity_search_with_min_score(
85        &self,
86        query_embedding: &[f32],
87        k: usize,
88        min_score: Option<f32>,
89    ) -> Result<Vec<SearchResult>, VectorStoreError> {
90        let store = self.documents.read().await;
91
92        // compute similarity for all documents, filter by threshold first, then take top-k (Q2)
93        let mut results: Vec<SearchResult> = store
94            .values()
95            .filter_map(|vd| {
96                let score = cosine_similarity(query_embedding, &vd.embedding).unwrap_or(0.0);
97                if min_score.is_none_or(|t| score >= t) {
98                    Some(SearchResult {
99                        document: vd.document.clone(),
100                        score,
101                    })
102                } else {
103                    None
104                }
105            })
106            .collect();
107
108        // sort by similarity descending
109        results.sort_by(|a, b| {
110            b.score
111                .partial_cmp(&a.score)
112                .unwrap_or(std::cmp::Ordering::Equal)
113        });
114
115        // return the top k results
116        Ok(results.into_iter().take(k).collect())
117    }
118
119    async fn similarity_search_with_filter(
120        &self,
121        query_embedding: &[f32],
122        k: usize,
123        filter: Option<&MetadataFilter>,
124    ) -> Result<Vec<SearchResult>, VectorStoreError> {
125        let store = self.documents.read().await;
126
127        // S3: in-memory metadata filtering — filter documents by the condition first, then compute similarity and take top-k.
128        let mut results: Vec<SearchResult> = store
129            .values()
130            .filter(|vd| filter.is_none_or(|f| f.matches(&vd.document.metadata)))
131            .map(|vd| {
132                let score = cosine_similarity(query_embedding, &vd.embedding).unwrap_or(0.0);
133                SearchResult {
134                    document: vd.document.clone(),
135                    score,
136                }
137            })
138            .collect();
139
140        results.sort_by(|a, b| {
141            b.score
142                .partial_cmp(&a.score)
143                .unwrap_or(std::cmp::Ordering::Equal)
144        });
145
146        Ok(results.into_iter().take(k).collect())
147    }
148
149    async fn get_document(&self, id: &str) -> Result<Option<Document>, VectorStoreError> {
150        let store = self.documents.read().await;
151        Ok(store.get(id).map(|vd| vd.document.clone()))
152    }
153
154    async fn get_embedding(&self, id: &str) -> Result<Option<Vec<f32>>, VectorStoreError> {
155        let store = self.documents.read().await;
156        Ok(store.get(id).map(|vd| vd.embedding.clone()))
157    }
158
159    async fn delete_document(&self, id: &str) -> Result<(), VectorStoreError> {
160        let mut store = self.documents.write().await;
161        store.remove(id);
162        Ok(())
163    }
164
165    async fn count(&self) -> usize {
166        let store = self.documents.read().await;
167        store.len()
168    }
169
170    async fn clear(&self) -> Result<(), VectorStoreError> {
171        let mut store = self.documents.write().await;
172        store.clear();
173        Ok(())
174    }
175}
176
177#[cfg(test)]
178mod tests {
179    use super::*;
180
181    #[tokio::test]
182    async fn test_add_and_search() {
183        let store = InMemoryVectorStore::new();
184
185        // add documents
186        let docs = vec![
187            Document::new("Rust is a systems programming language"),
188            Document::new("Python is a scripting language"),
189            Document::new("JavaScript is used for web development"),
190        ];
191
192        // create simple mock embedding vectors
193        let embeddings = vec![
194            vec![1.0, 0.0, 0.0], // Rust-related
195            vec![0.0, 1.0, 0.0], // Python-related
196            vec![0.0, 0.0, 1.0], // JavaScript-related
197        ];
198
199        let ids = store.add_documents(docs, embeddings).await.unwrap();
200        assert_eq!(ids.len(), 3);
201        assert_eq!(store.count().await, 3);
202
203        // search for similar documents
204        let query = vec![0.9, 0.1, 0.0]; // closer to Rust
205        let results = store.similarity_search(&query, 2).await.unwrap();
206
207        assert_eq!(results.len(), 2);
208        assert!(results[0].document.content.contains("Rust"));
209        assert!(results[0].score > results[1].score);
210    }
211
212    #[tokio::test]
213    async fn test_get_and_delete() {
214        let store = InMemoryVectorStore::new();
215
216        let doc = Document::new("Test document").with_id("test-id");
217        let embeddings = vec![vec![1.0, 0.0, 0.0]];
218
219        store.add_documents(vec![doc], embeddings).await.unwrap();
220
221        // get the document
222        let retrieved = store.get_document("test-id").await.unwrap();
223        assert!(retrieved.is_some());
224        assert_eq!(retrieved.unwrap().content, "Test document");
225
226        // delete the document
227        store.delete_document("test-id").await.unwrap();
228        assert_eq!(store.count().await, 0);
229
230        // fetching again should return None
231        let retrieved = store.get_document("test-id").await.unwrap();
232        assert!(retrieved.is_none());
233    }
234
235    #[tokio::test]
236    async fn test_clear() {
237        let store = InMemoryVectorStore::new();
238
239        let docs = vec![Document::new("Doc 1"), Document::new("Doc 2")];
240        let embeddings = vec![vec![1.0, 0.0], vec![0.0, 1.0]];
241
242        store.add_documents(docs, embeddings).await.unwrap();
243        assert_eq!(store.count().await, 2);
244
245        store.clear().await.unwrap();
246        assert_eq!(store.count().await, 0);
247    }
248
249    /// Q1: without a configured embedder, similarity_search_text must report EmbeddingError
250    /// explicitly, rather than silently succeeding or panicking.
251    #[tokio::test]
252    async fn test_similarity_search_text_without_embedder_errors() {
253        let store = InMemoryVectorStore::new();
254        let err = store.similarity_search_text("hello", 3).await.unwrap_err();
255        assert!(matches!(err, VectorStoreError::EmbeddingError(_)));
256    }
257
258    /// Q2: under an all-non-positive-score corpus, similarity_search still returns the top-k
259    /// (no longer cleared by a score>0 hard filter); similarity_search_with_min_score filters
260    /// explicitly by threshold.
261    #[tokio::test]
262    async fn test_negative_scores_not_dropped() {
263        let store = InMemoryVectorStore::new();
264        store
265            .add_documents(
266                vec![
267                    Document::new("orthogonal-up"),
268                    Document::new("opposite"),
269                    Document::new("orthogonal-down"),
270                ],
271                vec![vec![0.0, 1.0], vec![-1.0, 0.0], vec![0.0, -1.0]],
272            )
273            .await
274            .unwrap();
275
276        let query = vec![1.0, 0.0];
277
278        // the old implementation hard-filtered score > 0.0, which would return empty here; now it returns the top-k (3 items, all non-positive).
279        let results = store.similarity_search(&query, 3).await.unwrap();
280        assert_eq!(results.len(), 3);
281        assert!(results.iter().all(|r| r.score <= 0.0));
282
283        // explicit threshold: score >= -0.5 excludes the score = -1.0 entry
284        let filtered = store
285            .similarity_search_with_min_score(&query, 3, Some(-0.5))
286            .await
287            .unwrap();
288        assert_eq!(filtered.len(), 2);
289
290        // min_score = None behaves identically to similarity_search
291        let all = store
292            .similarity_search_with_min_score(&query, 3, None)
293            .await
294            .unwrap();
295        assert_eq!(all.len(), 3);
296    }
297
298    #[test]
299    fn test_cosine_similarity() {
300        // Identical vectors
301        let a = vec![1.0, 0.0, 0.0];
302        let b = vec![1.0, 0.0, 0.0];
303        assert!((cosine_similarity(&a, &b).unwrap() - 1.0).abs() < 0.0001);
304
305        // Orthogonal vectors
306        let a = vec![1.0, 0.0, 0.0];
307        let b = vec![0.0, 1.0, 0.0];
308        assert!((cosine_similarity(&a, &b).unwrap() - 0.0).abs() < 0.0001);
309    }
310
311    /// S3: in-memory metadata filtering — single condition + AND/OR combination; `filter: None` matches the legacy path.
312    #[tokio::test]
313    async fn test_metadata_filter() {
314        use crate::FilterOp;
315
316        let store = InMemoryVectorStore::new();
317        store
318            .add_documents(
319                vec![
320                    Document::new("rust doc")
321                        .with_metadata("lang", "rust")
322                        .with_metadata("year", 2024),
323                    Document::new("python doc")
324                        .with_metadata("lang", "python")
325                        .with_metadata("year", 2023),
326                    Document::new("rust legacy")
327                        .with_metadata("lang", "rust")
328                        .with_metadata("year", 2020),
329                ],
330                vec![
331                    vec![1.0, 0.0, 0.0],
332                    vec![0.0, 1.0, 0.0],
333                    vec![0.9, 0.1, 0.0],
334                ],
335            )
336            .await
337            .unwrap();
338
339        let query = vec![1.0, 0.0, 0.0];
340
341        // single condition
342        let eq = MetadataFilter::field("lang", FilterOp::Eq, "rust");
343        let r = store
344            .similarity_search_with_filter(&query, 5, Some(&eq))
345            .await
346            .unwrap();
347        assert_eq!(r.len(), 2);
348        assert!(r
349            .iter()
350            .all(|s| s.document.metadata.get("lang").and_then(|v| v.as_str()) == Some("rust")));
351
352        // AND combination
353        let and = MetadataFilter::and(vec![
354            MetadataFilter::field("lang", FilterOp::Eq, "rust"),
355            MetadataFilter::field("year", FilterOp::Gte, 2021),
356        ]);
357        let r = store
358            .similarity_search_with_filter(&query, 5, Some(&and))
359            .await
360            .unwrap();
361        assert_eq!(r.len(), 1);
362        assert!(r[0].document.content.contains("rust doc"));
363
364        // OR combination
365        let or = MetadataFilter::or(vec![
366            MetadataFilter::field("lang", FilterOp::Eq, "python"),
367            MetadataFilter::field("year", FilterOp::Lt, 2021),
368        ]);
369        let r = store
370            .similarity_search_with_filter(&query, 5, Some(&or))
371            .await
372            .unwrap();
373        assert_eq!(r.len(), 2);
374
375        // filter: None behaves identically to similarity_search (regression)
376        let none = store
377            .similarity_search_with_filter(&query, 5, None)
378            .await
379            .unwrap();
380        let base = store.similarity_search(&query, 5).await.unwrap();
381        assert_eq!(none.len(), base.len());
382    }
383}