use crate::search::query_classifier::{FusionWeights, QueryType};
use crate::types::{Confidence, ScoredChunkId};
use std::collections::HashMap;
use std::sync::LazyLock;
pub const RRF_K: f32 = 60.0;
pub const AMBIGUOUS_GAP_THRESHOLD: f32 = 0.05;
#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)]
pub enum FusionMode {
#[default]
Rrf,
WeightedRrf,
Cc,
}
impl FusionMode {
pub fn from_env_value(s: &str) -> Self {
match s.trim().to_ascii_lowercase().as_str() {
"weighted-rrf" | "wrrf" => Self::WeightedRrf,
"cc" | "convex" | "weighted" => Self::Cc,
_ => Self::Rrf,
}
}
}
static FUSION_MODE: LazyLock<FusionMode> = LazyLock::new(|| {
std::env::var("SEMANTEX_FUSION")
.ok()
.map_or(FusionMode::default(), |v| FusionMode::from_env_value(&v))
});
pub fn active_fusion_mode() -> FusionMode {
*FUSION_MODE
}
#[derive(Debug, Clone, Copy)]
pub struct TripleFusionWeights {
pub w_dense: f32,
pub w_sparse: f32,
pub w_exact: f32,
}
impl TripleFusionWeights {
pub fn max_possible(&self) -> f32 {
self.w_dense + self.w_sparse + self.w_exact
}
}
struct CachedWeightOverrides {
identifier: Option<TripleFusionWeights>,
keyword: Option<TripleFusionWeights>,
semantic: Option<TripleFusionWeights>,
mixed: Option<TripleFusionWeights>,
}
fn parse_weight_override(env_key: &str) -> Option<TripleFusionWeights> {
let val = std::env::var(env_key).ok()?;
let raw: Vec<&str> = val.split(',').map(str::trim).collect();
let mut parts: Vec<f32> = Vec::with_capacity(raw.len());
let mut rejected = 0usize;
for token in &raw {
match token.parse::<f32>() {
Ok(x) if x.is_finite() => parts.push(x),
Ok(_) | Err(_) => rejected += 1,
}
}
if rejected > 0 {
tracing::warn!(
env_var = env_key,
value = %val,
rejected,
"ignoring non-finite or unparseable values in fusion weight override; \
falling back to defaults unless exactly 3 finite values remain"
);
}
if parts.len() == 3 {
Some(TripleFusionWeights {
w_dense: parts[0],
w_sparse: parts[1],
w_exact: parts[2],
})
} else {
None
}
}
static WEIGHT_OVERRIDES: LazyLock<CachedWeightOverrides> =
LazyLock::new(|| CachedWeightOverrides {
identifier: parse_weight_override("SEMANTEX_WEIGHTS_IDENTIFIER"),
keyword: parse_weight_override("SEMANTEX_WEIGHTS_KEYWORD"),
semantic: parse_weight_override("SEMANTEX_WEIGHTS_SEMANTIC"),
mixed: parse_weight_override("SEMANTEX_WEIGHTS_MIXED"),
});
impl QueryType {
pub fn triple_fusion_weights(self) -> TripleFusionWeights {
let cached = match self {
QueryType::Identifier => WEIGHT_OVERRIDES.identifier,
QueryType::Keyword => WEIGHT_OVERRIDES.keyword,
QueryType::Semantic => WEIGHT_OVERRIDES.semantic,
QueryType::Mixed => WEIGHT_OVERRIDES.mixed,
};
if let Some(weights) = cached {
return weights;
}
match self {
QueryType::Identifier => TripleFusionWeights {
w_dense: 0.2,
w_sparse: 0.6,
w_exact: 5.0,
},
QueryType::Keyword => TripleFusionWeights {
w_dense: 0.3,
w_sparse: 0.6,
w_exact: 2.0,
},
QueryType::Semantic => TripleFusionWeights {
w_dense: 0.4,
w_sparse: 0.5,
w_exact: 0.8,
},
QueryType::Mixed => TripleFusionWeights {
w_dense: 0.5,
w_sparse: 0.4,
w_exact: 0.8,
},
}
}
}
fn top_score_normalize(list: &[ScoredChunkId]) -> Vec<(u64, f32)> {
if list.is_empty() {
return Vec::new();
}
if list.len() == 1 {
return vec![(list[0].chunk_id, 1.0)];
}
let max = list
.iter()
.map(|s| s.score)
.fold(f32::NEG_INFINITY, f32::max);
if max <= f32::EPSILON {
return list.iter().map(|s| (s.chunk_id, 0.0)).collect();
}
list.iter().map(|s| (s.chunk_id, s.score / max)).collect()
}
fn adapt_weights(
weights: TripleFusionWeights,
dense_top: f32,
sparse_top: f32,
exact_top: f32,
) -> TripleFusionWeights {
let mut w = weights;
if exact_top > 0.8 {
w.w_exact *= 1.5;
}
if sparse_top > 0.0 && dense_top > 0.0 {
let ratio = sparse_top / dense_top;
if ratio > 2.5 {
w.w_sparse *= 1.3;
w.w_dense *= 0.8;
} else if ratio < 0.4 {
w.w_dense *= 1.3;
w.w_sparse *= 0.8;
}
}
w
}
pub fn triple_cc_fuse(
dense_list: &[ScoredChunkId],
sparse_list: &[ScoredChunkId],
exact_ids: &[u64],
weights: &TripleFusionWeights,
) -> Vec<ScoredChunkId> {
struct ChannelScores {
total: f32,
dense: f32,
sparse: f32,
exact: f32,
}
let dense_normalized = top_score_normalize(dense_list);
let sparse_normalized = top_score_normalize(sparse_list);
let dense_top = dense_normalized
.iter()
.map(|&(_, s)| s)
.fold(0.0_f32, f32::max);
let sparse_top = sparse_normalized
.iter()
.map(|&(_, s)| s)
.fold(0.0_f32, f32::max);
let exact_top = if exact_ids.is_empty() { 0.0 } else { 1.0 };
let weights = adapt_weights(*weights, dense_top, sparse_top, exact_top);
let mut scores: HashMap<u64, ChannelScores> = HashMap::new();
for (chunk_id, norm_score) in dense_normalized {
let entry = scores.entry(chunk_id).or_insert(ChannelScores {
total: 0.0,
dense: 0.0,
sparse: 0.0,
exact: 0.0,
});
entry.dense = norm_score;
entry.total += weights.w_dense * norm_score;
}
for (chunk_id, norm_score) in sparse_normalized {
let entry = scores.entry(chunk_id).or_insert(ChannelScores {
total: 0.0,
dense: 0.0,
sparse: 0.0,
exact: 0.0,
});
entry.sparse = norm_score;
entry.total += weights.w_sparse * norm_score;
}
for &chunk_id in exact_ids {
let entry = scores.entry(chunk_id).or_insert(ChannelScores {
total: 0.0,
dense: 0.0,
sparse: 0.0,
exact: 0.0,
});
entry.exact = 1.0;
entry.total += weights.w_exact;
}
let mut fused: Vec<ScoredChunkId> = scores
.into_iter()
.map(|(chunk_id, cs)| ScoredChunkId {
chunk_id,
score: cs.total,
score_dense: cs.dense,
score_sparse: cs.sparse,
score_exact: cs.exact,
})
.collect();
fused.sort_by(|a, b| b.score.total_cmp(&a.score));
fused
}
#[derive(Debug, Clone)]
pub struct RrfFusedResult {
pub scored: ScoredChunkId,
pub channels_hit: u32,
pub channels_fired: u32,
}
impl RrfFusedResult {
pub fn confidence(&self, next_score: Option<f32>) -> Confidence {
let extracted = self.channels_fired >= 2 && self.channels_hit == self.channels_fired;
if !extracted
&& let Some(next) = next_score
&& self.scored.score > 0.0
{
let gap = (self.scored.score - next).abs() / self.scored.score;
if gap < AMBIGUOUS_GAP_THRESHOLD {
return Confidence::Ambiguous;
}
}
if extracted {
Confidence::Extracted
} else {
Confidence::Inferred
}
}
pub fn confidence_score(&self) -> f32 {
if self.channels_fired == 0 {
0.0
} else {
self.channels_hit as f32 / self.channels_fired as f32
}
}
}
struct RrfAccum {
total: f32,
dense: f32,
sparse: f32,
exact: f32,
channels_hit_mask: u32,
}
impl RrfAccum {
fn new() -> Self {
Self {
total: 0.0,
dense: 0.0,
sparse: 0.0,
exact: 0.0,
channels_hit_mask: 0,
}
}
}
fn accumulate_rrf_channel(
scores: &mut HashMap<u64, RrfAccum>,
ranked: &[ScoredChunkId],
channel_bit: u32,
score_field: ChannelKind,
) {
for (rank, item) in ranked.iter().enumerate() {
let contribution = 1.0 / (RRF_K + rank as f32 + 1.0);
let entry = scores.entry(item.chunk_id).or_insert_with(RrfAccum::new);
entry.total += contribution;
entry.channels_hit_mask |= channel_bit;
match score_field {
ChannelKind::Dense => entry.dense += contribution,
ChannelKind::Sparse => entry.sparse += contribution,
}
}
}
fn accumulate_rrf_exact(scores: &mut HashMap<u64, RrfAccum>, ids: &[u64], channel_bit: u32) {
for (rank, &id) in ids.iter().enumerate() {
let contribution = 1.0 / (RRF_K + rank as f32 + 1.0);
let entry = scores.entry(id).or_insert_with(RrfAccum::new);
entry.total += contribution;
entry.exact += contribution;
entry.channels_hit_mask |= channel_bit;
}
}
fn accumulate_weighted_rrf_channel(
scores: &mut HashMap<u64, RrfAccum>,
ranked: &[ScoredChunkId],
channel_bit: u32,
score_field: ChannelKind,
weight: f32,
k: f32,
) {
for (rank, item) in ranked.iter().enumerate() {
let contribution = weight * (1.0 / (k + rank as f32 + 1.0));
let entry = scores.entry(item.chunk_id).or_insert_with(RrfAccum::new);
entry.total += contribution;
entry.channels_hit_mask |= channel_bit;
match score_field {
ChannelKind::Dense => entry.dense += contribution,
ChannelKind::Sparse => entry.sparse += contribution,
}
}
}
fn accumulate_weighted_rrf_exact(
scores: &mut HashMap<u64, RrfAccum>,
ids: &[u64],
channel_bit: u32,
exact_weight: f32,
k: f32,
) {
for (rank, &id) in ids.iter().enumerate() {
let contribution = exact_weight * (1.0 / (k + rank as f32 + 1.0));
let entry = scores.entry(id).or_insert_with(RrfAccum::new);
entry.total += contribution;
entry.exact += contribution;
entry.channels_hit_mask |= channel_bit;
}
}
#[derive(Debug, Clone, Copy)]
enum ChannelKind {
Dense,
Sparse,
}
pub fn triple_rrf_fuse(
dense_list: &[ScoredChunkId],
sparse_list: &[ScoredChunkId],
exact_ids: &[u64],
) -> Vec<RrfFusedResult> {
let mut scores: HashMap<u64, RrfAccum> = HashMap::new();
let dense_fired = !dense_list.is_empty();
let sparse_fired = !sparse_list.is_empty();
let exact_fired = !exact_ids.is_empty();
if dense_fired {
accumulate_rrf_channel(&mut scores, dense_list, 0b001, ChannelKind::Dense);
}
if sparse_fired {
accumulate_rrf_channel(&mut scores, sparse_list, 0b010, ChannelKind::Sparse);
}
if exact_fired {
accumulate_rrf_exact(&mut scores, exact_ids, 0b100);
}
let channels_fired = u32::from(dense_fired) + u32::from(sparse_fired) + u32::from(exact_fired);
finalize_rrf(scores, channels_fired)
}
pub fn exp4_rrf_fuse(
orig_dense: &[ScoredChunkId],
orig_sparse: &[ScoredChunkId],
exp_dense: &[ScoredChunkId],
exp_sparse: &[ScoredChunkId],
exact_ids: &[u64],
) -> Vec<RrfFusedResult> {
let mut scores: HashMap<u64, RrfAccum> = HashMap::new();
let orig_dense_active = !orig_dense.is_empty();
let orig_sparse_active = !orig_sparse.is_empty();
let exp_dense_active = !exp_dense.is_empty();
let exp_sparse_active = !exp_sparse.is_empty();
let exact_active = !exact_ids.is_empty();
if orig_dense_active {
accumulate_rrf_channel(&mut scores, orig_dense, 0b0_0001, ChannelKind::Dense);
}
if orig_sparse_active {
accumulate_rrf_channel(&mut scores, orig_sparse, 0b0_0010, ChannelKind::Sparse);
}
if exp_dense_active {
accumulate_rrf_channel(&mut scores, exp_dense, 0b0_0100, ChannelKind::Dense);
}
if exp_sparse_active {
accumulate_rrf_channel(&mut scores, exp_sparse, 0b0_1000, ChannelKind::Sparse);
}
if exact_active {
accumulate_rrf_exact(&mut scores, exact_ids, 0b1_0000);
}
let channels_fired = u32::from(orig_dense_active)
+ u32::from(orig_sparse_active)
+ u32::from(exp_dense_active)
+ u32::from(exp_sparse_active)
+ u32::from(exact_active);
finalize_rrf(scores, channels_fired)
}
pub fn triple_weighted_rrf_fuse(
dense_list: &[ScoredChunkId],
sparse_list: &[ScoredChunkId],
exact_ids: &[u64],
weights: FusionWeights,
k: f32,
) -> Vec<RrfFusedResult> {
let mut scores: HashMap<u64, RrfAccum> = HashMap::new();
let dense_fired = !dense_list.is_empty();
let sparse_fired = !sparse_list.is_empty();
let exact_fired = !exact_ids.is_empty();
if dense_fired {
accumulate_weighted_rrf_channel(
&mut scores,
dense_list,
0b001,
ChannelKind::Dense,
weights.w_dense,
k,
);
}
if sparse_fired {
accumulate_weighted_rrf_channel(
&mut scores,
sparse_list,
0b010,
ChannelKind::Sparse,
weights.w_sparse,
k,
);
}
if exact_fired {
accumulate_weighted_rrf_exact(&mut scores, exact_ids, 0b100, 1.0, k);
}
let channels_fired = u32::from(dense_fired) + u32::from(sparse_fired) + u32::from(exact_fired);
finalize_rrf(scores, channels_fired)
}
pub fn exp4_weighted_rrf_fuse(
orig_dense: &[ScoredChunkId],
orig_sparse: &[ScoredChunkId],
exp_dense: &[ScoredChunkId],
exp_sparse: &[ScoredChunkId],
exact_ids: &[u64],
weights: FusionWeights,
k: f32,
) -> Vec<RrfFusedResult> {
let mut scores: HashMap<u64, RrfAccum> = HashMap::new();
let orig_dense_active = !orig_dense.is_empty();
let orig_sparse_active = !orig_sparse.is_empty();
let exp_dense_active = !exp_dense.is_empty();
let exp_sparse_active = !exp_sparse.is_empty();
let exact_active = !exact_ids.is_empty();
if orig_dense_active {
accumulate_weighted_rrf_channel(
&mut scores,
orig_dense,
0b0_0001,
ChannelKind::Dense,
weights.w_dense,
k,
);
}
if orig_sparse_active {
accumulate_weighted_rrf_channel(
&mut scores,
orig_sparse,
0b0_0010,
ChannelKind::Sparse,
weights.w_sparse,
k,
);
}
if exp_dense_active {
accumulate_weighted_rrf_channel(
&mut scores,
exp_dense,
0b0_0100,
ChannelKind::Dense,
weights.w_dense,
k,
);
}
if exp_sparse_active {
accumulate_weighted_rrf_channel(
&mut scores,
exp_sparse,
0b0_1000,
ChannelKind::Sparse,
weights.w_sparse,
k,
);
}
if exact_active {
accumulate_weighted_rrf_exact(&mut scores, exact_ids, 0b1_0000, 1.0, k);
}
let channels_fired = u32::from(orig_dense_active)
+ u32::from(orig_sparse_active)
+ u32::from(exp_dense_active)
+ u32::from(exp_sparse_active)
+ u32::from(exact_active);
finalize_rrf(scores, channels_fired)
}
fn finalize_rrf(scores: HashMap<u64, RrfAccum>, channels_fired: u32) -> Vec<RrfFusedResult> {
let mut fused: Vec<RrfFusedResult> = scores
.into_iter()
.map(|(chunk_id, acc)| {
let channels_hit = acc.channels_hit_mask.count_ones();
RrfFusedResult {
scored: ScoredChunkId {
chunk_id,
score: acc.total,
score_dense: acc.dense,
score_sparse: acc.sparse,
score_exact: acc.exact,
},
channels_hit,
channels_fired,
}
})
.collect();
fused.sort_by(|a, b| b.scored.score.total_cmp(&a.scored.score));
fused
}
pub fn assign_confidence(fused: &[RrfFusedResult]) -> Vec<(Confidence, f32)> {
fused
.iter()
.enumerate()
.map(|(i, r)| {
let next = fused.get(i + 1).map(|n| n.scored.score);
(r.confidence(next), r.confidence_score())
})
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
fn s(id: u64, score: f32) -> ScoredChunkId {
ScoredChunkId::new(id, score)
}
#[test]
fn test_weighted_accumulate_scales_by_weight_and_k() {
let mut scores: HashMap<u64, RrfAccum> = HashMap::new();
let dense = vec![s(1, 0.9), s(2, 0.5)];
accumulate_weighted_rrf_channel(&mut scores, &dense, 0b001, ChannelKind::Dense, 2.0, 30.0);
let c1 = &scores[&1];
let c2 = &scores[&2];
assert!(
(c1.total - 2.0 / 31.0).abs() < 1e-6,
"rank0 total = {}",
c1.total
);
assert!(
(c1.dense - 2.0 / 31.0).abs() < 1e-6,
"rank0 dense = {}",
c1.dense
);
assert!(
(c2.total - 2.0 / 32.0).abs() < 1e-6,
"rank1 total = {}",
c2.total
);
assert_eq!(c1.channels_hit_mask, 0b001);
}
#[test]
fn test_weighted_accumulate_exact_scales_by_weight() {
let mut scores: HashMap<u64, RrfAccum> = HashMap::new();
accumulate_weighted_rrf_exact(&mut scores, &[42, 7], 0b100, 3.0, 60.0);
assert!((scores[&42].total - 3.0 / 61.0).abs() < 1e-6);
assert!((scores[&42].exact - 3.0 / 61.0).abs() < 1e-6);
}
#[test]
fn test_triple_weighted_rrf_sparse_weight_lifts_sparse_only_chunk() {
let weights = FusionWeights {
w_dense: 0.2,
w_sparse: 1.0,
};
let dense = vec![s(1, 0.9)]; let sparse = vec![s(2, 10.0)]; let k = 60.0;
let fused = triple_weighted_rrf_fuse(&dense, &sparse, &[], weights, k);
assert_eq!(
fused[0].scored.chunk_id, 2,
"sparse-weighted chunk must win"
);
assert!(fused[0].scored.score > fused[1].scored.score);
}
#[test]
fn test_triple_weighted_rrf_consensus_and_confidence_preserved() {
let weights = FusionWeights {
w_dense: 0.5,
w_sparse: 0.5,
};
let dense = vec![s(5, 0.9), s(10, 0.5)];
let sparse = vec![s(5, 10.0), s(20, 5.0)];
let exact = vec![5u64, 30];
let fused = triple_weighted_rrf_fuse(&dense, &sparse, &exact, weights, 60.0);
assert_eq!(fused[0].scored.chunk_id, 5);
assert_eq!(fused[0].channels_hit, 3);
assert_eq!(fused[0].channels_fired, 3);
assert_eq!(fused[0].confidence(None), Confidence::Extracted);
}
#[test]
fn test_exp4_weighted_rrf_falls_back_on_empty_expansion() {
let weights = FusionWeights {
w_dense: 0.4,
w_sparse: 0.8,
};
let dense = vec![s(5, 0.9)];
let sparse = vec![s(7, 10.0)];
let exact = vec![5u64];
let triple = triple_weighted_rrf_fuse(&dense, &sparse, &exact, weights, 60.0);
let exp4 = exp4_weighted_rrf_fuse(&dense, &sparse, &[], &[], &exact, weights, 60.0);
assert_eq!(triple.len(), exp4.len());
assert_eq!(triple[0].channels_fired, exp4[0].channels_fired);
assert_eq!(triple[0].channels_fired, 3);
}
#[test]
fn test_top_score_normalize_empty() {
let result = top_score_normalize(&[]);
assert!(result.is_empty());
}
#[test]
fn test_top_score_normalize_single() {
let result = top_score_normalize(&[s(42, 0.5)]);
assert_eq!(result.len(), 1);
assert_eq!(result[0].0, 42);
assert!((result[0].1 - 1.0).abs() < f32::EPSILON);
}
#[test]
fn test_top_score_normalize_all_same() {
let list = vec![s(1, 5.0), s(2, 5.0), s(3, 5.0)];
let result = top_score_normalize(&list);
assert_eq!(result.len(), 3);
for (_, norm) in &result {
assert!((*norm - 1.0).abs() < f32::EPSILON);
}
}
#[test]
fn test_top_score_normalize_normal() {
let list = vec![s(1, 10.0), s(2, 5.0), s(3, 0.0)];
let result = top_score_normalize(&list);
assert_eq!(result.len(), 3);
assert_eq!(result[0].0, 1);
assert!((result[0].1 - 1.0).abs() < f32::EPSILON);
assert_eq!(result[1].0, 2);
assert!((result[1].1 - 0.5).abs() < f32::EPSILON);
assert_eq!(result[2].0, 3);
assert!(result[2].1.abs() < f32::EPSILON);
}
#[test]
fn test_triple_cc_fuse_basic() {
let weights = TripleFusionWeights {
w_dense: 1.0,
w_sparse: 1.0,
w_exact: 1.0,
};
let dense = vec![s(1, 0.9), s(2, 0.1)];
let sparse = vec![s(2, 10.0), s(1, 1.0)];
let result = triple_cc_fuse(&dense, &sparse, &[], &weights);
assert_eq!(result.len(), 2);
assert_eq!(result[0].chunk_id, 2);
assert!(result[0].score > result[1].score);
}
#[test]
fn test_triple_cc_fuse_exact_boost() {
let weights = TripleFusionWeights {
w_dense: 0.5,
w_sparse: 0.5,
w_exact: 2.0,
};
let dense = vec![s(1, 0.9), s(2, 0.8)];
let result = triple_cc_fuse(&dense, &[], &[2], &weights);
assert_eq!(result[0].chunk_id, 2);
let expected_chunk2 = 0.5 * (0.8_f32 / 0.9) + 3.0;
assert!((result[0].score - expected_chunk2).abs() < 1e-5);
assert_eq!(result[1].chunk_id, 1);
assert!((result[1].score - 0.5).abs() < f32::EPSILON);
}
#[test]
fn test_triple_cc_fuse_multi_source() {
let weights = TripleFusionWeights {
w_dense: 1.0,
w_sparse: 1.0,
w_exact: 1.0,
};
let dense = vec![s(5, 0.9), s(10, 0.5)];
let sparse = vec![s(5, 10.0), s(20, 5.0)];
let result = triple_cc_fuse(&dense, &sparse, &[5], &weights);
assert_eq!(result[0].chunk_id, 5);
assert!((result[0].score - 3.5).abs() < 1e-5);
assert!(result[0].score_dense > 0.0, "chunk 5 should have dense > 0");
assert!(
result[0].score_sparse > 0.0,
"chunk 5 should have sparse > 0"
);
assert!(
(result[0].score_exact - 1.0).abs() < f32::EPSILON,
"chunk 5 exact should be 1.0"
);
}
#[test]
fn test_triple_cc_fuse_empty_sources() {
let weights = TripleFusionWeights {
w_dense: 1.0,
w_sparse: 1.0,
w_exact: 1.0,
};
let result = triple_cc_fuse(&[], &[], &[], &weights);
assert!(result.is_empty());
let result = triple_cc_fuse(&[], &[], &[10, 20], &weights);
assert_eq!(result.len(), 2);
assert!((result[0].score - 1.5).abs() < f32::EPSILON);
let result = triple_cc_fuse(&[s(1, 0.5)], &[], &[], &weights);
assert_eq!(result.len(), 1);
assert!((result[0].score - 1.0).abs() < f32::EPSILON);
}
#[test]
fn test_identifier_weights_exact_dominates() {
let weights = QueryType::Identifier.triple_fusion_weights();
let dense = vec![s(1, 0.95)];
let sparse = vec![s(2, 15.0)];
let exact = vec![3];
let result = triple_cc_fuse(&dense, &sparse, &exact, &weights);
assert_eq!(result[0].chunk_id, 3);
}
#[test]
fn test_semantic_weights_exact_dominates() {
let weights = QueryType::Semantic.triple_fusion_weights();
let dense = vec![s(1, 0.95)];
let sparse = vec![s(2, 15.0)];
let exact = vec![3];
let result = triple_cc_fuse(&dense, &sparse, &exact, &weights);
assert_eq!(result[0].chunk_id, 3);
}
#[test]
fn test_consensus_wins() {
let weights = TripleFusionWeights {
w_dense: 1.0,
w_sparse: 1.0,
w_exact: 1.0,
};
let dense = vec![s(5, 0.9), s(10, 0.7)];
let sparse = vec![s(5, 10.0), s(20, 8.0)];
let exact = vec![5, 30];
let result = triple_cc_fuse(&dense, &sparse, &exact, &weights);
assert_eq!(result[0].chunk_id, 5);
}
#[test]
fn test_mixed_weights() {
let weights = QueryType::Mixed.triple_fusion_weights();
assert!((weights.w_dense - 0.5).abs() < f32::EPSILON);
assert!((weights.w_sparse - 0.4).abs() < f32::EPSILON);
assert!((weights.w_exact - 0.8).abs() < f32::EPSILON);
}
#[test]
fn test_keyword_weights() {
let weights = QueryType::Keyword.triple_fusion_weights();
assert!((weights.w_dense - 0.3).abs() < f32::EPSILON);
assert!((weights.w_sparse - 0.6).abs() < f32::EPSILON);
assert!((weights.w_exact - 2.0).abs() < f32::EPSILON);
}
#[test]
fn test_per_channel_scores_preserved() {
let weights = TripleFusionWeights {
w_dense: 0.4,
w_sparse: 0.5,
w_exact: 0.8,
};
let dense = vec![s(1, 0.8), s(4, 0.6)];
let sparse = vec![s(2, 5.0), s(4, 3.0)];
let exact = vec![3, 4];
let result = triple_cc_fuse(&dense, &sparse, &exact, &weights);
let by_id: HashMap<u64, &ScoredChunkId> = result.iter().map(|r| (r.chunk_id, r)).collect();
let c1 = by_id[&1];
assert!((c1.score_dense - 1.0).abs() < f32::EPSILON); assert!((c1.score_dense - 1.0).abs() < f32::EPSILON);
assert!(c1.score_sparse.abs() < f32::EPSILON);
assert!(c1.score_exact.abs() < f32::EPSILON);
let c2 = by_id[&2];
assert!(c2.score_dense.abs() < f32::EPSILON);
assert!((c2.score_sparse - 1.0).abs() < f32::EPSILON); assert!(c2.score_exact.abs() < f32::EPSILON);
let c3 = by_id[&3];
assert!(c3.score_dense.abs() < f32::EPSILON);
assert!(c3.score_sparse.abs() < f32::EPSILON);
assert!((c3.score_exact - 1.0).abs() < f32::EPSILON);
let c4 = by_id[&4];
assert!(c4.score_dense > 0.0);
assert!(c4.score_sparse > 0.0);
assert!((c4.score_exact - 1.0).abs() < f32::EPSILON);
}
#[test]
fn test_max_possible_identifier() {
let mp = QueryType::Identifier.triple_fusion_weights().max_possible();
assert!(
(mp - 5.8).abs() < f32::EPSILON,
"Identifier max_possible should be 5.8, got {mp}"
);
}
#[test]
fn test_max_possible_semantic() {
let mp = QueryType::Semantic.triple_fusion_weights().max_possible();
assert!(
(mp - 1.7).abs() < f32::EPSILON,
"Semantic max_possible should be 1.7, got {mp}"
);
}
#[test]
fn test_adapt_weights_high_exact_boosts_exact() {
let base = TripleFusionWeights {
w_dense: 0.4,
w_sparse: 0.5,
w_exact: 0.8,
};
let adapted = adapt_weights(base, 0.5, 0.5, 0.9);
assert!(
(adapted.w_exact - 1.2).abs() < 1e-5,
"w_exact should be boosted to 1.2, got {}",
adapted.w_exact
);
assert!((adapted.w_dense - 0.4).abs() < 1e-5);
assert!((adapted.w_sparse - 0.5).abs() < 1e-5);
}
#[test]
fn test_adapt_weights_low_exact_no_boost() {
let base = TripleFusionWeights {
w_dense: 0.4,
w_sparse: 0.5,
w_exact: 0.8,
};
let adapted = adapt_weights(base, 0.5, 0.5, 0.0);
assert!((adapted.w_exact - 0.8).abs() < 1e-5);
}
#[test]
fn test_adapt_weights_sparse_dominant_boosts_sparse() {
let base = TripleFusionWeights {
w_dense: 0.4,
w_sparse: 0.5,
w_exact: 0.8,
};
let adapted = adapt_weights(base, 0.3, 0.9, 0.0);
assert!(
(adapted.w_sparse - 0.5 * 1.3).abs() < 1e-5,
"w_sparse should be boosted, got {}",
adapted.w_sparse
);
assert!(
(adapted.w_dense - 0.4 * 0.8).abs() < 1e-5,
"w_dense should be dampened, got {}",
adapted.w_dense
);
}
#[test]
fn test_adapt_weights_dense_dominant_boosts_dense() {
let base = TripleFusionWeights {
w_dense: 0.4,
w_sparse: 0.5,
w_exact: 0.8,
};
let adapted = adapt_weights(base, 0.9, 0.1, 0.0);
assert!(
(adapted.w_dense - 0.4 * 1.3).abs() < 1e-5,
"w_dense should be boosted, got {}",
adapted.w_dense
);
assert!(
(adapted.w_sparse - 0.5 * 0.8).abs() < 1e-5,
"w_sparse should be dampened, got {}",
adapted.w_sparse
);
}
#[test]
fn test_adapt_weights_balanced_no_adjustment() {
let base = TripleFusionWeights {
w_dense: 0.4,
w_sparse: 0.5,
w_exact: 0.8,
};
let adapted = adapt_weights(base, 0.6, 0.6, 0.0);
assert!((adapted.w_dense - 0.4).abs() < 1e-5);
assert!((adapted.w_sparse - 0.5).abs() < 1e-5);
assert!((adapted.w_exact - 0.8).abs() < 1e-5);
}
#[test]
fn test_triple_rrf_basic() {
let dense = vec![s(1, 0.9), s(2, 0.1)];
let sparse = vec![s(2, 10.0), s(1, 1.0)];
let fused = triple_rrf_fuse(&dense, &sparse, &[]);
assert_eq!(fused.len(), 2);
assert_eq!(fused[0].channels_fired, 2);
assert_eq!(fused[0].channels_hit, 2);
assert_eq!(fused[1].channels_hit, 2);
let expected = 1.0 / 61.0 + 1.0 / 62.0;
assert!((fused[0].scored.score - expected).abs() < 1e-5);
assert!((fused[1].scored.score - expected).abs() < 1e-5);
}
#[test]
fn test_triple_rrf_consensus_wins() {
let dense = vec![s(5, 0.9), s(10, 0.5)];
let sparse = vec![s(5, 10.0), s(20, 5.0)];
let exact = vec![5, 30];
let fused = triple_rrf_fuse(&dense, &sparse, &exact);
assert_eq!(fused[0].scored.chunk_id, 5);
assert_eq!(fused[0].channels_hit, 3);
assert_eq!(fused[0].channels_fired, 3);
}
#[test]
fn test_triple_rrf_parameter_free_k_named_60() {
assert!((RRF_K - 60.0).abs() < f32::EPSILON);
}
#[test]
fn test_triple_rrf_empty_inputs() {
let fused = triple_rrf_fuse(&[], &[], &[]);
assert!(fused.is_empty());
}
#[test]
fn test_triple_rrf_only_exact() {
let fused = triple_rrf_fuse(&[], &[], &[42, 7]);
assert_eq!(fused.len(), 2);
assert_eq!(fused[0].scored.chunk_id, 42);
assert_eq!(fused[0].channels_fired, 1);
assert_eq!(fused[0].channels_hit, 1);
}
#[test]
fn test_triple_rrf_no_per_channel_weighting() {
let dense = vec![s(100, 999.0), s(200, 0.1)];
let sparse = vec![s(200, 0.1), s(100, 0.01)];
let fused = triple_rrf_fuse(&dense, &sparse, &[]);
let chunk_100 = fused.iter().find(|r| r.scored.chunk_id == 100).unwrap();
let chunk_200 = fused.iter().find(|r| r.scored.chunk_id == 200).unwrap();
assert!((chunk_100.scored.score - chunk_200.scored.score).abs() < 1e-6);
}
#[test]
fn test_rrf_vs_cc_perfect_agreement_both_pick_same_winner() {
let dense = vec![s(1, 0.9), s(2, 0.5), s(3, 0.1)];
let sparse = vec![s(1, 10.0), s(2, 5.0), s(3, 1.0)];
let exact: Vec<u64> = vec![];
let cc_weights = TripleFusionWeights {
w_dense: 1.0,
w_sparse: 1.0,
w_exact: 1.0,
};
let cc = triple_cc_fuse(&dense, &sparse, &exact, &cc_weights);
let rrf = triple_rrf_fuse(&dense, &sparse, &exact);
assert_eq!(cc[0].chunk_id, 1);
assert_eq!(rrf[0].scored.chunk_id, 1);
}
#[test]
fn test_exp4_basic_five_channels() {
let orig_dense = vec![s(1, 0.9), s(2, 0.5)];
let orig_sparse = vec![s(1, 10.0)];
let exp_dense = vec![s(1, 0.8)];
let exp_sparse = vec![s(1, 8.0)];
let exact = vec![1u64];
let fused = exp4_rrf_fuse(&orig_dense, &orig_sparse, &exp_dense, &exp_sparse, &exact);
assert_eq!(fused[0].scored.chunk_id, 1);
assert_eq!(fused[0].channels_fired, 5);
assert_eq!(fused[0].channels_hit, 5);
}
#[test]
fn test_exp4_with_no_expansion_falls_back() {
let dense = vec![s(5, 0.9)];
let sparse = vec![s(7, 10.0)];
let exact = vec![5u64];
let triple = triple_rrf_fuse(&dense, &sparse, &exact);
let exp4 = exp4_rrf_fuse(&dense, &sparse, &[], &[], &exact);
assert_eq!(triple.len(), exp4.len());
assert_eq!(triple[0].channels_fired, exp4[0].channels_fired);
assert_eq!(triple[0].channels_fired, 3);
}
#[test]
fn test_exp4_dual_route_finds_more_than_single_route() {
let orig_dense = vec![s(1, 0.9)];
let orig_sparse = vec![s(1, 10.0)];
let exp_dense = vec![s(99, 0.7)];
let exp_sparse = vec![s(99, 5.0)];
let exact: Vec<u64> = vec![];
let single = triple_rrf_fuse(&orig_dense, &orig_sparse, &exact);
let dual = exp4_rrf_fuse(&orig_dense, &orig_sparse, &exp_dense, &exp_sparse, &exact);
let single_ids: std::collections::HashSet<u64> =
single.iter().map(|r| r.scored.chunk_id).collect();
let dual_ids: std::collections::HashSet<u64> =
dual.iter().map(|r| r.scored.chunk_id).collect();
assert!(
!single_ids.contains(&99),
"Chunk 99 should be missing from single-route"
);
assert!(
dual_ids.contains(&99),
"Chunk 99 should be found by dual-route via expansion"
);
}
#[test]
fn test_confidence_extracted_when_all_channels_agree() {
let dense = vec![s(1, 0.9)];
let sparse = vec![s(1, 10.0)];
let exact = vec![1u64];
let fused = triple_rrf_fuse(&dense, &sparse, &exact);
let confidence = fused[0].confidence(None);
assert_eq!(confidence, Confidence::Extracted);
assert!((fused[0].confidence_score() - 1.0).abs() < f32::EPSILON);
}
#[test]
fn test_confidence_inferred_when_single_channel() {
let dense = vec![s(1, 0.9), s(2, 0.8)];
let sparse = vec![s(1, 10.0)]; let fused = triple_rrf_fuse(&dense, &sparse, &[]);
let chunk_2 = fused.iter().find(|r| r.scored.chunk_id == 2).unwrap();
assert_eq!(chunk_2.channels_hit, 1);
assert_eq!(chunk_2.channels_fired, 2);
assert_eq!(chunk_2.confidence(None), Confidence::Inferred);
assert!((chunk_2.confidence_score() - 0.5).abs() < f32::EPSILON);
}
#[test]
fn test_confidence_ambiguous_threshold() {
let r1 = RrfFusedResult {
scored: ScoredChunkId::new(1, 1.000),
channels_hit: 1,
channels_fired: 2,
};
let r2 = RrfFusedResult {
scored: ScoredChunkId::new(2, 0.970), channels_hit: 1,
channels_fired: 2,
};
assert_eq!(r1.confidence(Some(r2.scored.score)), Confidence::Ambiguous);
let r3 = RrfFusedResult {
scored: ScoredChunkId::new(3, 0.900), channels_hit: 1,
channels_fired: 2,
};
assert_eq!(r1.confidence(Some(r3.scored.score)), Confidence::Inferred);
}
#[test]
fn test_confidence_extracted_overrides_ambiguous() {
let extracted = RrfFusedResult {
scored: ScoredChunkId::new(1, 1.000),
channels_hit: 3,
channels_fired: 3,
};
let next_close = 0.999_f32;
assert_eq!(
extracted.confidence(Some(next_close)),
Confidence::Extracted
);
}
#[test]
fn test_confidence_single_channel_no_consensus_possible() {
let solo = RrfFusedResult {
scored: ScoredChunkId::new(1, 0.5),
channels_hit: 1,
channels_fired: 1,
};
assert_eq!(solo.confidence(None), Confidence::Inferred);
assert!((solo.confidence_score() - 1.0).abs() < f32::EPSILON);
}
#[test]
fn test_confidence_zero_channels_fired_safe() {
let empty = RrfFusedResult {
scored: ScoredChunkId::new(1, 0.0),
channels_hit: 0,
channels_fired: 0,
};
assert!(empty.confidence_score().abs() < f32::EPSILON);
}
#[test]
fn test_assign_confidence_propagates_to_list() {
let dense = vec![s(1, 0.9), s(2, 0.05)];
let sparse = vec![s(1, 10.0)];
let exact = vec![1u64];
let fused = triple_rrf_fuse(&dense, &sparse, &exact);
let labels = assign_confidence(&fused);
assert_eq!(labels.len(), fused.len());
assert_eq!(labels[0].0, Confidence::Extracted);
assert!((labels[0].1 - 1.0).abs() < f32::EPSILON);
let chunk_2_pos = fused.iter().position(|r| r.scored.chunk_id == 2).unwrap();
assert_eq!(labels[chunk_2_pos].0, Confidence::Inferred);
assert!((labels[chunk_2_pos].1 - 1.0 / 3.0).abs() < 1e-5);
}
#[test]
fn test_fusion_mode_default_is_rrf() {
assert_eq!(FusionMode::default(), FusionMode::Rrf);
}
#[test]
fn test_fusion_mode_parse_cc() {
assert_eq!(FusionMode::from_env_value("cc"), FusionMode::Cc);
assert_eq!(FusionMode::from_env_value("CC"), FusionMode::Cc);
assert_eq!(FusionMode::from_env_value(" cc "), FusionMode::Cc);
assert_eq!(FusionMode::from_env_value("convex"), FusionMode::Cc);
assert_eq!(FusionMode::from_env_value("weighted"), FusionMode::Cc);
}
#[test]
fn test_fusion_mode_parse_rrf() {
assert_eq!(FusionMode::from_env_value("rrf"), FusionMode::Rrf);
assert_eq!(FusionMode::from_env_value("RRF"), FusionMode::Rrf);
}
#[test]
fn test_fusion_mode_unknown_falls_back_to_default() {
assert_eq!(FusionMode::from_env_value(""), FusionMode::Rrf);
assert_eq!(FusionMode::from_env_value("zzz"), FusionMode::Rrf);
}
#[test]
fn test_fusion_mode_parse_weighted_rrf() {
assert_eq!(
FusionMode::from_env_value("weighted-rrf"),
FusionMode::WeightedRrf
);
assert_eq!(FusionMode::from_env_value("wrrf"), FusionMode::WeightedRrf);
assert_eq!(
FusionMode::from_env_value(" Weighted-RRF "),
FusionMode::WeightedRrf
);
}
#[test]
fn test_fusion_mode_weighted_does_not_collide_with_cc() {
assert_eq!(FusionMode::from_env_value("weighted"), FusionMode::Cc);
}
#[test]
fn test_parse_weight_override_rejects_all_non_finite() {
let key = "SEMANTEX_WEIGHTS_FINDING15_ALL_NAN";
unsafe {
std::env::set_var(key, "NaN,inf,-inf");
}
let parsed = parse_weight_override(key);
unsafe {
std::env::remove_var(key);
}
assert!(
parsed.is_none(),
"all non-finite weights must yield None, got {parsed:?}"
);
}
#[test]
fn test_parse_weight_override_rejects_mixed_nan() {
let key = "SEMANTEX_WEIGHTS_FINDING15_MIXED";
unsafe {
std::env::set_var(key, "NaN,1.0,0.5");
}
let parsed = parse_weight_override(key);
unsafe {
std::env::remove_var(key);
}
assert!(
parsed.is_none(),
"partial-finite weights must yield None, got {parsed:?}"
);
}
#[test]
fn test_parse_weight_override_accepts_three_finite() {
let key = "SEMANTEX_WEIGHTS_FINDING15_CLEAN";
unsafe {
std::env::set_var(key, "0.4,0.5,1.5");
}
let parsed = parse_weight_override(key);
unsafe {
std::env::remove_var(key);
}
let w = parsed.expect("three finite values must parse");
assert!((w.w_dense - 0.4).abs() < f32::EPSILON);
assert!((w.w_sparse - 0.5).abs() < f32::EPSILON);
assert!((w.w_exact - 1.5).abs() < f32::EPSILON);
assert!(w.w_dense.is_finite() && w.w_sparse.is_finite() && w.w_exact.is_finite());
}
#[test]
fn test_sort_does_not_panic_on_nan_scores() {
let weights = TripleFusionWeights {
w_dense: f32::NAN,
w_sparse: 0.5,
w_exact: 0.8,
};
let dense = vec![s(1, 0.9), s(2, 0.5)];
let sparse = vec![s(2, 5.0), s(3, 1.0)];
let cc = triple_cc_fuse(&dense, &sparse, &[1], &weights);
assert!(!cc.is_empty());
let rrf = triple_rrf_fuse(&dense, &sparse, &[1]);
assert!(!rrf.is_empty());
}
}