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)]
13pub enum RetrieverError {
14 StoreError(VectorStoreError),
16
17 EmbeddingError(String),
19
20 NoResults,
22}
23
24impl std::fmt::Display for RetrieverError {
25 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
26 match self {
27 RetrieverError::StoreError(e) => write!(f, "存储错误: {}", e),
28 RetrieverError::EmbeddingError(msg) => write!(f, "嵌入错误: {}", msg),
29 RetrieverError::NoResults => write!(f, "没有找到相关文档"),
30 }
31 }
32}
33
34impl std::error::Error for RetrieverError {}
35
36impl From<VectorStoreError> for RetrieverError {
37 fn from(e: VectorStoreError) -> Self {
38 RetrieverError::StoreError(e)
39 }
40}
41
42#[async_trait]
44pub trait RetrieverTrait: Send + Sync {
45 async fn retrieve(&self, query: &str, k: usize) -> Result<Vec<Document>, RetrieverError>;
54
55 async fn retrieve_with_scores(
57 &self,
58 query: &str,
59 k: usize,
60 ) -> Result<Vec<SearchResult>, RetrieverError>;
61
62 async fn add_documents(&self, documents: Vec<Document>) -> Result<(), RetrieverError>;
64}
65
66pub struct SimilarityRetriever {
68 store: Arc<dyn VectorStore>,
70
71 embeddings: Arc<dyn Embeddings>,
73}
74
75impl SimilarityRetriever {
76 pub fn new(store: Arc<dyn VectorStore>, embeddings: Arc<dyn Embeddings>) -> Self {
78 Self { store, embeddings }
79 }
80}
81
82#[async_trait]
83impl RetrieverTrait for SimilarityRetriever {
84 async fn retrieve(&self, query: &str, k: usize) -> Result<Vec<Document>, RetrieverError> {
85 let results = self.retrieve_with_scores(query, k).await?;
86 Ok(results.into_iter().map(|r| r.document).collect())
87 }
88
89 async fn retrieve_with_scores(
90 &self,
91 query: &str,
92 k: usize,
93 ) -> Result<Vec<SearchResult>, RetrieverError> {
94 let query_embedding = self
96 .embeddings
97 .embed_query(query)
98 .await
99 .map_err(|e| RetrieverError::EmbeddingError(e.to_string()))?;
100
101 let results = self.store.similarity_search(&query_embedding, k).await?;
103
104 Ok(results)
105 }
106
107 async fn add_documents(&self, documents: Vec<Document>) -> Result<(), RetrieverError> {
108 let texts: Vec<&str> = documents.iter().map(|d| d.content.as_str()).collect();
110 let embeddings = self
111 .embeddings
112 .embed_documents(&texts)
113 .await
114 .map_err(|e| RetrieverError::EmbeddingError(e.to_string()))?;
115
116 self.store.add_documents(documents, embeddings).await?;
118
119 Ok(())
120 }
121}
122
123pub type Retriever = SimilarityRetriever;
125
126#[cfg(test)]
127mod tests {
128 use super::*;
129 use crate::bm25::{AutoMergingConfig, BM25Retriever, ChunkedBM25Retriever};
130 use crate::chunked_hybrid::ChunkedHybridRetriever;
131 use crate::unified_hybrid::UnifiedHybridIndex;
132 use lc_embeddings::MockEmbeddings;
133 use lc_vector_stores::ChunkedDocumentStore;
134 use lc_vector_stores::InMemoryVectorStore;
135
136 #[allow(deprecated)] #[tokio::test]
140 async fn test_retriever_trait_object_hybrid_retrievers() {
141 let bm25: Arc<dyn RetrieverTrait> = Arc::new(BM25Retriever::new());
143 bm25.add_documents(vec![Document::new(
144 "Rust is a systems programming language",
145 )])
146 .await
147 .unwrap();
148 let results = bm25.retrieve("systems", 1).await.unwrap();
149 assert!(!results.is_empty());
150
151 let embeddings = Arc::new(MockEmbeddings::new(128));
153 let unified: Arc<dyn RetrieverTrait> =
154 Arc::new(UnifiedHybridIndex::new(embeddings.clone(), 128));
155 unified
156 .add_documents(vec![Document::new(
157 "Rust is a systems programming language",
158 )])
159 .await
160 .unwrap();
161 let results = unified.retrieve("systems", 1).await.unwrap();
162 assert!(!results.is_empty());
163
164 let store = Arc::new(ChunkedDocumentStore::new());
166 let bm25_retriever =
167 ChunkedBM25Retriever::with_config(store.clone(), AutoMergingConfig::new());
168 let chunked: Arc<dyn RetrieverTrait> = Arc::new(ChunkedHybridRetriever::new(
169 bm25_retriever,
170 store,
171 embeddings,
172 ));
173 chunked
174 .add_documents(vec![Document::new(
175 "Rust is a systems programming language",
176 )])
177 .await
178 .unwrap();
179 let results = chunked.retrieve("systems", 1).await.unwrap();
180 assert!(!results.is_empty());
181 }
182
183 #[tokio::test]
184 async fn test_retriever() {
185 let store = Arc::new(InMemoryVectorStore::new());
186 let embeddings = Arc::new(MockEmbeddings::new(128));
187
188 let retriever = SimilarityRetriever::new(store.clone(), embeddings.clone());
189
190 let docs = vec![
192 Document::new("Rust is a systems programming language"),
193 Document::new("Python is a scripting language"),
194 Document::new("JavaScript is used for web development"),
195 ];
196
197 retriever.add_documents(docs).await.unwrap();
198 assert_eq!(store.count().await, 3);
199
200 let results = retriever.retrieve("programming language", 2).await.unwrap();
202 assert!(
203 !results.is_empty(),
204 "expected at least 1 result, got {}",
205 results.len()
206 );
207 }
208
209 #[tokio::test]
210 async fn test_retriever_with_scores() {
211 let store = Arc::new(InMemoryVectorStore::new());
212 let embeddings = Arc::new(MockEmbeddings::new(64));
213
214 let retriever = SimilarityRetriever::new(store, embeddings);
215
216 let docs = vec![Document::new("Document A"), Document::new("Document B")];
217
218 retriever.add_documents(docs).await.unwrap();
219
220 let results = retriever.retrieve_with_scores("query", 2).await.unwrap();
221 assert_eq!(results.len(), 2);
222
223 assert!(results[0].score >= -1.0 && results[0].score <= 1.0);
225 }
226}