pub(crate) const RRF_K: f32 = 60.0;
pub(crate) const RETRIEVE_DEPTH: usize = 100;
pub(crate) type WeightedArm<'a> = (&'a [String], f32);
pub(crate) fn rrf_fuse_weighted(lists: &[WeightedArm<'_>], k: f32) -> Vec<(String, f32)> {
use std::collections::HashMap;
let mut scores: HashMap<&str, f32> = HashMap::new();
for (list, weight) in lists {
for (rank, id) in list.iter().enumerate() {
*scores.entry(id.as_str()).or_insert(0.0) += weight / (k + rank as f32);
}
}
let mut ranked: Vec<(String, f32)> = scores
.into_iter()
.map(|(id, score)| (id.to_string(), score))
.collect();
let len = ranked.len();
sort_and_truncate(&mut ranked, len);
ranked
}
pub(crate) fn sort_and_truncate(ranked: &mut Vec<(String, f32)>, top_k: usize) {
ranked.sort_by(|a, b| {
b.1.partial_cmp(&a.1)
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| a.0.cmp(&b.0))
});
ranked.truncate(top_k);
}
#[cfg(test)]
mod tests {
use super::*;
fn ids(v: &[&str]) -> Vec<String> {
v.iter().map(|s| s.to_string()).collect()
}
fn rrf_fuse(lists: &[&[String]], k: f32) -> Vec<(String, f32)> {
let weighted: Vec<WeightedArm<'_>> = lists.iter().map(|l| (*l, 1.0)).collect();
rrf_fuse_weighted(&weighted, k)
}
#[test]
fn empty_input_yields_no_hits() {
assert!(rrf_fuse(&[], RRF_K).is_empty());
let empty: Vec<String> = Vec::new();
assert!(rrf_fuse(&[empty.as_slice()], RRF_K).is_empty());
}
#[test]
fn single_list_preserves_order_with_reciprocal_scores() {
let list = ids(&["a", "b", "c"]);
let fused = rrf_fuse(&[list.as_slice()], RRF_K);
assert_eq!(
fused.iter().map(|(id, _)| id.as_str()).collect::<Vec<_>>(),
vec!["a", "b", "c"]
);
assert!((fused[0].1 - 1.0 / 60.0).abs() < 1e-6);
assert!((fused[1].1 - 1.0 / 61.0).abs() < 1e-6);
assert!((fused[2].1 - 1.0 / 62.0).abs() < 1e-6);
}
#[test]
fn doc_in_both_lists_outranks_doc_in_one() {
let bm25 = ids(&["solo", "shared"]);
let dense = ids(&["other", "shared"]);
let fused = rrf_fuse(&[bm25.as_slice(), dense.as_slice()], RRF_K);
assert_eq!(fused.first().map(|(id, _)| id.as_str()), Some("shared"));
}
#[test]
fn tied_scores_break_by_id_ascending() {
let l1 = ids(&["zeta"]);
let l2 = ids(&["alpha"]);
let l3 = ids(&["mid"]);
let fused = rrf_fuse(&[l1.as_slice(), l2.as_slice(), l3.as_slice()], RRF_K);
assert_eq!(
fused.iter().map(|(id, _)| id.as_str()).collect::<Vec<_>>(),
vec!["alpha", "mid", "zeta"]
);
}
#[test]
fn arg_order_does_not_change_the_result() {
let bm25 = ids(&["a", "b", "c"]);
let dense = ids(&["b", "c", "d"]);
let forward = rrf_fuse(&[bm25.as_slice(), dense.as_slice()], RRF_K);
let swapped = rrf_fuse(&[dense.as_slice(), bm25.as_slice()], RRF_K);
assert_eq!(forward, swapped);
}
#[test]
fn baseline_weights_reproduce_the_classic_rrf_scores() {
let bm25 = ids(&["a", "b"]);
let dense = ids(&["b", "a"]);
let fused = rrf_fuse_weighted(&[(bm25.as_slice(), 1.0), (dense.as_slice(), 1.0)], RRF_K);
let expected = 1.0 / 60.0 + 1.0 / 61.0;
for (id, score) in &fused {
assert!((score - expected).abs() < 1e-6, "{id} scored {score}");
}
}
#[test]
fn a_zero_weight_arm_contributes_nothing() {
let bm25 = ids(&["a", "b"]);
let muted = ids(&["z"]);
let fused = rrf_fuse_weighted(&[(bm25.as_slice(), 1.0), (muted.as_slice(), 0.0)], RRF_K);
assert_eq!(fused.last().map(|(id, _)| id.as_str()), Some("z"));
assert_eq!(fused.last().map(|(_, s)| *s), Some(0.0));
}
#[test]
fn a_heavier_arm_outvotes_a_lighter_one_at_the_same_rank() {
let light = ids(&["lex"]);
let heavy = ids(&["usage"]);
let fused = rrf_fuse_weighted(&[(light.as_slice(), 1.0), (heavy.as_slice(), 2.0)], RRF_K);
assert_eq!(fused.first().map(|(id, _)| id.as_str()), Some("usage"));
}
#[test]
fn sub_unit_usage_weight_lets_the_lexical_arm_win_at_equal_rank() {
let lexical = ids(&["a", "b", "d"]);
let usage = ids(&["a", "b", "c"]);
let fused = rrf_fuse_weighted(&[(lexical.as_slice(), 1.0), (usage.as_slice(), 0.5)], RRF_K);
let order: Vec<&str> = fused.iter().map(|(id, _)| id.as_str()).collect();
assert_eq!(order, vec!["a", "b", "d", "c"]);
}
#[test]
fn a_sub_unit_arm_still_promotes_a_deeply_ranked_id_past_the_lexical_top_hit() {
let mut lexical = vec!["docker_build".to_string()];
lexical.extend((0..49).map(|i| format!("filler{i:02}")));
lexical.push("gh_run_list".to_string()); let usage = ids(&["gh_run_list"]);
let fused = rrf_fuse_weighted(&[(lexical.as_slice(), 1.0), (usage.as_slice(), 0.5)], RRF_K);
assert_eq!(
fused.first().map(|(id, _)| id.as_str()),
Some("gh_run_list")
);
}
#[test]
fn sort_and_truncate_keeps_stable_membership_across_a_tie_boundary() {
let mut ranked = vec![
("zeta".to_string(), 1.0_f32),
("alpha".to_string(), 1.0),
("mid".to_string(), 1.0),
];
sort_and_truncate(&mut ranked, 2);
assert_eq!(
ranked.iter().map(|(id, _)| id.as_str()).collect::<Vec<_>>(),
vec!["alpha", "mid"]
);
}
}