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)]
13pub enum RetrieverError {
14    /// 向量存储错误
15    StoreError(VectorStoreError),
16
17    /// 嵌入错误
18    EmbeddingError(String),
19
20    /// 无结果
21    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/// 检索器 trait
43#[async_trait]
44pub trait RetrieverTrait: Send + Sync {
45    /// 检索相关文档
46    ///
47    /// # 参数
48    /// * `query` - 查询文本
49    /// * `k` - 返回的文档数量
50    ///
51    /// # 返回
52    /// 相关文档列表
53    async fn retrieve(&self, query: &str, k: usize) -> Result<Vec<Document>, RetrieverError>;
54
55    /// 检索相关文档(带分数)
56    async fn retrieve_with_scores(
57        &self,
58        query: &str,
59        k: usize,
60    ) -> Result<Vec<SearchResult>, RetrieverError>;
61
62    /// 添加文档
63    async fn add_documents(&self, documents: Vec<Document>) -> Result<(), RetrieverError>;
64}
65
66/// 基于相似度的检索器
67pub struct SimilarityRetriever {
68    /// 向量存储
69    store: Arc<dyn VectorStore>,
70
71    /// 嵌入模型
72    embeddings: Arc<dyn Embeddings>,
73}
74
75impl SimilarityRetriever {
76    /// 创建新的相似度检索器
77    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        // 生成查询向量
95        let query_embedding = self
96            .embeddings
97            .embed_query(query)
98            .await
99            .map_err(|e| RetrieverError::EmbeddingError(e.to_string()))?;
100
101        // 检索相似文档
102        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        // 生成文档嵌入
109        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        // 添加到存储
117        self.store.add_documents(documents, embeddings).await?;
118
119        Ok(())
120    }
121}
122
123/// 简化的 Retriever 类型别名(用于快速使用)
124pub type Retriever = SimilarityRetriever;
125
126#[cfg(test)]
127mod tests {
128    use super::*;
129    use crate::bm25::BM25Retriever;
130    use crate::unified_hybrid::UnifiedHybridIndex;
131    use lc_embeddings::MockEmbeddings;
132    use lc_vector_stores::InMemoryVectorStore;
133
134    /// P0-1: 验证 BM25 / UnifiedHybrid 均可作为
135    /// `Arc<dyn RetrieverTrait>` 使用,完成 add + retrieve 全流程。
136    #[tokio::test]
137    async fn test_retriever_trait_object_hybrid_retrievers() {
138        // BM25Retriever 作为 trait object
139        let bm25: Arc<dyn RetrieverTrait> = Arc::new(BM25Retriever::new());
140        bm25.add_documents(vec![Document::new(
141            "Rust is a systems programming language",
142        )])
143        .await
144        .unwrap();
145        let results = bm25.retrieve("systems", 1).await.unwrap();
146        assert!(!results.is_empty());
147
148        // UnifiedHybridIndex 作为 trait object
149        let embeddings = Arc::new(MockEmbeddings::new(128));
150        let vector_store: Arc<dyn VectorStore> = Arc::new(InMemoryVectorStore::new());
151        let unified: Arc<dyn RetrieverTrait> =
152            Arc::new(UnifiedHybridIndex::new(embeddings.clone(), vector_store, 128));
153        unified
154            .add_documents(vec![Document::new(
155                "Rust is a systems programming language",
156            )])
157            .await
158            .unwrap();
159        let results = unified.retrieve("systems", 1).await.unwrap();
160        assert!(!results.is_empty());
161    }
162
163    #[tokio::test]
164    async fn test_retriever() {
165        let store = Arc::new(InMemoryVectorStore::new());
166        let embeddings = Arc::new(MockEmbeddings::new(128));
167
168        let retriever = SimilarityRetriever::new(store.clone(), embeddings.clone());
169
170        // 添加文档
171        let docs = vec![
172            Document::new("Rust is a systems programming language"),
173            Document::new("Python is a scripting language"),
174            Document::new("JavaScript is used for web development"),
175        ];
176
177        retriever.add_documents(docs).await.unwrap();
178        assert_eq!(store.count().await, 3);
179
180        // 检索文档
181        let results = retriever.retrieve("programming language", 2).await.unwrap();
182        assert!(
183            !results.is_empty(),
184            "expected at least 1 result, got {}",
185            results.len()
186        );
187    }
188
189    #[tokio::test]
190    async fn test_retriever_with_scores() {
191        let store = Arc::new(InMemoryVectorStore::new());
192        let embeddings = Arc::new(MockEmbeddings::new(64));
193
194        let retriever = SimilarityRetriever::new(store, embeddings);
195
196        let docs = vec![Document::new("Document A"), Document::new("Document B")];
197
198        retriever.add_documents(docs).await.unwrap();
199
200        let results = retriever.retrieve_with_scores("query", 2).await.unwrap();
201        assert_eq!(results.len(), 2);
202
203        // 结果应该包含分数
204        assert!(results[0].score >= -1.0 && results[0].score <= 1.0);
205    }
206}