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            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]); // dim=2
268        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        // Upsert should keep count at 1
286        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    /// M-16 测试:top_k = 0 应被拒绝
415    #[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    /// M-16 测试:top_k 超过 MAX_TOP_K 应被拒绝
434    #[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    /// M-16 测试:top_k = MAX_TOP_K 应允许
452    #[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        // top_k = MAX_TOP_K 不应触发错误(即使记录数远少于 MAX_TOP_K)
461        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    /// M-16 测试:validate_top_k 函数单元测试
469    #[test]
470    fn test_m16_validate_top_k_function() {
471        use crate::validate_top_k;
472        // 有效值
473        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        // 无效值
477        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}