Skip to main content

sz_orm_vector/
memory.rs

1//! 内存实现:纯 Rust 向量计算(不连接数据库)
2//!
3//! 适用于:
4//! - 单元测试
5//! - 不需要真实 pgvector 的场景(如原型开发)
6//! - 性能基准(无 I/O 开销)
7
8use 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
15/// 内存 Vector Store 实现
16pub 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                // 将距离转为 [0, 1] 相似度
53                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
66/// 余弦相似度计算
67fn 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            // Upsert
133            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        // M-16 修复:校验 top_k 范围
155        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            // 将存储的 metadata 填充到搜索结果中(用于后续过滤)
183            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]); // dim=2
272        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        // Upsert should keep count at 1
290        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        // Fresh store: count for unknown collection must be 0
414        let count = store.count("nonexistent").await.unwrap();
415        assert_eq!(count, 0, "fresh store should have 0 records for unknown collection");
416        // Cosine similarity contract
417        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    /// M-16 测试:top_k = 0 应被拒绝
423    #[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    /// M-16 测试:top_k 超过 MAX_TOP_K 应被拒绝
442    #[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    /// M-16 测试:top_k = MAX_TOP_K 应允许
460    #[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        // top_k = MAX_TOP_K 不应触发错误(即使记录数远少于 MAX_TOP_K)
469        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    /// M-16 测试:validate_top_k 函数单元测试
477    #[test]
478    fn test_m16_validate_top_k_function() {
479        use crate::validate_top_k;
480        // 有效值
481        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        // 无效值
485        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}