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