1use 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
16pub struct InMemoryVectorStore {
18 documents: Arc<RwLock<HashMap<String, VectorDocument>>>,
20}
21
22impl InMemoryVectorStore {
23 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 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 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 results.sort_by(|a, b| {
110 b.score
111 .partial_cmp(&a.score)
112 .unwrap_or(std::cmp::Ordering::Equal)
113 });
114
115 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 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 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 let embeddings = vec![
194 vec![1.0, 0.0, 0.0], vec![0.0, 1.0, 0.0], vec![0.0, 0.0, 1.0], ];
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 let query = vec![0.9, 0.1, 0.0]; 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 let retrieved = store.get_document("test-id").await.unwrap();
223 assert!(retrieved.is_some());
224 assert_eq!(retrieved.unwrap().content, "Test document");
225
226 store.delete_document("test-id").await.unwrap();
228 assert_eq!(store.count().await, 0);
229
230 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 #[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 #[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 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 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 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 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 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 #[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 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 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 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 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}