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    /// 过滤器引用了 `allowed_attributes` 白名单外的字段(SelfQuery)。
25    /// 显式报错,绝不静默丢弃后回落无过滤检索——那会让本该被过滤排除的
26    /// 数据被返回(数据面过曝)。
27    InvalidFilter(String),
28
29    /// 无结果
30    NoResults,
31}
32
33impl std::fmt::Display for RetrieverError {
34    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
35        match self {
36            RetrieverError::StoreError(e) => write!(f, "storage error: {}", e),
37            RetrieverError::EmbeddingError(msg) => write!(f, "embedding error: {}", msg),
38            RetrieverError::LlmError(msg) => write!(f, "LLM error: {}", msg),
39            RetrieverError::InvalidFilter(msg) => write!(f, "invalid filter: {}", msg),
40            RetrieverError::NoResults => write!(f, "no relevant documents found"),
41        }
42    }
43}
44
45impl std::error::Error for RetrieverError {}
46
47impl From<VectorStoreError> for RetrieverError {
48    fn from(e: VectorStoreError) -> Self {
49        RetrieverError::StoreError(e)
50    }
51}
52
53/// 检索器 trait
54#[async_trait]
55pub trait RetrieverTrait: Send + Sync {
56    /// 检索相关文档
57    ///
58    /// # 参数
59    /// * `query` - 查询文本
60    /// * `k` - 返回的文档数量
61    ///
62    /// # 返回
63    /// 相关文档列表
64    async fn retrieve(&self, query: &str, k: usize) -> Result<Vec<Document>, RetrieverError>;
65
66    /// 检索相关文档(带分数)
67    async fn retrieve_with_scores(
68        &self,
69        query: &str,
70        k: usize,
71    ) -> Result<Vec<SearchResult>, RetrieverError>;
72
73    /// 添加文档
74    async fn add_documents(&self, documents: Vec<Document>) -> Result<(), RetrieverError>;
75}
76
77/// 基于相似度的检索器
78pub struct SimilarityRetriever {
79    /// 向量存储
80    store: Arc<dyn VectorStore>,
81
82    /// 嵌入模型
83    embeddings: Arc<dyn Embeddings>,
84}
85
86impl SimilarityRetriever {
87    /// 创建新的相似度检索器
88    pub fn new(store: Arc<dyn VectorStore>, embeddings: Arc<dyn Embeddings>) -> Self {
89        Self { store, embeddings }
90    }
91}
92
93#[async_trait]
94impl RetrieverTrait for SimilarityRetriever {
95    async fn retrieve(&self, query: &str, k: usize) -> Result<Vec<Document>, RetrieverError> {
96        let results = self.retrieve_with_scores(query, k).await?;
97        Ok(results.into_iter().map(|r| r.document).collect())
98    }
99
100    async fn retrieve_with_scores(
101        &self,
102        query: &str,
103        k: usize,
104    ) -> Result<Vec<SearchResult>, RetrieverError> {
105        // 生成查询向量
106        let query_embedding = self
107            .embeddings
108            .embed_query(query)
109            .await
110            .map_err(|e| RetrieverError::EmbeddingError(e.to_string()))?;
111
112        // 检索相似文档
113        let results = self.store.similarity_search(&query_embedding, k).await?;
114
115        Ok(results)
116    }
117
118    async fn add_documents(&self, documents: Vec<Document>) -> Result<(), RetrieverError> {
119        // 生成文档嵌入
120        let texts: Vec<&str> = documents.iter().map(|d| d.content.as_str()).collect();
121        let embeddings = self
122            .embeddings
123            .embed_documents(&texts)
124            .await
125            .map_err(|e| RetrieverError::EmbeddingError(e.to_string()))?;
126
127        // 添加到存储
128        self.store.add_documents(documents, embeddings).await?;
129
130        Ok(())
131    }
132}
133
134/// 简化的 Retriever 类型别名(用于快速使用)
135pub type Retriever = SimilarityRetriever;
136
137#[cfg(test)]
138mod tests {
139    use super::*;
140    use crate::bm25::BM25Retriever;
141    use crate::unified_hybrid::UnifiedHybridIndex;
142    use lc_embeddings::MockEmbeddings;
143    use lc_vector_stores::InMemoryVectorStore;
144
145    /// P0-1: 验证 BM25 / UnifiedHybrid 均可作为
146    /// `Arc<dyn RetrieverTrait>` 使用,完成 add + retrieve 全流程。
147    #[tokio::test]
148    async fn test_retriever_trait_object_hybrid_retrievers() {
149        // BM25Retriever 作为 trait object
150        let bm25: Arc<dyn RetrieverTrait> = Arc::new(BM25Retriever::new());
151        bm25.add_documents(vec![Document::new(
152            "Rust is a systems programming language",
153        )])
154        .await
155        .unwrap();
156        let results = bm25.retrieve("systems", 1).await.unwrap();
157        assert!(!results.is_empty());
158
159        // UnifiedHybridIndex 作为 trait object
160        let embeddings = Arc::new(MockEmbeddings::new(128));
161        let vector_store: Arc<dyn VectorStore> = Arc::new(InMemoryVectorStore::new());
162        let unified: Arc<dyn RetrieverTrait> = Arc::new(UnifiedHybridIndex::new(
163            embeddings.clone(),
164            vector_store,
165            128,
166        ));
167        unified
168            .add_documents(vec![Document::new(
169                "Rust is a systems programming language",
170            )])
171            .await
172            .unwrap();
173        let results = unified.retrieve("systems", 1).await.unwrap();
174        assert!(!results.is_empty());
175    }
176
177    #[tokio::test]
178    async fn test_retriever() {
179        let store = Arc::new(InMemoryVectorStore::new());
180        let embeddings = Arc::new(MockEmbeddings::new(128));
181
182        let retriever = SimilarityRetriever::new(store.clone(), embeddings.clone());
183
184        // 添加文档
185        let docs = vec![
186            Document::new("Rust is a systems programming language"),
187            Document::new("Python is a scripting language"),
188            Document::new("JavaScript is used for web development"),
189        ];
190
191        retriever.add_documents(docs).await.unwrap();
192        assert_eq!(store.count().await, 3);
193
194        // 检索文档
195        let results = retriever.retrieve("programming language", 2).await.unwrap();
196        assert!(
197            !results.is_empty(),
198            "expected at least 1 result, got {}",
199            results.len()
200        );
201    }
202
203    #[tokio::test]
204    async fn test_retriever_with_scores() {
205        let store = Arc::new(InMemoryVectorStore::new());
206        let embeddings = Arc::new(MockEmbeddings::new(64));
207
208        let retriever = SimilarityRetriever::new(store, embeddings);
209
210        let docs = vec![Document::new("Document A"), Document::new("Document B")];
211
212        retriever.add_documents(docs).await.unwrap();
213
214        let results = retriever.retrieve_with_scores("query", 2).await.unwrap();
215        assert_eq!(results.len(), 2);
216
217        // 结果应该包含分数
218        assert!(results[0].score >= -1.0 && results[0].score <= 1.0);
219    }
220}