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).await.map_err(|e| {
59                VectorStoreError::StorageError(format!("failed to read file: {}", e))
60            })?;
61            serde_json::from_str::<FileStoreData>(&content).map_err(|e| {
62                VectorStoreError::StorageError(format!("failed to parse file: {}", e))
63            })?
64        } else {
65            // 确保父目录存在
66            if let Some(parent) = path.parent() {
67                if !parent.as_os_str().is_empty() {
68                    tokio::fs::create_dir_all(parent).await.map_err(|e| {
69                        VectorStoreError::StorageError(format!("failed to create directory: {}", e))
70                    })?;
71                }
72            }
73            let data = FileStoreData {
74                dimension,
75                documents: HashMap::new(),
76            };
77            // 首次创建时也持久化空文件,确保 path.exists() 为 true
78            Self::persist(&data, &path).await?;
79            data
80        };
81
82        Ok(Self {
83            path,
84            dimension: data.dimension,
85            data: RwLock::new(data),
86        })
87    }
88
89    /// 持久化当前数据到磁盘
90    async fn persist(data: &FileStoreData, path: &PathBuf) -> Result<(), VectorStoreError> {
91        let json = serde_json::to_string(data)
92            .map_err(|e| VectorStoreError::StorageError(format!("failed to serialize: {}", e)))?;
93        // 先写临时文件,再 rename,避免写一半断电损坏
94        let tmp_path = path.with_extension("json.tmp");
95        tokio::fs::write(&tmp_path, &json).await.map_err(|e| {
96            VectorStoreError::StorageError(format!("failed to write temporary file: {}", e))
97        })?;
98        tokio::fs::rename(&tmp_path, path)
99            .await
100            .map_err(|e| VectorStoreError::StorageError(format!("failed to rename file: {}", e)))?;
101        Ok(())
102    }
103
104    /// 返回向量维度
105    pub fn dimension(&self) -> usize {
106        self.dimension
107    }
108
109    /// 返回存储文件路径
110    pub fn path(&self) -> &PathBuf {
111        &self.path
112    }
113}
114
115#[async_trait]
116impl VectorStore for FileVectorStore {
117    async fn add_documents(
118        &self,
119        documents: Vec<Document>,
120        embeddings: Vec<Vec<f32>>,
121    ) -> Result<Vec<String>, VectorStoreError> {
122        if documents.len() != embeddings.len() {
123            return Err(VectorStoreError::StorageError(
124                "document count and embedding count mismatch".to_string(),
125            ));
126        }
127
128        let mut data = self.data.write().await;
129        let mut ids = Vec::new();
130
131        for (doc, embedding) in documents.into_iter().zip(embeddings.into_iter()) {
132            // 校验维度
133            if !embedding.is_empty() && embedding.len() != data.dimension {
134                return Err(VectorStoreError::StorageError(format!(
135                    "embedding dimension {} does not match storage dimension {}",
136                    embedding.len(),
137                    data.dimension
138                )));
139            }
140
141            let id = doc.id.clone().unwrap_or_else(|| Uuid::new_v4().to_string());
142            let vector_doc = VectorDocument {
143                document: Document {
144                    id: Some(id.clone()),
145                    content: doc.content,
146                    metadata: doc.metadata,
147                },
148                embedding,
149            };
150            data.documents.insert(id.clone(), vector_doc);
151            ids.push(id);
152        }
153
154        Self::persist(&data, &self.path).await?;
155        Ok(ids)
156    }
157
158    async fn similarity_search(
159        &self,
160        query_embedding: &[f32],
161        k: usize,
162    ) -> Result<Vec<SearchResult>, VectorStoreError> {
163        let data = self.data.read().await;
164
165        let mut results: Vec<SearchResult> = data
166            .documents
167            .values()
168            .map(|vd| {
169                let score = cosine_similarity(query_embedding, &vd.embedding).unwrap_or(0.0);
170                SearchResult {
171                    document: vd.document.clone(),
172                    score,
173                }
174            })
175            .collect();
176
177        results.sort_by(|a, b| {
178            b.score
179                .partial_cmp(&a.score)
180                .unwrap_or(std::cmp::Ordering::Equal)
181        });
182
183        Ok(results.into_iter().take(k).collect())
184    }
185
186    async fn get_document(&self, id: &str) -> Result<Option<Document>, VectorStoreError> {
187        let data = self.data.read().await;
188        Ok(data.documents.get(id).map(|vd| vd.document.clone()))
189    }
190
191    async fn get_embedding(&self, id: &str) -> Result<Option<Vec<f32>>, VectorStoreError> {
192        let data = self.data.read().await;
193        Ok(data.documents.get(id).map(|vd| vd.embedding.clone()))
194    }
195
196    async fn delete_document(&self, id: &str) -> Result<(), VectorStoreError> {
197        let mut data = self.data.write().await;
198        data.documents
199            .remove(id)
200            .ok_or_else(|| VectorStoreError::DocumentNotFound(id.to_string()))?;
201        Self::persist(&data, &self.path).await?;
202        Ok(())
203    }
204
205    async fn count(&self) -> usize {
206        let data = self.data.read().await;
207        data.documents.len()
208    }
209
210    async fn clear(&self) -> Result<(), VectorStoreError> {
211        let mut data = self.data.write().await;
212        data.documents.clear();
213        Self::persist(&data, &self.path).await?;
214        Ok(())
215    }
216}
217
218#[cfg(test)]
219mod tests {
220    use super::*;
221    use tempfile::TempDir;
222
223    fn test_store_path(dir: &TempDir) -> PathBuf {
224        dir.path().join("test_vectors.json")
225    }
226
227    #[tokio::test]
228    async fn test_new_creates_empty_store() {
229        let dir = TempDir::new().unwrap();
230        let path = test_store_path(&dir);
231        let store = FileVectorStore::new(path.clone(), 3).await.unwrap();
232        assert_eq!(store.count().await, 0);
233        assert_eq!(store.dimension(), 3);
234        assert!(path.exists());
235    }
236
237    #[tokio::test]
238    async fn test_add_and_search() {
239        let dir = TempDir::new().unwrap();
240        let path = test_store_path(&dir);
241        let store = FileVectorStore::new(path, 3).await.unwrap();
242
243        let docs = vec![
244            Document::new("Rust is a systems programming language"),
245            Document::new("Python is a scripting language"),
246            Document::new("JavaScript is used for web development"),
247        ];
248        let embeddings = vec![
249            vec![1.0, 0.0, 0.0],
250            vec![0.0, 1.0, 0.0],
251            vec![0.0, 0.0, 1.0],
252        ];
253
254        let ids = store.add_documents(docs, embeddings).await.unwrap();
255        assert_eq!(ids.len(), 3);
256        assert_eq!(store.count().await, 3);
257
258        let query = vec![0.9, 0.1, 0.0];
259        let results = store.similarity_search(&query, 2).await.unwrap();
260        assert_eq!(results.len(), 2);
261        assert!(results[0].document.content.contains("Rust"));
262        assert!(results[0].score > results[1].score);
263    }
264
265    #[tokio::test]
266    async fn test_persistence_across_instances() {
267        let dir = TempDir::new().unwrap();
268        let path = test_store_path(&dir);
269
270        // 第一个实例:写入
271        {
272            let store = FileVectorStore::new(path.clone(), 3).await.unwrap();
273            let doc = Document::new("persistent doc").with_id("p1");
274            store
275                .add_documents(vec![doc], vec![vec![1.0, 0.0, 0.0]])
276                .await
277                .unwrap();
278        }
279
280        // 第二个实例:加载并验证
281        {
282            let store = FileVectorStore::new(path.clone(), 3).await.unwrap();
283            assert_eq!(store.count().await, 1);
284            let doc = store.get_document("p1").await.unwrap().unwrap();
285            assert_eq!(doc.content, "persistent doc");
286        }
287    }
288
289    #[tokio::test]
290    async fn test_delete_persists() {
291        let dir = TempDir::new().unwrap();
292        let path = test_store_path(&dir);
293
294        {
295            let store = FileVectorStore::new(path.clone(), 3).await.unwrap();
296            let doc = Document::new("to delete").with_id("d1");
297            store
298                .add_documents(vec![doc], vec![vec![1.0, 0.0, 0.0]])
299                .await
300                .unwrap();
301            store.delete_document("d1").await.unwrap();
302        }
303
304        let store = FileVectorStore::new(path.clone(), 3).await.unwrap();
305        assert_eq!(store.count().await, 0);
306    }
307
308    #[tokio::test]
309    async fn test_clear_persists() {
310        let dir = TempDir::new().unwrap();
311        let path = test_store_path(&dir);
312
313        {
314            let store = FileVectorStore::new(path.clone(), 3).await.unwrap();
315            let docs = vec![Document::new("a"), Document::new("b")];
316            let embeddings = vec![vec![1.0, 0.0, 0.0], vec![0.0, 1.0, 0.0]];
317            store.add_documents(docs, embeddings).await.unwrap();
318            store.clear().await.unwrap();
319        }
320
321        let store = FileVectorStore::new(path.clone(), 3).await.unwrap();
322        assert_eq!(store.count().await, 0);
323    }
324
325    #[tokio::test]
326    async fn test_dimension_mismatch() {
327        let dir = TempDir::new().unwrap();
328        let path = test_store_path(&dir);
329        let store = FileVectorStore::new(path, 3).await.unwrap();
330
331        let doc = Document::new("wrong dim");
332        let wrong_embedding = vec![1.0, 0.0]; // 维度 2,存储维度 3
333        let result = store.add_documents(vec![doc], vec![wrong_embedding]).await;
334        assert!(result.is_err());
335    }
336
337    #[tokio::test]
338    async fn test_get_embedding() {
339        let dir = TempDir::new().unwrap();
340        let path = test_store_path(&dir);
341        let store = FileVectorStore::new(path, 3).await.unwrap();
342
343        let doc = Document::new("embed test").with_id("e1");
344        store
345            .add_documents(vec![doc], vec![vec![0.5, 0.5, 0.0]])
346            .await
347            .unwrap();
348
349        let emb = store.get_embedding("e1").await.unwrap().unwrap();
350        assert_eq!(emb, vec![0.5, 0.5, 0.0]);
351    }
352
353    #[tokio::test]
354    async fn test_delete_nonexistent() {
355        let dir = TempDir::new().unwrap();
356        let path = test_store_path(&dir);
357        let store = FileVectorStore::new(path, 3).await.unwrap();
358
359        let result = store.delete_document("no-such-id").await;
360        assert!(result.is_err());
361    }
362
363    #[tokio::test]
364    async fn test_cosine_similarity() {
365        let a = vec![1.0, 0.0, 0.0];
366        let b = vec![1.0, 0.0, 0.0];
367        assert!((cosine_similarity(&a, &b).unwrap() - 1.0).abs() < 0.0001);
368
369        let a = vec![1.0, 0.0, 0.0];
370        let b = vec![0.0, 1.0, 0.0];
371        assert!((cosine_similarity(&a, &b).unwrap() - 0.0).abs() < 0.0001);
372    }
373}