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;
27
28use std::sync::Arc;
29
30/// RAG Pipeline — 切分 + 嵌入 + 存储 + 检索 + 生成
31///
32/// 将 LLM、Embeddings、VectorStore 组装成完整的 RAG 管线,
33/// 提供 `index_documents()`、`query()`、`query_with_sources()` 三个核心方法。
34pub struct RAGPipeline {
35    llm: Arc<dyn BaseChatModel<Error = ProviderError> + Send + Sync>,
36    embeddings: Arc<dyn Embeddings + Send + Sync>,
37    vector_store: Arc<dyn VectorStore + Send + Sync>,
38    /// 检索文档数量
39    retrieve_k: usize,
40    /// 系统提示词
41    system_prompt: String,
42}
43
44impl std::fmt::Debug for RAGPipeline {
45    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
46        f.debug_struct("RAGPipeline")
47            .field("model_name", &self.llm.model_name())
48            .field("retrieve_k", &self.retrieve_k)
49            .field("system_prompt", &self.system_prompt)
50            .finish()
51    }
52}
53
54impl RAGPipeline {
55    /// 索引文档:嵌入 + 存储
56    ///
57    /// 将文档列表嵌入向量并添加到 VectorStore。
58    pub async fn index_documents(&self, documents: Vec<Document>) -> Result<(), RetrieverError> {
59        let texts: Vec<&str> = documents.iter().map(|d| d.page_content()).collect();
60        let embeddings = self
61            .embeddings
62            .embed_documents(&texts)
63            .await
64            .map_err(|e| RetrieverError::EmbeddingError(e.to_string()))?;
65
66        self.vector_store
67            .add_documents(documents, embeddings)
68            .await
69            .map_err(RetrieverError::StoreError)?;
70
71        Ok(())
72    }
73
74    /// 查询:检索 + 生成回答
75    ///
76    /// 1. 将问题嵌入向量
77    /// 2. 从 VectorStore 检索相似文档
78    /// 3. 将检索结果作为上下文,让 LLM 生成回答
79    pub async fn query(&self, question: &str) -> Result<String, RetrieverError> {
80        let result = self.query_with_sources(question).await?;
81        Ok(result.answer)
82    }
83
84    /// 查询并返回来源文档
85    ///
86    /// 返回生成的答案和检索到的源文档列表。
87    pub async fn query_with_sources(
88        &self,
89        question: &str,
90    ) -> Result<RAGQueryResult, RetrieverError> {
91        // 1. 嵌入问题
92        let query_embedding = self
93            .embeddings
94            .embed_query(question)
95            .await
96            .map_err(|e| RetrieverError::EmbeddingError(e.to_string()))?;
97
98        // 2. 检索相似文档
99        let search_results = self
100            .vector_store
101            .similarity_search(&query_embedding, self.retrieve_k)
102            .await
103            .map_err(RetrieverError::StoreError)?;
104
105        let sources: Vec<Document> = search_results.iter().map(|r| r.document.clone()).collect();
106
107        // 3. 构建上下文
108        let context = if sources.is_empty() {
109            "No relevant documents found.".to_string()
110        } else {
111            sources
112                .iter()
113                .enumerate()
114                .map(|(i, doc)| format!("[{}] {}", i + 1, doc.page_content()))
115                .collect::<Vec<_>>()
116                .join("\n\n")
117        };
118
119        // 4. 生成回答
120        let messages = vec![
121            Message::system(format!(
122                "{}\n\nUse the following context to answer the question. If the context doesn't contain the answer, say so.",
123                self.system_prompt
124            )),
125            Message::human(format!("Context:\n{}\n\nQuestion: {}", context, question)),
126        ];
127
128        let llm_result = self
129            .llm
130            .chat(messages, None)
131            .await
132            .map_err(|e| RetrieverError::EmbeddingError(format!("LLM 调用失败: {}", e)))?;
133
134        Ok(RAGQueryResult {
135            answer: llm_result.content,
136            sources,
137        })
138    }
139}
140
141/// RAG 查询结果
142#[derive(Debug, Clone)]
143pub struct RAGQueryResult {
144    /// 生成的答案
145    pub answer: String,
146    /// 检索到的源文档
147    pub sources: Vec<Document>,
148}
149
150// ---------------------------------------------------------------------------
151// RAGPipelineBuilder
152// ---------------------------------------------------------------------------
153
154/// RAG Pipeline Builder — 流畅 API 创建 RAG 管线
155///
156/// # Example
157///
158/// ```ignore
159/// let rag = RAGPipelineBuilder::new()
160///     .llm(OpenAIChat::new(OpenAIConfig::new("sk-...")))
161///     .embeddings(OpenAIEmbeddings::new(config))
162///     .vector_store(InMemoryVectorStore::new())
163///     .build()?;
164/// ```
165pub struct RAGPipelineBuilder {
166    llm: Option<Arc<dyn BaseChatModel<Error = ProviderError> + Send + Sync>>,
167    embeddings: Option<Arc<dyn Embeddings + Send + Sync>>,
168    vector_store: Option<Arc<dyn VectorStore + Send + Sync>>,
169    retrieve_k: usize,
170    system_prompt: Option<String>,
171}
172
173impl RAGPipelineBuilder {
174    /// 创建新的 RAGPipelineBuilder
175    pub fn new() -> Self {
176        Self {
177            llm: None,
178            embeddings: None,
179            vector_store: None,
180            retrieve_k: 4,
181            system_prompt: None,
182        }
183    }
184
185    /// 设置 LLM(任何实现了 `BaseChatModel` 的类型)
186    pub fn llm<L>(mut self, llm: L) -> Self
187    where
188        L: BaseChatModel + Send + Sync + 'static,
189        L::Error: Into<ProviderError>,
190    {
191        self.llm = Some(lc_providers::wrap_chat_model(llm));
192        self
193    }
194
195    /// 设置 LLM(从已包装的 `Arc<dyn BaseChatModel>`)
196    pub fn llm_from_arc(
197        mut self,
198        llm: Arc<dyn BaseChatModel<Error = ProviderError> + Send + Sync>,
199    ) -> Self {
200        self.llm = Some(llm);
201        self
202    }
203
204    /// 设置 LLM(从 `LLMClient`)
205    pub fn llm_client(mut self, client: lc_providers::LLMClient) -> Self {
206        let provider_arc = client.into_inner();
207        self.llm = Some(provider_arc);
208        self
209    }
210
211    /// 设置 Embeddings
212    pub fn embeddings<E: Embeddings + Send + Sync + 'static>(mut self, embeddings: E) -> Self {
213        self.embeddings = Some(Arc::new(embeddings));
214        self
215    }
216
217    /// 设置 VectorStore
218    pub fn vector_store<V: VectorStore + Send + Sync + 'static>(mut self, store: V) -> Self {
219        self.vector_store = Some(Arc::new(store));
220        self
221    }
222
223    /// 设置检索文档数量
224    pub fn retrieve_k(mut self, k: usize) -> Self {
225        self.retrieve_k = k;
226        self
227    }
228
229    /// 设置系统提示词
230    pub fn system(mut self, prompt: impl Into<String>) -> Self {
231        self.system_prompt = Some(prompt.into());
232        self
233    }
234
235    /// 构建 RAGPipeline
236    ///
237    /// # Errors
238    ///
239    /// 如果缺少 LLM、Embeddings 或 VectorStore,返回错误。
240    pub fn build(self) -> Result<RAGPipeline, RetrieverError> {
241        let llm = self.llm.ok_or_else(|| {
242            RetrieverError::EmbeddingError(
243                "RAGPipelineBuilder: LLM is required. Call .llm() first.".into(),
244            )
245        })?;
246
247        let embeddings = self.embeddings.ok_or_else(|| {
248            RetrieverError::EmbeddingError(
249                "RAGPipelineBuilder: Embeddings is required. Call .embeddings() first.".into(),
250            )
251        })?;
252
253        let vector_store = self.vector_store.ok_or_else(|| {
254            RetrieverError::StoreError(VectorStoreError::StorageError(
255                "RAGPipelineBuilder: VectorStore is required. Call .vector_store() first.".into(),
256            ))
257        })?;
258
259        Ok(RAGPipeline {
260            llm,
261            embeddings,
262            vector_store,
263            retrieve_k: self.retrieve_k,
264            system_prompt: self.system_prompt.unwrap_or_else(|| {
265                "You are a helpful assistant that answers questions based on the provided context.".to_string()
266            }),
267        })
268    }
269}
270
271impl Default for RAGPipelineBuilder {
272    fn default() -> Self {
273        Self::new()
274    }
275}
276
277#[cfg(test)]
278mod tests {
279    use super::*;
280    use lc_embeddings::MockEmbeddings;
281    use lc_providers::{OpenAIChat, OpenAIConfig};
282    use lc_vector_stores::InMemoryVectorStore;
283
284    #[test]
285    fn test_builder_missing_llm() {
286        let result = RAGPipelineBuilder::new()
287            .embeddings(MockEmbeddings::new(3))
288            .vector_store(InMemoryVectorStore::new())
289            .build();
290
291        assert!(result.is_err());
292        assert!(result.unwrap_err().to_string().contains("LLM is required"));
293    }
294
295    #[test]
296    fn test_builder_missing_embeddings() {
297        let config = OpenAIConfig::new("test_key").with_base_url("http://localhost:8080/v1");
298        let result = RAGPipelineBuilder::new()
299            .llm(OpenAIChat::new(config))
300            .vector_store(InMemoryVectorStore::new())
301            .build();
302
303        assert!(result.is_err());
304        assert!(result
305            .unwrap_err()
306            .to_string()
307            .contains("Embeddings is required"));
308    }
309
310    #[test]
311    fn test_builder_missing_vector_store() {
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            .embeddings(MockEmbeddings::new(3))
316            .build();
317
318        assert!(result.is_err());
319        assert!(result
320            .unwrap_err()
321            .to_string()
322            .contains("VectorStore is required"));
323    }
324
325    #[test]
326    fn test_builder_success() {
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            .vector_store(InMemoryVectorStore::new())
332            .system("You are a test assistant.")
333            .retrieve_k(5)
334            .build();
335
336        assert!(result.is_ok());
337        let pipeline = result.unwrap();
338        assert_eq!(pipeline.retrieve_k, 5);
339        assert_eq!(pipeline.system_prompt, "You are a test assistant.");
340    }
341
342    #[test]
343    fn test_builder_default() {
344        let builder = RAGPipelineBuilder::default();
345        assert_eq!(builder.retrieve_k, 4);
346        assert!(builder.llm.is_none());
347    }
348}