#![cfg_attr(not(feature = "embeddings"), allow(dead_code))]
use std::collections::HashMap;
use crate::context::Candidate;
pub const DEFAULT_RRF_K: f32 = 60.0;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Fusion {
Linear,
Rrf,
}
impl Fusion {
pub fn from_config(value: &str) -> Self {
match value.trim().to_ascii_lowercase().as_str() {
"rrf" => Fusion::Rrf,
_ => Fusion::Linear,
}
}
pub fn as_str(self) -> &'static str {
match self {
Fusion::Linear => "linear",
Fusion::Rrf => "rrf",
}
}
}
fn candidate_key(candidate: &Candidate) -> &str {
&candidate.capsule.expansion_handle
}
pub(crate) fn union_max(lists: Vec<Vec<Candidate>>) -> Vec<Candidate> {
let mut seen: HashMap<String, usize> = HashMap::new();
let mut merged: Vec<Candidate> = Vec::new();
for candidate in lists.into_iter().flatten() {
let key = candidate_key(&candidate).to_string();
match seen.get(&key) {
Some(&idx) => {
if candidate.raw_relevance > merged[idx].raw_relevance {
merged[idx] = candidate;
}
}
None => {
seen.insert(key, merged.len());
merged.push(candidate);
}
}
}
merged
}
pub(crate) fn rrf_fuse(lists: Vec<Vec<Candidate>>, k: f32) -> Vec<Candidate> {
let k = if k > 0.0 { k } else { DEFAULT_RRF_K };
let mut fused: HashMap<String, f32> = HashMap::new();
for list in &lists {
for (rank, candidate) in list.iter().enumerate() {
let contribution = 1.0 / (k + (rank as f32 + 1.0));
*fused
.entry(candidate_key(candidate).to_string())
.or_insert(0.0) += contribution;
}
}
let mut best = union_max(lists);
let max_fused = fused.values().copied().fold(0.0_f32, f32::max);
for candidate in &mut best {
let score = fused
.get(candidate_key(candidate))
.copied()
.unwrap_or_default();
candidate.raw_relevance = if max_fused > 0.0 {
score / max_fused
} else {
0.0
};
}
best.sort_by(|a, b| {
b.raw_relevance
.partial_cmp(&a.raw_relevance)
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| candidate_key(a).cmp(candidate_key(b)))
});
best
}
pub(crate) fn fuse(mode: Fusion, lists: Vec<Vec<Candidate>>) -> Vec<Candidate> {
match mode {
Fusion::Linear => union_max(lists),
Fusion::Rrf => rrf_fuse(lists, DEFAULT_RRF_K),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::context::{Candidate, ContextCapsule};
fn candidate(handle: &str, relevance: f32) -> Candidate {
Candidate {
capsule: ContextCapsule {
id: String::new(),
kind: "memory".to_string(),
summary: format!("project:fact - {handle}"),
token_estimate: 10,
expansion_handle: handle.to_string(),
provenance: Vec::new(),
confidence: 0.9,
freshness: 0.5,
relevance: 0.0,
scope_weight: 0.9,
score: 0.0,
superseded_hint: false,
rerank_policy_tier: 0,
claim_revision: None,
facts: vec![],
rerank_usefulness: None,
rerank_trust: None,
},
raw_relevance: relevance,
embedding: None,
cosine: None,
created_at: None,
}
}
fn handles(candidates: &[Candidate]) -> Vec<&str> {
candidates
.iter()
.map(|c| c.capsule.expansion_handle.as_str())
.collect()
}
#[test]
fn fusion_parses_and_falls_back_to_linear() {
assert_eq!(Fusion::from_config("rrf"), Fusion::Rrf);
assert_eq!(Fusion::from_config("RRF"), Fusion::Rrf);
assert_eq!(Fusion::from_config("linear"), Fusion::Linear);
assert_eq!(
Fusion::from_config("reciprocal-rank"),
Fusion::Linear,
"an unknown value degrades to the previous behaviour, it does not fail retrieval"
);
}
#[test]
fn union_max_keeps_the_best_instance_of_each_candidate() {
let merged = union_max(vec![
vec![candidate("memory:a", 0.4), candidate("memory:b", 0.9)],
vec![candidate("memory:a", 0.7)],
]);
assert_eq!(merged.len(), 2);
let a = merged
.iter()
.find(|c| c.capsule.expansion_handle == "memory:a")
.unwrap();
assert!((a.raw_relevance - 0.7).abs() < f32::EPSILON);
}
#[test]
fn rrf_rewards_agreement_between_lists() {
let lexical = vec![
candidate("memory:lexical_only", 0.95),
candidate("memory:both", 0.90),
];
let semantic = vec![
candidate("memory:semantic_only", 0.95),
candidate("memory:both", 0.90),
];
let fused = rrf_fuse(vec![lexical.clone(), semantic.clone()], DEFAULT_RRF_K);
assert_eq!(
handles(&fused)[0],
"memory:both",
"agreement across lists must win: {:?}",
handles(&fused)
);
let unioned = union_max(vec![lexical, semantic]);
let both_idx = handles(&unioned)
.iter()
.position(|h| *h == "memory:both")
.unwrap();
assert!(
both_idx > 0,
"union-max cannot see the agreement: {:?}",
handles(&unioned)
);
}
#[test]
fn rrf_normalizes_the_top_score_to_one() {
let fused = rrf_fuse(
vec![
vec![candidate("memory:a", 0.5), candidate("memory:b", 0.4)],
vec![candidate("memory:a", 0.3)],
],
DEFAULT_RRF_K,
);
assert!((fused[0].raw_relevance - 1.0).abs() < 1e-5);
assert!(fused[1].raw_relevance < 1.0);
}
#[test]
fn rrf_is_deterministic_on_ties() {
let build = || {
vec![
vec![candidate("memory:b", 0.5)],
vec![candidate("memory:a", 0.5)],
]
};
let first = handles(&rrf_fuse(build(), DEFAULT_RRF_K))
.iter()
.map(|s| s.to_string())
.collect::<Vec<_>>();
let second = handles(&rrf_fuse(build(), DEFAULT_RRF_K))
.iter()
.map(|s| s.to_string())
.collect::<Vec<_>>();
assert_eq!(first, second);
assert_eq!(first[0], "memory:a", "ties break on the stable key");
}
#[test]
fn rrf_over_one_list_preserves_its_order() {
let single = vec![
candidate("memory:first", 0.9),
candidate("memory:second", 0.5),
candidate("memory:third", 0.1),
];
let fused = rrf_fuse(vec![single], DEFAULT_RRF_K);
assert_eq!(
handles(&fused),
vec!["memory:first", "memory:second", "memory:third"]
);
}
#[test]
fn fuse_dispatches_on_the_mode() {
let lists = || {
vec![
vec![candidate("memory:a", 0.4)],
vec![candidate("memory:a", 0.8), candidate("memory:b", 0.2)],
]
};
let linear = fuse(Fusion::Linear, lists());
let a = linear
.iter()
.find(|c| c.capsule.expansion_handle == "memory:a")
.unwrap();
assert!(
(a.raw_relevance - 0.8).abs() < f32::EPSILON,
"linear keeps the raw blended score"
);
let rrf = fuse(Fusion::Rrf, lists());
let a = rrf
.iter()
.find(|c| c.capsule.expansion_handle == "memory:a")
.unwrap();
assert!(
(a.raw_relevance - 1.0).abs() < 1e-5,
"rrf rewrites relevance to the normalized fused score"
);
}
}