use std::collections::HashMap;
use crate::search::types::SearchResult;
pub fn rrf_fuse(result_sets: &[Vec<SearchResult>], k: u32) -> Vec<SearchResult> {
let weights = vec![1.0_f32; result_sets.len()];
rrf_fuse_weighted(result_sets, &weights, k)
}
pub fn rrf_fuse_weighted(
result_sets: &[Vec<SearchResult>],
weights: &[f32],
k: u32,
) -> Vec<SearchResult> {
assert_eq!(
result_sets.len(),
weights.len(),
"rrf_fuse_weighted: result_sets and weights length must match"
);
let mut agg: HashMap<String, (f32, String)> = HashMap::new();
for (results, weight) in result_sets.iter().zip(weights.iter().copied()) {
for (rank, result) in results.iter().enumerate() {
let entry = agg
.entry(result.id.clone())
.or_insert_with(|| (0.0, result.text.clone()));
entry.0 += weight * (1.0 / (k as f32 + rank as f32 + 1.0));
}
}
let mut fused: Vec<SearchResult> = agg
.into_iter()
.map(|(id, (score, text))| SearchResult { text, id, score })
.collect();
fused.sort_by(|a, b| {
b.score
.partial_cmp(&a.score)
.unwrap_or(std::cmp::Ordering::Equal)
});
fused
}
#[cfg(test)]
mod tests {
use super::*;
fn make_results(ids: &[(&str, f32)]) -> Vec<SearchResult> {
ids.iter()
.map(|(id, score)| SearchResult {
id: id.to_string(),
score: *score,
text: format!("text for {id}"),
})
.collect()
}
#[test]
fn empty_inputs() {
let result = rrf_fuse(&[], 60);
assert!(result.is_empty());
}
#[test]
fn single_list() {
let list = make_results(&[("a", 0.9), ("b", 0.5)]);
let fused = rrf_fuse(&[list], 60);
assert_eq!(fused.len(), 2);
assert_eq!(fused[0].id, "a"); }
#[test]
fn two_lists_overlap() {
let bm25 = make_results(&[("a", 0.9), ("b", 0.7)]);
let vector = make_results(&[("b", 0.95), ("c", 0.6)]);
let fused = rrf_fuse(&[bm25, vector], 60);
assert_eq!(fused.len(), 3);
assert_eq!(fused[0].id, "b");
}
#[test]
fn rrf_score_formula() {
let list = make_results(&[("x", 1.0)]);
let fused = rrf_fuse(&[list], 60);
let expected = 1.0 / 61.0;
assert!((fused[0].score - expected).abs() < 1e-6);
}
#[test]
fn three_lists() {
let a = make_results(&[("doc1", 1.0), ("doc2", 0.8)]);
let b = make_results(&[("doc2", 0.9), ("doc3", 0.7)]);
let c = make_results(&[("doc1", 0.8), ("doc3", 0.9)]);
let fused = rrf_fuse(&[a, b, c], 60);
assert_eq!(fused.len(), 3);
let doc3_score = fused.iter().find(|r| r.id == "doc3").unwrap().score;
let doc1_score = fused.iter().find(|r| r.id == "doc1").unwrap().score;
assert!(doc1_score > doc3_score);
}
#[test]
fn weighted_higher_weight_dominates() {
let lexical = make_results(&[("a", 0.9), ("b", 0.3)]);
let vector = make_results(&[("b", 0.95), ("a", 0.2)]);
let fused = rrf_fuse_weighted(&[lexical, vector], &[2.0, 0.5], 60);
assert_eq!(fused.len(), 2);
let a_score = fused.iter().find(|r| r.id == "a").unwrap().score;
let b_score = fused.iter().find(|r| r.id == "b").unwrap().score;
assert!(
a_score > b_score,
"higher-weighted source must dominate: a={a_score}, b={b_score}"
);
}
#[test]
fn weighted_uniform_matches_rrf_fuse() {
let a = make_results(&[("x", 1.0), ("y", 0.5)]);
let b = make_results(&[("y", 0.9), ("z", 0.4)]);
let plain = rrf_fuse(&[a.clone(), b.clone()], 60);
let weighted = rrf_fuse_weighted(&[a, b], &[1.0, 1.0], 60);
assert_eq!(plain.len(), weighted.len());
for (p, w) in plain.iter().zip(weighted.iter()) {
assert_eq!(p.id, w.id);
assert!((p.score - w.score).abs() < 1e-9);
}
}
#[test]
#[should_panic(expected = "result_sets and weights length must match")]
fn weighted_length_mismatch_panics() {
let a = make_results(&[("x", 1.0)]);
rrf_fuse_weighted(&[a], &[1.0, 2.0], 60);
}
}