1use lc_core::language_models::BaseChatModel;
11use lc_core::tools::ToolDefinition;
12use lc_rag::RetrieverTrait;
13use lc_schema::Message;
14use lc_vector_stores::Document;
15use serde_json::json;
16
17mod prompts;
18#[cfg(test)]
19mod tests;
20mod types;
21
22pub use types::{AdaptiveRAGError, AdaptiveRAGResult, RagDecision};
23
24use prompts::{GENERATE_SYSTEM_PROMPT, MULTI_QUERY_PROMPT, ROUTING_PROMPT};
25
26pub struct AdaptiveRAG<M: BaseChatModel, R: RetrieverTrait> {
50 llm: M,
51 retriever: R,
52 retrieve_k: usize,
54 multi_query_count: usize,
56}
57
58impl<M: BaseChatModel, R: RetrieverTrait> AdaptiveRAG<M, R> {
59 pub fn new(llm: M, retriever: R) -> Self {
61 Self {
62 llm,
63 retriever,
64 retrieve_k: 4,
65 multi_query_count: 3,
66 }
67 }
68
69 pub fn with_retrieve_k(mut self, k: usize) -> Self {
71 self.retrieve_k = k;
72 self
73 }
74
75 pub fn with_multi_query_count(mut self, count: usize) -> Self {
77 self.multi_query_count = count;
78 self
79 }
80
81 pub async fn invoke(&self, query: &str) -> Result<AdaptiveRAGResult, AdaptiveRAGError> {
83 let decision = self.route(query).await?;
85
86 match decision {
87 RagDecision::NoRetrieval => self.generate_no_retrieval(query).await,
88 RagDecision::SingleSearch => self.generate_single_search(query).await,
89 RagDecision::MultiQuery => self.generate_multi_query(query).await,
90 }
91 }
92
93 pub async fn stream(
98 &self,
99 query: &str,
100 ) -> Result<
101 std::pin::Pin<
102 Box<dyn futures_util::Stream<Item = crate::streaming::AgentStreamEvent> + Send>,
103 >,
104 AdaptiveRAGError,
105 > {
106 use crate::streaming::AgentStreamEvent;
107
108 let mut events: Vec<AgentStreamEvent> = Vec::new();
109
110 events.push(AgentStreamEvent::PipelineStep {
112 step: "routing".to_string(),
113 detail: Some("Classifying query...".to_string()),
114 });
115
116 let decision = self.route(query).await?;
117
118 events.push(AgentStreamEvent::PipelineStep {
119 step: "routed".to_string(),
120 detail: Some(format!("Decision: {}", decision)),
121 });
122
123 let result = match decision {
125 RagDecision::NoRetrieval => {
126 events.push(AgentStreamEvent::PipelineStep {
127 step: "generating".to_string(),
128 detail: Some("No retrieval needed, generating directly...".to_string()),
129 });
130 self.generate_no_retrieval(query).await?
131 }
132 RagDecision::SingleSearch => {
133 events.push(AgentStreamEvent::PipelineStep {
134 step: "retrieving".to_string(),
135 detail: Some("Single search retrieval...".to_string()),
136 });
137 let result = self.generate_single_search(query).await?;
138 events.push(AgentStreamEvent::PipelineStep {
139 step: "generating".to_string(),
140 detail: Some(format!("Sources: {} documents", result.sources.len())),
141 });
142 result
143 }
144 RagDecision::MultiQuery => {
145 events.push(AgentStreamEvent::PipelineStep {
146 step: "multi_query".to_string(),
147 detail: Some("Generating multiple queries...".to_string()),
148 });
149 let result = self.generate_multi_query(query).await?;
150 events.push(AgentStreamEvent::PipelineStep {
151 step: "generating".to_string(),
152 detail: Some(format!("Sources: {} documents", result.sources.len())),
153 });
154 result
155 }
156 };
157
158 events.push(AgentStreamEvent::FinalAnswer {
160 content: result.answer,
161 });
162
163 Ok(Box::pin(futures_util::stream::iter(events)))
164 }
165
166 async fn route(&self, query: &str) -> Result<RagDecision, AdaptiveRAGError> {
170 let prompt = ROUTING_PROMPT.replace("{query}", query);
171 let messages = vec![Message::human(&prompt)];
172
173 let structured = crate::structured::chat_structured(
175 &self.llm,
176 Some(route_tool()),
177 messages,
178 None,
179 &crate::retry::RetryConfig::default(),
180 )
181 .await
182 .map_err(|e| AdaptiveRAGError::Llm(e.to_string()))?;
183
184 if let Some(args) = &structured.tool_args {
185 if let Some(decision) = args.get("decision").and_then(|v| v.as_str()) {
186 return parse_decision(decision);
187 }
188 }
189 parse_decision(&structured.content)
190 }
191
192 async fn generate_no_retrieval(
196 &self,
197 query: &str,
198 ) -> Result<AdaptiveRAGResult, AdaptiveRAGError> {
199 let messages = vec![Message::human(query)];
200 let result = self
201 .llm
202 .chat_with_system(GENERATE_SYSTEM_PROMPT.to_string(), messages)
203 .await
204 .map_err(|e| AdaptiveRAGError::Llm(e.to_string()))?;
205
206 Ok(AdaptiveRAGResult {
207 answer: result.content,
208 decision: RagDecision::NoRetrieval,
209 sources: Vec::new(),
210 })
211 }
212
213 async fn generate_single_search(
217 &self,
218 query: &str,
219 ) -> Result<AdaptiveRAGResult, AdaptiveRAGError> {
220 let docs = self.retriever.retrieve(query, self.retrieve_k).await?;
221
222 let context = build_context(&docs);
223 let user_msg = format!("Context:\n{}\n\nQuestion: {}\n\nAnswer:", context, query);
224 let messages = vec![Message::human(&user_msg)];
225
226 let result = self
227 .llm
228 .chat_with_system(GENERATE_SYSTEM_PROMPT.to_string(), messages)
229 .await
230 .map_err(|e| AdaptiveRAGError::Llm(e.to_string()))?;
231
232 Ok(AdaptiveRAGResult {
233 answer: result.content,
234 decision: RagDecision::SingleSearch,
235 sources: docs,
236 })
237 }
238
239 async fn generate_multi_query(
244 &self,
245 query: &str,
246 ) -> Result<AdaptiveRAGResult, AdaptiveRAGError> {
247 let alternative_queries = self.generate_queries(query).await?;
248
249 let all_queries: Vec<String> = std::iter::once(query.to_string())
251 .chain(alternative_queries)
252 .collect();
253
254 let docs = self.retrieve_and_merge(&all_queries).await?;
256
257 let context = build_context(&docs);
258 let user_msg = format!("Context:\n{}\n\nQuestion: {}\n\nAnswer:", context, query);
259 let messages = vec![Message::human(&user_msg)];
260
261 let result = self
262 .llm
263 .chat_with_system(GENERATE_SYSTEM_PROMPT.to_string(), messages)
264 .await
265 .map_err(|e| AdaptiveRAGError::Llm(e.to_string()))?;
266
267 Ok(AdaptiveRAGResult {
268 answer: result.content,
269 decision: RagDecision::MultiQuery,
270 sources: docs,
271 })
272 }
273
274 async fn generate_queries(&self, query: &str) -> Result<Vec<String>, AdaptiveRAGError> {
276 let prompt = MULTI_QUERY_PROMPT
277 .replace("{count}", &self.multi_query_count.to_string())
278 .replace("{question}", query);
279
280 let messages = vec![Message::human(&prompt)];
281 let result = crate::retry::retry_chat(
282 &self.llm,
283 messages,
284 None,
285 &crate::retry::RetryConfig::default(),
286 )
287 .await
288 .map_err(|e| AdaptiveRAGError::Llm(e.to_string()))?;
289
290 let queries: Vec<String> = result
291 .content
292 .lines()
293 .filter(|line| !line.trim().is_empty())
294 .map(|line| line.trim().to_string())
295 .collect();
296
297 Ok(queries)
298 }
299
300 async fn retrieve_and_merge(
302 &self,
303 queries: &[String],
304 ) -> Result<Vec<Document>, AdaptiveRAGError> {
305 let futures: Vec<_> = queries
307 .iter()
308 .map(|q| self.retriever.retrieve(q, self.retrieve_k))
309 .collect();
310 let all_results = futures_util::future::join_all(futures).await;
311
312 let mut seen_content: std::collections::HashSet<String> = std::collections::HashSet::new();
313 let mut merged: Vec<Document> = Vec::new();
314
315 for result in all_results {
316 let docs = result?;
317 for doc in docs {
318 let key = {
321 use std::hash::Hasher;
322 let mut hasher = std::collections::hash_map::DefaultHasher::new();
323 hasher.write(doc.content.as_bytes());
324 format!("{:016x}", hasher.finish())
325 };
326 if seen_content.insert(key) {
327 merged.push(doc);
328 }
329 }
330 }
331
332 Ok(merged)
333 }
334}
335
336fn route_tool() -> ToolDefinition {
342 ToolDefinition::new(
343 "route_decision",
344 "判断查询是否需要检索及检索策略:no_retrieval / single_search / multi_query",
345 )
346 .with_parameters(json!({
347 "type": "object",
348 "properties": {
349 "decision": {
350 "type": "string",
351 "enum": ["no_retrieval", "single_search", "multi_query"]
352 }
353 },
354 "required": ["decision"]
355 }))
356}
357
358fn parse_decision(response: &str) -> Result<RagDecision, AdaptiveRAGError> {
360 let lower = response.to_lowercase();
361
362 if lower.contains("no_retrieval") {
364 return Ok(RagDecision::NoRetrieval);
365 }
366 if lower.contains("multi_query") {
367 return Ok(RagDecision::MultiQuery);
368 }
369 if lower.contains("single_search") {
370 return Ok(RagDecision::SingleSearch);
371 }
372
373 Err(AdaptiveRAGError::DecisionParse(response.to_string()))
374}
375
376fn build_context(docs: &[Document]) -> String {
378 docs.iter()
379 .enumerate()
380 .map(|(i, doc)| format!("[Document {}]: {}", i + 1, doc.content))
381 .collect::<Vec<_>>()
382 .join("\n\n")
383}