use anyhow::Result;
use std::collections::HashMap;
use std::path::PathBuf;
use super::budget::TokenBudget;
use super::parser::{parse, Language, ParsedFile, Symbol};
use crate::token_count::estimate_content_tokens;
const BM25_K1: f64 = 1.2;
const BM25_B: f64 = 0.75;
pub struct CodeQueryEngine {
doc_freq: HashMap<String, usize>,
total_docs: usize,
avg_doc_len: f64,
}
#[derive(Debug, Clone)]
pub struct RankedResult {
pub path: PathBuf,
pub relevance: f64,
pub matched_symbols: Vec<Symbol>,
}
#[derive(Debug, Clone)]
pub struct QueryResults {
pub results: Vec<RankedResult>,
pub tokens_used: usize,
pub total_matches: usize,
}
impl CodeQueryEngine {
pub fn new() -> Self {
Self {
doc_freq: HashMap::new(),
total_docs: 0,
avg_doc_len: 0.0,
}
}
pub async fn rank_files(&self, files: &[(PathBuf, f64)], query: &str) -> Vec<(PathBuf, f64)> {
let query_terms = tokenize(query);
let mut ranked = Vec::new();
for (path, base_score) in files {
if let Ok(content) = tokio::fs::read_to_string(path).await {
let language = Language::detect(path, None);
let parsed = parse(&content, language);
let score = self.bm25_score(&parsed, &query_terms, content.len());
let combined_score = score * base_score;
ranked.push((path.clone(), combined_score));
} else {
ranked.push((path.clone(), *base_score));
}
}
ranked.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
ranked
}
pub async fn search(
&self,
query: &str,
files: &[PathBuf],
budget: &TokenBudget,
) -> Result<QueryResults> {
let query_terms = tokenize(query);
let mut results = Vec::new();
let mut tokens_used = 0;
for path in files {
if tokens_used >= budget.remaining() {
break;
}
if let Ok(content) = tokio::fs::read_to_string(path).await {
let language = Language::detect(path, None);
let parsed = parse(&content, language);
let score = self.bm25_score(&parsed, &query_terms, content.len());
let matched: Vec<Symbol> = parsed
.symbols
.into_iter()
.filter(|s| self.symbol_matches(s, &query_terms))
.collect();
if !matched.is_empty() {
let symbol_tokens: usize = matched
.iter()
.map(|s| estimate_content_tokens(&s.render()))
.sum();
if tokens_used + symbol_tokens <= budget.remaining() {
results.push(RankedResult {
path: path.clone(),
relevance: score,
matched_symbols: matched,
});
tokens_used += symbol_tokens;
}
}
}
}
let total_matches = results.len();
Ok(QueryResults {
results,
tokens_used,
total_matches,
})
}
fn bm25_score(&self, parsed: &ParsedFile, query_terms: &[String], doc_len: usize) -> f64 {
let mut score = 0.0;
let doc_len_norm = doc_len as f64 / self.avg_doc_len.max(1.0);
for term in query_terms {
let tf = self.term_frequency(parsed, term);
let idf = self.idf(term);
let numerator = tf * (BM25_K1 + 1.0);
let denominator = tf + BM25_K1 * (1.0 - BM25_B + BM25_B * doc_len_norm);
score += idf * numerator / denominator;
}
if parsed.module_doc.is_some() {
score *= 1.1;
}
score
}
fn idf(&self, term: &str) -> f64 {
let doc_freq = self.doc_freq.get(term).copied().unwrap_or(1);
let n = self.total_docs.max(doc_freq);
((n as f64 - doc_freq as f64 + 0.5) / (doc_freq as f64 + 0.5) + 1.0).ln()
}
fn term_frequency(&self, parsed: &ParsedFile, term: &str) -> f64 {
let mut count = 0.0;
let term_lower = term.to_lowercase();
for sym in &parsed.symbols {
if sym.name.to_lowercase().contains(&term_lower) {
count += 3.0; }
if sym.signature.to_lowercase().contains(&term_lower) {
count += 1.0;
}
if let Some(doc) = &sym.documentation {
if doc.to_lowercase().contains(&term_lower) {
count += 0.5;
}
}
}
if let Some(doc) = &parsed.module_doc {
let matches = doc.to_lowercase().matches(&term_lower).count();
count += matches as f64 * 0.5;
}
count
}
fn symbol_matches(&self, symbol: &Symbol, query_terms: &[String]) -> bool {
let name_lower = symbol.name.to_lowercase();
let sig_lower = symbol.signature.to_lowercase();
query_terms.iter().any(|term| {
let term_lower = term.to_lowercase();
name_lower.contains(&term_lower) || sig_lower.contains(&term_lower)
})
}
pub async fn build_index(&mut self, files: &[PathBuf]) -> Result<()> {
self.total_docs = files.len();
let mut total_len = 0;
let mut term_doc_freq: HashMap<String, usize> = HashMap::new();
for path in files {
if let Ok(content) = tokio::fs::read_to_string(path).await {
total_len += content.len();
let terms = tokenize(&content);
let mut seen = std::collections::HashSet::new();
for term in terms {
if seen.insert(term.clone()) {
*term_doc_freq.entry(term).or_default() += 1;
}
}
}
}
self.avg_doc_len = if self.total_docs > 0 {
total_len as f64 / self.total_docs as f64
} else {
0.0
};
self.doc_freq = term_doc_freq;
Ok(())
}
}
impl Default for CodeQueryEngine {
fn default() -> Self {
Self::new()
}
}
fn tokenize(text: &str) -> Vec<String> {
text.to_lowercase()
.split(|c: char| !c.is_alphanumeric() && c != '_')
.filter(|s| !s.is_empty() && s.len() > 1)
.map(|s| s.to_string())
.collect()
}
pub async fn find_related_symbols(
query: &str,
files: &[PathBuf],
max_results: usize,
) -> Result<Vec<(PathBuf, Symbol)>> {
let engine = CodeQueryEngine::new();
let budget = TokenBudget::new(usize::MAX);
let results = engine.search(query, files, &budget).await?;
let mut all_symbols = Vec::new();
for result in results.results {
for symbol in result.matched_symbols {
all_symbols.push((result.path.clone(), symbol));
}
}
all_symbols.truncate(max_results);
Ok(all_symbols)
}
pub fn extract_keywords(query: &str) -> Vec<String> {
let stop_words: std::collections::HashSet<&str> = [
"the", "a", "an", "is", "are", "was", "were", "be", "been", "being", "have", "has", "had",
"do", "does", "did", "will", "would", "could", "should", "may", "might", "must", "shall",
"can", "need", "dare", "ought", "used", "to", "of", "in", "for", "on", "with", "at", "by",
"from", "as", "into", "through", "during", "before", "after", "above", "below", "between",
"under", "and", "but", "or", "yet", "so", "if", "because", "although", "though", "while",
"where", "when", "that", "which", "who", "whom", "whose", "what", "this", "these", "those",
"i", "you", "he", "she", "it", "we", "they", "me", "him", "her", "us", "them", "my",
"your", "his", "its", "our", "their", "how", "does", "work", "use", "using", "get",
]
.iter()
.cloned()
.collect();
tokenize(query)
.into_iter()
.filter(|t| !stop_words.contains(t.as_str()))
.collect()
}
#[cfg(test)]
#[path = "../../../tests/unit/tools/introspect/query/query_test.rs"]
mod tests;