Skip to main content

lc_vector_stores/
memory.rs

1// lc-vector-stores/src/memory.rs
2//! 内存向量存储
3//!
4//! 将文档和向量存储在内存中,适用于小规模数据和测试。
5
6use crate::{
7    cosine_similarity, Document, SearchResult, VectorDocument, VectorStore, VectorStoreError,
8};
9use async_trait::async_trait;
10use std::collections::HashMap;
11use std::sync::Arc;
12use tokio::sync::RwLock;
13use uuid::Uuid;
14
15/// 内存向量存储
16pub struct InMemoryVectorStore {
17    /// 文档存储
18    documents: Arc<RwLock<HashMap<String, VectorDocument>>>,
19}
20
21impl InMemoryVectorStore {
22    /// 创建新的内存向量存储
23    pub fn new() -> Self {
24        Self {
25            documents: Arc::new(RwLock::new(HashMap::new())),
26        }
27    }
28}
29
30impl Default for InMemoryVectorStore {
31    fn default() -> Self {
32        Self::new()
33    }
34}
35
36#[async_trait]
37impl VectorStore for InMemoryVectorStore {
38    async fn add_documents(
39        &self,
40        documents: Vec<Document>,
41        embeddings: Vec<Vec<f32>>,
42    ) -> Result<Vec<String>, VectorStoreError> {
43        if documents.len() != embeddings.len() {
44            return Err(VectorStoreError::StorageError(
45                "文档数量和嵌入向量数量不匹配".to_string(),
46            ));
47        }
48
49        let mut store = self.documents.write().await;
50        let mut ids = Vec::new();
51
52        for (doc, embedding) in documents.into_iter().zip(embeddings.into_iter()) {
53            let id = doc.id.clone().unwrap_or_else(|| Uuid::new_v4().to_string());
54
55            let vector_doc = VectorDocument {
56                document: Document {
57                    id: Some(id.clone()),
58                    content: doc.content,
59                    metadata: doc.metadata,
60                },
61                embedding,
62            };
63
64            store.insert(id.clone(), vector_doc);
65            ids.push(id);
66        }
67
68        Ok(ids)
69    }
70
71    async fn similarity_search(
72        &self,
73        query_embedding: &[f32],
74        k: usize,
75    ) -> Result<Vec<SearchResult>, VectorStoreError> {
76        // Q2: 不再硬过滤 score > 0 —— 全负分语料下也应返回 top-k;
77        // 是否设阈值由调用方通过 similarity_search_with_min_score 显式决定。
78        self.similarity_search_with_min_score(query_embedding, k, None)
79            .await
80    }
81
82    async fn similarity_search_with_min_score(
83        &self,
84        query_embedding: &[f32],
85        k: usize,
86        min_score: Option<f32>,
87    ) -> Result<Vec<SearchResult>, VectorStoreError> {
88        let store = self.documents.read().await;
89
90        // 计算所有文档的相似度,先按阈值过滤再取 top-k (Q2)
91        let mut results: Vec<SearchResult> = store
92            .values()
93            .filter_map(|vd| {
94                let score = cosine_similarity(query_embedding, &vd.embedding).unwrap_or(0.0);
95                if min_score.is_none_or(|t| score >= t) {
96                    Some(SearchResult {
97                        document: vd.document.clone(),
98                        score,
99                    })
100                } else {
101                    None
102                }
103            })
104            .collect();
105
106        // 按相似度降序排序
107        results.sort_by(|a, b| {
108            b.score
109                .partial_cmp(&a.score)
110                .unwrap_or(std::cmp::Ordering::Equal)
111        });
112
113        // 返回前 k 个结果
114        Ok(results.into_iter().take(k).collect())
115    }
116
117    async fn get_document(&self, id: &str) -> Result<Option<Document>, VectorStoreError> {
118        let store = self.documents.read().await;
119        Ok(store.get(id).map(|vd| vd.document.clone()))
120    }
121
122    async fn get_embedding(&self, id: &str) -> Result<Option<Vec<f32>>, VectorStoreError> {
123        let store = self.documents.read().await;
124        Ok(store.get(id).map(|vd| vd.embedding.clone()))
125    }
126
127    async fn delete_document(&self, id: &str) -> Result<(), VectorStoreError> {
128        let mut store = self.documents.write().await;
129        store.remove(id);
130        Ok(())
131    }
132
133    async fn count(&self) -> usize {
134        let store = self.documents.read().await;
135        store.len()
136    }
137
138    async fn clear(&self) -> Result<(), VectorStoreError> {
139        let mut store = self.documents.write().await;
140        store.clear();
141        Ok(())
142    }
143}
144
145#[cfg(test)]
146mod tests {
147    use super::*;
148
149    #[tokio::test]
150    async fn test_add_and_search() {
151        let store = InMemoryVectorStore::new();
152
153        // 添加文档
154        let docs = vec![
155            Document::new("Rust is a systems programming language"),
156            Document::new("Python is a scripting language"),
157            Document::new("JavaScript is used for web development"),
158        ];
159
160        // 创建简单的模拟嵌入向量
161        let embeddings = vec![
162            vec![1.0, 0.0, 0.0], // Rust 相关
163            vec![0.0, 1.0, 0.0], // Python 相关
164            vec![0.0, 0.0, 1.0], // JavaScript 相关
165        ];
166
167        let ids = store.add_documents(docs, embeddings).await.unwrap();
168        assert_eq!(ids.len(), 3);
169        assert_eq!(store.count().await, 3);
170
171        // 搜索相似文档
172        let query = vec![0.9, 0.1, 0.0]; // 更接近 Rust
173        let results = store.similarity_search(&query, 2).await.unwrap();
174
175        assert_eq!(results.len(), 2);
176        assert!(results[0].document.content.contains("Rust"));
177        assert!(results[0].score > results[1].score);
178    }
179
180    #[tokio::test]
181    async fn test_get_and_delete() {
182        let store = InMemoryVectorStore::new();
183
184        let doc = Document::new("Test document").with_id("test-id");
185        let embeddings = vec![vec![1.0, 0.0, 0.0]];
186
187        store.add_documents(vec![doc], embeddings).await.unwrap();
188
189        // 获取文档
190        let retrieved = store.get_document("test-id").await.unwrap();
191        assert!(retrieved.is_some());
192        assert_eq!(retrieved.unwrap().content, "Test document");
193
194        // 删除文档
195        store.delete_document("test-id").await.unwrap();
196        assert_eq!(store.count().await, 0);
197
198        // 再次获取应该返回 None
199        let retrieved = store.get_document("test-id").await.unwrap();
200        assert!(retrieved.is_none());
201    }
202
203    #[tokio::test]
204    async fn test_clear() {
205        let store = InMemoryVectorStore::new();
206
207        let docs = vec![Document::new("Doc 1"), Document::new("Doc 2")];
208        let embeddings = vec![vec![1.0, 0.0], vec![0.0, 1.0]];
209
210        store.add_documents(docs, embeddings).await.unwrap();
211        assert_eq!(store.count().await, 2);
212
213        store.clear().await.unwrap();
214        assert_eq!(store.count().await, 0);
215    }
216
217    /// Q1: 未配置嵌入器时,similarity_search_text 应显式报 EmbeddingError,
218    /// 而不是静默成功或 panic。
219    #[tokio::test]
220    async fn test_similarity_search_text_without_embedder_errors() {
221        let store = InMemoryVectorStore::new();
222        let err = store.similarity_search_text("hello", 3).await.unwrap_err();
223        assert!(matches!(err, VectorStoreError::EmbeddingError(_)));
224    }
225
226    /// Q2: 全非正分语料下 similarity_search 仍返回 top-k(不再被 score>0 硬过滤清空);
227    /// similarity_search_with_min_score 按阈值显式过滤。
228    #[tokio::test]
229    async fn test_negative_scores_not_dropped() {
230        let store = InMemoryVectorStore::new();
231        store
232            .add_documents(
233                vec![
234                    Document::new("orthogonal-up"),
235                    Document::new("opposite"),
236                    Document::new("orthogonal-down"),
237                ],
238                vec![vec![0.0, 1.0], vec![-1.0, 0.0], vec![0.0, -1.0]],
239            )
240            .await
241            .unwrap();
242
243        let query = vec![1.0, 0.0];
244
245        // 旧实现 score > 0.0 硬过滤,该语料下会返回空;现在返回 top-k(3 条,全部非正分)。
246        let results = store.similarity_search(&query, 3).await.unwrap();
247        assert_eq!(results.len(), 3);
248        assert!(results.iter().all(|r| r.score <= 0.0));
249
250        // 显式阈值:score >= -0.5 → 排除 score = -1.0 的那条
251        let filtered = store
252            .similarity_search_with_min_score(&query, 3, Some(-0.5))
253            .await
254            .unwrap();
255        assert_eq!(filtered.len(), 2);
256
257        // min_score = None 时与 similarity_search 行为一致
258        let all = store
259            .similarity_search_with_min_score(&query, 3, None)
260            .await
261            .unwrap();
262        assert_eq!(all.len(), 3);
263    }
264
265    #[test]
266    fn test_cosine_similarity() {
267        // Identical vectors
268        let a = vec![1.0, 0.0, 0.0];
269        let b = vec![1.0, 0.0, 0.0];
270        assert!((cosine_similarity(&a, &b).unwrap() - 1.0).abs() < 0.0001);
271
272        // Orthogonal vectors
273        let a = vec![1.0, 0.0, 0.0];
274        let b = vec![0.0, 1.0, 0.0];
275        assert!((cosine_similarity(&a, &b).unwrap() - 0.0).abs() < 0.0001);
276    }
277}