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    /// 无结果
22    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/// 检索器 trait
44#[async_trait]
45pub trait RetrieverTrait: Send + Sync {
46    /// 检索相关文档
47    ///
48    /// # 参数
49    /// * `query` - 查询文本
50    /// * `k` - 返回的文档数量
51    ///
52    /// # 返回
53    /// 相关文档列表
54    async fn retrieve(&self, query: &str, k: usize) -> Result<Vec<Document>, RetrieverError>;
55
56    /// 检索相关文档(带分数)
57    async fn retrieve_with_scores(
58        &self,
59        query: &str,
60        k: usize,
61    ) -> Result<Vec<SearchResult>, RetrieverError>;
62
63    /// 添加文档
64    async fn add_documents(&self, documents: Vec<Document>) -> Result<(), RetrieverError>;
65}
66
67/// 基于相似度的检索器
68pub struct SimilarityRetriever {
69    /// 向量存储
70    store: Arc<dyn VectorStore>,
71
72    /// 嵌入模型
73    embeddings: Arc<dyn Embeddings>,
74}
75
76impl SimilarityRetriever {
77    /// 创建新的相似度检索器
78    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        // 生成查询向量
96        let query_embedding = self
97            .embeddings
98            .embed_query(query)
99            .await
100            .map_err(|e| RetrieverError::EmbeddingError(e.to_string()))?;
101
102        // 检索相似文档
103        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        // 生成文档嵌入
110        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        // 添加到存储
118        self.store.add_documents(documents, embeddings).await?;
119
120        Ok(())
121    }
122}
123
124/// 简化的 Retriever 类型别名(用于快速使用)
125pub 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    /// P0-1: 验证 BM25 / UnifiedHybrid 均可作为
136    /// `Arc<dyn RetrieverTrait>` 使用,完成 add + retrieve 全流程。
137    #[tokio::test]
138    async fn test_retriever_trait_object_hybrid_retrievers() {
139        // BM25Retriever 作为 trait object
140        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        // UnifiedHybridIndex 作为 trait object
150        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        // 添加文档
175        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        // 检索文档
185        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        // 结果应该包含分数
208        assert!(results[0].score >= -1.0 && results[0].score <= 1.0);
209    }
210}