pub mod provider;
pub mod templates;
use std::collections::HashSet;
#[derive(Debug)]
pub struct QueryRewriter {
prompt_template: String,
}
pub enum RewriteTemplate {
KeywordExtraction,
MultiPerspective { n: usize },
TranslateToEnglish,
Decompose { max_subqueries: usize },
}
const STOP_WORDS: &[&str] = &[
"的",
"了",
"在",
"是",
"我",
"有",
"和",
"就",
"不",
"人",
"都",
"一",
"the",
"a",
"an",
"is",
"are",
"was",
"were",
"be",
"been",
"being",
"have",
"has",
"had",
"do",
"does",
"did",
"will",
"would",
"could",
"should",
"may",
"might",
"can",
"shall",
"to",
"of",
"in",
"for",
"on",
"with",
"at",
"by",
"from",
"as",
"into",
"about",
"what",
"which",
"who",
"whom",
"this",
"that",
"these",
"those",
"it",
"its",
"and",
"but",
"or",
"not",
"no",
"if",
"then",
"else",
"when",
"how",
"why",
"where",
"我",
"你",
"他",
"她",
"它",
"们",
"这",
"那",
"哪",
"什么",
"怎么",
"为什么",
"哪",
"谁",
"吗",
"吧",
"呢",
];
impl QueryRewriter {
pub fn new(prompt_template: &str) -> Self {
Self { prompt_template: prompt_template.to_string() }
}
pub async fn rewrite_with_template(
&self,
query: &str,
template: RewriteTemplate,
) -> Result<Vec<String>, crate::error::SearchError> {
match template {
RewriteTemplate::KeywordExtraction => Ok(self.extract_keywords(query)),
RewriteTemplate::MultiPerspective { n } => Ok(self.generate_perspectives(query, n)),
RewriteTemplate::TranslateToEnglish => {
Ok(vec![query.to_string()])
}
RewriteTemplate::Decompose { max_subqueries } => {
Ok(self.decompose_query(query, max_subqueries))
}
}
}
pub async fn rewrite_with_llm(
&self,
query: &str,
template: RewriteTemplate,
provider: Option<&dyn provider::QueryRewriteProvider>,
) -> Result<Vec<String>, crate::error::SearchError> {
match provider {
Some(p) => {
let system_prompt = match &template {
RewriteTemplate::KeywordExtraction => {
"Extract the core keywords from the query. Return only keywords."
}
RewriteTemplate::MultiPerspective { n } => {
"Rephrase the query from multiple perspectives."
}
RewriteTemplate::TranslateToEnglish => "Translate the query to English.",
RewriteTemplate::Decompose { max_subqueries } => {
"Break the query into sub-queries."
}
};
let rewritten = p.rewrite(query, system_prompt).await?;
Ok(vec![rewritten])
}
None => self.rewrite_with_template(query, template).await,
}
}
pub async fn multi_perspective(
&self,
query: &str,
n: usize,
) -> Result<Vec<String>, crate::error::SearchError> {
Ok(self.generate_perspectives(query, n))
}
fn extract_keywords(&self, text: &str) -> Vec<String> {
let stop_words: HashSet<&str> = STOP_WORDS.iter().copied().collect();
let tokens: Vec<String> = text
.split(|c: char| {
c.is_ascii_punctuation()
|| c.is_whitespace()
|| c == ','
|| c == '。'
|| c == '?'
|| c == '!'
|| c == '、'
})
.map(|s| s.trim().to_lowercase())
.filter(|s| !s.is_empty() && s.len() >= 2)
.collect();
if tokens.is_empty() {
return vec![text.to_string()];
}
let filtered: Vec<String> = tokens
.iter()
.filter(|t| {
if stop_words.contains(t.as_str()) {
return false;
}
if t.chars().all(|c| c.is_ascii_alphabetic()) && t.len() <= 2 {
return false;
}
true
})
.cloned()
.collect();
if filtered.is_empty() {
return vec![text.to_string()];
}
let keywords = filtered.join(" ");
if keywords.len() < text.len() / 2 {
vec![text.to_string(), keywords]
} else {
vec![keywords]
}
}
fn generate_perspectives(&self, query: &str, n: usize) -> Vec<String> {
let mut results = vec![query.to_string()];
let modifiers =
["best", "top", "guide", "tutorial", "how to", "最新", "指南", "教程", "最佳", "推荐"];
for i in 1..n {
let modifier = modifiers[i % modifiers.len()];
results.push(format!("{} {}", query, modifier));
}
results
}
fn decompose_query(&self, query: &str, max: usize) -> Vec<String> {
let separators = [" and ", " or ", " vs ", " 和 ", " 或 ", " 与 ", " 对比 "];
for sep in &separators {
if query.contains(sep) {
let parts: Vec<String> = query
.splitn(max + 1, sep)
.map(|s| s.trim().to_string())
.filter(|s| !s.is_empty())
.collect();
if parts.len() > 1 {
return parts;
}
}
}
vec![query.to_string()]
}
}