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).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 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 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 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 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 pub fn dimension(&self) -> usize {
106 self.dimension
107 }
108
109 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 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 {
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 {
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]; 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}