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