use ahash::{AHashMap, AHashSet};
use probe_code::search::elastic_query::Expr;
use probe_code::search::tokenization;
use rust_stemmers::{Algorithm, Stemmer};
use std::sync::OnceLock;
type HashMap<K, V> = AHashMap<K, V>;
type HashSet<T> = AHashSet<T>;
pub type QueryTokenMap = HashMap<String, u8>;
pub struct TfDfResult {
pub term_frequencies: Vec<HashMap<u8, usize>>,
pub document_frequencies: HashMap<String, usize>,
pub document_lengths: Vec<usize>,
}
pub struct RankingParams<'a> {
pub documents: &'a [&'a str],
pub query: &'a str,
pub pre_tokenized: Option<&'a [Vec<String>]>,
}
pub fn get_stemmer() -> &'static Stemmer {
static STEMMER: OnceLock<Stemmer> = OnceLock::new();
STEMMER.get_or_init(|| Stemmer::create(Algorithm::English))
}
pub fn tokenize(text: &str) -> Vec<String> {
tokenization::tokenize(text)
}
pub fn preprocess_text_with_filename(text: &str, filename: &str) -> Vec<String> {
let mut tokens = tokenize(text);
let filename_tokens = tokenize(filename);
tokens.extend(filename_tokens);
tokens
}
pub fn compute_avgdl(lengths: &[usize]) -> f64 {
if lengths.is_empty() {
return 0.0;
}
let sum: f64 = lengths.iter().map(|&x| x as f64).sum();
sum / lengths.len() as f64
}
pub struct PrecomputedBm25Params<'a> {
pub doc_tf: &'a HashMap<u8, usize>,
pub doc_len: usize,
pub avgdl: f64,
pub idfs: &'a HashMap<String, f64>,
pub query_token_map: &'a QueryTokenMap,
pub k1: f64,
pub b: f64,
}
pub fn extract_query_terms(expr: &Expr) -> HashSet<String> {
use Expr::*;
let mut terms = HashSet::new();
match expr {
Term { keywords, .. } => {
terms.extend(keywords.iter().cloned());
}
And(left, right) | Or(left, right) => {
terms.extend(extract_query_terms(left));
terms.extend(extract_query_terms(right));
}
}
terms
}
pub fn precompute_idfs(
terms: &HashSet<String>,
dfs: &HashMap<String, usize>,
n_docs: usize,
) -> HashMap<String, f64> {
let debug_mode = std::env::var("DEBUG").unwrap_or_default() == "1";
if debug_mode {
println!(
"DEBUG: Precomputing IDF values for {terms_len} terms",
terms_len = terms.len()
);
}
terms
.iter()
.filter_map(|term| {
let df = *dfs.get(term).unwrap_or(&0);
if df > 0 {
let numerator = (n_docs as f64 - df as f64) + 0.5;
let denominator = df as f64 + 0.5;
let idf = (1.0 + (numerator / denominator)).ln();
Some((term.as_str(), idf))
} else {
None
}
})
.map(|(term, idf)| (term.to_string(), idf))
.collect()
}
fn generate_query_token_map(query_terms: &HashSet<String>) -> Result<QueryTokenMap, &'static str> {
if query_terms.len() > 256 {
return Err("Query exceeds the 256 unique token limit for u8 mapping");
}
let mut token_map = QueryTokenMap::new();
let mut index: u8 = 0;
let mut sorted_terms: Vec<&str> = query_terms.iter().map(|s| s.as_str()).collect();
sorted_terms.sort();
for term in sorted_terms {
token_map.insert(term.to_string(), index);
index = index.wrapping_add(1); }
Ok(token_map)
}
fn bm25_single_token_optimized(token: &str, params: &PrecomputedBm25Params) -> f64 {
let Some(&token_index) = params.query_token_map.get(token) else {
return 0.0;
};
let freq_in_doc = *params.doc_tf.get(&token_index).unwrap_or(&0) as f64;
if freq_in_doc <= 0.0 {
return 0.0;
}
let idf = *params.idfs.get(token).unwrap_or(&0.0);
let tf_part = (freq_in_doc * (params.k1 + 1.0))
/ (freq_in_doc
+ params.k1 * (1.0 - params.b + params.b * (params.doc_len as f64 / params.avgdl)));
idf * tf_part
}
fn score_term_bm25_optimized(keywords: &[String], params: &PrecomputedBm25Params) -> f64 {
let mut total = 0.0;
for kw in keywords {
total += bm25_single_token_optimized(kw, params);
}
total
}
pub fn score_expr_bm25_optimized(expr: &Expr, params: &PrecomputedBm25Params) -> Option<f64> {
use Expr::*;
match expr {
Term {
keywords,
required,
excluded,
..
} => {
let score = score_term_bm25_optimized(keywords, params);
if *excluded {
if score > 0.0 {
None
} else {
Some(0.0)
}
} else if *required {
if score > 0.0 {
Some(score)
} else {
None
}
} else {
Some(score)
}
}
And(left, right) => {
let lscore = score_expr_bm25_optimized(left, params)?;
let rscore = score_expr_bm25_optimized(right, params)?;
Some(lscore + rscore)
}
Or(left, right) => {
let l = score_expr_bm25_optimized(left, params);
let r = score_expr_bm25_optimized(right, params);
match (l, r) {
(None, None) => None,
(None, Some(rs)) => Some(rs),
(Some(ls), None) => Some(ls),
(Some(ls), Some(rs)) => Some(ls + rs),
}
}
}
}
pub fn rank_documents(params: &RankingParams) -> Vec<(usize, f64)> {
use rayon::prelude::*;
use std::cmp::Ordering;
let debug_mode = std::env::var("DEBUG").unwrap_or_default() == "1";
let parsed_expr = match crate::search::elastic_query::parse_query(params.query, false) {
Ok(expr) => expr,
Err(e) => {
if debug_mode {
eprintln!("DEBUG: parse_query failed: {e:?}");
}
eprintln!("WARNING: Query parsing failed: {e:?}. Returning empty results.");
return vec![];
}
};
let query_terms = extract_query_terms(&parsed_expr);
let query_token_map = match generate_query_token_map(&query_terms) {
Ok(map) => map,
Err(e) => {
if debug_mode {
eprintln!("DEBUG: Failed to generate query token map: {e}");
}
eprintln!("WARNING: {e}");
return vec![];
}
};
if debug_mode {
println!(
"DEBUG: Generated query token map with {} entries",
query_token_map.len()
);
}
let tf_df_result = if let Some(pre_tokenized) = ¶ms.pre_tokenized {
if debug_mode {
println!("DEBUG: Using pre-tokenized content for ranking");
}
compute_tf_df_from_tokenized(pre_tokenized, &query_token_map)
} else {
if debug_mode {
println!("DEBUG: Tokenizing documents for ranking");
}
let tokenized_docs: Vec<Vec<String>> =
params.documents.iter().map(|doc| tokenize(doc)).collect();
compute_tf_df_from_tokenized(&tokenized_docs, &query_token_map)
};
let n_docs = params.documents.len();
let avgdl = compute_avgdl(&tf_df_result.document_lengths);
let precomputed_idfs =
precompute_idfs(&query_terms, &tf_df_result.document_frequencies, n_docs);
if debug_mode {
println!(
"DEBUG: Precomputed IDF values for {} unique query terms",
precomputed_idfs.len()
);
}
let k1 = 1.2;
let b = 0.75;
if debug_mode {
println!("DEBUG: Starting parallel document scoring for {n_docs} documents");
}
let scored_docs: Vec<(usize, Option<f64>)> = (0..tf_df_result.term_frequencies.len())
.collect::<Vec<_>>() .par_iter() .map(|&i| {
let doc_tf = &tf_df_result.term_frequencies[i];
let doc_len = tf_df_result.document_lengths[i];
let precomputed_bm25_params = PrecomputedBm25Params {
doc_tf,
doc_len,
avgdl,
idfs: &precomputed_idfs,
query_token_map: &query_token_map,
k1,
b,
};
let bm25_score_opt = score_expr_bm25_optimized(&parsed_expr, &precomputed_bm25_params);
(i, bm25_score_opt)
})
.collect();
if debug_mode {
println!("DEBUG: Parallel document scoring completed");
}
let mut filtered_docs: Vec<(usize, f64)> = scored_docs
.into_iter()
.filter_map(|(i, score_opt)| score_opt.map(|score| (i, score)))
.collect();
filtered_docs.sort_by(|a, b| {
match b.1.partial_cmp(&a.1).unwrap_or(Ordering::Equal) {
Ordering::Equal => {
a.0.cmp(&b.0)
}
other => other,
}
});
if debug_mode {
println!(
"DEBUG: Sorted {} matching documents by score",
filtered_docs.len()
);
}
filtered_docs
}
pub fn compute_tf_df_from_tokenized(
tokenized_docs: &[Vec<String>],
query_token_map: &QueryTokenMap,
) -> TfDfResult {
use rayon::prelude::*;
let debug_mode = std::env::var("DEBUG").unwrap_or_default() == "1";
if debug_mode {
println!("DEBUG: Starting parallel TF-DF computation from pre-tokenized content for {docs_len} documents", docs_len = tokenized_docs.len());
}
#[allow(clippy::type_complexity)]
let doc_results: Vec<(
HashMap<u8, usize>,
HashMap<String, usize>,
usize,
HashSet<String>,
)> = tokenized_docs
.par_iter()
.map(|tokens| {
let mut tf_u8 = HashMap::new(); let mut tf_str = HashMap::new();
for token in tokens.iter() {
*tf_str.entry(token.clone()).or_insert(0) += 1;
if let Some(&token_index) = query_token_map.get(token) {
*tf_u8.entry(token_index).or_insert(0) += 1;
}
}
let unique_terms: HashSet<String> = tf_str.keys().cloned().collect();
(tf_u8, tf_str, tokens.len(), unique_terms)
})
.collect();
let mut term_frequencies = Vec::with_capacity(tokenized_docs.len());
let mut document_lengths = Vec::with_capacity(tokenized_docs.len());
let min_chunk_size = tokenized_docs
.len()
.checked_div(rayon::current_num_threads())
.unwrap_or(1)
.max(1);
let document_frequencies = doc_results
.par_iter()
.with_min_len(min_chunk_size) .map(|(_, _, _, unique_terms)| {
let mut local_df = HashMap::new();
for term in unique_terms {
*local_df.entry(term.clone()).or_insert(0) += 1;
}
local_df
})
.reduce(HashMap::new, |mut acc, local_df| {
for (term, count) in local_df {
*acc.entry(term).or_insert(0) += count;
}
acc
});
if debug_mode {
println!(
"DEBUG: Parallel DF computation completed with {} unique terms",
document_frequencies.len()
);
}
for (tf_u8, _, doc_len, _) in doc_results {
term_frequencies.push(tf_u8);
document_lengths.push(doc_len);
}
if debug_mode {
println!("DEBUG: Parallel TF-DF computation from pre-tokenized content completed");
println!("DEBUG: Using u8 indices for term frequencies (optimized storage)");
}
TfDfResult {
term_frequencies,
document_frequencies,
document_lengths,
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_basic_bm25_scoring() {
let docs = vec!["api process load", "another random text with process"];
let query = "+api +process +load";
let params = RankingParams {
documents: &docs,
query,
pre_tokenized: None,
};
let results = rank_documents(¶ms);
assert_eq!(results.len(), 1);
assert_eq!(results[0].0, 0);
assert!(results[0].1 > 0.0);
assert!(results[0].1 < 10.0); }
#[test]
fn test_bm25_scoring_with_pre_tokenized() {
let docs = vec!["api process load", "another random text with process"];
let query = "+api +process +load";
let pre_tokenized = vec![
vec!["api".to_string(), "process".to_string(), "load".to_string()],
vec![
"another".to_string(),
"random".to_string(),
"text".to_string(),
"with".to_string(),
"process".to_string(),
],
];
let params = RankingParams {
documents: &docs,
query,
pre_tokenized: Some(&pre_tokenized),
};
let results = rank_documents(¶ms);
assert_eq!(results.len(), 1);
assert_eq!(results[0].0, 0);
assert!(results[0].1 > 0.0);
assert!(results[0].1 < 10.0); }
#[test]
fn test_relative_bm25_scoring() {
let docs = vec![
"api process load data", "api process load", "api process", "api", ];
let query = "api process load data";
let params = RankingParams {
documents: &docs,
query,
pre_tokenized: None,
};
let results = rank_documents(¶ms);
assert_eq!(results.len(), 4);
assert_eq!(results[0].0, 0); assert_eq!(results[1].0, 1); assert_eq!(results[2].0, 2); assert_eq!(results[3].0, 3);
assert!(results[0].1 > results[1].1); assert!(results[1].1 > results[2].1); assert!(results[2].1 > results[3].1); }
#[test]
fn test_generate_query_token_map_basic() {
let mut query_terms = HashSet::new();
query_terms.insert("apple".to_string());
query_terms.insert("banana".to_string());
query_terms.insert("cherry".to_string());
let token_map = generate_query_token_map(&query_terms).unwrap();
assert_eq!(token_map.len(), 3);
let mut indices = HashSet::new();
for (_, &idx) in &token_map {
assert!(indices.insert(idx), "Duplicate index found");
}
assert_eq!(indices.len(), 3);
assert!(indices.contains(&0));
assert!(indices.contains(&1));
assert!(indices.contains(&2));
}
#[test]
fn test_generate_query_token_map_empty() {
let query_terms = HashSet::new();
let token_map = generate_query_token_map(&query_terms).unwrap();
assert!(token_map.is_empty());
}
#[test]
fn test_generate_query_token_map_deterministic() {
let mut query_terms1 = HashSet::new();
query_terms1.insert("apple".to_string());
query_terms1.insert("banana".to_string());
query_terms1.insert("cherry".to_string());
let mut query_terms2 = HashSet::new();
query_terms2.insert("cherry".to_string());
query_terms2.insert("apple".to_string());
query_terms2.insert("banana".to_string());
let token_map1 = generate_query_token_map(&query_terms1).unwrap();
let token_map2 = generate_query_token_map(&query_terms2).unwrap();
assert_eq!(token_map1.len(), token_map2.len());
for (term, &idx1) in &token_map1 {
assert_eq!(
Some(&idx1),
token_map2.get(term),
"Term '{term}' has different indices in the two maps"
);
}
}
#[test]
fn test_generate_query_token_map_too_many_terms() {
let query_terms: HashSet<String> = (0..257).map(|i| format!("term{i}")).collect();
let result = generate_query_token_map(&query_terms);
assert!(result.is_err());
assert_eq!(
result.unwrap_err(),
"Query exceeds the 256 unique token limit for u8 mapping"
);
}
#[test]
fn test_compute_tf_df_with_u8_indices() {
let docs = vec![
vec![
"apple".to_string(),
"banana".to_string(),
"cherry".to_string(),
],
vec!["apple".to_string(), "banana".to_string()],
vec!["apple".to_string()],
];
let mut query_token_map = QueryTokenMap::new();
query_token_map.insert("apple".to_string(), 0);
query_token_map.insert("banana".to_string(), 1);
query_token_map.insert("cherry".to_string(), 2);
let tf_df_result = compute_tf_df_from_tokenized(&docs, &query_token_map);
assert_eq!(tf_df_result.document_lengths[0], 3);
assert_eq!(tf_df_result.document_lengths[1], 2);
assert_eq!(tf_df_result.document_lengths[2], 1);
assert_eq!(*tf_df_result.term_frequencies[0].get(&0).unwrap(), 1); assert_eq!(*tf_df_result.term_frequencies[0].get(&1).unwrap(), 1); assert_eq!(*tf_df_result.term_frequencies[0].get(&2).unwrap(), 1);
assert_eq!(*tf_df_result.term_frequencies[1].get(&0).unwrap(), 1); assert_eq!(*tf_df_result.term_frequencies[1].get(&1).unwrap(), 1); assert!(tf_df_result.term_frequencies[1].get(&2).is_none());
assert_eq!(*tf_df_result.term_frequencies[2].get(&0).unwrap(), 1); assert!(tf_df_result.term_frequencies[2].get(&1).is_none()); assert!(tf_df_result.term_frequencies[2].get(&2).is_none());
assert_eq!(*tf_df_result.document_frequencies.get("apple").unwrap(), 3); assert_eq!(*tf_df_result.document_frequencies.get("banana").unwrap(), 2); assert_eq!(*tf_df_result.document_frequencies.get("cherry").unwrap(), 1);
}
#[test]
fn test_bm25_scoring_with_u8_indices() {
let _doc_content = "apple banana cherry";
let mut query_token_map = QueryTokenMap::new();
query_token_map.insert("apple".to_string(), 0);
query_token_map.insert("banana".to_string(), 1);
query_token_map.insert("cherry".to_string(), 2);
let mut doc_tf = HashMap::new();
doc_tf.insert(0u8, 1); doc_tf.insert(1u8, 1); doc_tf.insert(2u8, 1);
let mut doc_freqs = HashMap::new();
doc_freqs.insert("apple".to_string(), 1);
doc_freqs.insert("banana".to_string(), 1);
doc_freqs.insert("cherry".to_string(), 1);
let mut idfs = HashMap::new();
idfs.insert("apple".to_string(), 1.0);
idfs.insert("banana".to_string(), 1.0);
idfs.insert("cherry".to_string(), 1.0);
let params = PrecomputedBm25Params {
doc_tf: &doc_tf,
doc_len: 3,
avgdl: 3.0,
idfs: &idfs,
query_token_map: &query_token_map,
k1: 1.2,
b: 0.75,
};
let apple_score = bm25_single_token_optimized("apple", ¶ms);
let banana_score = bm25_single_token_optimized("banana", ¶ms);
let cherry_score = bm25_single_token_optimized("cherry", ¶ms);
assert!(apple_score > 0.0);
assert_eq!(apple_score, banana_score);
assert_eq!(banana_score, cherry_score);
let unknown_score = bm25_single_token_optimized("unknown", ¶ms);
assert_eq!(unknown_score, 0.0);
let keywords = vec!["apple".to_string(), "banana".to_string()];
let term_score = score_term_bm25_optimized(&keywords, ¶ms);
assert_eq!(term_score, apple_score + banana_score);
}
}