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 if let Some(ref metadata) = record.metadata {
184 result = result.with_metadata(metadata.clone());
185 }
186 results.push(result);
187 }
188 Ok(results)
189 }
190
191 async fn get(&self, collection: &str, id: &str) -> Result<Option<VectorRecord>, VectorError> {
192 let collections = self
193 .collections
194 .read()
195 .map_err(|e| VectorError::Query(format!("lock error: {}", e)))?;
196 let state = collections
197 .get(collection)
198 .ok_or_else(|| VectorError::CollectionNotFound(collection.to_string()))?;
199 Ok(state
200 .records
201 .iter()
202 .find(|r| r.id == id)
203 .map(|r| VectorRecord {
204 id: r.id.clone(),
205 vector: r.vector.clone(),
206 score: None,
207 metadata: r.metadata.clone(),
208 }))
209 }
210
211 async fn delete(&self, collection: &str, ids: Vec<String>) -> Result<u64, VectorError> {
212 let mut collections = self
213 .collections
214 .write()
215 .map_err(|e| VectorError::Query(format!("lock error: {}", e)))?;
216 let state = collections
217 .get_mut(collection)
218 .ok_or_else(|| VectorError::CollectionNotFound(collection.to_string()))?;
219 let before = state.records.len();
220 state.records.retain(|r| !ids.contains(&r.id));
221 let removed = (before - state.records.len()) as u64;
222 Ok(removed)
223 }
224
225 async fn count(&self, collection: &str) -> Result<usize, VectorError> {
226 let collections = self
227 .collections
228 .read()
229 .map_err(|e| VectorError::Query(format!("lock error: {}", e)))?;
230 Ok(collections
231 .get(collection)
232 .map(|s| s.records.len())
233 .unwrap_or(0))
234 }
235}
236
237#[cfg(test)]
238mod tests {
239 use super::*;
240 use crate::VectorMetric;
241
242 #[tokio::test]
243 async fn test_create_and_delete_collection() {
244 let store = InMemoryVectorStore::new();
245 store.create_collection("docs", 4, None).await.unwrap();
246 assert_eq!(store.count("docs").await.unwrap(), 0);
247
248 store.delete_collection("docs").await.unwrap();
249 assert_eq!(store.count("docs").await.unwrap(), 0);
250 }
251
252 #[tokio::test]
253 async fn test_insert_and_get() {
254 let store = InMemoryVectorStore::new();
255 store.create_collection("docs", 3, None).await.unwrap();
256 let rec = VectorRecord::new("r1", vec![1.0, 0.0, 0.0]);
257 store.insert("docs", vec![rec]).await.unwrap();
258 assert_eq!(store.count("docs").await.unwrap(), 1);
259
260 let fetched = store.get("docs", "r1").await.unwrap().unwrap();
261 assert_eq!(fetched.id, "r1");
262 assert_eq!(fetched.vector, vec![1.0, 0.0, 0.0]);
263
264 assert!(store.get("docs", "missing").await.unwrap().is_none());
265 }
266
267 #[tokio::test]
268 async fn test_insert_dimension_mismatch() {
269 let store = InMemoryVectorStore::new();
270 store.create_collection("docs", 3, None).await.unwrap();
271 let rec = VectorRecord::new("r1", vec![1.0, 0.0]); let err = store.insert("docs", vec![rec]).await;
273 assert!(err.is_err());
274 assert!(matches!(err, Err(VectorError::DimensionMismatch { .. })));
275 }
276
277 #[tokio::test]
278 async fn test_insert_upsert() {
279 let store = InMemoryVectorStore::new();
280 store.create_collection("docs", 2, None).await.unwrap();
281 store
282 .insert("docs", vec![VectorRecord::new("r1", vec![1.0, 0.0])])
283 .await
284 .unwrap();
285 store
286 .insert("docs", vec![VectorRecord::new("r1", vec![0.0, 1.0])])
287 .await
288 .unwrap();
289 assert_eq!(store.count("docs").await.unwrap(), 1);
291 let fetched = store.get("docs", "r1").await.unwrap().unwrap();
292 assert_eq!(fetched.vector, vec![0.0, 1.0]);
293 }
294
295 #[tokio::test]
296 async fn test_search_cosine_returns_closest_first() {
297 let store = InMemoryVectorStore::new();
298 store
299 .create_collection("docs", 3, Some(VectorMetric::Cosine))
300 .await
301 .unwrap();
302 let records = vec![
303 VectorRecord::new("a", vec![1.0, 0.0, 0.0]),
304 VectorRecord::new("b", vec![0.0, 1.0, 0.0]),
305 VectorRecord::new("c", vec![1.0, 1.0, 0.0]),
306 ];
307 store.insert("docs", records).await.unwrap();
308
309 let results = store.search("docs", &[1.0, 0.0, 0.0], 2).await.unwrap();
310 assert_eq!(results.len(), 2);
311 assert_eq!(results[0].id, "a");
312 assert!(results[0].score > results[1].score);
313 }
314
315 #[tokio::test]
316 async fn test_search_top_k_limit() {
317 let store = InMemoryVectorStore::new();
318 store.create_collection("docs", 2, None).await.unwrap();
319 for i in 0..5 {
320 store
321 .insert(
322 "docs",
323 vec![VectorRecord::new(format!("r{}", i), vec![i as f32, 1.0])],
324 )
325 .await
326 .unwrap();
327 }
328 let results = store.search("docs", &[0.0, 1.0], 3).await.unwrap();
329 assert_eq!(results.len(), 3);
330 }
331
332 #[tokio::test]
333 async fn test_delete_records() {
334 let store = InMemoryVectorStore::new();
335 store.create_collection("docs", 2, None).await.unwrap();
336 store
337 .insert(
338 "docs",
339 vec![
340 VectorRecord::new("a", vec![1.0, 0.0]),
341 VectorRecord::new("b", vec![0.0, 1.0]),
342 VectorRecord::new("c", vec![1.0, 1.0]),
343 ],
344 )
345 .await
346 .unwrap();
347 let removed = store
348 .delete("docs", vec!["a".to_string(), "c".to_string()])
349 .await
350 .unwrap();
351 assert_eq!(removed, 2);
352 assert_eq!(store.count("docs").await.unwrap(), 1);
353 }
354
355 #[tokio::test]
356 async fn test_search_euclidean() {
357 let store = InMemoryVectorStore::new();
358 store
359 .create_collection("docs", 2, Some(VectorMetric::Euclidean))
360 .await
361 .unwrap();
362 let records = vec![
363 VectorRecord::new("near", vec![0.0, 0.0]),
364 VectorRecord::new("far", vec![10.0, 10.0]),
365 ];
366 store.insert("docs", records).await.unwrap();
367
368 let results = store.search("docs", &[0.0, 0.0], 2).await.unwrap();
369 assert_eq!(results.len(), 2);
370 assert_eq!(results[0].id, "near");
371 assert!(results[0].score > results[1].score);
372 }
373
374 #[tokio::test]
375 async fn test_search_dot_product() {
376 let store = InMemoryVectorStore::new();
377 store
378 .create_collection("docs", 2, Some(VectorMetric::DotProduct))
379 .await
380 .unwrap();
381 let records = vec![
382 VectorRecord::new("high", vec![2.0, 3.0]),
383 VectorRecord::new("low", vec![0.0, 0.0]),
384 ];
385 store.insert("docs", records).await.unwrap();
386
387 let results = store.search("docs", &[1.0, 1.0], 2).await.unwrap();
388 assert_eq!(results.len(), 2);
389 assert_eq!(results[0].id, "high");
390 }
391
392 #[tokio::test]
393 async fn test_collection_not_found() {
394 let store = InMemoryVectorStore::new();
395 let result = store.count("nonexistent").await;
396 assert_eq!(result.unwrap(), 0);
397
398 let err = store.search("nonexistent", &[1.0, 0.0], 5).await;
399 assert!(matches!(err, Err(VectorError::CollectionNotFound(_))));
400 }
401
402 #[tokio::test]
403 async fn test_get_nonexistent_record() {
404 let store = InMemoryVectorStore::new();
405 store.create_collection("docs", 2, None).await.unwrap();
406 let result = store.get("docs", "nonexistent").await.unwrap();
407 assert!(result.is_none());
408 }
409
410 #[tokio::test]
411 async fn test_helpers_compile() {
412 let store = InMemoryVectorStore::new();
413 let count = store.count("nonexistent").await.unwrap();
415 assert_eq!(count, 0, "fresh store should have 0 records for unknown collection");
416 assert_eq!(cosine_similarity(&[1.0, 0.0], &[1.0, 0.0]), 1.0);
418 assert!((cosine_similarity(&[1.0, 0.0], &[0.0, 1.0])).abs() < 1e-6);
419 assert_eq!(cosine_similarity(&[], &[]), 0.0);
420 }
421
422 #[tokio::test]
424 async fn test_m16_top_k_zero_rejected() {
425 let store = InMemoryVectorStore::new();
426 store.create_collection("docs", 2, None).await.unwrap();
427 store
428 .insert("docs", vec![VectorRecord::new("a", vec![1.0, 0.0])])
429 .await
430 .unwrap();
431 let err = store.search("docs", &[1.0, 0.0], 0).await;
432 assert!(matches!(
433 err,
434 Err(VectorError::TopKExceeded {
435 requested: 0,
436 max: crate::MAX_TOP_K
437 })
438 ));
439 }
440
441 #[tokio::test]
443 async fn test_m16_top_k_exceeded_rejected() {
444 let store = InMemoryVectorStore::new();
445 store.create_collection("docs", 2, None).await.unwrap();
446 store
447 .insert("docs", vec![VectorRecord::new("a", vec![1.0, 0.0])])
448 .await
449 .unwrap();
450 let err = store
451 .search("docs", &[1.0, 0.0], crate::MAX_TOP_K + 1)
452 .await;
453 assert!(matches!(
454 err,
455 Err(VectorError::TopKExceeded { requested, max }) if requested == crate::MAX_TOP_K + 1 && max == crate::MAX_TOP_K
456 ));
457 }
458
459 #[tokio::test]
461 async fn test_m16_top_k_max_allowed() {
462 let store = InMemoryVectorStore::new();
463 store.create_collection("docs", 2, None).await.unwrap();
464 store
465 .insert("docs", vec![VectorRecord::new("a", vec![1.0, 0.0])])
466 .await
467 .unwrap();
468 let results = store
470 .search("docs", &[1.0, 0.0], crate::MAX_TOP_K)
471 .await
472 .unwrap();
473 assert_eq!(results.len(), 1);
474 }
475
476 #[test]
478 fn test_m16_validate_top_k_function() {
479 use crate::validate_top_k;
480 assert_eq!(validate_top_k(1).unwrap(), 1);
482 assert_eq!(validate_top_k(100).unwrap(), 100);
483 assert_eq!(validate_top_k(crate::MAX_TOP_K).unwrap(), crate::MAX_TOP_K);
484 assert!(validate_top_k(0).is_err());
486 assert!(validate_top_k(crate::MAX_TOP_K + 1).is_err());
487 assert!(validate_top_k(usize::MAX).is_err());
488 }
489}