Skip to main content

lc_chains/
retrieval_qa.rs

1// lc-chains/src/retrieval_qa.rs
2//! RetrievalQA Chain
3//!
4//! One-stop retrieval QA chain that encapsulates the complete RAG workflow.
5
6use async_trait::async_trait;
7use lc_core::language_models::LLMResult;
8use lc_core::{BaseChatModel, Runnable};
9use lc_rag::retriever::RetrieverTrait;
10use lc_schema::Message;
11use lc_shared::document::Document;
12use serde_json::Value;
13use std::collections::HashMap;
14use std::sync::Arc;
15
16use crate::base::{BaseChain, ChainError, ChainResult};
17
18/// Default QA prompt template.
19const DEFAULT_QA_PROMPT: &str = "Answer the question based on the following context. If the context does not contain relevant information, say 'I don't know'.
20
21Context:
22{context}
23
24Question: {question}
25
26Answer:";
27
28/// RetrievalQA Chain
29///
30/// One-stop retrieval QA chain that automatically:
31/// 1. Retrieves relevant documents
32/// 2. Assembles prompt (context + question)
33/// 3. LLM generates answer
34pub struct RetrievalQA<M: BaseChatModel> {
35    llm: M,
36    retriever: Arc<dyn RetrieverTrait>,
37
38    prompt_template: String,
39    input_key: String,
40    output_key: String,
41    name: String,
42
43    k: usize,
44    verbose: bool,
45
46    return_source_documents: bool,
47    source_document_key: String,
48}
49
50impl<M: BaseChatModel + 'static> RetrievalQA<M> {
51    pub fn new(llm: M, retriever: Arc<dyn RetrieverTrait>) -> Self {
52        Self {
53            llm,
54            retriever,
55            prompt_template: DEFAULT_QA_PROMPT.to_string(),
56            input_key: "query".to_string(),
57            output_key: "result".to_string(),
58            name: "retrieval_qa".to_string(),
59            k: 4,
60            verbose: false,
61            return_source_documents: false,
62            source_document_key: "source_documents".to_string(),
63        }
64    }
65
66    pub fn with_prompt_template(mut self, template: impl Into<String>) -> Self {
67        self.prompt_template = template.into();
68        self
69    }
70
71    pub fn with_input_key(mut self, key: impl Into<String>) -> Self {
72        self.input_key = key.into();
73        self
74    }
75
76    pub fn with_output_key(mut self, key: impl Into<String>) -> Self {
77        self.output_key = key.into();
78        self
79    }
80
81    pub fn with_name(mut self, name: impl Into<String>) -> Self {
82        self.name = name.into();
83        self
84    }
85
86    pub fn with_k(mut self, k: usize) -> Self {
87        self.k = k;
88        self
89    }
90
91    pub fn with_verbose(mut self, verbose: bool) -> Self {
92        self.verbose = verbose;
93        self
94    }
95
96    pub fn with_return_source_documents(mut self, return_source: bool) -> Self {
97        self.return_source_documents = return_source;
98        self
99    }
100
101    pub fn with_source_document_key(mut self, key: impl Into<String>) -> Self {
102        self.source_document_key = key.into();
103        self
104    }
105
106    pub fn retriever(&self) -> &Arc<dyn RetrieverTrait> {
107        &self.retriever
108    }
109
110    pub fn k(&self) -> usize {
111        self.k
112    }
113
114    fn format_context(&self, documents: &[Document]) -> String {
115        documents
116            .iter()
117            .map(|doc| doc.content.clone())
118            .collect::<Vec<_>>()
119            .join("\n\n")
120    }
121
122    fn build_prompt(&self, context: &str, question: &str) -> String {
123        self.prompt_template
124            .replace("{context}", context)
125            .replace("{question}", question)
126    }
127
128    pub async fn query(&self, question: impl Into<String>) -> Result<String, ChainError> {
129        let inputs = HashMap::from([(self.input_key.clone(), Value::String(question.into()))]);
130
131        let result = self.invoke(inputs).await?;
132
133        result
134            .get(&self.output_key)
135            .and_then(|v| v.as_str())
136            .map(|s| s.to_string())
137            .ok_or_else(|| ChainError::OutputError("Missing output result".to_string()))
138    }
139
140    pub async fn query_with_sources(
141        &self,
142        question: impl Into<String>,
143    ) -> Result<(String, Vec<Document>), ChainError> {
144        let inputs = HashMap::from([(self.input_key.clone(), Value::String(question.into()))]);
145
146        let was_returning_sources = self.return_source_documents;
147        if !was_returning_sources {
148            let question_str = inputs
149                .get(&self.input_key)
150                .and_then(|v| v.as_str())
151                .ok_or_else(|| ChainError::MissingInput(self.input_key.clone()))?;
152
153            let documents = self
154                .retriever
155                .retrieve(question_str, self.k)
156                .await
157                .map_err(|e| ChainError::ExecutionError(format!("Retrieval failed: {}", e)))?;
158
159            let context = self.format_context(&documents);
160            let prompt = self.build_prompt(&context, question_str);
161            let messages = vec![Message::human(&prompt)];
162            let response = self
163                .llm
164                .invoke(messages, None)
165                .await
166                .map_err(|e| ChainError::ExecutionError(format!("LLM call failed: {}", e)))?;
167
168            return Ok((response.content, documents));
169        }
170
171        let result = self.invoke(inputs).await?;
172
173        let answer = result
174            .get(&self.output_key)
175            .and_then(|v| v.as_str())
176            .map(|s| s.to_string())
177            .ok_or_else(|| ChainError::OutputError("Missing output result".to_string()))?;
178
179        let sources: Vec<Document> = result
180            .get(&self.source_document_key)
181            .and_then(|v| v.as_array())
182            .map(|arr| {
183                arr.iter()
184                    .filter_map(|v| serde_json::from_value(v.clone()).ok())
185                    .collect()
186            })
187            .unwrap_or_default();
188
189        Ok((answer, sources))
190    }
191}
192
193#[async_trait]
194impl<M: BaseChatModel + Send + Sync + 'static> BaseChain for RetrievalQA<M>
195where
196    <M as Runnable<Vec<Message>, LLMResult>>::Error: std::fmt::Display,
197{
198    fn input_keys(&self) -> Vec<&str> {
199        vec![&self.input_key]
200    }
201
202    fn output_keys(&self) -> Vec<&str> {
203        if self.return_source_documents {
204            vec![&self.output_key, &self.source_document_key]
205        } else {
206            vec![&self.output_key]
207        }
208    }
209
210    async fn invoke(&self, inputs: HashMap<String, Value>) -> Result<ChainResult, ChainError> {
211        self.validate_inputs(&inputs)?;
212
213        let question = inputs
214            .get(&self.input_key)
215            .and_then(|v| v.as_str())
216            .ok_or_else(|| ChainError::MissingInput(self.input_key.clone()))?;
217
218        if self.verbose {
219            println!("\n=== RetrievalQA Execution ===");
220            println!("Question: {}", question);
221            println!("Retrieval count (k): {}", self.k);
222        }
223
224        if self.verbose {
225            println!("\n--- Step 1: Retrieve relevant documents ---");
226        }
227
228        let documents = self
229            .retriever
230            .retrieve(question, self.k)
231            .await
232            .map_err(|e| ChainError::ExecutionError(format!("Retrieval failed: {}", e)))?;
233
234        if self.verbose {
235            println!("Retrieved {} documents", documents.len());
236            for (i, doc) in documents.iter().enumerate() {
237                let preview: String = doc.content.chars().take(100).collect();
238                println!("Document {}: {}", i + 1, preview);
239            }
240        }
241
242        if documents.is_empty() && self.verbose {
243            println!("Warning: No relevant documents retrieved");
244        }
245
246        if self.verbose {
247            println!("\n--- Step 2: Assemble Prompt ---");
248        }
249
250        let context = self.format_context(&documents);
251        let prompt = self.build_prompt(&context, question);
252
253        if self.verbose {
254            println!("Context length: {} characters", context.len());
255            println!("Prompt length: {} characters", prompt.len());
256        }
257
258        if self.verbose {
259            println!("\n--- Step 3: LLM generates answer ---");
260        }
261
262        let messages = vec![Message::human(&prompt)];
263        let response = self
264            .llm
265            .invoke(messages, None)
266            .await
267            .map_err(|e| ChainError::ExecutionError(format!("LLM call failed: {}", e)))?;
268
269        let answer = response.content;
270
271        if self.verbose {
272            println!("Answer: {}", answer);
273            println!("=== RetrievalQA Complete ===\n");
274        }
275
276        let mut result = HashMap::new();
277        result.insert(self.output_key.clone(), Value::String(answer));
278
279        if self.return_source_documents {
280            let sources: Vec<Value> = documents
281                .iter()
282                .map(|doc| serde_json::to_value(doc).unwrap_or(Value::Null))
283                .collect();
284            result.insert(self.source_document_key.clone(), Value::Array(sources));
285        }
286
287        Ok(result)
288    }
289
290    fn name(&self) -> &str {
291        &self.name
292    }
293}