1use crate::error::VectorError;
9use crate::PgVectorStore;
10use crate::{SearchResult, VectorMetric, VectorRecord};
11use async_trait::async_trait;
12use std::collections::HashMap;
13use std::sync::RwLock;
14
15pub struct InMemoryVectorStore {
17 collections: RwLock<HashMap<String, CollectionState>>,
18}
19
20#[derive(Debug, Clone)]
21struct CollectionState {
22 dimension: usize,
23 metric: VectorMetric,
24 records: Vec<StoredRecord>,
25}
26
27#[derive(Debug, Clone)]
28struct StoredRecord {
29 id: String,
30 vector: Vec<f32>,
31 metadata: Option<HashMap<String, serde_json::Value>>,
32 text: Option<String>,
33}
34
35impl InMemoryVectorStore {
36 pub fn new() -> Self {
37 Self {
38 collections: RwLock::new(HashMap::new()),
39 }
40 }
41
42 fn metric_value(metric: VectorMetric, a: &[f32], b: &[f32]) -> f32 {
43 match metric {
44 VectorMetric::Cosine => cosine_similarity(a, b),
45 VectorMetric::Euclidean => {
46 let dist: f32 = a
47 .iter()
48 .zip(b.iter())
49 .map(|(x, y)| (x - y) * (x - y))
50 .sum::<f32>()
51 .sqrt();
52 1.0 / (1.0 + dist)
54 }
55 VectorMetric::DotProduct => a.iter().zip(b.iter()).map(|(x, y)| x * y).sum(),
56 }
57 }
58}
59
60impl Default for InMemoryVectorStore {
61 fn default() -> Self {
62 Self::new()
63 }
64}
65
66fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
68 if a.len() != b.len() || a.is_empty() {
69 return 0.0;
70 }
71 let dot: f32 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
72 let na: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
73 let nb: f32 = b.iter().map(|x| x * x).sum::<f32>().sqrt();
74 if na == 0.0 || nb == 0.0 {
75 return 0.0;
76 }
77 dot / (na * nb)
78}
79
80#[async_trait]
81impl PgVectorStore for InMemoryVectorStore {
82 async fn create_collection(
83 &self,
84 name: &str,
85 dimension: usize,
86 metric: Option<VectorMetric>,
87 ) -> Result<(), VectorError> {
88 let mut collections = self
89 .collections
90 .write()
91 .map_err(|e| VectorError::Query(format!("lock error: {}", e)))?;
92 collections.insert(
93 name.to_string(),
94 CollectionState {
95 dimension,
96 metric: metric.unwrap_or_default(),
97 records: Vec::new(),
98 },
99 );
100 Ok(())
101 }
102
103 async fn delete_collection(&self, name: &str) -> Result<(), VectorError> {
104 let mut collections = self
105 .collections
106 .write()
107 .map_err(|e| VectorError::Query(format!("lock error: {}", e)))?;
108 collections.remove(name);
109 Ok(())
110 }
111
112 async fn insert(
113 &self,
114 collection: &str,
115 records: Vec<VectorRecord>,
116 ) -> Result<(), VectorError> {
117 let mut collections = self
118 .collections
119 .write()
120 .map_err(|e| VectorError::Query(format!("lock error: {}", e)))?;
121 let state = collections
122 .get_mut(collection)
123 .ok_or_else(|| VectorError::CollectionNotFound(collection.to_string()))?;
124
125 for record in records {
126 if record.vector.len() != state.dimension {
127 return Err(VectorError::DimensionMismatch {
128 expected: state.dimension,
129 actual: record.vector.len(),
130 });
131 }
132 if let Some(existing) = state.records.iter_mut().find(|r| r.id == record.id) {
134 existing.vector = record.vector;
135 existing.metadata = record.metadata;
136 continue;
137 }
138 state.records.push(StoredRecord {
139 id: record.id,
140 vector: record.vector,
141 metadata: record.metadata,
142 text: None,
143 });
144 }
145 Ok(())
146 }
147
148 async fn search(
149 &self,
150 collection: &str,
151 query: &[f32],
152 top_k: usize,
153 ) -> Result<Vec<SearchResult>, VectorError> {
154 let top_k = crate::validate_top_k(top_k)?;
156
157 let collections = self
158 .collections
159 .read()
160 .map_err(|e| VectorError::Query(format!("lock error: {}", e)))?;
161 let state = collections
162 .get(collection)
163 .ok_or_else(|| VectorError::CollectionNotFound(collection.to_string()))?;
164
165 let mut scored: Vec<(usize, f32)> = state
166 .records
167 .iter()
168 .enumerate()
169 .map(|(i, r)| (i, Self::metric_value(state.metric, query, &r.vector)))
170 .collect();
171
172 scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
173
174 let k = top_k.min(scored.len());
175 let mut results = Vec::with_capacity(k);
176 for (idx, score) in scored.into_iter().take(k) {
177 let record = &state.records[idx];
178 let mut result = SearchResult::new(record.id.clone(), score, record.vector.clone());
179 if let Some(ref text) = record.text {
180 result = result.with_text(text.clone());
181 }
182 results.push(result);
183 }
184 Ok(results)
185 }
186
187 async fn get(&self, collection: &str, id: &str) -> Result<Option<VectorRecord>, VectorError> {
188 let collections = self
189 .collections
190 .read()
191 .map_err(|e| VectorError::Query(format!("lock error: {}", e)))?;
192 let state = collections
193 .get(collection)
194 .ok_or_else(|| VectorError::CollectionNotFound(collection.to_string()))?;
195 Ok(state
196 .records
197 .iter()
198 .find(|r| r.id == id)
199 .map(|r| VectorRecord {
200 id: r.id.clone(),
201 vector: r.vector.clone(),
202 score: None,
203 metadata: r.metadata.clone(),
204 }))
205 }
206
207 async fn delete(&self, collection: &str, ids: Vec<String>) -> Result<u64, VectorError> {
208 let mut collections = self
209 .collections
210 .write()
211 .map_err(|e| VectorError::Query(format!("lock error: {}", e)))?;
212 let state = collections
213 .get_mut(collection)
214 .ok_or_else(|| VectorError::CollectionNotFound(collection.to_string()))?;
215 let before = state.records.len();
216 state.records.retain(|r| !ids.contains(&r.id));
217 let removed = (before - state.records.len()) as u64;
218 Ok(removed)
219 }
220
221 async fn count(&self, collection: &str) -> Result<usize, VectorError> {
222 let collections = self
223 .collections
224 .read()
225 .map_err(|e| VectorError::Query(format!("lock error: {}", e)))?;
226 Ok(collections
227 .get(collection)
228 .map(|s| s.records.len())
229 .unwrap_or(0))
230 }
231}
232
233#[cfg(test)]
234mod tests {
235 use super::*;
236 use crate::VectorMetric;
237
238 #[tokio::test]
239 async fn test_create_and_delete_collection() {
240 let store = InMemoryVectorStore::new();
241 store.create_collection("docs", 4, None).await.unwrap();
242 assert_eq!(store.count("docs").await.unwrap(), 0);
243
244 store.delete_collection("docs").await.unwrap();
245 assert_eq!(store.count("docs").await.unwrap(), 0);
246 }
247
248 #[tokio::test]
249 async fn test_insert_and_get() {
250 let store = InMemoryVectorStore::new();
251 store.create_collection("docs", 3, None).await.unwrap();
252 let rec = VectorRecord::new("r1", vec![1.0, 0.0, 0.0]);
253 store.insert("docs", vec![rec]).await.unwrap();
254 assert_eq!(store.count("docs").await.unwrap(), 1);
255
256 let fetched = store.get("docs", "r1").await.unwrap().unwrap();
257 assert_eq!(fetched.id, "r1");
258 assert_eq!(fetched.vector, vec![1.0, 0.0, 0.0]);
259
260 assert!(store.get("docs", "missing").await.unwrap().is_none());
261 }
262
263 #[tokio::test]
264 async fn test_insert_dimension_mismatch() {
265 let store = InMemoryVectorStore::new();
266 store.create_collection("docs", 3, None).await.unwrap();
267 let rec = VectorRecord::new("r1", vec![1.0, 0.0]); let err = store.insert("docs", vec![rec]).await;
269 assert!(err.is_err());
270 assert!(matches!(err, Err(VectorError::DimensionMismatch { .. })));
271 }
272
273 #[tokio::test]
274 async fn test_insert_upsert() {
275 let store = InMemoryVectorStore::new();
276 store.create_collection("docs", 2, None).await.unwrap();
277 store
278 .insert("docs", vec![VectorRecord::new("r1", vec![1.0, 0.0])])
279 .await
280 .unwrap();
281 store
282 .insert("docs", vec![VectorRecord::new("r1", vec![0.0, 1.0])])
283 .await
284 .unwrap();
285 assert_eq!(store.count("docs").await.unwrap(), 1);
287 let fetched = store.get("docs", "r1").await.unwrap().unwrap();
288 assert_eq!(fetched.vector, vec![0.0, 1.0]);
289 }
290
291 #[tokio::test]
292 async fn test_search_cosine_returns_closest_first() {
293 let store = InMemoryVectorStore::new();
294 store
295 .create_collection("docs", 3, Some(VectorMetric::Cosine))
296 .await
297 .unwrap();
298 let records = vec![
299 VectorRecord::new("a", vec![1.0, 0.0, 0.0]),
300 VectorRecord::new("b", vec![0.0, 1.0, 0.0]),
301 VectorRecord::new("c", vec![1.0, 1.0, 0.0]),
302 ];
303 store.insert("docs", records).await.unwrap();
304
305 let results = store.search("docs", &[1.0, 0.0, 0.0], 2).await.unwrap();
306 assert_eq!(results.len(), 2);
307 assert_eq!(results[0].id, "a");
308 assert!(results[0].score > results[1].score);
309 }
310
311 #[tokio::test]
312 async fn test_search_top_k_limit() {
313 let store = InMemoryVectorStore::new();
314 store.create_collection("docs", 2, None).await.unwrap();
315 for i in 0..5 {
316 store
317 .insert(
318 "docs",
319 vec![VectorRecord::new(format!("r{}", i), vec![i as f32, 1.0])],
320 )
321 .await
322 .unwrap();
323 }
324 let results = store.search("docs", &[0.0, 1.0], 3).await.unwrap();
325 assert_eq!(results.len(), 3);
326 }
327
328 #[tokio::test]
329 async fn test_delete_records() {
330 let store = InMemoryVectorStore::new();
331 store.create_collection("docs", 2, None).await.unwrap();
332 store
333 .insert(
334 "docs",
335 vec![
336 VectorRecord::new("a", vec![1.0, 0.0]),
337 VectorRecord::new("b", vec![0.0, 1.0]),
338 VectorRecord::new("c", vec![1.0, 1.0]),
339 ],
340 )
341 .await
342 .unwrap();
343 let removed = store
344 .delete("docs", vec!["a".to_string(), "c".to_string()])
345 .await
346 .unwrap();
347 assert_eq!(removed, 2);
348 assert_eq!(store.count("docs").await.unwrap(), 1);
349 }
350
351 #[tokio::test]
352 async fn test_search_euclidean() {
353 let store = InMemoryVectorStore::new();
354 store
355 .create_collection("docs", 2, Some(VectorMetric::Euclidean))
356 .await
357 .unwrap();
358 let records = vec![
359 VectorRecord::new("near", vec![0.0, 0.0]),
360 VectorRecord::new("far", vec![10.0, 10.0]),
361 ];
362 store.insert("docs", records).await.unwrap();
363
364 let results = store.search("docs", &[0.0, 0.0], 2).await.unwrap();
365 assert_eq!(results.len(), 2);
366 assert_eq!(results[0].id, "near");
367 assert!(results[0].score > results[1].score);
368 }
369
370 #[tokio::test]
371 async fn test_search_dot_product() {
372 let store = InMemoryVectorStore::new();
373 store
374 .create_collection("docs", 2, Some(VectorMetric::DotProduct))
375 .await
376 .unwrap();
377 let records = vec![
378 VectorRecord::new("high", vec![2.0, 3.0]),
379 VectorRecord::new("low", vec![0.0, 0.0]),
380 ];
381 store.insert("docs", records).await.unwrap();
382
383 let results = store.search("docs", &[1.0, 1.0], 2).await.unwrap();
384 assert_eq!(results.len(), 2);
385 assert_eq!(results[0].id, "high");
386 }
387
388 #[tokio::test]
389 async fn test_collection_not_found() {
390 let store = InMemoryVectorStore::new();
391 let result = store.count("nonexistent").await;
392 assert_eq!(result.unwrap(), 0);
393
394 let err = store.search("nonexistent", &[1.0, 0.0], 5).await;
395 assert!(matches!(err, Err(VectorError::CollectionNotFound(_))));
396 }
397
398 #[tokio::test]
399 async fn test_get_nonexistent_record() {
400 let store = InMemoryVectorStore::new();
401 store.create_collection("docs", 2, None).await.unwrap();
402 let result = store.get("docs", "nonexistent").await.unwrap();
403 assert!(result.is_none());
404 }
405
406 #[tokio::test]
407 async fn test_helpers_compile() {
408 let _ = InMemoryVectorStore::new();
409 assert_eq!(cosine_similarity(&[1.0, 0.0], &[1.0, 0.0]), 1.0);
410 assert!((cosine_similarity(&[1.0, 0.0], &[0.0, 1.0])).abs() < 1e-6);
411 assert_eq!(cosine_similarity(&[], &[]), 0.0);
412 }
413
414 #[tokio::test]
416 async fn test_m16_top_k_zero_rejected() {
417 let store = InMemoryVectorStore::new();
418 store.create_collection("docs", 2, None).await.unwrap();
419 store
420 .insert("docs", vec![VectorRecord::new("a", vec![1.0, 0.0])])
421 .await
422 .unwrap();
423 let err = store.search("docs", &[1.0, 0.0], 0).await;
424 assert!(matches!(
425 err,
426 Err(VectorError::TopKExceeded {
427 requested: 0,
428 max: crate::MAX_TOP_K
429 })
430 ));
431 }
432
433 #[tokio::test]
435 async fn test_m16_top_k_exceeded_rejected() {
436 let store = InMemoryVectorStore::new();
437 store.create_collection("docs", 2, None).await.unwrap();
438 store
439 .insert("docs", vec![VectorRecord::new("a", vec![1.0, 0.0])])
440 .await
441 .unwrap();
442 let err = store
443 .search("docs", &[1.0, 0.0], crate::MAX_TOP_K + 1)
444 .await;
445 assert!(matches!(
446 err,
447 Err(VectorError::TopKExceeded { requested, max }) if requested == crate::MAX_TOP_K + 1 && max == crate::MAX_TOP_K
448 ));
449 }
450
451 #[tokio::test]
453 async fn test_m16_top_k_max_allowed() {
454 let store = InMemoryVectorStore::new();
455 store.create_collection("docs", 2, None).await.unwrap();
456 store
457 .insert("docs", vec![VectorRecord::new("a", vec![1.0, 0.0])])
458 .await
459 .unwrap();
460 let results = store
462 .search("docs", &[1.0, 0.0], crate::MAX_TOP_K)
463 .await
464 .unwrap();
465 assert_eq!(results.len(), 1);
466 }
467
468 #[test]
470 fn test_m16_validate_top_k_function() {
471 use crate::validate_top_k;
472 assert_eq!(validate_top_k(1).unwrap(), 1);
474 assert_eq!(validate_top_k(100).unwrap(), 100);
475 assert_eq!(validate_top_k(crate::MAX_TOP_K).unwrap(), crate::MAX_TOP_K);
476 assert!(validate_top_k(0).is_err());
478 assert!(validate_top_k(crate::MAX_TOP_K + 1).is_err());
479 assert!(validate_top_k(usize::MAX).is_err());
480 }
481}