pub use csm_retrieval::{HybridResult, RetrievalAbstention};
use std::collections::HashMap;
pub const fn compute_weights(token_count: usize) -> (f32, f32) {
match token_count {
1..=2 => (0.9, 0.1),
3..=4 => (0.7, 0.3),
5..=8 => (0.4, 0.6),
_ => (0.2, 0.8),
}
}
pub fn normalize_scores(scores: &[(String, f32)]) -> Vec<(String, f32)> {
if scores.is_empty() {
return Vec::new();
}
let mut min = f32::INFINITY;
let mut max = f32::NEG_INFINITY;
for (_, s) in scores {
let s = *s;
min = min.min(s);
max = max.max(s);
}
let range = max - min;
let epsilon = 1e-10;
if range < epsilon {
return scores.iter().map(|(id, _)| (id.clone(), 1.0)).collect();
}
let inv_range = 1.0 / range;
scores
.iter()
.map(|(id, score)| {
let normalized = (score - min) * inv_range;
(id.clone(), normalized)
})
.collect()
}
pub fn merge_results(
bm25_results: &[(String, f32)],
hdc_results: &[(String, f32)],
weights: (f32, f32),
top_k: usize,
) -> Vec<(String, f32)> {
if top_k == 0 {
return Vec::new();
}
let (kw_weight, sem_weight) = weights;
let mut combined: HashMap<&str, f32> =
HashMap::with_capacity(bm25_results.len() + hdc_results.len());
if !bm25_results.is_empty() {
let mut min = f32::INFINITY;
let mut max = f32::NEG_INFINITY;
for (_, s) in bm25_results {
let s = *s;
min = min.min(s);
max = max.max(s);
}
let range = max - min;
if range < 1e-10 {
for (id, _) in bm25_results {
combined.insert(id.as_str(), kw_weight);
}
} else {
let inv_range = 1.0 / range;
for (id, score) in bm25_results {
combined.insert(id.as_str(), kw_weight * (score - min) * inv_range);
}
}
}
if !hdc_results.is_empty() {
let mut min = f32::INFINITY;
let mut max = f32::NEG_INFINITY;
for (_, s) in hdc_results {
let s = *s;
min = min.min(s);
max = max.max(s);
}
let range = max - min;
if range < 1e-10 {
for (id, _) in hdc_results {
combined
.entry(id.as_str())
.and_modify(|s| *s += sem_weight)
.or_insert(sem_weight);
}
} else {
let inv_range = 1.0 / range;
for (id, score) in hdc_results {
let weighted_norm = sem_weight * (score - min) * inv_range;
combined
.entry(id.as_str())
.and_modify(|s| *s += weighted_norm)
.or_insert(weighted_norm);
}
}
}
let mut ref_results: Vec<(&str, f32)> = combined.into_iter().collect();
if ref_results.len() > top_k {
ref_results.select_nth_unstable_by(top_k, |a, b| b.1.total_cmp(&a.1));
ref_results.truncate(top_k);
}
ref_results.sort_unstable_by(|a, b| b.1.total_cmp(&a.1));
ref_results
.into_iter()
.map(|(id, score)| (id.to_string(), score))
.collect()
}
#[derive(Debug, Clone, Copy, PartialEq, Default)]
pub enum HybridMode {
#[default]
Auto,
SemanticOnly,
KeywordOnly,
Custom(f32),
}
#[derive(Debug, Clone)]
pub struct HybridConfig {
pub mode: HybridMode,
pub min_score: f32,
}
impl Default for HybridConfig {
fn default() -> Self {
Self {
mode: HybridMode::Auto,
min_score: 0.0,
}
}
}
#[cfg(test)]
mod tests {
#![allow(clippy::unwrap_used, clippy::expect_used, clippy::panic)]
use super::*;
#[test]
fn test_compute_weights_short_query() {
let (kw, sem) = compute_weights(1);
assert!((kw - (0.9)).abs() < 1e-6);
assert!((sem - (0.1)).abs() < 1e-6);
let (kw, sem) = compute_weights(2);
assert!((kw - (0.9)).abs() < 1e-6);
assert!((sem - (0.1)).abs() < 1e-6);
}
#[test]
fn test_compute_weights_medium_query() {
let (kw, sem) = compute_weights(3);
assert!((kw - (0.7)).abs() < 1e-6);
assert!((sem - (0.3)).abs() < 1e-6);
let (kw, sem) = compute_weights(4);
assert!((kw - (0.7)).abs() < 1e-6);
assert!((sem - (0.3)).abs() < 1e-6);
}
#[test]
fn test_compute_weights_long_query() {
let (kw, sem) = compute_weights(5);
assert!((kw - (0.4)).abs() < 1e-6);
assert!((sem - (0.6)).abs() < 1e-6);
let (kw, sem) = compute_weights(8);
assert!((kw - (0.4)).abs() < 1e-6);
assert!((sem - (0.6)).abs() < 1e-6);
}
#[test]
fn test_compute_weights_very_long_query() {
let (kw, sem) = compute_weights(9);
assert!((kw - (0.2)).abs() < 1e-6);
assert!((sem - (0.8)).abs() < 1e-6);
let (kw, sem) = compute_weights(100);
assert!((kw - (0.2)).abs() < 1e-6);
assert!((sem - (0.8)).abs() < 1e-6);
}
#[test]
fn test_normalize_scores_basic() {
let scores = vec![
("a".to_string(), 10.0),
("b".to_string(), 15.0),
("c".to_string(), 20.0),
];
let normalized = normalize_scores(&scores);
assert!((normalized[0].1 - 0.0).abs() < 1e-6);
assert!((normalized[1].1 - 0.5).abs() < 1e-6);
assert!((normalized[2].1 - 1.0).abs() < 1e-6);
}
#[test]
fn test_normalize_scores_range_two() {
let scores = vec![("a".to_string(), 2.0), ("b".to_string(), 0.0)];
let normalized = normalize_scores(&scores);
assert!((normalized[0].1 - 1.0).abs() < 1e-6);
assert!((normalized[1].1 - 0.0).abs() < 1e-6);
}
#[test]
fn test_normalize_scores_empty() {
let normalized = normalize_scores(&[]);
assert!(normalized.is_empty());
}
#[test]
fn test_normalize_scores_equal() {
let scores = vec![("a".to_string(), 5.0), ("b".to_string(), 5.0)];
let normalized = normalize_scores(&scores);
assert!((normalized[0].1 - 1.0).abs() < 1e-6);
assert!((normalized[1].1 - 1.0).abs() < 1e-6);
}
#[test]
fn test_merge_results_basic() {
let bm25 = vec![("doc1".to_string(), 1.0), ("doc2".to_string(), 0.5)];
let hdc = vec![("doc1".to_string(), 0.5), ("doc3".to_string(), 1.0)];
let merged = merge_results(&bm25, &hdc, (0.5, 0.5), 10);
assert!(merged.iter().any(|(id, _)| id == "doc1"));
assert!(merged.iter().any(|(id, _)| id == "doc2"));
assert!(merged.iter().any(|(id, _)| id == "doc3"));
}
#[test]
fn test_merge_results_weighted() {
let bm25 = vec![("doc1".to_string(), 1.0)];
let hdc = vec![("doc1".to_string(), 1.0)];
let merged = merge_results(&bm25, &hdc, (0.9, 0.1), 10);
assert!(merged.iter().any(|(id, s)| id == "doc1" && *s > 0.0));
}
#[test]
fn test_merge_results_empty() {
let merged = merge_results(&[], &[], (0.5, 0.5), 10);
assert!(merged.is_empty());
let merged = merge_results(&[("a".to_string(), 1.0)], &[], (0.5, 0.5), 10);
assert_eq!(merged.len(), 1);
let merged = merge_results(&[], &[("a".to_string(), 1.0)], (0.5, 0.5), 10);
assert_eq!(merged.len(), 1);
}
#[test]
fn test_exact_score_calculation() {
let weights = (0.6, 0.4);
let bm25 = vec![
("d1".to_string(), 12.0),
("d2".to_string(), 2.0),
("d4".to_string(), 7.0),
];
let hdc = vec![
("d1".to_string(), 1.2),
("d3".to_string(), 0.2),
("d4".to_string(), 0.7),
];
let merged = merge_results(&bm25, &hdc, weights, 10);
let d1_score = merged.iter().find(|(id, _)| id == "d1").unwrap().1;
assert!((d1_score - 1.0).abs() < 1e-6);
let d2_score = merged.iter().find(|(id, _)| id == "d2").unwrap().1;
assert!((d2_score - 0.0).abs() < 1e-6);
let d3_score = merged.iter().find(|(id, _)| id == "d3").unwrap().1;
assert!((d3_score - 0.0).abs() < 1e-6);
let d4_score = merged.iter().find(|(id, _)| id == "d4").unwrap().1;
assert!((d4_score - 0.5).abs() < 1e-6);
}
#[test]
fn test_range_epsilon_boundary() {
let epsilon = 1e-10;
let scores = vec![("a".to_string(), epsilon), ("b".to_string(), 0.0)];
let normalized = normalize_scores(&scores);
assert!((normalized[0].1 - 1.0).abs() < 1e-6);
assert!((normalized[1].1 - 0.0).abs() < 1e-6);
let just_below = epsilon * 0.9;
let scores_below = vec![("a".to_string(), just_below), ("b".to_string(), 0.0)];
let normalized_below = normalize_scores(&scores_below);
assert!((normalized_below[0].1 - 1.0).abs() < 1e-6);
assert!((normalized_below[1].1 - 1.0).abs() < 1e-6);
}
#[test]
fn test_merge_results_epsilon_boundary() {
let epsilon = 1e-10;
let weights = (0.5, 0.5);
let bm25 = vec![("d1".to_string(), epsilon), ("d2".to_string(), 0.0)];
let hdc = vec![("d1".to_string(), epsilon), ("d2".to_string(), 0.0)];
let merged = merge_results(&bm25, &hdc, weights, 10);
let d1_score = merged.iter().find(|(id, _)| id == "d1").unwrap().1;
let d2_score = merged.iter().find(|(id, _)| id == "d2").unwrap().1;
assert!((d1_score - 1.0).abs() < 1e-6);
assert!((d2_score - 0.0).abs() < 1e-6);
let just_below = epsilon * 0.9;
let bm25_small = vec![("d1".to_string(), just_below), ("d2".to_string(), 0.0)];
let hdc_small = vec![("d1".to_string(), just_below), ("d2".to_string(), 0.0)];
let merged_small = merge_results(&bm25_small, &hdc_small, weights, 10);
for (_, score) in merged_small {
assert!((score - 1.0).abs() < 1e-6);
}
}
#[test]
fn test_merge_results_top_k() {
let bm25 = vec![
("d1".to_string(), 10.0),
("d2".to_string(), 8.0),
("d3".to_string(), 6.0),
];
let hdc = vec![
("d1".to_string(), 10.0),
("d4".to_string(), 4.0),
("d5".to_string(), 2.0),
];
let weights = (0.5, 0.5);
let merged = merge_results(&bm25, &hdc, weights, 2);
assert_eq!(merged.len(), 2);
assert_eq!(merged[0].0, "d1");
assert_eq!(merged[1].0, "d2");
let merged_one = merge_results(&bm25, &hdc, weights, 1);
assert_eq!(merged_one.len(), 1);
assert_eq!(merged_one[0].0, "d1");
let merged_zero = merge_results(&bm25, &hdc, weights, 0);
assert!(merged_zero.is_empty());
}
#[test]
fn test_merge_results_top_k_exact_boundary() {
let bm25 = vec![("d1".to_string(), 10.0), ("d2".to_string(), 8.0)];
let hdc = vec![("d1".to_string(), 10.0), ("d2".to_string(), 8.0)];
let weights = (0.5, 0.5);
let merged = merge_results(&bm25, &hdc, weights, 2);
assert_eq!(merged.len(), 2);
assert_eq!(merged[0].0, "d1");
assert_eq!(merged[1].0, "d2");
}
}