lc_agents/crag/
rewriter.rs1use lc_core::language_models::BaseChatModel;
8use lc_schema::Message;
9
10#[derive(Debug, thiserror::Error)]
12#[non_exhaustive]
13pub enum RewriterError {
14 #[error("LLM error during query rewriting: {0}")]
16 LLMError(String),
17
18 #[error("Failed to extract rewritten query from LLM response: {0}")]
20 ParseError(String),
21}
22
23pub struct QueryRewriter<'a, M: BaseChatModel> {
25 llm: &'a M,
26}
27
28impl<'a, M: BaseChatModel> QueryRewriter<'a, M> {
29 pub fn new(llm: &'a M) -> Self {
31 Self { llm }
32 }
33
34 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 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 return Ok(vec![query.to_string()]);
83 }
84
85 Ok(alternatives)
86 }
87}
88
89fn 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
102fn 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
117fn extract_rewritten_query(response: &str) -> String {
122 let trimmed = response.trim();
123
124 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 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
146fn 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 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
172const 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
187const 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}