1use 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#[derive(serde::Serialize, serde::Deserialize)]
28struct FileStoreData {
29 dimension: usize,
31 documents: HashMap<String, VectorDocument>,
33}
34
35pub struct FileVectorStore {
40 path: PathBuf,
42 dimension: usize,
44 data: RwLock<FileStoreData>,
46}
47
48impl FileVectorStore {
49 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 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 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 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 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 pub fn dimension(&self) -> usize {
105 self.dimension
106 }
107
108 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 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 {
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 {
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]; 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}