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;
27
28use std::sync::Arc;
29
30pub 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 retrieve_k: usize,
40 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 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 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 pub async fn query_with_sources(
88 &self,
89 question: &str,
90 ) -> Result<RAGQueryResult, RetrieverError> {
91 let query_embedding = self
93 .embeddings
94 .embed_query(question)
95 .await
96 .map_err(|e| RetrieverError::EmbeddingError(e.to_string()))?;
97
98 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 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 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#[derive(Debug, Clone)]
143pub struct RAGQueryResult {
144 pub answer: String,
146 pub sources: Vec<Document>,
148}
149
150pub 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 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 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 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 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 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 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 pub fn retrieve_k(mut self, k: usize) -> Self {
225 self.retrieve_k = k;
226 self
227 }
228
229 pub fn system(mut self, prompt: impl Into<String>) -> Self {
231 self.system_prompt = Some(prompt.into());
232 self
233 }
234
235 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}