lc_chains/
retrieval_qa.rs1use 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
18const 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
28pub 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}