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