1pub mod provider;
2pub mod templates;
3
4use std::collections::HashSet;
5
6#[derive(Debug)]
13pub struct QueryRewriter {
14 prompt_template: String,
15}
16
17pub enum RewriteTemplate {
19 KeywordExtraction,
21 MultiPerspective { n: usize },
23 TranslateToEnglish,
25 Decompose { max_subqueries: usize },
27}
28
29const STOP_WORDS: &[&str] = &[
31 "的",
32 "了",
33 "在",
34 "是",
35 "我",
36 "有",
37 "和",
38 "就",
39 "不",
40 "人",
41 "都",
42 "一",
43 "the",
44 "a",
45 "an",
46 "is",
47 "are",
48 "was",
49 "were",
50 "be",
51 "been",
52 "being",
53 "have",
54 "has",
55 "had",
56 "do",
57 "does",
58 "did",
59 "will",
60 "would",
61 "could",
62 "should",
63 "may",
64 "might",
65 "can",
66 "shall",
67 "to",
68 "of",
69 "in",
70 "for",
71 "on",
72 "with",
73 "at",
74 "by",
75 "from",
76 "as",
77 "into",
78 "about",
79 "what",
80 "which",
81 "who",
82 "whom",
83 "this",
84 "that",
85 "these",
86 "those",
87 "it",
88 "its",
89 "and",
90 "but",
91 "or",
92 "not",
93 "no",
94 "if",
95 "then",
96 "else",
97 "when",
98 "how",
99 "why",
100 "where",
101 "我",
102 "你",
103 "他",
104 "她",
105 "它",
106 "们",
107 "这",
108 "那",
109 "哪",
110 "什么",
111 "怎么",
112 "为什么",
113 "哪",
114 "谁",
115 "吗",
116 "吧",
117 "呢",
118];
119
120impl QueryRewriter {
121 pub fn new(prompt_template: &str) -> Self {
122 Self { prompt_template: prompt_template.to_string() }
123 }
124
125 pub async fn rewrite_with_template(
127 &self,
128 query: &str,
129 template: RewriteTemplate,
130 ) -> Result<Vec<String>, crate::error::SearchError> {
131 match template {
132 RewriteTemplate::KeywordExtraction => Ok(self.extract_keywords(query)),
133 RewriteTemplate::MultiPerspective { n } => Ok(self.generate_perspectives(query, n)),
134 RewriteTemplate::TranslateToEnglish => {
135 Ok(vec![query.to_string()])
137 }
138 RewriteTemplate::Decompose { max_subqueries } => {
139 Ok(self.decompose_query(query, max_subqueries))
140 }
141 }
142 }
143
144 pub async fn rewrite_with_llm(
146 &self,
147 query: &str,
148 template: RewriteTemplate,
149 provider: Option<&dyn provider::QueryRewriteProvider>,
150 ) -> Result<Vec<String>, crate::error::SearchError> {
151 match provider {
152 Some(p) => {
153 let system_prompt = match &template {
154 RewriteTemplate::KeywordExtraction => {
155 "Extract the core keywords from the query. Return only keywords."
156 }
157 RewriteTemplate::MultiPerspective { n } => {
158 "Rephrase the query from multiple perspectives."
159 }
160 RewriteTemplate::TranslateToEnglish => "Translate the query to English.",
161 RewriteTemplate::Decompose { max_subqueries } => {
162 "Break the query into sub-queries."
163 }
164 };
165 let rewritten = p.rewrite(query, system_prompt).await?;
166 Ok(vec![rewritten])
167 }
168 None => self.rewrite_with_template(query, template).await,
169 }
170 }
171
172 pub async fn multi_perspective(
174 &self,
175 query: &str,
176 n: usize,
177 ) -> Result<Vec<String>, crate::error::SearchError> {
178 Ok(self.generate_perspectives(query, n))
179 }
180
181 fn extract_keywords(&self, text: &str) -> Vec<String> {
185 let stop_words: HashSet<&str> = STOP_WORDS.iter().copied().collect();
186
187 let tokens: Vec<String> = text
189 .split(|c: char| {
190 c.is_ascii_punctuation()
191 || c.is_whitespace()
192 || c == ','
193 || c == '。'
194 || c == '?'
195 || c == '!'
196 || c == '、'
197 })
198 .map(|s| s.trim().to_lowercase())
199 .filter(|s| !s.is_empty() && s.len() >= 2)
200 .collect();
201
202 if tokens.is_empty() {
203 return vec![text.to_string()];
204 }
205
206 let filtered: Vec<String> = tokens
208 .iter()
209 .filter(|t| {
210 if stop_words.contains(t.as_str()) {
211 return false;
212 }
213 if t.chars().all(|c| c.is_ascii_alphabetic()) && t.len() <= 2 {
215 return false;
216 }
217 true
218 })
219 .cloned()
220 .collect();
221
222 if filtered.is_empty() {
223 return vec![text.to_string()];
224 }
225
226 let keywords = filtered.join(" ");
228 if keywords.len() < text.len() / 2 {
229 vec![text.to_string(), keywords]
231 } else {
232 vec![keywords]
233 }
234 }
235
236 fn generate_perspectives(&self, query: &str, n: usize) -> Vec<String> {
238 let mut results = vec![query.to_string()];
239
240 let modifiers =
241 ["best", "top", "guide", "tutorial", "how to", "最新", "指南", "教程", "最佳", "推荐"];
242
243 for i in 1..n {
244 let modifier = modifiers[i % modifiers.len()];
245 results.push(format!("{} {}", query, modifier));
246 }
247
248 results
249 }
250
251 fn decompose_query(&self, query: &str, max: usize) -> Vec<String> {
253 let separators = [" and ", " or ", " vs ", " 和 ", " 或 ", " 与 ", " 对比 "];
254
255 for sep in &separators {
256 if query.contains(sep) {
257 let parts: Vec<String> = query
258 .splitn(max + 1, sep)
259 .map(|s| s.trim().to_string())
260 .filter(|s| !s.is_empty())
261 .collect();
262
263 if parts.len() > 1 {
264 return parts;
265 }
266 }
267 }
268
269 vec![query.to_string()]
270 }
271}