Skip to main content

lc_rag/
pipeline.rs

1// lc-rag/src/pipeline.rs
2//! RAGPipeline & RAGPipelineBuilder — 一行搞定 RAG 管线
3//!
4//! 提供流畅的 Builder API,将 LLM + Embeddings + VectorStore + Retriever
5//! 组装成完整的 RAG 管线。
6//!
7//! # Example
8//!
9//! ```ignore
10//! let rag = RAGPipelineBuilder::new()
11//!     .llm(OpenAIChat::new(OpenAIConfig::new("sk-...")))
12//!     .embeddings(OpenAIEmbeddings::new(config)?)
13//!     .vector_store(InMemoryVectorStore::new())
14//!     .build()?;
15//!
16//! rag.index_documents(docs).await?;
17//! let answer = rag.query("What is RustB?").await?;
18//! ```
19
20use lc_core::language_models::BaseChatModel;
21use lc_embeddings::Embeddings;
22use lc_providers::ProviderError;
23use lc_schema::Message;
24use lc_vector_stores::{Document, VectorStore, VectorStoreError};
25
26use crate::retriever::{RetrieverError, RetrieverTrait, SimilarityRetriever};
27
28use std::sync::Arc;
29
30/// RAG Pipeline — 切分 + 嵌入 + 存储 + 检索 + 生成
31///
32/// 将 LLM 与一个 `RetrieverTrait` 实现(BM25、向量相似度、混合检索等)组装成
33/// 完整的 RAG 管线,提供 `index_documents()`、`query()`、`query_with_sources()`
34/// 三个核心方法。
35///
36/// P0-2: 检索路径收敛到 `Arc<dyn RetrieverTrait>`,不再直接依赖
37/// `Embeddings + VectorStore`,可无缝切换任意检索器。
38pub struct RAGPipeline {
39    llm: Arc<dyn BaseChatModel<Error = ProviderError> + Send + Sync>,
40    retriever: Arc<dyn RetrieverTrait>,
41    /// 检索文档数量
42    retrieve_k: usize,
43    /// 系统提示词
44    system_prompt: String,
45}
46
47impl std::fmt::Debug for RAGPipeline {
48    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
49        f.debug_struct("RAGPipeline")
50            .field("model_name", &self.llm.model_name())
51            .field("retrieve_k", &self.retrieve_k)
52            .field("system_prompt", &self.system_prompt)
53            .finish()
54    }
55}
56
57impl RAGPipeline {
58    /// 索引文档
59    ///
60    /// P0-2: 委托给 `RetrieverTrait::add_documents`(嵌入 + 存储由检索器内部完成)。
61    pub async fn index_documents(&self, documents: Vec<Document>) -> Result<(), RetrieverError> {
62        self.retriever.add_documents(documents).await
63    }
64
65    /// 查询:检索 + 生成回答
66    ///
67    /// 1. 将问题嵌入向量
68    /// 2. 从 VectorStore 检索相似文档
69    /// 3. 将检索结果作为上下文,让 LLM 生成回答
70    pub async fn query(&self, question: &str) -> Result<String, RetrieverError> {
71        let result = self.query_with_sources(question).await?;
72        Ok(result.answer)
73    }
74
75    /// 查询并返回来源文档
76    ///
77    /// 返回生成的答案和检索到的源文档列表。
78    pub async fn query_with_sources(
79        &self,
80        question: &str,
81    ) -> Result<RAGQueryResult, RetrieverError> {
82        // 1. 检索相关文档(P0-2: 委托给 RetrieverTrait)
83        let search_results = self
84            .retriever
85            .retrieve_with_scores(question, self.retrieve_k)
86            .await?;
87
88        let sources: Vec<Document> = search_results.iter().map(|r| r.document.clone()).collect();
89
90        // 3. 构建上下文
91        let context = if sources.is_empty() {
92            "No relevant documents found.".to_string()
93        } else {
94            sources
95                .iter()
96                .enumerate()
97                .map(|(i, doc)| format!("[{}] {}", i + 1, doc.page_content()))
98                .collect::<Vec<_>>()
99                .join("\n\n")
100        };
101
102        // 4. 生成回答
103        let messages = vec![
104            Message::system(format!(
105                "{}\n\nUse the following context to answer the question. If the context doesn't contain the answer, say so.",
106                self.system_prompt
107            )),
108            Message::human(format!("Context:\n{}\n\nQuestion: {}", context, question)),
109        ];
110
111        let llm_result = self
112            .llm
113            .chat(messages, None)
114            .await
115            .map_err(|e| RetrieverError::EmbeddingError(format!("LLM call failed: {}", e)))?;
116
117        Ok(RAGQueryResult {
118            answer: llm_result.content,
119            sources,
120        })
121    }
122}
123
124/// RAG 查询结果
125#[derive(Debug, Clone)]
126pub struct RAGQueryResult {
127    /// 生成的答案
128    pub answer: String,
129    /// 检索到的源文档
130    pub sources: Vec<Document>,
131}
132
133// ---------------------------------------------------------------------------
134// RAGPipelineBuilder
135// ---------------------------------------------------------------------------
136
137/// RAG Pipeline Builder — 流畅 API 创建 RAG 管线
138///
139/// # Example
140///
141/// ```ignore
142/// let rag = RAGPipelineBuilder::new()
143///     .llm(OpenAIChat::new(OpenAIConfig::new("sk-...")))
144///     .embeddings(OpenAIEmbeddings::new(config)?)
145///     .vector_store(InMemoryVectorStore::new())
146///     .build()?;
147/// ```
148pub struct RAGPipelineBuilder {
149    llm: Option<Arc<dyn BaseChatModel<Error = ProviderError> + Send + Sync>>,
150    embeddings: Option<Arc<dyn Embeddings + Send + Sync>>,
151    vector_store: Option<Arc<dyn VectorStore + Send + Sync>>,
152    /// P0-2: 显式传入的检索器(优先);未提供时用 embeddings + vector_store 构建
153    retriever: Option<Arc<dyn RetrieverTrait>>,
154    retrieve_k: usize,
155    system_prompt: Option<String>,
156}
157
158impl RAGPipelineBuilder {
159    /// 创建新的 RAGPipelineBuilder
160    pub fn new() -> Self {
161        Self {
162            llm: None,
163            embeddings: None,
164            vector_store: None,
165            retriever: None,
166            retrieve_k: 4,
167            system_prompt: None,
168        }
169    }
170
171    /// 设置 LLM(任何实现了 `BaseChatModel` 的类型)
172    pub fn llm<L>(mut self, llm: L) -> Self
173    where
174        L: BaseChatModel + Send + Sync + 'static,
175        L::Error: Into<ProviderError>,
176    {
177        self.llm = Some(lc_providers::wrap_chat_model(llm));
178        self
179    }
180
181    /// 设置 LLM(从已包装的 `Arc<dyn BaseChatModel>`)
182    pub fn llm_from_arc(
183        mut self,
184        llm: Arc<dyn BaseChatModel<Error = ProviderError> + Send + Sync>,
185    ) -> Self {
186        self.llm = Some(llm);
187        self
188    }
189
190    /// 设置 LLM(从 `LLMClient`)
191    pub fn llm_client(mut self, client: lc_providers::LLMClient) -> Self {
192        let provider_arc = client.into_inner();
193        self.llm = Some(provider_arc);
194        self
195    }
196
197    /// 设置 Embeddings
198    pub fn embeddings<E: Embeddings + Send + Sync + 'static>(mut self, embeddings: E) -> Self {
199        self.embeddings = Some(Arc::new(embeddings));
200        self
201    }
202
203    /// 设置 VectorStore
204    pub fn vector_store<V: VectorStore + Send + Sync + 'static>(mut self, store: V) -> Self {
205        self.vector_store = Some(Arc::new(store));
206        self
207    }
208
209    /// 设置自定义检索器(任何实现了 `RetrieverTrait` 的类型,
210    /// 如 BM25、UnifiedHybridIndex 等)
211    ///
212    /// P0-2: 显式检索器优先于 `.embeddings() + .vector_store()` 构建的
213    /// 相似度检索器。
214    pub fn retriever<R>(mut self, retriever: R) -> Self
215    where
216        R: RetrieverTrait + Send + Sync + 'static,
217    {
218        self.retriever = Some(Arc::new(retriever));
219        self
220    }
221
222    /// 设置检索器(从已包装的 `Arc<dyn RetrieverTrait>`)
223    pub fn retriever_from_arc(mut self, retriever: Arc<dyn RetrieverTrait>) -> Self {
224        self.retriever = Some(retriever);
225        self
226    }
227
228    /// 设置检索文档数量
229    pub fn retrieve_k(mut self, k: usize) -> Self {
230        self.retrieve_k = k;
231        self
232    }
233
234    /// 设置系统提示词
235    pub fn system(mut self, prompt: impl Into<String>) -> Self {
236        self.system_prompt = Some(prompt.into());
237        self
238    }
239
240    /// 构建 RAGPipeline
241    ///
242    /// # Errors
243    ///
244    /// 如果缺少 LLM、Embeddings 或 VectorStore,返回错误。
245    pub fn build(self) -> Result<RAGPipeline, RetrieverError> {
246        let llm = self.llm.ok_or_else(|| {
247            RetrieverError::EmbeddingError(
248                "RAGPipelineBuilder: LLM is required. Call .llm() first.".into(),
249            )
250        })?;
251
252        // P0-2: 优先使用显式检索器;否则回退到 embeddings + vector_store
253        // 构建 SimilarityRetriever,兼容旧用法。
254        let retriever = match self.retriever {
255            Some(r) => r,
256            None => {
257                let embeddings = self.embeddings.ok_or_else(|| {
258                    RetrieverError::EmbeddingError(
259                        "RAGPipelineBuilder: Embeddings is required (or use .retriever()). Call .embeddings() first."
260                            .into(),
261                    )
262                })?;
263
264                let vector_store = self.vector_store.ok_or_else(|| {
265                    RetrieverError::StoreError(VectorStoreError::StorageError(
266                        "RAGPipelineBuilder: VectorStore is required (or use .retriever()). Call .vector_store() first."
267                            .into(),
268                    ))
269                })?;
270
271                Arc::new(SimilarityRetriever::new(vector_store, embeddings))
272            }
273        };
274
275        Ok(RAGPipeline {
276            llm,
277            retriever,
278            retrieve_k: self.retrieve_k,
279            system_prompt: self.system_prompt.unwrap_or_else(|| {
280                "You are a helpful assistant that answers questions based on the provided context.".to_string()
281            }),
282        })
283    }
284}
285
286impl Default for RAGPipelineBuilder {
287    fn default() -> Self {
288        Self::new()
289    }
290}
291
292#[cfg(test)]
293mod tests {
294    use super::*;
295    use lc_embeddings::MockEmbeddings;
296    use lc_providers::{OpenAIChat, OpenAIConfig};
297    use lc_vector_stores::InMemoryVectorStore;
298
299    #[test]
300    fn test_builder_missing_llm() {
301        let result = RAGPipelineBuilder::new()
302            .embeddings(MockEmbeddings::new(3))
303            .vector_store(InMemoryVectorStore::new())
304            .build();
305
306        assert!(result.is_err());
307        assert!(result.unwrap_err().to_string().contains("LLM is required"));
308    }
309
310    #[test]
311    fn test_builder_missing_embeddings() {
312        let config = OpenAIConfig::new("test_key").with_base_url("http://localhost:8080/v1");
313        let result = RAGPipelineBuilder::new()
314            .llm(OpenAIChat::new(config))
315            .vector_store(InMemoryVectorStore::new())
316            .build();
317
318        assert!(result.is_err());
319        assert!(result
320            .unwrap_err()
321            .to_string()
322            .contains("Embeddings is required"));
323    }
324
325    #[test]
326    fn test_builder_missing_vector_store() {
327        let config = OpenAIConfig::new("test_key").with_base_url("http://localhost:8080/v1");
328        let result = RAGPipelineBuilder::new()
329            .llm(OpenAIChat::new(config))
330            .embeddings(MockEmbeddings::new(3))
331            .build();
332
333        assert!(result.is_err());
334        assert!(result
335            .unwrap_err()
336            .to_string()
337            .contains("VectorStore is required"));
338    }
339
340    #[test]
341    fn test_builder_success() {
342        let config = OpenAIConfig::new("test_key").with_base_url("http://localhost:8080/v1");
343        let result = RAGPipelineBuilder::new()
344            .llm(OpenAIChat::new(config))
345            .embeddings(MockEmbeddings::new(3))
346            .vector_store(InMemoryVectorStore::new())
347            .system("You are a test assistant.")
348            .retrieve_k(5)
349            .build();
350
351        assert!(result.is_ok());
352        let pipeline = result.unwrap();
353        assert_eq!(pipeline.retrieve_k, 5);
354        assert_eq!(pipeline.system_prompt, "You are a test assistant.");
355    }
356
357    #[test]
358    fn test_builder_default() {
359        let builder = RAGPipelineBuilder::default();
360        assert_eq!(builder.retrieve_k, 4);
361        assert!(builder.llm.is_none());
362    }
363
364    /// P0-2: 支持通过 `.retriever()` 注入自定义 `RetrieverTrait` 实现(BM25),
365    /// `index_documents` 委托给该检索器,无需 Embeddings/VectorStore。
366    #[tokio::test]
367    async fn test_builder_with_custom_retriever() {
368        use crate::bm25::BM25Retriever;
369
370        let config = OpenAIConfig::new("test_key").with_base_url("http://localhost:8080/v1");
371        let rag = RAGPipelineBuilder::new()
372            .llm(OpenAIChat::new(config))
373            .retriever(BM25Retriever::new())
374            .build()
375            .expect("build should succeed");
376
377        rag.index_documents(vec![Document::new("Rust is a systems language")])
378            .await
379            .expect("index_documents should delegate to BM25 retriever successfully");
380    }
381}