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