Skip to main content

sz_orm_ai/vector/
mod.rs

1use crate::embedding::EmbeddingRecord;
2use crate::error::AiError;
3use async_trait::async_trait;
4use parking_lot::RwLock;
5use std::collections::HashMap;
6
7// HNSW 向量索引子模块(近似最近邻搜索)
8pub 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
225/// In-memory VectorStore backed by `HashMap` + `Vec`.
226///
227/// Stores records per collection and supports cosine/euclidean/dot-product
228/// similarity search. Suitable for unit tests and small in-process workloads.
229pub 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                // Convert distance to similarity score in [0, 1].
260                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            // Replace existing record if id is the same (upsert semantics).
333            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        // Sort by score descending (stable for ties).
369        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
422/// Very small filter expression parser: `{"field": {"eq": value}}`.
423/// Returns true if metadata matches; false (or true if no filter) otherwise.
424fn 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        // After deletion, count is 0 (collection does not exist).
485        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]); // dim=2
508        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        // Upsert should keep count at 1.
525        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        // Cosine similarity of [1,0,0] and [1,1,0] is 1/sqrt(2) ~= 0.707
551        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        // Smoke-test the helper: verify store constructs and is empty for unknown collection.
622        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        // Cosine similarity contract: identical vectors -> 1.0, orthogonal -> 0.0
629        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}