Skip to main content

lc_agents/crag/
rewriter.rs

1// src/agents/crag/rewriter.rs
2//! Query rewriting for CRAG.
3//!
4//! When retrieved documents score below the threshold, the query is
5//! rewritten to improve retrieval quality.
6
7use lc_core::language_models::BaseChatModel;
8use lc_schema::Message;
9
10/// Query rewriting error types.
11#[derive(Debug, thiserror::Error)]
12#[non_exhaustive]
13pub enum RewriterError {
14    /// LLM invocation failed.
15    #[error("LLM error during query rewriting: {0}")]
16    LLMError(String),
17
18    /// Failed to parse the rewritten query.
19    #[error("Failed to extract rewritten query from LLM response: {0}")]
20    ParseError(String),
21}
22
23/// Rewrites queries to improve retrieval quality.
24pub struct QueryRewriter<'a, M: BaseChatModel> {
25    llm: &'a M,
26}
27
28impl<'a, M: BaseChatModel> QueryRewriter<'a, M> {
29    /// Creates a new query rewriter.
30    pub fn new(llm: &'a M) -> Self {
31        Self { llm }
32    }
33
34    /// Rewrites the query to be more effective for document retrieval.
35    ///
36    /// The rewriter generates alternative phrasings that may match
37    /// documents the original query missed.
38    pub async fn rewrite(&self, query: &str) -> Result<String, RewriterError> {
39        let prompt = build_rewrite_prompt(query);
40
41        let messages = vec![Message::human(&prompt)];
42        let result = crate::retry::retry_chat(
43            self.llm,
44            messages,
45            None,
46            &crate::retry::RetryConfig::default(),
47        )
48        .await
49        .map_err(|e| RewriterError::LLMError(e.to_string()))?;
50
51        let rewritten = extract_rewritten_query(&result.content);
52        if rewritten.is_empty() {
53            return Err(RewriterError::ParseError(result.content));
54        }
55
56        Ok(rewritten)
57    }
58
59    /// Generates multiple alternative queries for broader retrieval.
60    ///
61    /// Returns a list of rewritten queries including the original.
62    pub async fn generate_alternatives(
63        &self,
64        query: &str,
65        count: usize,
66    ) -> Result<Vec<String>, RewriterError> {
67        let prompt = build_alternatives_prompt(query, count);
68
69        let messages = vec![Message::human(&prompt)];
70        let result = crate::retry::retry_chat(
71            self.llm,
72            messages,
73            None,
74            &crate::retry::RetryConfig::default(),
75        )
76        .await
77        .map_err(|e| RewriterError::LLMError(e.to_string()))?;
78
79        let alternatives = parse_alternatives(&result.content);
80        if alternatives.is_empty() {
81            // Fallback: return the original query
82            return Ok(vec![query.to_string()]);
83        }
84
85        Ok(alternatives)
86    }
87}
88
89/// Builds the single query rewrite prompt.
90fn build_rewrite_prompt(query: &str) -> String {
91    use lc_prompts::PromptTemplate;
92    use std::collections::HashMap;
93
94    let template = PromptTemplate::new(REWRITE_PROMPT);
95    let mut vars = HashMap::new();
96    vars.insert("query", query);
97    template
98        .format(&vars)
99        .unwrap_or_else(|_| REWRITE_PROMPT.to_string())
100}
101
102/// Builds the alternatives generation prompt.
103fn build_alternatives_prompt(query: &str, count: usize) -> String {
104    use lc_prompts::PromptTemplate;
105    use std::collections::HashMap;
106
107    let template = PromptTemplate::new(ALTERNATIVES_PROMPT);
108    let mut vars = HashMap::new();
109    vars.insert("query", query);
110    let count_str = count.to_string();
111    vars.insert("count", &count_str);
112    template
113        .format(&vars)
114        .unwrap_or_else(|_| ALTERNATIVES_PROMPT.to_string())
115}
116
117/// Extracts the rewritten query from the LLM response.
118///
119/// Looks for "Rewritten query:" prefix, or takes the first non-empty line
120/// if no prefix is found.
121fn extract_rewritten_query(response: &str) -> String {
122    let trimmed = response.trim();
123
124    // Try to find "Rewritten query:" prefix
125    if let Some(pos) = trimmed.find("Rewritten query:") {
126        let after = trimmed[pos + "Rewritten query:".len()..].trim();
127        if let Some(line) = after.lines().next() {
128            let cleaned = line.trim().trim_start_matches('-').trim();
129            if !cleaned.is_empty() {
130                return cleaned.to_string();
131            }
132        }
133    }
134
135    // Fallback: take the first non-empty line
136    for line in trimmed.lines() {
137        let cleaned = line.trim();
138        if !cleaned.is_empty() && !cleaned.starts_with('#') {
139            return cleaned.to_string();
140        }
141    }
142
143    String::new()
144}
145
146/// Parses multiple alternative queries from the LLM response.
147///
148/// Expects numbered or bulleted list format.
149fn parse_alternatives(response: &str) -> Vec<String> {
150    let mut alternatives = Vec::new();
151
152    for line in response.lines() {
153        let trimmed = line.trim();
154        if trimmed.is_empty() || trimmed.starts_with('#') {
155            continue;
156        }
157
158        // Strip numbered prefix like "1. " or "1) "
159        let cleaned = trimmed
160            .trim_start_matches(|c: char| c.is_ascii_digit())
161            .trim_start_matches(['.', ')', '-'])
162            .trim();
163
164        if !cleaned.is_empty() {
165            alternatives.push(cleaned.to_string());
166        }
167    }
168
169    alternatives
170}
171
172/// Prompt template for single query rewriting.
173const REWRITE_PROMPT: &str = r#"You are a query rewriter. Your task is to rewrite the given query to improve document retrieval results.
174
175Original query: {query}
176
177Instructions:
1781. Analyze the original query and identify the core information need.
1792. Rewrite the query using different phrasing, synonyms, or more specific terms.
1803. The rewritten query should be optimized for semantic search retrieval.
1814. Keep the rewritten query concise and focused.
182
183Respond with ONLY the rewritten query, no explanation needed.
184
185Rewritten query:"#;
186
187/// Prompt template for generating multiple alternative queries.
188const ALTERNATIVES_PROMPT: &str = r#"You are a query rewriter. Generate {count} alternative versions of the following query to improve document retrieval coverage.
189
190Original query: {query}
191
192Instructions:
1931. Each alternative should approach the information need from a different angle.
1942. Use synonyms, related terms, or more specific phrasing.
1953. Keep each alternative concise and focused.
1964. Number each alternative.
197
198Alternative queries:"#;
199
200#[cfg(test)]
201mod tests {
202    use super::*;
203
204    #[test]
205    fn test_extract_rewritten_query_with_prefix() {
206        let response = "Rewritten query: What are the key features of Rust programming language?";
207        let result = extract_rewritten_query(response);
208        assert_eq!(
209            result,
210            "What are the key features of Rust programming language?"
211        );
212    }
213
214    #[test]
215    fn test_extract_rewritten_query_without_prefix() {
216        let response = "What are the main characteristics of the Rust language?";
217        let result = extract_rewritten_query(response);
218        assert_eq!(
219            result,
220            "What are the main characteristics of the Rust language?"
221        );
222    }
223
224    #[test]
225    fn test_extract_rewritten_query_multiline() {
226        let response = "Here is the rewritten query:\nRewritten query: How does Rust ensure memory safety?\n\nThis focuses on the safety aspect.";
227        let result = extract_rewritten_query(response);
228        assert_eq!(result, "How does Rust ensure memory safety?");
229    }
230
231    #[test]
232    fn test_extract_rewritten_query_empty() {
233        let result = extract_rewritten_query("");
234        assert!(result.is_empty());
235    }
236
237    #[test]
238    fn test_parse_alternatives_numbered() {
239        let response = "1. What is Rust?\n2. How does Rust work?\n3. Rust language overview";
240        let result = parse_alternatives(response);
241        assert_eq!(result.len(), 3);
242        assert_eq!(result[0], "What is Rust?");
243        assert_eq!(result[1], "How does Rust work?");
244        assert_eq!(result[2], "Rust language overview");
245    }
246
247    #[test]
248    fn test_parse_alternatives_bulleted() {
249        let response = "- First alternative\n- Second alternative";
250        let result = parse_alternatives(response);
251        assert_eq!(result.len(), 2);
252    }
253
254    #[test]
255    fn test_parse_alternatives_empty() {
256        let result = parse_alternatives("");
257        assert!(result.is_empty());
258    }
259
260    #[test]
261    fn test_build_rewrite_prompt() {
262        let prompt = build_rewrite_prompt("What is Rust?");
263        assert!(prompt.contains("What is Rust?"));
264        assert!(prompt.contains("Rewritten query:"));
265    }
266
267    #[test]
268    fn test_build_alternatives_prompt() {
269        let prompt = build_alternatives_prompt("What is Rust?", 3);
270        assert!(prompt.contains("What is Rust?"));
271        assert!(prompt.contains("3"));
272    }
273}