1pub 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#[derive(Clone)]
17pub struct Retriever {
18 store: Arc<dyn VectorStore>,
19 embedder: Arc<dyn Embedder>,
20 chat: Option<Arc<dyn ChatModel>>,
22 rrf_k: f32,
23 multiquery_n: usize,
24 bm25_params: Bm25Params,
25 bm25_cache: Arc<Bm25Cache>,
28}
29
30impl Retriever {
31 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 pub fn with_rrf_k(mut self, k: f32) -> Self {
50 self.rrf_k = k;
51 self
52 }
53
54 pub fn with_multiquery_n(mut self, n: usize) -> Self {
56 self.multiquery_n = n.max(1);
57 self
58 }
59
60 pub fn with_bm25_params(mut self, params: Bm25Params) -> Self {
62 self.bm25_params = params;
63 self
64 }
65
66 pub fn with_bm25_cache(mut self, cache: Arc<Bm25Cache>) -> Self {
70 self.bm25_cache = cache;
71 self
72 }
73
74 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 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 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 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 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 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 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 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 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}