1use async_trait::async_trait;
10use chrono::Utc;
11use ragfs_core::{
12 Chunk, DirectoryScope, FileRecord, SearchQuery, SearchResult, StoreError, StoreStats,
13 VectorStore, dir_path_matches_scope,
14};
15use std::collections::HashMap;
16use std::path::{Path, PathBuf};
17use std::sync::Arc;
18use tokio::sync::RwLock;
19use tracing::debug;
20use uuid::Uuid;
21
22pub struct MemoryStore {
45 dimension: usize,
46 chunks: Arc<RwLock<HashMap<Uuid, Chunk>>>,
47 files: Arc<RwLock<HashMap<PathBuf, FileRecord>>>,
48 initialized: Arc<RwLock<bool>>,
49}
50
51impl MemoryStore {
52 #[must_use]
54 pub fn new(dimension: usize) -> Self {
55 Self {
56 dimension,
57 chunks: Arc::new(RwLock::new(HashMap::new())),
58 files: Arc::new(RwLock::new(HashMap::new())),
59 initialized: Arc::new(RwLock::new(false)),
60 }
61 }
62
63 fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
65 if a.len() != b.len() {
66 return 0.0;
67 }
68
69 let dot: f32 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
70 let norm_a: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
71 let norm_b: f32 = b.iter().map(|x| x * x).sum::<f32>().sqrt();
72
73 if norm_a == 0.0 || norm_b == 0.0 {
74 return 0.0;
75 }
76
77 dot / (norm_a * norm_b)
78 }
79}
80
81impl Default for MemoryStore {
82 fn default() -> Self {
83 Self::new(384)
84 }
85}
86
87#[async_trait]
88impl VectorStore for MemoryStore {
89 async fn init(&self) -> Result<(), StoreError> {
90 let mut initialized = self.initialized.write().await;
91 *initialized = true;
92 debug!("MemoryStore initialized (dimension: {})", self.dimension);
93 Ok(())
94 }
95
96 async fn upsert_chunks(&self, chunks: &[Chunk]) -> Result<(), StoreError> {
97 let mut store = self.chunks.write().await;
98 for chunk in chunks {
99 store.insert(chunk.id, chunk.clone());
100 }
101 debug!("Upserted {} chunks", chunks.len());
102 Ok(())
103 }
104
105 async fn search(&self, query: SearchQuery) -> Result<Vec<SearchResult>, StoreError> {
106 let chunks = self.chunks.read().await;
107 let mut results: Vec<(f32, &Chunk)> = Vec::new();
108
109 for chunk in chunks.values() {
111 if let Some(ref scope) = query.scope_prefix
112 && !dir_path_matches_scope(&chunk.dir_path, scope)
113 {
114 continue;
115 }
116 if let Some(embedding) = &chunk.embedding {
117 let score = Self::cosine_similarity(&query.embedding, embedding);
118 results.push((score, chunk));
119 }
120 }
121
122 results.sort_by(|a, b| b.0.partial_cmp(&a.0).unwrap_or(std::cmp::Ordering::Equal));
124
125 let top_k = results
127 .into_iter()
128 .take(query.limit)
129 .map(|(score, chunk)| SearchResult {
130 chunk_id: chunk.id,
131 file_path: chunk.file_path.clone(),
132 content: chunk.content.clone(),
133 score,
134 byte_range: chunk.byte_range.clone(),
135 line_range: chunk.line_range.clone(),
136 metadata: chunk.metadata.extra.clone(),
137 })
138 .collect();
139
140 Ok(top_k)
141 }
142
143 async fn hybrid_search(&self, query: SearchQuery) -> Result<Vec<SearchResult>, StoreError> {
144 self.search(query).await
147 }
148
149 async fn delete_by_file_path(&self, path: &Path) -> Result<u64, StoreError> {
150 let mut chunks = self.chunks.write().await;
151 let mut files = self.files.write().await;
152
153 let before = chunks.len();
154 chunks.retain(|_, chunk| chunk.file_path != path);
155 let deleted = (before - chunks.len()) as u64;
156
157 files.remove(path);
158
159 debug!("Deleted {} chunks for {:?}", deleted, path);
160 Ok(deleted)
161 }
162
163 async fn update_file_path(&self, from: &Path, to: &Path) -> Result<u64, StoreError> {
164 let mut chunks = self.chunks.write().await;
165 let mut files = self.files.write().await;
166 let mut updated = 0u64;
167
168 for chunk in chunks.values_mut() {
170 if chunk.file_path == from {
171 let root = DirectoryScope::infer_root(&chunk.file_path, &chunk.dir_path);
172 let scope = DirectoryScope::from_paths(to, root.as_deref());
173 chunk.file_path = to.to_path_buf();
174 chunk.dir_path = scope.dir_path;
175 chunk.dir_depth = scope.dir_depth;
176 chunk.path_components = scope.path_components;
177 updated += 1;
178 }
179 }
180
181 if let Some(mut record) = files.remove(from) {
183 record.path = to.to_path_buf();
184 files.insert(to.to_path_buf(), record);
185 }
186
187 debug!("Updated {} chunks from {:?} to {:?}", updated, from, to);
188 Ok(updated)
189 }
190
191 async fn get_chunks_for_file(&self, path: &Path) -> Result<Vec<Chunk>, StoreError> {
192 let chunks = self.chunks.read().await;
193 let file_chunks: Vec<Chunk> = chunks
194 .values()
195 .filter(|chunk| chunk.file_path == path)
196 .cloned()
197 .collect();
198 Ok(file_chunks)
199 }
200
201 async fn get_file(&self, path: &Path) -> Result<Option<FileRecord>, StoreError> {
202 let files = self.files.read().await;
203 Ok(files.get(path).cloned())
204 }
205
206 async fn upsert_file(&self, record: &FileRecord) -> Result<(), StoreError> {
207 let mut files = self.files.write().await;
208 files.insert(record.path.clone(), record.clone());
209 debug!("Upserted file record for {:?}", record.path);
210 Ok(())
211 }
212
213 async fn stats(&self) -> Result<StoreStats, StoreError> {
214 let chunks = self.chunks.read().await;
215 let files = self.files.read().await;
216
217 Ok(StoreStats {
218 total_chunks: chunks.len() as u64,
219 total_files: files.len() as u64,
220 index_size_bytes: 0, last_updated: Some(Utc::now()),
222 })
223 }
224
225 async fn get_all_chunks(&self) -> Result<Vec<Chunk>, StoreError> {
226 let chunks = self.chunks.read().await;
227 Ok(chunks.values().cloned().collect())
228 }
229
230 async fn get_all_files(&self) -> Result<Vec<FileRecord>, StoreError> {
231 let files = self.files.read().await;
232 Ok(files.values().cloned().collect())
233 }
234}
235
236#[cfg(test)]
237mod tests {
238 use super::*;
239 use ragfs_core::{ChunkMetadata, ContentType, DistanceMetric};
240
241 fn create_test_chunk(id: Uuid, file_id: Uuid, path: &str, embedding: Vec<f32>) -> Chunk {
242 let file_path = PathBuf::from(path);
243 let scope = DirectoryScope::from_file_path(&file_path);
244 Chunk {
245 id,
246 file_id,
247 file_path,
248 content: "test content".to_string(),
249 content_type: ContentType::Text,
250 mime_type: Some("text/plain".to_string()),
251 chunk_index: 0,
252 byte_range: 0..12,
253 line_range: Some(0..1),
254 parent_chunk_id: None,
255 depth: 0,
256 embedding: Some(embedding),
257 dir_path: scope.dir_path,
258 dir_depth: scope.dir_depth,
259 path_components: scope.path_components,
260 metadata: ChunkMetadata::default(),
261 }
262 }
263
264 fn create_scoped_chunk(
265 path: &str,
266 root: Option<&Path>,
267 embedding: Vec<f32>,
268 content: &str,
269 ) -> Chunk {
270 let file_path = PathBuf::from(path);
271 let scope = DirectoryScope::from_paths(&file_path, root);
272 Chunk {
273 id: Uuid::new_v4(),
274 file_id: Uuid::new_v4(),
275 file_path,
276 content: content.to_string(),
277 content_type: ContentType::Text,
278 mime_type: Some("text/plain".to_string()),
279 chunk_index: 0,
280 byte_range: 0..content.len() as u64,
281 line_range: Some(0..1),
282 parent_chunk_id: None,
283 depth: 0,
284 embedding: Some(embedding),
285 dir_path: scope.dir_path,
286 dir_depth: scope.dir_depth,
287 path_components: scope.path_components,
288 metadata: ChunkMetadata::default(),
289 }
290 }
291
292 #[tokio::test]
293 async fn test_memory_store_new() {
294 let store = MemoryStore::new(384);
295 assert_eq!(store.dimension, 384);
296 }
297
298 #[tokio::test]
299 async fn test_memory_store_init() {
300 let store = MemoryStore::new(384);
301 let result = store.init().await;
302 assert!(result.is_ok());
303 }
304
305 #[tokio::test]
306 async fn test_memory_store_upsert_and_stats() {
307 let store = MemoryStore::new(3);
308 store.init().await.unwrap();
309
310 let file_id = Uuid::new_v4();
311 let chunks = vec![
312 create_test_chunk(
313 Uuid::new_v4(),
314 file_id,
315 "/test/file.txt",
316 vec![1.0, 0.0, 0.0],
317 ),
318 create_test_chunk(
319 Uuid::new_v4(),
320 file_id,
321 "/test/file.txt",
322 vec![0.0, 1.0, 0.0],
323 ),
324 ];
325
326 store.upsert_chunks(&chunks).await.unwrap();
327
328 let stats = store.stats().await.unwrap();
329 assert_eq!(stats.total_chunks, 2);
330 }
331
332 #[tokio::test]
333 async fn test_memory_store_search() {
334 let store = MemoryStore::new(3);
335 store.init().await.unwrap();
336
337 let file_id = Uuid::new_v4();
338 let chunk1_id = Uuid::new_v4();
339 let chunks = vec![
340 create_test_chunk(chunk1_id, file_id, "/test/file.txt", vec![1.0, 0.0, 0.0]),
341 create_test_chunk(
342 Uuid::new_v4(),
343 file_id,
344 "/test/file.txt",
345 vec![0.0, 1.0, 0.0],
346 ),
347 create_test_chunk(
348 Uuid::new_v4(),
349 file_id,
350 "/test/file.txt",
351 vec![0.0, 0.0, 1.0],
352 ),
353 ];
354
355 store.upsert_chunks(&chunks).await.unwrap();
356
357 let query = SearchQuery {
358 embedding: vec![1.0, 0.0, 0.0],
359 text: None,
360 limit: 2,
361 filters: vec![],
362 metric: Default::default(),
363 scope_prefix: None,
364 };
365
366 let results = store.search(query).await.unwrap();
367 assert_eq!(results.len(), 2);
368 assert_eq!(results[0].chunk_id, chunk1_id);
369 assert!((results[0].score - 1.0).abs() < 0.001);
370 }
371
372 #[tokio::test]
373 async fn test_memory_store_delete_by_file_path() {
374 let store = MemoryStore::new(3);
375 store.init().await.unwrap();
376
377 let chunks = vec![
378 create_test_chunk(
379 Uuid::new_v4(),
380 Uuid::new_v4(),
381 "/test/file1.txt",
382 vec![1.0, 0.0, 0.0],
383 ),
384 create_test_chunk(
385 Uuid::new_v4(),
386 Uuid::new_v4(),
387 "/test/file2.txt",
388 vec![0.0, 1.0, 0.0],
389 ),
390 ];
391
392 store.upsert_chunks(&chunks).await.unwrap();
393
394 let deleted = store
395 .delete_by_file_path(Path::new("/test/file1.txt"))
396 .await
397 .unwrap();
398 assert_eq!(deleted, 1);
399
400 let stats = store.stats().await.unwrap();
401 assert_eq!(stats.total_chunks, 1);
402 }
403
404 #[tokio::test]
405 async fn test_memory_store_get_all_chunks() {
406 let store = MemoryStore::new(3);
407 store.init().await.unwrap();
408
409 let file_id = Uuid::new_v4();
410 let chunks = vec![
411 create_test_chunk(
412 Uuid::new_v4(),
413 file_id,
414 "/test/file.txt",
415 vec![1.0, 0.0, 0.0],
416 ),
417 create_test_chunk(
418 Uuid::new_v4(),
419 file_id,
420 "/test/file.txt",
421 vec![0.0, 1.0, 0.0],
422 ),
423 ];
424
425 store.upsert_chunks(&chunks).await.unwrap();
426
427 let all_chunks = store.get_all_chunks().await.unwrap();
428 assert_eq!(all_chunks.len(), 2);
429 }
430
431 #[tokio::test]
432 async fn test_memory_store_scoped_search() {
433 let store = MemoryStore::new(3);
434 store.init().await.unwrap();
435 let root = Path::new("/project");
436 let embedding = vec![1.0, 0.0, 0.0];
437 store
438 .upsert_chunks(&[
439 create_scoped_chunk(
440 "/project/src/auth/login.rs",
441 Some(root),
442 embedding.clone(),
443 "login",
444 ),
445 create_scoped_chunk(
446 "/project/src/auth/oauth/token.rs",
447 Some(root),
448 embedding.clone(),
449 "oauth",
450 ),
451 create_scoped_chunk("/project/src/db.rs", Some(root), embedding.clone(), "db"),
452 create_scoped_chunk(
453 "/project/docs/readme.md",
454 Some(root),
455 embedding.clone(),
456 "docs",
457 ),
458 ])
459 .await
460 .unwrap();
461
462 let stored = store
463 .get_chunks_for_file(Path::new("/project/src/auth/login.rs"))
464 .await
465 .unwrap();
466 assert_eq!(stored[0].dir_path, "src/auth");
467 assert!(!stored[0].dir_path.starts_with('/'));
468
469 let exact = store
470 .search(SearchQuery {
471 embedding: embedding.clone(),
472 text: None,
473 limit: 10,
474 filters: vec![],
475 metric: DistanceMetric::Cosine,
476 scope_prefix: Some("src/auth".to_string()),
477 })
478 .await
479 .unwrap();
480 assert_eq!(exact.len(), 2);
481 assert!(
482 exact
483 .iter()
484 .any(|r| r.file_path.ends_with("src/auth/login.rs"))
485 );
486 assert!(
487 exact
488 .iter()
489 .any(|r| r.file_path.ends_with("src/auth/oauth/token.rs"))
490 );
491
492 let src = store
493 .search(SearchQuery {
494 embedding: embedding.clone(),
495 text: None,
496 limit: 10,
497 filters: vec![],
498 metric: DistanceMetric::Cosine,
499 scope_prefix: Some("src".to_string()),
500 })
501 .await
502 .unwrap();
503 assert_eq!(src.len(), 3);
504 assert!(src.iter().any(|r| r.file_path.ends_with("src/db.rs")));
505
506 let docs = store
507 .search(SearchQuery {
508 embedding,
509 text: None,
510 limit: 10,
511 filters: vec![],
512 metric: DistanceMetric::Cosine,
513 scope_prefix: Some("docs".to_string()),
514 })
515 .await
516 .unwrap();
517 assert_eq!(docs.len(), 1);
518 assert!(docs[0].file_path.ends_with("docs/readme.md"));
519 }
520
521 #[test]
522 fn test_cosine_similarity() {
523 let sim = MemoryStore::cosine_similarity(&[1.0, 0.0, 0.0], &[1.0, 0.0, 0.0]);
525 assert!((sim - 1.0).abs() < 0.001);
526
527 let sim = MemoryStore::cosine_similarity(&[1.0, 0.0, 0.0], &[0.0, 1.0, 0.0]);
529 assert!(sim.abs() < 0.001);
530
531 let sim = MemoryStore::cosine_similarity(&[1.0, 0.0, 0.0], &[-1.0, 0.0, 0.0]);
533 assert!((sim - (-1.0)).abs() < 0.001);
534 }
535}