docling_rag/store/
memory.rs1use super::{top_k_by_cosine, VectorStore};
4use crate::model::{Chunk, Document, Scored};
5use crate::Result;
6use async_trait::async_trait;
7use std::sync::RwLock;
8
9#[derive(Default)]
11pub struct MemoryStore {
12 docs: RwLock<Vec<Document>>,
13 chunks: RwLock<Vec<Chunk>>,
14}
15
16impl MemoryStore {
17 pub fn new() -> Self {
19 MemoryStore::default()
20 }
21}
22
23#[async_trait]
24impl VectorStore for MemoryStore {
25 async fn migrate(&self) -> Result<()> {
26 Ok(())
27 }
28
29 async fn upsert_document(&self, doc: &Document) -> Result<()> {
30 let mut docs = self.docs.write().unwrap();
31 if let Some(existing) = docs.iter_mut().find(|d| d.id == doc.id) {
32 *existing = doc.clone();
33 } else {
34 docs.push(doc.clone());
35 }
36 Ok(())
37 }
38
39 async fn find_document_by_hash(&self, hash: &str) -> Result<Option<String>> {
40 Ok(self
41 .docs
42 .read()
43 .unwrap()
44 .iter()
45 .find(|d| d.hash == hash)
46 .map(|d| d.id.clone()))
47 }
48
49 async fn insert_chunks(&self, chunks: &[Chunk]) -> Result<()> {
50 self.chunks.write().unwrap().extend_from_slice(chunks);
51 Ok(())
52 }
53
54 async fn vector_search(&self, query: &[f32], k: usize) -> Result<Vec<Scored>> {
55 let chunks = self.chunks.read().unwrap();
56 let candidates = chunks.iter().filter_map(|c| {
57 c.embedding.as_ref().map(|e| {
58 let mut bare = c.clone();
59 bare.embedding = None;
60 (bare, e.clone())
61 })
62 });
63 Ok(top_k_by_cosine(query, candidates.collect::<Vec<_>>(), k))
64 }
65
66 async fn all_chunks(&self) -> Result<Vec<Chunk>> {
67 Ok(self
68 .chunks
69 .read()
70 .unwrap()
71 .iter()
72 .map(|c| {
73 let mut bare = c.clone();
74 bare.embedding = None;
75 bare
76 })
77 .collect())
78 }
79
80 async fn count_chunks(&self) -> Result<usize> {
81 Ok(self.chunks.read().unwrap().len())
82 }
83
84 async fn count_chunks_for(&self, doc_id: &str) -> Result<usize> {
85 Ok(self
86 .chunks
87 .read()
88 .unwrap()
89 .iter()
90 .filter(|c| c.doc_id == doc_id)
91 .count())
92 }
93
94 async fn chunk_neighborhood(&self, doc_id: &str, ordinal: i64) -> Result<Vec<Chunk>> {
95 let mut out: Vec<Chunk> = self
96 .chunks
97 .read()
98 .unwrap()
99 .iter()
100 .filter(|c| c.doc_id == doc_id && (c.ordinal - ordinal).abs() <= 1)
101 .cloned()
102 .collect();
103 out.sort_by_key(|c| c.ordinal);
104 Ok(out)
105 }
106
107 async fn count_documents(&self) -> Result<usize> {
108 Ok(self.docs.read().unwrap().len())
109 }
110
111 async fn list_documents(&self) -> Result<Vec<Document>> {
112 Ok(self.docs.read().unwrap().clone())
113 }
114
115 async fn delete_document(&self, doc_id: &str) -> Result<()> {
116 self.docs.write().unwrap().retain(|d| d.id != doc_id);
117 self.chunks.write().unwrap().retain(|c| c.doc_id != doc_id);
118 Ok(())
119 }
120
121 async fn delete_documents_by_source(&self, source_uri: &str) -> Result<()> {
122 let ids: Vec<String> = {
123 let docs = self.docs.read().unwrap();
124 docs.iter()
125 .filter(|d| d.source_uri == source_uri)
126 .map(|d| d.id.clone())
127 .collect()
128 };
129 self.docs
130 .write()
131 .unwrap()
132 .retain(|d| d.source_uri != source_uri);
133 self.chunks
134 .write()
135 .unwrap()
136 .retain(|c| !ids.contains(&c.doc_id));
137 Ok(())
138 }
139
140 async fn clear(&self) -> Result<()> {
141 self.docs.write().unwrap().clear();
142 self.chunks.write().unwrap().clear();
143 Ok(())
144 }
145}