Skip to main content

lc_rag/
retriever.rs

1// lc-rag/src/retriever.rs
2//! 检索器实现
3//!
4//! 提供基于相似度的文档检索功能。
5
6use async_trait::async_trait;
7use lc_embeddings::Embeddings;
8use lc_vector_stores::{Document, SearchResult, VectorStore, VectorStoreError};
9use std::sync::Arc;
10
11/// 检索器错误类型
12#[derive(Debug)]
13#[non_exhaustive]
14pub enum RetrieverError {
15    /// 向量存储错误
16    StoreError(VectorStoreError),
17
18    /// 嵌入错误
19    EmbeddingError(String),
20
21    /// LLM 拆解失败(调用失败 / 输出无法解析)。SelfQuery(S4) 用。
22    LlmError(String),
23
24    /// 无结果
25    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/// 检索器 trait
48#[async_trait]
49pub trait RetrieverTrait: Send + Sync {
50    /// 检索相关文档
51    ///
52    /// # 参数
53    /// * `query` - 查询文本
54    /// * `k` - 返回的文档数量
55    ///
56    /// # 返回
57    /// 相关文档列表
58    async fn retrieve(&self, query: &str, k: usize) -> Result<Vec<Document>, RetrieverError>;
59
60    /// 检索相关文档(带分数)
61    async fn retrieve_with_scores(
62        &self,
63        query: &str,
64        k: usize,
65    ) -> Result<Vec<SearchResult>, RetrieverError>;
66
67    /// 添加文档
68    async fn add_documents(&self, documents: Vec<Document>) -> Result<(), RetrieverError>;
69}
70
71/// 基于相似度的检索器
72pub struct SimilarityRetriever {
73    /// 向量存储
74    store: Arc<dyn VectorStore>,
75
76    /// 嵌入模型
77    embeddings: Arc<dyn Embeddings>,
78}
79
80impl SimilarityRetriever {
81    /// 创建新的相似度检索器
82    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        // 生成查询向量
100        let query_embedding = self
101            .embeddings
102            .embed_query(query)
103            .await
104            .map_err(|e| RetrieverError::EmbeddingError(e.to_string()))?;
105
106        // 检索相似文档
107        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        // 生成文档嵌入
114        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        // 添加到存储
122        self.store.add_documents(documents, embeddings).await?;
123
124        Ok(())
125    }
126}
127
128/// 简化的 Retriever 类型别名(用于快速使用)
129pub 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    /// P0-1: 验证 BM25 / UnifiedHybrid 均可作为
140    /// `Arc<dyn RetrieverTrait>` 使用,完成 add + retrieve 全流程。
141    #[tokio::test]
142    async fn test_retriever_trait_object_hybrid_retrievers() {
143        // BM25Retriever 作为 trait object
144        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        // UnifiedHybridIndex 作为 trait object
154        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        // 添加文档
179        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        // 检索文档
189        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        // 结果应该包含分数
212        assert!(results[0].score >= -1.0 && results[0].score <= 1.0);
213    }
214}