Skip to main content

docling_rag/retrieve/
mod.rs

1//! Retrieval: dense vector, sparse BM25, and the advanced modes that combine or
2//! rewrite queries (Hybrid, Multi-Query fusion, HyDE).
3
4pub mod bm25;
5pub mod fusion;
6
7use crate::embed::Embedder;
8use crate::llm::ChatModel;
9use crate::model::{RetrievalMode, Scored};
10use crate::store::VectorStore;
11use crate::{RagError, Result};
12use std::sync::Arc;
13
14/// Orchestrates the retrieval modes over a store + embedder (+ optional LLM).
15#[derive(Clone)]
16pub struct Retriever {
17    store: Arc<dyn VectorStore>,
18    embedder: Arc<dyn Embedder>,
19    /// Required by the Multi-Query and HyDE modes; `None` disables them.
20    chat: Option<Arc<dyn ChatModel>>,
21    rrf_k: f32,
22    multiquery_n: usize,
23}
24
25impl Retriever {
26    /// Build a retriever. Pass `chat = None` to run without an LLM (vector/bm25/hybrid only).
27    pub fn new(
28        store: Arc<dyn VectorStore>,
29        embedder: Arc<dyn Embedder>,
30        chat: Option<Arc<dyn ChatModel>>,
31    ) -> Self {
32        Retriever {
33            store,
34            embedder,
35            chat,
36            rrf_k: fusion::DEFAULT_RRF_K,
37            multiquery_n: 4,
38        }
39    }
40
41    /// Override the RRF constant (default 60).
42    pub fn with_rrf_k(mut self, k: f32) -> Self {
43        self.rrf_k = k;
44        self
45    }
46
47    /// Override the number of Multi-Query rewrites (default 4).
48    pub fn with_multiquery_n(mut self, n: usize) -> Self {
49        self.multiquery_n = n.max(1);
50        self
51    }
52
53    /// Retrieve the top `k` chunks for `query` using `mode`.
54    pub async fn retrieve(
55        &self,
56        mode: RetrievalMode,
57        query: &str,
58        k: usize,
59    ) -> Result<Vec<Scored>> {
60        match mode {
61            RetrievalMode::Vector => self.vector(query, k).await,
62            RetrievalMode::Bm25 => self.bm25(query, k).await,
63            RetrievalMode::Hybrid => self.hybrid(query, k).await,
64            RetrievalMode::MultiQuery => self.multi_query(query, k).await,
65            RetrievalMode::Hyde => self.hyde(query, k).await,
66        }
67    }
68
69    /// Dense vector search.
70    pub async fn vector(&self, query: &str, k: usize) -> Result<Vec<Scored>> {
71        let emb = self.embedder.embed_one(query).await?;
72        self.store.vector_search(&emb, k).await
73    }
74
75    /// Sparse BM25 keyword search over the whole chunk corpus.
76    pub async fn bm25(&self, query: &str, k: usize) -> Result<Vec<Scored>> {
77        let chunks = self.store.all_chunks().await?;
78        let index = bm25::Bm25Index::build(chunks);
79        Ok(index.search(query, k))
80    }
81
82    /// Hybrid RAG: fuse dense + sparse results with RRF. Over-fetches each arm so
83    /// fusion has depth to work with.
84    pub async fn hybrid(&self, query: &str, k: usize) -> Result<Vec<Scored>> {
85        let depth = (k * 4).max(20);
86        let vec_hits = self.vector(query, depth).await?;
87        let bm_hits = self.bm25(query, depth).await?;
88        Ok(fusion::rrf(&[vec_hits, bm_hits], self.rrf_k, k))
89    }
90
91    /// Multi-Query (Fusion) RAG: the LLM rewrites the question into several diverse
92    /// queries; each is retrieved (hybrid) and the results are fused with RRF.
93    pub async fn multi_query(&self, query: &str, k: usize) -> Result<Vec<Scored>> {
94        let chat = self.require_chat()?;
95        let system = "You rewrite a user's question into diverse search queries that \
96                      surface relevant documents. Output only the queries, one per line, \
97                      with no numbering or commentary.";
98        let user = format!(
99            "Rewrite this question into {} diverse search queries:\n\n{}",
100            self.multiquery_n, query
101        );
102        let raw = chat.ask(system, &user).await?;
103        let mut queries: Vec<String> = raw
104            .lines()
105            .map(|l| {
106                l.trim()
107                    .trim_start_matches(|c: char| {
108                        c.is_ascii_digit() || c == '.' || c == '-' || c == ')'
109                    })
110                    .trim()
111            })
112            .filter(|l| !l.is_empty())
113            .map(|l| l.to_string())
114            .take(self.multiquery_n)
115            .collect();
116        // Always include the original query.
117        queries.push(query.to_string());
118
119        let depth = (k * 4).max(20);
120        let mut rankings = Vec::with_capacity(queries.len());
121        for q in &queries {
122            rankings.push(self.hybrid(q, depth).await?);
123        }
124        Ok(fusion::rrf(&rankings, self.rrf_k, k))
125    }
126
127    /// HyDE: the LLM writes a hypothetical answer; its embedding drives the search.
128    pub async fn hyde(&self, query: &str, k: usize) -> Result<Vec<Scored>> {
129        let chat = self.require_chat()?;
130        let system = "You are helping a search system. Write a short, factual passage \
131                      (2-4 sentences) that could plausibly answer the user's question, as \
132                      if quoted from a relevant document. Do not hedge or mention that it \
133                      is hypothetical.";
134        let hypothetical = chat.ask(system, query).await?;
135        // Search with the hypothetical document's embedding; fall back to the raw
136        // query text if the model returned nothing usable.
137        let search_text = if hypothetical.trim().is_empty() {
138            query
139        } else {
140            &hypothetical
141        };
142        let emb = self.embedder.embed_one(search_text).await?;
143        self.store.vector_search(&emb, k).await
144    }
145
146    fn require_chat(&self) -> Result<&Arc<dyn ChatModel>> {
147        self.chat.as_ref().ok_or_else(|| {
148            RagError::Llm("this retrieval mode needs an LLM; set OPENROUTER_API_KEY".into())
149        })
150    }
151}
152
153#[cfg(test)]
154mod tests {
155    use super::*;
156    use crate::embed::HashEmbedder;
157    use crate::llm::Message;
158    use crate::model::{Chunk, Document};
159    use crate::store::memory::MemoryStore;
160    use async_trait::async_trait;
161
162    async fn seeded_store() -> Arc<dyn VectorStore> {
163        let store = Arc::new(MemoryStore::new());
164        let embedder = HashEmbedder::new(512);
165        let doc = Document::new("mem://t", "T", "h");
166        store.upsert_document(&doc).await.unwrap();
167        let texts = [
168            "postgres vector database stores embeddings for semantic search",
169            "a banana smoothie recipe with yogurt and honey",
170            "rust async runtime tokio spawns tasks on a thread pool",
171        ];
172        let mut chunks = Vec::new();
173        for (i, t) in texts.iter().enumerate() {
174            let mut c = Chunk::new(&doc.id, i as i64, *t, 0);
175            c.embedding = Some(
176                crate::embed::Embedder::embed(&embedder, &[t.to_string()])
177                    .await
178                    .unwrap()
179                    .pop()
180                    .unwrap(),
181            );
182            chunks.push(c);
183        }
184        store.insert_chunks(&chunks).await.unwrap();
185        store
186    }
187
188    #[tokio::test]
189    async fn vector_bm25_hybrid_find_relevant_chunk() {
190        let store = seeded_store().await;
191        let embedder: Arc<dyn Embedder> = Arc::new(HashEmbedder::new(512));
192        let r = Retriever::new(store, embedder, None);
193
194        for mode in RetrievalMode::OFFLINE {
195            let hits = r
196                .retrieve(mode, "semantic search vector database", 3)
197                .await
198                .unwrap();
199            assert!(!hits.is_empty(), "{mode} returned nothing");
200            assert!(
201                hits[0].chunk.text.contains("vector database"),
202                "{mode} ranked the wrong chunk first: {}",
203                hits[0].chunk.text
204            );
205        }
206    }
207
208    struct MockChat;
209    #[async_trait]
210    impl ChatModel for MockChat {
211        async fn complete(&self, messages: &[Message]) -> Result<String> {
212            let user = messages.last().map(|m| m.content.as_str()).unwrap_or("");
213            if user.contains("Rewrite") {
214                Ok("vector database\nsemantic search embeddings\npostgres storage".into())
215            } else {
216                // HyDE hypothetical answer.
217                Ok("A vector database stores embeddings and performs semantic search.".into())
218            }
219        }
220    }
221
222    #[tokio::test]
223    async fn multiquery_and_hyde_use_the_llm() {
224        let store = seeded_store().await;
225        let embedder: Arc<dyn Embedder> = Arc::new(HashEmbedder::new(512));
226        let chat: Arc<dyn ChatModel> = Arc::new(MockChat);
227        let r = Retriever::new(store, embedder, Some(chat));
228
229        let mq = r
230            .retrieve(RetrievalMode::MultiQuery, "how are embeddings stored?", 3)
231            .await
232            .unwrap();
233        assert!(mq.iter().any(|h| h.chunk.text.contains("vector database")));
234
235        let hyde = r
236            .retrieve(RetrievalMode::Hyde, "how are embeddings stored?", 3)
237            .await
238            .unwrap();
239        assert!(hyde
240            .iter()
241            .any(|h| h.chunk.text.contains("vector database")));
242    }
243
244    #[tokio::test]
245    async fn llm_modes_error_without_chat() {
246        let store = seeded_store().await;
247        let embedder: Arc<dyn Embedder> = Arc::new(HashEmbedder::new(512));
248        let r = Retriever::new(store, embedder, None);
249        assert!(r.retrieve(RetrievalMode::Hyde, "q", 3).await.is_err());
250    }
251}