use std::path::Path;
use crate::error::Result;
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct SearchHit {
pub id: u64,
pub score: f32,
}
impl PartialOrd for SearchHit {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
Some(
other
.score
.total_cmp(&self.score)
.then_with(|| self.id.cmp(&other.id)),
)
}
}
#[derive(Debug, Clone, PartialEq)]
pub enum Fusion {
Rrf {
k: u32,
},
Weighted {
weights: Vec<f32>,
},
}
impl Fusion {
pub fn rrf() -> Self {
Fusion::Rrf { k: 60 }
}
pub fn fuse(&self, hit_sets: &[Vec<SearchHit>]) -> Vec<SearchHit> {
use std::collections::HashMap;
let mut scores: HashMap<u64, f32> = HashMap::new();
match self {
Fusion::Rrf { k } => {
let k = *k as f32;
for set in hit_sets {
for (rank, hit) in set.iter().enumerate() {
*scores.entry(hit.id).or_insert(0.0) += 1.0 / (k + rank as f32 + 1.0);
}
}
}
Fusion::Weighted { weights } => {
assert_eq!(
hit_sets.len(),
weights.len(),
"Fusion::Weighted: hit_sets.len() ({}) must equal weights.len() ({})",
hit_sets.len(),
weights.len(),
);
for (set, w) in hit_sets.iter().zip(weights) {
for hit in set {
*scores.entry(hit.id).or_insert(0.0) += w * hit.score;
}
}
}
}
let mut fused: Vec<SearchHit> = scores
.into_iter()
.map(|(id, score)| SearchHit { id, score })
.collect();
fused.sort_unstable_by(|a, b| b.score.total_cmp(&a.score).then_with(|| a.id.cmp(&b.id)));
fused
}
}
pub trait VectorIndex: Send + Sync {
fn add(&mut self, vectors: &[Vec<f32>]) -> Result<()>;
fn add_with_ids(&mut self, vectors: &[Vec<f32>], ids: &[u64]) -> Result<()>;
fn remove(&mut self, ids: &[u64]) -> Result<()>;
fn search(&self, query: &[f32], k: usize) -> Result<Vec<SearchHit>>;
fn search_filtered(&self, query: &[f32], k: usize, allowlist: &[u64])
-> Result<Vec<SearchHit>>;
fn len(&self) -> usize;
fn is_empty(&self) -> bool;
fn dim(&self) -> usize;
fn save(&self, path: &Path) -> Result<()>;
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn search_hit_fields() {
let hit = SearchHit {
id: 42,
score: 0.95,
};
assert_eq!(hit.id, 42);
assert!((hit.score - 0.95).abs() < f32::EPSILON);
}
#[test]
fn search_hit_copy() {
let hit = SearchHit { id: 1, score: 0.5 };
let copied = hit; assert_eq!(copied.id, hit.id);
assert_eq!(copied.score, hit.score);
}
#[test]
fn search_hit_sort_descending_by_score() {
let mut hits = [
SearchHit { id: 1, score: 0.3 },
SearchHit { id: 2, score: 0.9 },
SearchHit { id: 3, score: 0.5 },
];
hits.sort_by(|a, b| a.partial_cmp(b).unwrap());
assert_eq!(hits[0].id, 2); assert_eq!(hits[1].id, 3);
assert_eq!(hits[2].id, 1);
}
#[test]
fn search_hit_tie_break_by_id() {
let mut hits = [
SearchHit { id: 30, score: 0.5 },
SearchHit { id: 10, score: 0.5 },
SearchHit { id: 20, score: 0.5 },
];
hits.sort_by(|a, b| a.partial_cmp(b).unwrap());
assert_eq!(hits[0].id, 10);
assert_eq!(hits[1].id, 20);
assert_eq!(hits[2].id, 30);
}
fn hit(id: u64, score: f32) -> SearchHit {
SearchHit { id, score }
}
#[test]
fn rrf_default_uses_k_60() {
assert_eq!(Fusion::rrf(), Fusion::Rrf { k: 60 });
}
#[test]
fn rrf_rewards_agreement_across_branches() {
let dense = vec![hit(10, 0.99), hit(20, 0.80)];
let sparse = vec![hit(20, 12.0), hit(30, 9.0)];
let fused = Fusion::rrf().fuse(&[dense, sparse]);
assert_eq!(fused[0].id, 20);
assert_eq!(fused.len(), 3);
}
#[test]
fn rrf_ignores_raw_score_scale() {
let dense = vec![hit(1, 0.9), hit(2, 0.8)];
let sparse = vec![hit(1, 900.0), hit(2, 800.0)];
let a = Fusion::rrf().fuse(&[dense.clone(), sparse]);
let sparse_small = vec![hit(1, 0.009), hit(2, 0.008)];
let b = Fusion::rrf().fuse(&[dense, sparse_small]);
assert_eq!(
a.iter().map(|h| h.id).collect::<Vec<_>>(),
b.iter().map(|h| h.id).collect::<Vec<_>>()
);
}
#[test]
fn weighted_sums_scores_per_branch() {
let dense = vec![hit(1, 1.0)];
let sparse = vec![hit(2, 1.0)];
let fused = Fusion::Weighted {
weights: vec![1.0, 2.0],
}
.fuse(&[dense, sparse]);
assert_eq!(fused[0].id, 2);
assert!((fused[0].score - 2.0).abs() < 1e-6);
assert!((fused[1].score - 1.0).abs() < 1e-6);
}
#[test]
#[should_panic(expected = "must equal weights.len()")]
fn weighted_panics_on_weight_count_mismatch() {
Fusion::Weighted { weights: vec![1.0] }.fuse(&[vec![hit(1, 1.0)], vec![hit(2, 1.0)]]);
}
#[test]
fn fuse_of_nothing_is_empty() {
assert!(Fusion::rrf().fuse(&[]).is_empty());
}
}