1use crate::embedding::EmbeddingRecord;
2use crate::error::AiError;
3use async_trait::async_trait;
4use parking_lot::RwLock;
5use std::collections::HashMap;
6
7pub mod hnsw;
9
10#[derive(Debug, Clone)]
11pub struct VectorError {
12 pub message: String,
13 pub collection: Option<String>,
14}
15
16impl VectorError {
17 pub fn new(message: impl Into<String>) -> Self {
18 Self {
19 message: message.into(),
20 collection: None,
21 }
22 }
23
24 pub fn with_collection(message: impl Into<String>, collection: impl Into<String>) -> Self {
25 Self {
26 message: message.into(),
27 collection: Some(collection.into()),
28 }
29 }
30}
31
32impl std::fmt::Display for VectorError {
33 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
34 write!(f, "VectorError: {}", self.message)?;
35 if let Some(ref coll) = self.collection {
36 write!(f, " (collection: {})", coll)?;
37 }
38 Ok(())
39 }
40}
41
42impl std::error::Error for VectorError {}
43
44#[derive(Debug, Clone)]
45pub struct VectorRecord {
46 pub id: String,
47 pub vector: Vec<f32>,
48 pub score: Option<f32>,
49 pub metadata: Option<std::collections::HashMap<String, serde_json::Value>>,
50}
51
52impl VectorRecord {
53 pub fn new(id: impl Into<String>, vector: Vec<f32>) -> Self {
54 Self {
55 id: id.into(),
56 vector,
57 score: None,
58 metadata: None,
59 }
60 }
61
62 pub fn with_score(mut self, score: f32) -> Self {
63 self.score = Some(score);
64 self
65 }
66
67 pub fn with_metadata(
68 mut self,
69 metadata: std::collections::HashMap<String, serde_json::Value>,
70 ) -> Self {
71 self.metadata = Some(metadata);
72 self
73 }
74
75 pub fn from_embedding(record: &EmbeddingRecord) -> Self {
76 Self {
77 id: record.id.clone(),
78 vector: record.vector.clone(),
79 score: None,
80 metadata: record.metadata.clone(),
81 }
82 }
83}
84
85#[derive(Debug, Clone)]
86pub struct SearchResult {
87 pub id: String,
88 pub score: f32,
89 pub vector: Vec<f32>,
90 pub text: Option<String>,
91 pub metadata: Option<std::collections::HashMap<String, serde_json::Value>>,
92}
93
94impl SearchResult {
95 pub fn new(id: impl Into<String>, score: f32, vector: Vec<f32>) -> Self {
96 Self {
97 id: id.into(),
98 score,
99 vector,
100 text: None,
101 metadata: None,
102 }
103 }
104
105 pub fn with_text(mut self, text: impl Into<String>) -> Self {
106 self.text = Some(text.into());
107 self
108 }
109}
110
111#[derive(Debug, Clone, Default)]
112pub struct VectorFilter {
113 pub field: Option<String>,
114 pub operator: Option<String>,
115 pub value: Option<serde_json::Value>,
116}
117
118impl VectorFilter {
119 pub fn new() -> Self {
120 Self::default()
121 }
122
123 pub fn field(mut self, field: impl Into<String>) -> Self {
124 self.field = Some(field.into());
125 self
126 }
127
128 pub fn eq(mut self, value: impl Into<serde_json::Value>) -> Self {
129 self.operator = Some("eq".to_string());
130 self.value = Some(value.into());
131 self
132 }
133
134 pub fn gt(mut self, value: impl Into<serde_json::Value>) -> Self {
135 self.operator = Some("gt".to_string());
136 self.value = Some(value.into());
137 self
138 }
139
140 pub fn lt(mut self, value: impl Into<serde_json::Value>) -> Self {
141 self.operator = Some("lt".to_string());
142 self.value = Some(value.into());
143 self
144 }
145
146 pub fn build(&self) -> Option<String> {
147 match (&self.field, &self.operator, &self.value) {
148 (Some(field), Some(op), Some(value)) => {
149 Some(format!(r#"{{"{}": {{"{}": {}}}}}"#, field, op, value))
150 }
151 _ => None,
152 }
153 }
154}
155
156#[async_trait]
157pub trait VectorStore: Send + Sync {
158 async fn create_collection(
159 &self,
160 name: &str,
161 dimension: usize,
162 metric: Option<VectorMetric>,
163 ) -> Result<(), AiError>;
164
165 async fn delete_collection(&self, name: &str) -> Result<(), AiError>;
166
167 async fn insert(&self, collection: &str, records: Vec<VectorRecord>) -> Result<(), AiError>;
168
169 async fn search(
170 &self,
171 collection: &str,
172 query: &[f32],
173 top_k: usize,
174 filter: Option<&str>,
175 ) -> Result<Vec<SearchResult>, AiError>;
176
177 async fn get(&self, collection: &str, id: &str) -> Result<Option<VectorRecord>, AiError>;
178
179 async fn delete(&self, collection: &str, ids: Vec<String>) -> Result<u64, AiError>;
180
181 async fn count(&self, collection: &str) -> Result<usize, AiError>;
182}
183
184#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
185pub enum VectorMetric {
186 #[default]
187 Cosine,
188 Euclidean,
189 DotProduct,
190}
191
192impl VectorMetric {
193 pub fn as_str(&self) -> &str {
194 match self {
195 VectorMetric::Cosine => "cosine",
196 VectorMetric::Euclidean => "euclidean",
197 VectorMetric::DotProduct => "dotproduct",
198 }
199 }
200}
201
202pub struct CollectionMeta {
203 pub name: String,
204 pub dimension: usize,
205 pub metric: VectorMetric,
206 pub count: usize,
207}
208
209impl CollectionMeta {
210 pub fn new(name: impl Into<String>, dimension: usize) -> Self {
211 Self {
212 name: name.into(),
213 dimension,
214 metric: VectorMetric::default(),
215 count: 0,
216 }
217 }
218
219 pub fn with_metric(mut self, metric: VectorMetric) -> Self {
220 self.metric = metric;
221 self
222 }
223}
224
225pub struct InMemoryVectorStore {
230 collections: RwLock<HashMap<String, CollectionState>>,
231}
232
233#[derive(Debug, Clone)]
234struct CollectionState {
235 dimension: usize,
236 metric: VectorMetric,
237 records: Vec<StoredRecord>,
238}
239
240#[derive(Debug, Clone)]
241struct StoredRecord {
242 id: String,
243 vector: Vec<f32>,
244 metadata: Option<HashMap<String, serde_json::Value>>,
245 text: Option<String>,
246}
247
248impl InMemoryVectorStore {
249 pub fn new() -> Self {
250 Self {
251 collections: RwLock::new(HashMap::new()),
252 }
253 }
254
255 fn metric_value(metric: VectorMetric, a: &[f32], b: &[f32]) -> f32 {
256 match metric {
257 VectorMetric::Cosine => cosine_similarity(a, b),
258 VectorMetric::Euclidean => {
259 let dist: f32 = a
261 .iter()
262 .zip(b.iter())
263 .map(|(x, y)| (x - y) * (x - y))
264 .sum::<f32>()
265 .sqrt();
266 1.0 / (1.0 + dist)
267 }
268 VectorMetric::DotProduct => a.iter().zip(b.iter()).map(|(x, y)| x * y).sum(),
269 }
270 }
271}
272
273impl Default for InMemoryVectorStore {
274 fn default() -> Self {
275 Self::new()
276 }
277}
278
279fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
280 if a.len() != b.len() || a.is_empty() {
281 return 0.0;
282 }
283 let dot: f32 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
284 let na: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
285 let nb: f32 = b.iter().map(|x| x * x).sum::<f32>().sqrt();
286 if na == 0.0 || nb == 0.0 {
287 return 0.0;
288 }
289 dot / (na * nb)
290}
291
292#[async_trait]
293impl VectorStore for InMemoryVectorStore {
294 async fn create_collection(
295 &self,
296 name: &str,
297 dimension: usize,
298 metric: Option<VectorMetric>,
299 ) -> Result<(), AiError> {
300 let mut collections = self.collections.write();
301 collections.insert(
302 name.to_string(),
303 CollectionState {
304 dimension,
305 metric: metric.unwrap_or_default(),
306 records: Vec::new(),
307 },
308 );
309 Ok(())
310 }
311
312 async fn delete_collection(&self, name: &str) -> Result<(), AiError> {
313 let mut collections = self.collections.write();
314 collections.remove(name);
315 Ok(())
316 }
317
318 async fn insert(&self, collection: &str, records: Vec<VectorRecord>) -> Result<(), AiError> {
319 let mut collections = self.collections.write();
320 let state = collections
321 .get_mut(collection)
322 .ok_or_else(|| AiError::Vector(format!("collection not found: {}", collection)))?;
323
324 for record in records {
325 if record.vector.len() != state.dimension {
326 return Err(AiError::Vector(format!(
327 "dimension mismatch: expected {}, got {}",
328 state.dimension,
329 record.vector.len()
330 )));
331 }
332 if let Some(existing) = state.records.iter_mut().find(|r| r.id == record.id) {
334 existing.vector = record.vector;
335 existing.metadata = record.metadata;
336 continue;
337 }
338 state.records.push(StoredRecord {
339 id: record.id,
340 vector: record.vector,
341 metadata: record.metadata,
342 text: None,
343 });
344 }
345 Ok(())
346 }
347
348 async fn search(
349 &self,
350 collection: &str,
351 query: &[f32],
352 top_k: usize,
353 filter: Option<&str>,
354 ) -> Result<Vec<SearchResult>, AiError> {
355 let collections = self.collections.read();
356 let state = collections
357 .get(collection)
358 .ok_or_else(|| AiError::Vector(format!("collection not found: {}", collection)))?;
359
360 let mut scored: Vec<(usize, f32)> = state
361 .records
362 .iter()
363 .enumerate()
364 .filter(|(_, r)| match_filter(r.metadata.as_ref(), filter))
365 .map(|(i, r)| (i, Self::metric_value(state.metric, query, &r.vector)))
366 .collect();
367
368 scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
370
371 let k = top_k.min(scored.len());
372 let mut results = Vec::with_capacity(k);
373 for (idx, score) in scored.into_iter().take(k) {
374 let record = &state.records[idx];
375 let mut search_result =
376 SearchResult::new(record.id.clone(), score, record.vector.clone());
377 if let Some(ref text) = record.text {
378 search_result = search_result.with_text(text.clone());
379 }
380 results.push(search_result);
381 }
382 Ok(results)
383 }
384
385 async fn get(&self, collection: &str, id: &str) -> Result<Option<VectorRecord>, AiError> {
386 let collections = self.collections.read();
387 let state = collections
388 .get(collection)
389 .ok_or_else(|| AiError::Vector(format!("collection not found: {}", collection)))?;
390 Ok(state
391 .records
392 .iter()
393 .find(|r| r.id == id)
394 .map(|r| VectorRecord {
395 id: r.id.clone(),
396 vector: r.vector.clone(),
397 score: None,
398 metadata: r.metadata.clone(),
399 }))
400 }
401
402 async fn delete(&self, collection: &str, ids: Vec<String>) -> Result<u64, AiError> {
403 let mut collections = self.collections.write();
404 let state = collections
405 .get_mut(collection)
406 .ok_or_else(|| AiError::Vector(format!("collection not found: {}", collection)))?;
407 let before = state.records.len();
408 state.records.retain(|r| !ids.contains(&r.id));
409 let removed = (before - state.records.len()) as u64;
410 Ok(removed)
411 }
412
413 async fn count(&self, collection: &str) -> Result<usize, AiError> {
414 let collections = self.collections.read();
415 Ok(collections
416 .get(collection)
417 .map(|s| s.records.len())
418 .unwrap_or(0))
419 }
420}
421
422fn match_filter(
425 metadata: Option<&HashMap<String, serde_json::Value>>,
426 filter: Option<&str>,
427) -> bool {
428 let Some(expr) = filter else { return true };
429 let Some(metadata) = metadata else {
430 return false;
431 };
432 let Ok(parsed) = serde_json::from_str::<serde_json::Value>(expr) else {
433 return false;
434 };
435 let Some(obj) = parsed.as_object() else {
436 return false;
437 };
438 for (field, cond) in obj {
439 let Some(actual) = metadata.get(field) else {
440 return false;
441 };
442 let Some(cond_obj) = cond.as_object() else {
443 return false;
444 };
445 for (op, val) in cond_obj {
446 match op.as_str() {
447 "eq" if actual == val => continue,
448 "gt" => {
449 let greater = match (actual.as_f64(), val.as_f64()) {
450 (Some(a), Some(b)) => a > b,
451 _ => false,
452 };
453 if !greater {
454 return false;
455 }
456 }
457 "lt" => {
458 let less = match (actual.as_f64(), val.as_f64()) {
459 (Some(a), Some(b)) => a < b,
460 _ => false,
461 };
462 if !less {
463 return false;
464 }
465 }
466 _ => return false,
467 }
468 }
469 }
470 true
471}
472
473#[cfg(test)]
474mod tests {
475 use super::*;
476
477 #[tokio::test]
478 async fn test_create_and_delete_collection() {
479 let store = InMemoryVectorStore::new();
480 store.create_collection("docs", 4, None).await.unwrap();
481 assert_eq!(store.count("docs").await.unwrap(), 0);
482
483 store.delete_collection("docs").await.unwrap();
484 assert_eq!(store.count("docs").await.unwrap(), 0);
486 }
487
488 #[tokio::test]
489 async fn test_insert_and_get() {
490 let store = InMemoryVectorStore::new();
491 store.create_collection("docs", 3, None).await.unwrap();
492 let rec = VectorRecord::new("r1", vec![1.0, 0.0, 0.0]);
493 store.insert("docs", vec![rec]).await.unwrap();
494 assert_eq!(store.count("docs").await.unwrap(), 1);
495
496 let fetched = store.get("docs", "r1").await.unwrap().unwrap();
497 assert_eq!(fetched.id, "r1");
498 assert_eq!(fetched.vector, vec![1.0, 0.0, 0.0]);
499
500 assert!(store.get("docs", "missing").await.unwrap().is_none());
501 }
502
503 #[tokio::test]
504 async fn test_insert_dimension_mismatch() {
505 let store = InMemoryVectorStore::new();
506 store.create_collection("docs", 3, None).await.unwrap();
507 let rec = VectorRecord::new("r1", vec![1.0, 0.0]); let err = store.insert("docs", vec![rec]).await;
509 assert!(err.is_err());
510 }
511
512 #[tokio::test]
513 async fn test_insert_upsert() {
514 let store = InMemoryVectorStore::new();
515 store.create_collection("docs", 2, None).await.unwrap();
516 store
517 .insert("docs", vec![VectorRecord::new("r1", vec![1.0, 0.0])])
518 .await
519 .unwrap();
520 store
521 .insert("docs", vec![VectorRecord::new("r1", vec![0.0, 1.0])])
522 .await
523 .unwrap();
524 assert_eq!(store.count("docs").await.unwrap(), 1);
526 let fetched = store.get("docs", "r1").await.unwrap().unwrap();
527 assert_eq!(fetched.vector, vec![0.0, 1.0]);
528 }
529
530 #[tokio::test]
531 async fn test_search_cosine_returns_closest_first() {
532 let store = InMemoryVectorStore::new();
533 store
534 .create_collection("docs", 3, Some(VectorMetric::Cosine))
535 .await
536 .unwrap();
537 let records = vec![
538 VectorRecord::new("a", vec![1.0, 0.0, 0.0]),
539 VectorRecord::new("b", vec![0.0, 1.0, 0.0]),
540 VectorRecord::new("c", vec![1.0, 1.0, 0.0]),
541 ];
542 store.insert("docs", records).await.unwrap();
543
544 let results = store
545 .search("docs", &[1.0, 0.0, 0.0], 2, None)
546 .await
547 .unwrap();
548 assert_eq!(results.len(), 2);
549 assert_eq!(results[0].id, "a");
550 assert!(results[0].score > results[1].score);
552 }
553
554 #[tokio::test]
555 async fn test_search_top_k_limit() {
556 let store = InMemoryVectorStore::new();
557 store.create_collection("docs", 2, None).await.unwrap();
558 for i in 0..5 {
559 store
560 .insert(
561 "docs",
562 vec![VectorRecord::new(format!("r{}", i), vec![i as f32, 1.0])],
563 )
564 .await
565 .unwrap();
566 }
567 let results = store.search("docs", &[0.0, 1.0], 3, None).await.unwrap();
568 assert_eq!(results.len(), 3);
569 }
570
571 #[tokio::test]
572 async fn test_delete_records() {
573 let store = InMemoryVectorStore::new();
574 store.create_collection("docs", 2, None).await.unwrap();
575 store
576 .insert(
577 "docs",
578 vec![
579 VectorRecord::new("a", vec![1.0, 0.0]),
580 VectorRecord::new("b", vec![0.0, 1.0]),
581 VectorRecord::new("c", vec![1.0, 1.0]),
582 ],
583 )
584 .await
585 .unwrap();
586 let removed = store
587 .delete("docs", vec!["a".to_string(), "c".to_string()])
588 .await
589 .unwrap();
590 assert_eq!(removed, 2);
591 assert_eq!(store.count("docs").await.unwrap(), 1);
592 }
593
594 #[tokio::test]
595 async fn test_search_with_filter() {
596 let store = InMemoryVectorStore::new();
597 store.create_collection("docs", 2, None).await.unwrap();
598 let mut md = HashMap::new();
599 md.insert("kind".to_string(), serde_json::json!("alpha"));
600 let r1 = VectorRecord::new("a", vec![1.0, 0.0]).with_metadata(md);
601 let mut md2 = HashMap::new();
602 md2.insert("kind".to_string(), serde_json::json!("beta"));
603 let r2 = VectorRecord::new("b", vec![1.0, 0.0]).with_metadata(md2);
604 store.insert("docs", vec![r1, r2]).await.unwrap();
605
606 let results = store
607 .search(
608 "docs",
609 &[1.0, 0.0],
610 10,
611 Some(r#"{"kind": {"eq": "alpha"}}"#),
612 )
613 .await
614 .unwrap();
615 assert_eq!(results.len(), 1);
616 assert_eq!(results[0].id, "a");
617 }
618
619 #[tokio::test]
620 async fn test_helpers_compile() {
621 let store = InMemoryVectorStore::new();
623 let count = store.count("nonexistent").await.unwrap();
624 assert_eq!(
625 count, 0,
626 "fresh store should have 0 records for unknown collection"
627 );
628 assert_eq!(cosine_similarity(&[1.0, 0.0], &[1.0, 0.0]), 1.0);
630 assert!((cosine_similarity(&[1.0, 0.0], &[0.0, 1.0])).abs() < 1e-6);
631 }
632}