use rustc_hash::FxHashMap;
use super::vector::MultiValueCombiner;
use super::{ScoredPosition, SearchResult, compare_search_results_desc};
pub const DEFAULT_RRF_K: f32 = 60.0;
pub const MAX_FUSION_SUB_QUERIES: usize = 16;
pub const MAX_FUSION_CANDIDATE_SLOTS: usize = 200_000;
pub const MAX_FUSION_CHUNK_SLOTS: usize = 500_000;
#[derive(Debug, Clone, Copy, PartialEq)]
pub enum FusionMethod {
Rrf { k: f32 },
NormalizedWeightedSum,
}
impl Default for FusionMethod {
fn default() -> Self {
FusionMethod::Rrf { k: DEFAULT_RRF_K }
}
}
#[inline]
pub(crate) fn rrf_contribution(k: f32, rank: usize) -> f32 {
1.0 / (k + rank as f32)
}
pub fn fuse_ranked_lists(
lists: Vec<(Vec<SearchResult>, f32)>,
method: FusionMethod,
limit: usize,
) -> Vec<SearchResult> {
const MAX_INITIAL_FUSION_CAPACITY: usize = 200_000;
let capacity = lists
.iter()
.map(|(list, _)| list.len())
.fold(0usize, usize::saturating_add)
.min(MAX_INITIAL_FUSION_CAPACITY);
let mut fused: FxHashMap<(u128, u32), SearchResult> =
FxHashMap::with_capacity_and_hasher(capacity, Default::default());
for (list, weight) in lists {
let (min_score, inv_range) = match method {
FusionMethod::NormalizedWeightedSum if !list.is_empty() => {
let mut min = f32::INFINITY;
let mut max = f32::NEG_INFINITY;
for r in &list {
min = min.min(r.score);
max = max.max(r.score);
}
let range = max - min;
(min, if range > 0.0 { 1.0 / range } else { 0.0 })
}
_ => (0.0, 0.0),
};
for (idx, result) in list.into_iter().enumerate() {
let contribution = match method {
FusionMethod::Rrf { k } => weight * rrf_contribution(k, idx + 1),
FusionMethod::NormalizedWeightedSum => {
if inv_range > 0.0 {
weight * (result.score - min_score) * inv_range
} else {
weight
}
}
};
fused
.entry((result.segment_id, result.doc_id))
.and_modify(|r| r.score += contribution)
.or_insert_with(|| SearchResult {
score: contribution,
..result
});
}
}
let mut results: Vec<SearchResult> = fused.into_values().collect();
if results.len() > limit {
results.select_nth_unstable_by(limit, compare_search_results_desc);
results.truncate(limit);
}
results.sort_unstable_by(compare_search_results_desc);
results
}
pub fn fuse_ranked_lists_chunked(
lists: Vec<(Vec<SearchResult>, f32)>,
method: FusionMethod,
combiner: MultiValueCombiner,
limit: usize,
) -> Vec<SearchResult> {
type ChunkKey = (u128, u32, u32);
let mut contributions: Vec<(ChunkKey, u16, f32)> = Vec::new();
let mut chunks: Vec<(ChunkKey, f32)> = Vec::new();
for (list_index, (list, weight)) in lists.into_iter().enumerate() {
chunks.clear();
for result in &list {
let mut had_positions = false;
for (_field_id, scored_positions) in &result.positions {
for sp in scored_positions {
had_positions = true;
chunks.push(((result.segment_id, result.doc_id, sp.position), sp.score));
}
}
if !had_positions {
chunks.push(((result.segment_id, result.doc_id, 0), result.score));
}
}
if chunks.is_empty() {
continue;
}
chunks.sort_unstable_by(|a, b| b.1.total_cmp(&a.1).then_with(|| a.0.cmp(&b.0)));
let (min_score, inv_range) = match method {
FusionMethod::NormalizedWeightedSum => {
let max = chunks.first().map(|c| c.1).unwrap_or(0.0);
let min = chunks.last().map(|c| c.1).unwrap_or(0.0);
let range = max - min;
(min, if range > 0.0 { 1.0 / range } else { 0.0 })
}
_ => (0.0, 0.0),
};
contributions.reserve(chunks.len());
let list_index = list_index as u16;
for (rank, &(key, score)) in chunks.iter().enumerate() {
let contribution = match method {
FusionMethod::Rrf { k } => weight * rrf_contribution(k, rank + 1),
FusionMethod::NormalizedWeightedSum => {
if inv_range > 0.0 {
weight * (score - min_score) * inv_range
} else {
weight
}
}
};
contributions.push((key, list_index, contribution));
}
}
contributions.sort_unstable_by(|a, b| a.0.cmp(&b.0).then_with(|| a.1.cmp(&b.1)));
let mut results: Vec<SearchResult> = Vec::new();
let mut ordinals: Vec<(u32, f32)> = Vec::new();
let mut index = 0;
while index < contributions.len() {
let (segment_id, doc_id, _) = contributions[index].0;
ordinals.clear();
while index < contributions.len() {
let (seg, doc, ordinal) = contributions[index].0;
if seg != segment_id || doc != doc_id {
break;
}
let mut fused = 0.0f32;
while index < contributions.len()
&& contributions[index].0 == (segment_id, doc_id, ordinal)
{
fused += contributions[index].2;
index += 1;
}
ordinals.push((ordinal, fused));
}
let score = combiner.combine(&ordinals);
let scored_positions: Vec<ScoredPosition> = ordinals
.iter()
.map(|&(ord, s)| ScoredPosition::new(ord, s))
.collect();
results.push(SearchResult {
doc_id,
score,
segment_id,
positions: vec![(0, scored_positions)],
});
}
if results.len() > limit {
results.select_nth_unstable_by(limit, compare_search_results_desc);
results.truncate(limit);
}
results.sort_unstable_by(compare_search_results_desc);
results
}
pub fn try_fuse_ranked_lists_chunked(
lists: Vec<(Vec<SearchResult>, f32)>,
method: FusionMethod,
combiner: MultiValueCombiner,
limit: usize,
) -> Result<Vec<SearchResult>, String> {
if lists.is_empty() {
return Err("fusion requires at least one ranked list".to_string());
}
if lists.len() > MAX_FUSION_SUB_QUERIES {
return Err(format!(
"fusion supports at most {MAX_FUSION_SUB_QUERIES} ranked lists"
));
}
if let FusionMethod::Rrf { k } = method
&& (!k.is_finite() || k < 0.0)
{
return Err(format!(
"fusion RRF k must be finite and non-negative, got {k}"
));
}
combiner.validate()?;
let mut candidates = 0usize;
let mut chunks = 0usize;
for (list_index, (list, weight)) in lists.iter().enumerate() {
if !weight.is_finite() || *weight < 0.0 {
return Err(format!(
"fusion list weight at index {list_index} must be finite and non-negative, \
got {weight}"
));
}
candidates = candidates
.checked_add(list.len())
.ok_or_else(|| "fusion candidate count overflow".to_string())?;
if candidates > MAX_FUSION_CANDIDATE_SLOTS {
return Err(format!(
"fusion contains more than {MAX_FUSION_CANDIDATE_SLOTS} candidate slots"
));
}
for result in list {
let position_count = result
.positions
.iter()
.try_fold(0usize, |count, (_, positions)| {
count.checked_add(positions.len())
})
.ok_or_else(|| "fusion chunk count overflow".to_string())?;
chunks = chunks
.checked_add(position_count.max(1))
.ok_or_else(|| "fusion chunk count overflow".to_string())?;
if chunks > MAX_FUSION_CHUNK_SLOTS {
return Err(format!(
"fusion expands to more than {MAX_FUSION_CHUNK_SLOTS} ordinal chunks"
));
}
}
}
Ok(fuse_ranked_lists_chunked(lists, method, combiner, limit))
}
#[cfg(test)]
mod tests {
use super::*;
fn result(doc_id: u32, score: f32) -> SearchResult {
SearchResult {
doc_id,
score,
segment_id: 1,
positions: Vec::new(),
}
}
#[test]
fn test_rrf_union_includes_single_list_docs() {
let sparse = vec![result(1, 10.0), result(2, 5.0)];
let dense = vec![result(3, 0.9), result(1, 0.8)];
let fused = fuse_ranked_lists(
vec![(sparse, 1.0), (dense, 1.0)],
FusionMethod::Rrf { k: 60.0 },
10,
);
assert_eq!(fused.len(), 3);
assert_eq!(fused[0].doc_id, 1);
let expected = 1.0 / 61.0 + 1.0 / 62.0;
assert!((fused[0].score - expected).abs() < 1e-6);
let ids: Vec<u32> = fused.iter().map(|r| r.doc_id).collect();
assert!(ids.contains(&2) && ids.contains(&3));
}
#[test]
fn test_rrf_weights_scale_contribution() {
let a = vec![result(1, 1.0)];
let b = vec![result(2, 1.0)];
let fused = fuse_ranked_lists(vec![(a, 1.0), (b, 2.0)], FusionMethod::Rrf { k: 60.0 }, 10);
assert_eq!(fused[0].doc_id, 2);
assert!((fused[0].score - 2.0 / 61.0).abs() < 1e-6);
}
#[test]
fn test_normalized_weighted_sum() {
let sparse = vec![result(1, 20.0), result(2, 10.0), result(3, 0.0)];
let dense = vec![result(2, 0.99), result(1, 0.55), result(3, 0.11)];
let fused = fuse_ranked_lists(
vec![(sparse, 0.5), (dense, 0.5)],
FusionMethod::NormalizedWeightedSum,
10,
);
assert_eq!(fused.len(), 3);
assert_eq!(fused[0].doc_id, 1);
assert!((fused[0].score - 0.75).abs() < 1e-6);
assert!((fused[1].score - 0.75).abs() < 1e-6);
assert_eq!(fused[2].doc_id, 3);
assert!(fused[2].score.abs() < 1e-6);
}
#[test]
fn test_limit_truncation() {
let list: Vec<SearchResult> = (0..100).map(|i| result(i, 100.0 - i as f32)).collect();
let fused = fuse_ranked_lists(vec![(list, 1.0)], FusionMethod::default(), 5);
assert_eq!(fused.len(), 5);
assert_eq!(fused[0].doc_id, 0);
}
fn chunked(doc_id: u32, chunks: &[(u32, f32)]) -> SearchResult {
let positions = vec![(
0u32,
chunks
.iter()
.map(|&(ord, s)| ScoredPosition::new(ord, s))
.collect(),
)];
SearchResult {
doc_id,
score: chunks.iter().map(|&(_, s)| s).fold(0.0, f32::max),
segment_id: 1,
positions,
}
}
#[test]
fn test_chunked_fusion_junk_vertical_does_not_outvote() {
let sparse = vec![
chunked(1, &[(0, 10.0)]),
chunked(2, &[(0, 5.0)]),
chunked(3, &[(0, 4.0)]),
chunked(4, &[(0, 3.0)]),
chunked(9, &[(2, 2.0)]),
];
let dense = vec![
chunked(7, &[(0, 0.31)]),
chunked(8, &[(1, 0.30)]),
chunked(6, &[(0, 0.29)]),
chunked(5, &[(3, 0.28)]),
chunked(9, &[(5, 0.27)]),
];
let fused = fuse_ranked_lists_chunked(
vec![(sparse, 1.0), (dense, 1.0)],
FusionMethod::Rrf { k: 60.0 },
MultiValueCombiner::Max,
10,
);
assert_eq!(
fused[0].doc_id, 1,
"sparse rank-1 doc must win over doc 9 (present in both lists on different chunks)"
);
}
#[test]
fn test_chunked_fusion_same_chunk_corroboration_wins() {
let sparse = vec![chunked(1, &[(3, 9.0)]), chunked(2, &[(0, 8.0)])];
let dense = vec![chunked(1, &[(3, 0.9)]), chunked(2, &[(7, 0.8)])];
let fused = fuse_ranked_lists_chunked(
vec![(sparse, 1.0), (dense, 1.0)],
FusionMethod::Rrf { k: 60.0 },
MultiValueCombiner::Max,
10,
);
assert_eq!(fused[0].doc_id, 1);
let expected_doc1 = 2.0 / 61.0;
assert!((fused[0].score - expected_doc1).abs() < 1e-6);
assert!(fused[1].score < expected_doc1 / 1.9);
let (_, positions) = &fused[0].positions[0..1][0];
assert_eq!(positions.len(), 1);
assert_eq!(positions[0].position, 3, "fused chunk ordinal preserved");
}
#[test]
fn test_chunked_fusion_pseudo_chunk_for_docs_without_positions() {
let text = vec![result(1, 3.0), result(2, 2.0)]; let dense = vec![chunked(1, &[(0, 0.9)])];
let fused = fuse_ranked_lists_chunked(
vec![(text, 1.0), (dense, 1.0)],
FusionMethod::Rrf { k: 60.0 },
MultiValueCombiner::Max,
10,
);
assert_eq!(fused[0].doc_id, 1);
assert!((fused[0].score - 2.0 / 61.0).abs() < 1e-6);
assert_eq!(fused.len(), 2);
}
#[test]
fn test_validated_chunked_fusion_rejects_invalid_parameters() {
assert!(
try_fuse_ranked_lists_chunked(
vec![(vec![result(1, 1.0)], -1.0)],
FusionMethod::default(),
MultiValueCombiner::Max,
10,
)
.is_err()
);
assert!(
try_fuse_ranked_lists_chunked(
vec![(vec![result(1, 1.0)], 1.0)],
FusionMethod::Rrf { k: f32::NAN },
MultiValueCombiner::Max,
10,
)
.is_err()
);
}
#[test]
fn chunked_fusion_sort_grouping_matches_hash_map_reference() {
use rustc_hash::FxHashMap;
fn seg(mut result: SearchResult, segment_id: u128) -> SearchResult {
result.segment_id = segment_id;
result
}
let lists = vec![
(
vec![
chunked(1, &[(2, 9.0), (0, 8.5)]),
seg(chunked(1, &[(0, 7.0)]), 2),
chunked(5, &[(1, 6.0)]),
result(8, 5.0),
],
1.0,
),
(
vec![
chunked(5, &[(1, 0.9), (4, 0.8)]),
chunked(1, &[(0, 0.7)]),
seg(chunked(1, &[(0, 0.6)]), 2),
result(9, 0.5),
],
0.7,
),
(vec![chunked(1, &[(2, 3.0)]), chunked(8, &[(0, 2.0)])], 1.3),
];
let mut reference: FxHashMap<(u128, u32, u32), f32> = FxHashMap::default();
for (list, weight) in &lists {
let mut chunks: Vec<((u128, u32, u32), f32)> = Vec::new();
for r in list {
let mut had = false;
for (_, positions) in &r.positions {
for p in positions {
had = true;
chunks.push(((r.segment_id, r.doc_id, p.position), p.score));
}
}
if !had {
chunks.push(((r.segment_id, r.doc_id, 0), r.score));
}
}
chunks.sort_unstable_by(|a, b| b.1.total_cmp(&a.1).then_with(|| a.0.cmp(&b.0)));
for (rank, &(key, _)) in chunks.iter().enumerate() {
*reference.entry(key).or_insert(0.0) += weight * rrf_contribution(60.0, rank + 1);
}
}
let fused = fuse_ranked_lists_chunked(
lists,
FusionMethod::Rrf { k: 60.0 },
MultiValueCombiner::Max,
100,
);
let mut seen = 0;
for result in &fused {
let (_, positions) = &result.positions[0];
let ordinals: Vec<u32> = positions.iter().map(|p| p.position).collect();
let mut sorted = ordinals.clone();
sorted.sort_unstable();
assert_eq!(
ordinals, sorted,
"ordinals ascending for doc {}",
result.doc_id
);
let mut best = f32::NEG_INFINITY;
for p in positions {
let expected = reference[&(result.segment_id, result.doc_id, p.position)];
assert_eq!(
p.score.to_bits(),
expected.to_bits(),
"chunk ({}, {}, {}) fused score",
result.segment_id,
result.doc_id,
p.position
);
best = best.max(expected);
seen += 1;
}
assert_eq!(result.score.to_bits(), best.to_bits());
}
assert_eq!(seen, reference.len(), "every fused chunk is reported once");
assert_eq!(fused.len(), 5, "(1,seg1) (1,seg2) 5 8 9");
for pair in fused.windows(2) {
assert!(compare_search_results_desc(&pair[0], &pair[1]).is_le());
}
}
#[test]
fn test_duplicate_across_segments_not_merged() {
let mut a = result(1, 1.0);
a.segment_id = 1;
let mut b = result(1, 1.0);
b.segment_id = 2;
let fused = fuse_ranked_lists(
vec![(vec![a], 1.0), (vec![b], 1.0)],
FusionMethod::default(),
10,
);
assert_eq!(fused.len(), 2);
}
}