1pub 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#[derive(Clone)]
16pub struct Retriever {
17 store: Arc<dyn VectorStore>,
18 embedder: Arc<dyn Embedder>,
19 chat: Option<Arc<dyn ChatModel>>,
21 rrf_k: f32,
22 multiquery_n: usize,
23}
24
25impl Retriever {
26 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 pub fn with_rrf_k(mut self, k: f32) -> Self {
43 self.rrf_k = k;
44 self
45 }
46
47 pub fn with_multiquery_n(mut self, n: usize) -> Self {
49 self.multiquery_n = n.max(1);
50 self
51 }
52
53 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 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 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 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 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 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 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 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 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}