1use 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
30pub struct RAGPipeline {
39 llm: Arc<dyn BaseChatModel<Error = ProviderError> + Send + Sync>,
40 retriever: Arc<dyn RetrieverTrait>,
41 retrieve_k: usize,
43 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 pub async fn index_documents(&self, documents: Vec<Document>) -> Result<(), RetrieverError> {
62 self.retriever.add_documents(documents).await
63 }
64
65 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 pub async fn query_with_sources(
79 &self,
80 question: &str,
81 ) -> Result<RAGQueryResult, RetrieverError> {
82 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 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 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 调用失败: {}", e)))?;
116
117 Ok(RAGQueryResult {
118 answer: llm_result.content,
119 sources,
120 })
121 }
122}
123
124#[derive(Debug, Clone)]
126pub struct RAGQueryResult {
127 pub answer: String,
129 pub sources: Vec<Document>,
131}
132
133pub 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 retriever: Option<Arc<dyn RetrieverTrait>>,
154 retrieve_k: usize,
155 system_prompt: Option<String>,
156}
157
158impl RAGPipelineBuilder {
159 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 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 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 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 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 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 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 pub fn retriever_from_arc(mut self, retriever: Arc<dyn RetrieverTrait>) -> Self {
224 self.retriever = Some(retriever);
225 self
226 }
227
228 pub fn retrieve_k(mut self, k: usize) -> Self {
230 self.retrieve_k = k;
231 self
232 }
233
234 pub fn system(mut self, prompt: impl Into<String>) -> Self {
236 self.system_prompt = Some(prompt.into());
237 self
238 }
239
240 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 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 #[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 应成功");
376
377 rag.index_documents(vec![Document::new("Rust is a systems language")])
378 .await
379 .expect("index_documents 应委托给 BM25 检索器成功");
380 }
381}