Skip to main content

lc_vector_stores/
file_store.rs

1// lc-vector-stores/src/file_store.rs
2//! 文件持久化向量存储
3//!
4//! 将文档和向量持久化到本地文件(JSON 序列化),适用于个人知识库和离线场景。
5//! 填补 InMemory(不持久)与外部数据库(太重)之间的空缺,类似 SQLite 之于 MySQL。
6//!
7//! # 使用方式
8//! ```ignore
9//! use lc_vector_stores::FileVectorStore;
10//! use std::path::PathBuf;
11//!
12//! let store = FileVectorStore::new(PathBuf::from("./my_vectors.json"), 1536).unwrap();
13//! // add_documents / similarity_search 与 InMemoryVectorStore 接口一致
14//! // 每次增删后自动持久化到磁盘
15//! ```
16
17use crate::{
18    cosine_similarity, Document, SearchResult, VectorDocument, VectorStore, VectorStoreError,
19};
20use async_trait::async_trait;
21use std::collections::HashMap;
22use std::path::PathBuf;
23use tokio::sync::RwLock;
24use uuid::Uuid;
25
26/// 持久化数据结构(序列化为 JSON)
27#[derive(serde::Serialize, serde::Deserialize)]
28struct FileStoreData {
29    /// 向量维度(用于校验)
30    dimension: usize,
31    /// 文档 + 向量
32    documents: HashMap<String, VectorDocument>,
33}
34
35/// 文件持久化向量存储
36///
37/// 向量 + 元数据序列化为 JSON 写入磁盘,启动时加载,增删后自动写回。
38/// 使用 `RwLock` 保证并发安全:读操作用读锁,写操作用写锁 + 持久化。
39pub struct FileVectorStore {
40    /// 存储文件路径
41    path: PathBuf,
42    /// 向量维度
43    dimension: usize,
44    /// 内存中的数据(与文件同步)
45    data: RwLock<FileStoreData>,
46}
47
48impl FileVectorStore {
49    /// 创建或加载文件向量存储
50    ///
51    /// 如果文件已存在,加载其中数据;否则创建空存储。
52    ///
53    /// # Arguments
54    /// * `path` - 存储文件路径(建议 `.json` 后缀)
55    /// * `dimension` - 向量维度(用于校验,已有文件时以文件为准)
56    pub async fn new(path: PathBuf, dimension: usize) -> Result<Self, VectorStoreError> {
57        let data = if path.exists() {
58            let content = tokio::fs::read_to_string(&path)
59                .await
60                .map_err(|e| VectorStoreError::StorageError(format!("读取文件失败: {}", e)))?;
61            serde_json::from_str::<FileStoreData>(&content)
62                .map_err(|e| VectorStoreError::StorageError(format!("解析文件失败: {}", e)))?
63        } else {
64            // 确保父目录存在
65            if let Some(parent) = path.parent() {
66                if !parent.as_os_str().is_empty() {
67                    tokio::fs::create_dir_all(parent).await.map_err(|e| {
68                        VectorStoreError::StorageError(format!("创建目录失败: {}", e))
69                    })?;
70                }
71            }
72            let data = FileStoreData {
73                dimension,
74                documents: HashMap::new(),
75            };
76            // 首次创建时也持久化空文件,确保 path.exists() 为 true
77            Self::persist(&data, &path).await?;
78            data
79        };
80
81        Ok(Self {
82            path,
83            dimension: data.dimension,
84            data: RwLock::new(data),
85        })
86    }
87
88    /// 持久化当前数据到磁盘
89    async fn persist(data: &FileStoreData, path: &PathBuf) -> Result<(), VectorStoreError> {
90        let json = serde_json::to_string(data)
91            .map_err(|e| VectorStoreError::StorageError(format!("序列化失败: {}", e)))?;
92        // 先写临时文件,再 rename,避免写一半断电损坏
93        let tmp_path = path.with_extension("json.tmp");
94        tokio::fs::write(&tmp_path, &json)
95            .await
96            .map_err(|e| VectorStoreError::StorageError(format!("写入临时文件失败: {}", e)))?;
97        tokio::fs::rename(&tmp_path, path)
98            .await
99            .map_err(|e| VectorStoreError::StorageError(format!("重命名文件失败: {}", e)))?;
100        Ok(())
101    }
102
103    /// 返回向量维度
104    pub fn dimension(&self) -> usize {
105        self.dimension
106    }
107
108    /// 返回存储文件路径
109    pub fn path(&self) -> &PathBuf {
110        &self.path
111    }
112}
113
114#[async_trait]
115impl VectorStore for FileVectorStore {
116    async fn add_documents(
117        &self,
118        documents: Vec<Document>,
119        embeddings: Vec<Vec<f32>>,
120    ) -> Result<Vec<String>, VectorStoreError> {
121        if documents.len() != embeddings.len() {
122            return Err(VectorStoreError::StorageError(
123                "文档数量和嵌入向量数量不匹配".to_string(),
124            ));
125        }
126
127        let mut data = self.data.write().await;
128        let mut ids = Vec::new();
129
130        for (doc, embedding) in documents.into_iter().zip(embeddings.into_iter()) {
131            // 校验维度
132            if !embedding.is_empty() && embedding.len() != data.dimension {
133                return Err(VectorStoreError::StorageError(format!(
134                    "嵌入维度 {} 与存储维度 {} 不匹配",
135                    embedding.len(),
136                    data.dimension
137                )));
138            }
139
140            let id = doc.id.clone().unwrap_or_else(|| Uuid::new_v4().to_string());
141            let vector_doc = VectorDocument {
142                document: Document {
143                    id: Some(id.clone()),
144                    content: doc.content,
145                    metadata: doc.metadata,
146                },
147                embedding,
148            };
149            data.documents.insert(id.clone(), vector_doc);
150            ids.push(id);
151        }
152
153        Self::persist(&data, &self.path).await?;
154        Ok(ids)
155    }
156
157    async fn similarity_search(
158        &self,
159        query_embedding: &[f32],
160        k: usize,
161    ) -> Result<Vec<SearchResult>, VectorStoreError> {
162        let data = self.data.read().await;
163
164        let mut results: Vec<SearchResult> = data
165            .documents
166            .values()
167            .map(|vd| {
168                let score = cosine_similarity(query_embedding, &vd.embedding).unwrap_or(0.0);
169                SearchResult {
170                    document: vd.document.clone(),
171                    score,
172                }
173            })
174            .collect();
175
176        results.sort_by(|a, b| {
177            b.score
178                .partial_cmp(&a.score)
179                .unwrap_or(std::cmp::Ordering::Equal)
180        });
181
182        Ok(results.into_iter().take(k).collect())
183    }
184
185    async fn get_document(&self, id: &str) -> Result<Option<Document>, VectorStoreError> {
186        let data = self.data.read().await;
187        Ok(data.documents.get(id).map(|vd| vd.document.clone()))
188    }
189
190    async fn get_embedding(&self, id: &str) -> Result<Option<Vec<f32>>, VectorStoreError> {
191        let data = self.data.read().await;
192        Ok(data.documents.get(id).map(|vd| vd.embedding.clone()))
193    }
194
195    async fn delete_document(&self, id: &str) -> Result<(), VectorStoreError> {
196        let mut data = self.data.write().await;
197        data.documents
198            .remove(id)
199            .ok_or_else(|| VectorStoreError::DocumentNotFound(id.to_string()))?;
200        Self::persist(&data, &self.path).await?;
201        Ok(())
202    }
203
204    async fn count(&self) -> usize {
205        let data = self.data.read().await;
206        data.documents.len()
207    }
208
209    async fn clear(&self) -> Result<(), VectorStoreError> {
210        let mut data = self.data.write().await;
211        data.documents.clear();
212        Self::persist(&data, &self.path).await?;
213        Ok(())
214    }
215}
216
217#[cfg(test)]
218mod tests {
219    use super::*;
220    use tempfile::TempDir;
221
222    fn test_store_path(dir: &TempDir) -> PathBuf {
223        dir.path().join("test_vectors.json")
224    }
225
226    #[tokio::test]
227    async fn test_new_creates_empty_store() {
228        let dir = TempDir::new().unwrap();
229        let path = test_store_path(&dir);
230        let store = FileVectorStore::new(path.clone(), 3).await.unwrap();
231        assert_eq!(store.count().await, 0);
232        assert_eq!(store.dimension(), 3);
233        assert!(path.exists());
234    }
235
236    #[tokio::test]
237    async fn test_add_and_search() {
238        let dir = TempDir::new().unwrap();
239        let path = test_store_path(&dir);
240        let store = FileVectorStore::new(path, 3).await.unwrap();
241
242        let docs = vec![
243            Document::new("Rust is a systems programming language"),
244            Document::new("Python is a scripting language"),
245            Document::new("JavaScript is used for web development"),
246        ];
247        let embeddings = vec![
248            vec![1.0, 0.0, 0.0],
249            vec![0.0, 1.0, 0.0],
250            vec![0.0, 0.0, 1.0],
251        ];
252
253        let ids = store.add_documents(docs, embeddings).await.unwrap();
254        assert_eq!(ids.len(), 3);
255        assert_eq!(store.count().await, 3);
256
257        let query = vec![0.9, 0.1, 0.0];
258        let results = store.similarity_search(&query, 2).await.unwrap();
259        assert_eq!(results.len(), 2);
260        assert!(results[0].document.content.contains("Rust"));
261        assert!(results[0].score > results[1].score);
262    }
263
264    #[tokio::test]
265    async fn test_persistence_across_instances() {
266        let dir = TempDir::new().unwrap();
267        let path = test_store_path(&dir);
268
269        // 第一个实例:写入
270        {
271            let store = FileVectorStore::new(path.clone(), 3).await.unwrap();
272            let doc = Document::new("persistent doc").with_id("p1");
273            store
274                .add_documents(vec![doc], vec![vec![1.0, 0.0, 0.0]])
275                .await
276                .unwrap();
277        }
278
279        // 第二个实例:加载并验证
280        {
281            let store = FileVectorStore::new(path.clone(), 3).await.unwrap();
282            assert_eq!(store.count().await, 1);
283            let doc = store.get_document("p1").await.unwrap().unwrap();
284            assert_eq!(doc.content, "persistent doc");
285        }
286    }
287
288    #[tokio::test]
289    async fn test_delete_persists() {
290        let dir = TempDir::new().unwrap();
291        let path = test_store_path(&dir);
292
293        {
294            let store = FileVectorStore::new(path.clone(), 3).await.unwrap();
295            let doc = Document::new("to delete").with_id("d1");
296            store
297                .add_documents(vec![doc], vec![vec![1.0, 0.0, 0.0]])
298                .await
299                .unwrap();
300            store.delete_document("d1").await.unwrap();
301        }
302
303        let store = FileVectorStore::new(path.clone(), 3).await.unwrap();
304        assert_eq!(store.count().await, 0);
305    }
306
307    #[tokio::test]
308    async fn test_clear_persists() {
309        let dir = TempDir::new().unwrap();
310        let path = test_store_path(&dir);
311
312        {
313            let store = FileVectorStore::new(path.clone(), 3).await.unwrap();
314            let docs = vec![Document::new("a"), Document::new("b")];
315            let embeddings = vec![vec![1.0, 0.0, 0.0], vec![0.0, 1.0, 0.0]];
316            store.add_documents(docs, embeddings).await.unwrap();
317            store.clear().await.unwrap();
318        }
319
320        let store = FileVectorStore::new(path.clone(), 3).await.unwrap();
321        assert_eq!(store.count().await, 0);
322    }
323
324    #[tokio::test]
325    async fn test_dimension_mismatch() {
326        let dir = TempDir::new().unwrap();
327        let path = test_store_path(&dir);
328        let store = FileVectorStore::new(path, 3).await.unwrap();
329
330        let doc = Document::new("wrong dim");
331        let wrong_embedding = vec![1.0, 0.0]; // 维度 2,存储维度 3
332        let result = store.add_documents(vec![doc], vec![wrong_embedding]).await;
333        assert!(result.is_err());
334    }
335
336    #[tokio::test]
337    async fn test_get_embedding() {
338        let dir = TempDir::new().unwrap();
339        let path = test_store_path(&dir);
340        let store = FileVectorStore::new(path, 3).await.unwrap();
341
342        let doc = Document::new("embed test").with_id("e1");
343        store
344            .add_documents(vec![doc], vec![vec![0.5, 0.5, 0.0]])
345            .await
346            .unwrap();
347
348        let emb = store.get_embedding("e1").await.unwrap().unwrap();
349        assert_eq!(emb, vec![0.5, 0.5, 0.0]);
350    }
351
352    #[tokio::test]
353    async fn test_delete_nonexistent() {
354        let dir = TempDir::new().unwrap();
355        let path = test_store_path(&dir);
356        let store = FileVectorStore::new(path, 3).await.unwrap();
357
358        let result = store.delete_document("no-such-id").await;
359        assert!(result.is_err());
360    }
361
362    #[tokio::test]
363    async fn test_cosine_similarity() {
364        let a = vec![1.0, 0.0, 0.0];
365        let b = vec![1.0, 0.0, 0.0];
366        assert!((cosine_similarity(&a, &b).unwrap() - 1.0).abs() < 0.0001);
367
368        let a = vec![1.0, 0.0, 0.0];
369        let b = vec![0.0, 1.0, 0.0];
370        assert!((cosine_similarity(&a, &b).unwrap() - 0.0).abs() < 0.0001);
371    }
372}