Skip to main content

xz_search/rewrite/
mod.rs

1pub mod provider;
2pub mod templates;
3
4use std::collections::HashSet;
5
6/// 查询重写器 — 通过启发式规则或 LLM 优化搜索查询
7///
8/// 使用场景:
9/// - 用户输入自然语言 → 提取关键词
10/// - 拆分复杂查询为多个子查询
11/// - 多角度表述同一查询
12#[derive(Debug)]
13pub struct QueryRewriter {
14    prompt_template: String,
15}
16
17/// 预置重写模板
18pub enum RewriteTemplate {
19    /// 关键词提取
20    KeywordExtraction,
21    /// 多角度表述
22    MultiPerspective { n: usize },
23    /// 翻译为英文
24    TranslateToEnglish,
25    /// 分解查询
26    Decompose { max_subqueries: usize },
27}
28
29/// 常见停用词(中文 + 英文)
30const 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    /// 使用模板重写查询(无 LLM 的启发式实现)
126    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                // 无 LLM 时返回原始查询
136                Ok(vec![query.to_string()])
137            }
138            RewriteTemplate::Decompose { max_subqueries } => {
139                Ok(self.decompose_query(query, max_subqueries))
140            }
141        }
142    }
143
144    /// 使用 LLM 提供者改写查询。无 provider 时回退到启发式方法。
145    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    /// 多角度查询拓展
173    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    // ─── 启发式方法 ────────────────────────────────────────────
182
183    /// 从查询中提取关键词(去除停用词和标点)
184    fn extract_keywords(&self, text: &str) -> Vec<String> {
185        let stop_words: HashSet<&str> = STOP_WORDS.iter().copied().collect();
186
187        // 简单的分词:按空格和常见标点分割
188        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        // 英文停用词检查 + 中文短词过滤
207        let filtered: Vec<String> = tokens
208            .iter()
209            .filter(|t| {
210                if stop_words.contains(t.as_str()) {
211                    return false;
212                }
213                // 单中文字符通常有意义,保留
214                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        // 关键词连接成搜索词
227        let keywords = filtered.join(" ");
228        if keywords.len() < text.len() / 2 {
229            // 如果删太多,回退到原始查询
230            vec![text.to_string(), keywords]
231        } else {
232            vec![keywords]
233        }
234    }
235
236    /// 生成多角度查询变体
237    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    /// 按连接词分解复杂查询
252    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}