1use crate::{
7 cosine_similarity, Document, SearchResult, VectorDocument, VectorStore, VectorStoreError,
8};
9use async_trait::async_trait;
10use std::collections::HashMap;
11use std::sync::Arc;
12use tokio::sync::RwLock;
13use uuid::Uuid;
14
15pub struct InMemoryVectorStore {
17 documents: Arc<RwLock<HashMap<String, VectorDocument>>>,
19}
20
21impl InMemoryVectorStore {
22 pub fn new() -> Self {
24 Self {
25 documents: Arc::new(RwLock::new(HashMap::new())),
26 }
27 }
28}
29
30impl Default for InMemoryVectorStore {
31 fn default() -> Self {
32 Self::new()
33 }
34}
35
36#[async_trait]
37impl VectorStore for InMemoryVectorStore {
38 async fn add_documents(
39 &self,
40 documents: Vec<Document>,
41 embeddings: Vec<Vec<f32>>,
42 ) -> Result<Vec<String>, VectorStoreError> {
43 if documents.len() != embeddings.len() {
44 return Err(VectorStoreError::StorageError(
45 "document count and embedding count mismatch".to_string(),
46 ));
47 }
48
49 let mut store = self.documents.write().await;
50 let mut ids = Vec::new();
51
52 for (doc, embedding) in documents.into_iter().zip(embeddings.into_iter()) {
53 let id = doc.id.clone().unwrap_or_else(|| Uuid::new_v4().to_string());
54
55 let vector_doc = VectorDocument {
56 document: Document {
57 id: Some(id.clone()),
58 content: doc.content,
59 metadata: doc.metadata,
60 },
61 embedding,
62 };
63
64 store.insert(id.clone(), vector_doc);
65 ids.push(id);
66 }
67
68 Ok(ids)
69 }
70
71 async fn similarity_search(
72 &self,
73 query_embedding: &[f32],
74 k: usize,
75 ) -> Result<Vec<SearchResult>, VectorStoreError> {
76 self.similarity_search_with_min_score(query_embedding, k, None)
79 .await
80 }
81
82 async fn similarity_search_with_min_score(
83 &self,
84 query_embedding: &[f32],
85 k: usize,
86 min_score: Option<f32>,
87 ) -> Result<Vec<SearchResult>, VectorStoreError> {
88 let store = self.documents.read().await;
89
90 let mut results: Vec<SearchResult> = store
92 .values()
93 .filter_map(|vd| {
94 let score = cosine_similarity(query_embedding, &vd.embedding).unwrap_or(0.0);
95 if min_score.is_none_or(|t| score >= t) {
96 Some(SearchResult {
97 document: vd.document.clone(),
98 score,
99 })
100 } else {
101 None
102 }
103 })
104 .collect();
105
106 results.sort_by(|a, b| {
108 b.score
109 .partial_cmp(&a.score)
110 .unwrap_or(std::cmp::Ordering::Equal)
111 });
112
113 Ok(results.into_iter().take(k).collect())
115 }
116
117 async fn get_document(&self, id: &str) -> Result<Option<Document>, VectorStoreError> {
118 let store = self.documents.read().await;
119 Ok(store.get(id).map(|vd| vd.document.clone()))
120 }
121
122 async fn get_embedding(&self, id: &str) -> Result<Option<Vec<f32>>, VectorStoreError> {
123 let store = self.documents.read().await;
124 Ok(store.get(id).map(|vd| vd.embedding.clone()))
125 }
126
127 async fn delete_document(&self, id: &str) -> Result<(), VectorStoreError> {
128 let mut store = self.documents.write().await;
129 store.remove(id);
130 Ok(())
131 }
132
133 async fn count(&self) -> usize {
134 let store = self.documents.read().await;
135 store.len()
136 }
137
138 async fn clear(&self) -> Result<(), VectorStoreError> {
139 let mut store = self.documents.write().await;
140 store.clear();
141 Ok(())
142 }
143}
144
145#[cfg(test)]
146mod tests {
147 use super::*;
148
149 #[tokio::test]
150 async fn test_add_and_search() {
151 let store = InMemoryVectorStore::new();
152
153 let docs = vec![
155 Document::new("Rust is a systems programming language"),
156 Document::new("Python is a scripting language"),
157 Document::new("JavaScript is used for web development"),
158 ];
159
160 let embeddings = vec![
162 vec![1.0, 0.0, 0.0], vec![0.0, 1.0, 0.0], vec![0.0, 0.0, 1.0], ];
166
167 let ids = store.add_documents(docs, embeddings).await.unwrap();
168 assert_eq!(ids.len(), 3);
169 assert_eq!(store.count().await, 3);
170
171 let query = vec![0.9, 0.1, 0.0]; let results = store.similarity_search(&query, 2).await.unwrap();
174
175 assert_eq!(results.len(), 2);
176 assert!(results[0].document.content.contains("Rust"));
177 assert!(results[0].score > results[1].score);
178 }
179
180 #[tokio::test]
181 async fn test_get_and_delete() {
182 let store = InMemoryVectorStore::new();
183
184 let doc = Document::new("Test document").with_id("test-id");
185 let embeddings = vec![vec![1.0, 0.0, 0.0]];
186
187 store.add_documents(vec![doc], embeddings).await.unwrap();
188
189 let retrieved = store.get_document("test-id").await.unwrap();
191 assert!(retrieved.is_some());
192 assert_eq!(retrieved.unwrap().content, "Test document");
193
194 store.delete_document("test-id").await.unwrap();
196 assert_eq!(store.count().await, 0);
197
198 let retrieved = store.get_document("test-id").await.unwrap();
200 assert!(retrieved.is_none());
201 }
202
203 #[tokio::test]
204 async fn test_clear() {
205 let store = InMemoryVectorStore::new();
206
207 let docs = vec![Document::new("Doc 1"), Document::new("Doc 2")];
208 let embeddings = vec![vec![1.0, 0.0], vec![0.0, 1.0]];
209
210 store.add_documents(docs, embeddings).await.unwrap();
211 assert_eq!(store.count().await, 2);
212
213 store.clear().await.unwrap();
214 assert_eq!(store.count().await, 0);
215 }
216
217 #[tokio::test]
220 async fn test_similarity_search_text_without_embedder_errors() {
221 let store = InMemoryVectorStore::new();
222 let err = store.similarity_search_text("hello", 3).await.unwrap_err();
223 assert!(matches!(err, VectorStoreError::EmbeddingError(_)));
224 }
225
226 #[tokio::test]
229 async fn test_negative_scores_not_dropped() {
230 let store = InMemoryVectorStore::new();
231 store
232 .add_documents(
233 vec![
234 Document::new("orthogonal-up"),
235 Document::new("opposite"),
236 Document::new("orthogonal-down"),
237 ],
238 vec![vec![0.0, 1.0], vec![-1.0, 0.0], vec![0.0, -1.0]],
239 )
240 .await
241 .unwrap();
242
243 let query = vec![1.0, 0.0];
244
245 let results = store.similarity_search(&query, 3).await.unwrap();
247 assert_eq!(results.len(), 3);
248 assert!(results.iter().all(|r| r.score <= 0.0));
249
250 let filtered = store
252 .similarity_search_with_min_score(&query, 3, Some(-0.5))
253 .await
254 .unwrap();
255 assert_eq!(filtered.len(), 2);
256
257 let all = store
259 .similarity_search_with_min_score(&query, 3, None)
260 .await
261 .unwrap();
262 assert_eq!(all.len(), 3);
263 }
264
265 #[test]
266 fn test_cosine_similarity() {
267 let a = vec![1.0, 0.0, 0.0];
269 let b = vec![1.0, 0.0, 0.0];
270 assert!((cosine_similarity(&a, &b).unwrap() - 1.0).abs() < 0.0001);
271
272 let a = vec![1.0, 0.0, 0.0];
274 let b = vec![0.0, 1.0, 0.0];
275 assert!((cosine_similarity(&a, &b).unwrap() - 0.0).abs() < 0.0001);
276 }
277}