1use async_trait::async_trait;
7use lc_embeddings::Embeddings;
8use lc_vector_stores::{Document, SearchResult, VectorStore, VectorStoreError};
9use std::sync::Arc;
10
11#[derive(Debug)]
13#[non_exhaustive]
14pub enum RetrieverError {
15 StoreError(VectorStoreError),
17
18 EmbeddingError(String),
20
21 LlmError(String),
23
24 InvalidFilter(String),
28
29 NoResults,
31}
32
33impl std::fmt::Display for RetrieverError {
34 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
35 match self {
36 RetrieverError::StoreError(e) => write!(f, "storage error: {}", e),
37 RetrieverError::EmbeddingError(msg) => write!(f, "embedding error: {}", msg),
38 RetrieverError::LlmError(msg) => write!(f, "LLM error: {}", msg),
39 RetrieverError::InvalidFilter(msg) => write!(f, "invalid filter: {}", msg),
40 RetrieverError::NoResults => write!(f, "no relevant documents found"),
41 }
42 }
43}
44
45impl std::error::Error for RetrieverError {}
46
47impl From<VectorStoreError> for RetrieverError {
48 fn from(e: VectorStoreError) -> Self {
49 RetrieverError::StoreError(e)
50 }
51}
52
53#[async_trait]
55pub trait RetrieverTrait: Send + Sync {
56 async fn retrieve(&self, query: &str, k: usize) -> Result<Vec<Document>, RetrieverError>;
65
66 async fn retrieve_with_scores(
68 &self,
69 query: &str,
70 k: usize,
71 ) -> Result<Vec<SearchResult>, RetrieverError>;
72
73 async fn add_documents(&self, documents: Vec<Document>) -> Result<(), RetrieverError>;
75}
76
77pub struct SimilarityRetriever {
79 store: Arc<dyn VectorStore>,
81
82 embeddings: Arc<dyn Embeddings>,
84}
85
86impl SimilarityRetriever {
87 pub fn new(store: Arc<dyn VectorStore>, embeddings: Arc<dyn Embeddings>) -> Self {
89 Self { store, embeddings }
90 }
91}
92
93#[async_trait]
94impl RetrieverTrait for SimilarityRetriever {
95 async fn retrieve(&self, query: &str, k: usize) -> Result<Vec<Document>, RetrieverError> {
96 let results = self.retrieve_with_scores(query, k).await?;
97 Ok(results.into_iter().map(|r| r.document).collect())
98 }
99
100 async fn retrieve_with_scores(
101 &self,
102 query: &str,
103 k: usize,
104 ) -> Result<Vec<SearchResult>, RetrieverError> {
105 let query_embedding = self
107 .embeddings
108 .embed_query(query)
109 .await
110 .map_err(|e| RetrieverError::EmbeddingError(e.to_string()))?;
111
112 let results = self.store.similarity_search(&query_embedding, k).await?;
114
115 Ok(results)
116 }
117
118 async fn add_documents(&self, documents: Vec<Document>) -> Result<(), RetrieverError> {
119 let texts: Vec<&str> = documents.iter().map(|d| d.content.as_str()).collect();
121 let embeddings = self
122 .embeddings
123 .embed_documents(&texts)
124 .await
125 .map_err(|e| RetrieverError::EmbeddingError(e.to_string()))?;
126
127 self.store.add_documents(documents, embeddings).await?;
129
130 Ok(())
131 }
132}
133
134pub type Retriever = SimilarityRetriever;
136
137#[cfg(test)]
138mod tests {
139 use super::*;
140 use crate::bm25::BM25Retriever;
141 use crate::unified_hybrid::UnifiedHybridIndex;
142 use lc_embeddings::MockEmbeddings;
143 use lc_vector_stores::InMemoryVectorStore;
144
145 #[tokio::test]
148 async fn test_retriever_trait_object_hybrid_retrievers() {
149 let bm25: Arc<dyn RetrieverTrait> = Arc::new(BM25Retriever::new());
151 bm25.add_documents(vec![Document::new(
152 "Rust is a systems programming language",
153 )])
154 .await
155 .unwrap();
156 let results = bm25.retrieve("systems", 1).await.unwrap();
157 assert!(!results.is_empty());
158
159 let embeddings = Arc::new(MockEmbeddings::new(128));
161 let vector_store: Arc<dyn VectorStore> = Arc::new(InMemoryVectorStore::new());
162 let unified: Arc<dyn RetrieverTrait> = Arc::new(UnifiedHybridIndex::new(
163 embeddings.clone(),
164 vector_store,
165 128,
166 ));
167 unified
168 .add_documents(vec![Document::new(
169 "Rust is a systems programming language",
170 )])
171 .await
172 .unwrap();
173 let results = unified.retrieve("systems", 1).await.unwrap();
174 assert!(!results.is_empty());
175 }
176
177 #[tokio::test]
178 async fn test_retriever() {
179 let store = Arc::new(InMemoryVectorStore::new());
180 let embeddings = Arc::new(MockEmbeddings::new(128));
181
182 let retriever = SimilarityRetriever::new(store.clone(), embeddings.clone());
183
184 let docs = vec![
186 Document::new("Rust is a systems programming language"),
187 Document::new("Python is a scripting language"),
188 Document::new("JavaScript is used for web development"),
189 ];
190
191 retriever.add_documents(docs).await.unwrap();
192 assert_eq!(store.count().await, 3);
193
194 let results = retriever.retrieve("programming language", 2).await.unwrap();
196 assert!(
197 !results.is_empty(),
198 "expected at least 1 result, got {}",
199 results.len()
200 );
201 }
202
203 #[tokio::test]
204 async fn test_retriever_with_scores() {
205 let store = Arc::new(InMemoryVectorStore::new());
206 let embeddings = Arc::new(MockEmbeddings::new(64));
207
208 let retriever = SimilarityRetriever::new(store, embeddings);
209
210 let docs = vec![Document::new("Document A"), Document::new("Document B")];
211
212 retriever.add_documents(docs).await.unwrap();
213
214 let results = retriever.retrieve_with_scores("query", 2).await.unwrap();
215 assert_eq!(results.len(), 2);
216
217 assert!(results[0].score >= -1.0 && results[0].score <= 1.0);
219 }
220}