Skip to main content

docling_rag/store/
memory.rs

1//! In-memory vector store. Never persisted; used by tests and quick evals.
2
3use 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/// A store that keeps everything in process memory.
10#[derive(Default)]
11pub struct MemoryStore {
12    docs: RwLock<Vec<Document>>,
13    chunks: RwLock<Vec<Chunk>>,
14}
15
16impl MemoryStore {
17    /// Create an empty in-memory store.
18    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}