Skip to main content

lc_agents/adaptive_rag/
mod.rs

1// src/agents/adaptive_rag.rs → adaptive_rag/ (mod.rs + types.rs + prompts.rs + tests.rs)
2//! Adaptive RAG implementation.
3//!
4//! Uses an LLM to decide whether retrieval is needed and what strategy to use.
5//! Three decision branches:
6//! - **NoRetrieval**: The query can be answered from general knowledge.
7//! - **SingleSearch**: A single search is sufficient.
8//! - **MultiQuery**: The query is complex and needs multiple search angles.
9
10use 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
26/// Adaptive RAG that routes queries to the most appropriate strategy.
27///
28/// # Overview
29///
30/// 1. The LLM classifies the query into one of three buckets:
31///    `no_retrieval`, `single_search`, or `multi_query`.
32/// 2. Based on the decision:
33///    - **NoRetrieval**: call the LLM directly.
34///    - **SingleSearch**: retrieve documents, then generate.
35///    - **MultiQuery**: generate multiple query variants, retrieve for each,
36///      merge results, then generate.
37///
38/// # Example
39///
40/// ```ignore
41/// use langchainrust::agents::adaptive_rag::{AdaptiveRAG, RagDecision};
42/// use langchainrust::OpenAIChat;
43/// use langchainrust::retrieval::SimilarityRetriever;
44///
45/// let rag = AdaptiveRAG::new(llm, retriever);
46/// let result = rag.invoke("What is the capital of France?").await?;
47/// assert_eq!(result.decision, RagDecision::NoRetrieval);
48/// ```
49pub struct AdaptiveRAG<M: BaseChatModel, R: RetrieverTrait> {
50    llm: M,
51    retriever: R,
52    /// Number of documents to retrieve per query.
53    retrieve_k: usize,
54    /// Number of alternative queries to generate for multi-query mode.
55    multi_query_count: usize,
56}
57
58impl<M: BaseChatModel, R: RetrieverTrait> AdaptiveRAG<M, R> {
59    /// Creates a new `AdaptiveRAG` with the given LLM and retriever.
60    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    /// Sets the number of documents to retrieve per query.
70    pub fn with_retrieve_k(mut self, k: usize) -> Self {
71        self.retrieve_k = k;
72        self
73    }
74
75    /// Sets the number of alternative queries for multi-query mode.
76    pub fn with_multi_query_count(mut self, count: usize) -> Self {
77        self.multi_query_count = count;
78        self
79    }
80
81    /// Invokes the adaptive RAG pipeline for the given query.
82    pub async fn invoke(&self, query: &str) -> Result<AdaptiveRAGResult, AdaptiveRAGError> {
83        // Step 1: Route the query.
84        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    /// Streams the AdaptiveRAG execution, emitting pipeline step events.
94    ///
95    /// Emits `AgentStreamEvent::PipelineStep` events for routing and generation,
96    /// and `AgentStreamEvent::FinalAnswer` when the answer is ready.
97    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        // Step 1: Route
111        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        // Step 2: Execute based on decision
124        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        // Final answer
159        events.push(AgentStreamEvent::FinalAnswer {
160            content: result.answer,
161        });
162
163        Ok(Box::pin(futures_util::stream::iter(events)))
164    }
165
166    // -- Routing -----------------------------------------------------------
167
168    /// Asks the LLM to classify the query.
169    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        // P1-3: prefer structured routing via tool_calls, falling back to text parsing.
174        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    // -- No retrieval ------------------------------------------------------
193
194    /// Generates an answer without any retrieval.
195    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    // -- Single search -----------------------------------------------------
214
215    /// Retrieves documents with a single query and generates an answer.
216    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    // -- Multi query -------------------------------------------------------
240
241    /// Generates multiple query variants, retrieves for each, merges, and
242    /// generates an answer.
243    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        // Combine original + alternatives.
250        let all_queries: Vec<String> = std::iter::once(query.to_string())
251            .chain(alternative_queries)
252            .collect();
253
254        // Retrieve and merge.
255        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    /// Asks the LLM to generate alternative query variants.
275    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    /// Retrieves documents for each query and merges/deduplicates by content.
301    async fn retrieve_and_merge(
302        &self,
303        queries: &[String],
304    ) -> Result<Vec<Document>, AdaptiveRAGError> {
305        // M11: Parallel retrieval instead of sequential.
306        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                // M6 fix: deduplicate by full content hash instead of first 80 chars
319                // to avoid false collisions on documents with common prefixes.
320                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
336// ---------------------------------------------------------------------------
337// Helpers
338// ---------------------------------------------------------------------------
339
340/// Routing tool definition: forces the LLM to emit a three-way decision (P1-3).
341fn 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
358/// Parses the routing decision from the LLM response.
359fn parse_decision(response: &str) -> Result<RagDecision, AdaptiveRAGError> {
360    let lower = response.to_lowercase();
361
362    // Check for exact or contained keywords, with precedence.
363    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
376/// Builds a context string from source documents.
377fn 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}