use std::collections::HashMap;
use std::hash::Hash;
use khive_fusion::reciprocal_rank_fusion;
use khive_score::DeterministicScore;
use crate::hit::{SearchSignals, SearchSource};
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct HitLabel {
pub rank: usize,
pub signals: SearchSignals,
pub source: SearchSource,
pub title: Option<String>,
pub snippet: Option<String>,
}
pub fn fuse_labelled<Id, L, F>(
arms: Vec<Vec<(Id, L)>>,
k: usize,
combine: F,
) -> Vec<(Id, DeterministicScore, L)>
where
Id: Eq + Hash + Clone + Ord,
F: Fn(L, L) -> L,
{
let scored_arms: Vec<Vec<(Id, DeterministicScore)>> = arms
.iter()
.map(|arm| {
arm.iter()
.map(|(id, _)| (id.clone(), DeterministicScore::ZERO))
.collect()
})
.collect();
let ranked = reciprocal_rank_fusion(scored_arms, k);
let mut labels: HashMap<Id, L> = HashMap::new();
for (id, label) in arms.into_iter().flatten() {
let merged = match labels.remove(&id) {
Some(held) => combine(held, label),
None => label,
};
labels.insert(id, merged);
}
ranked
.into_iter()
.map(|(id, score)| {
let label = labels.remove(&id).expect("every ranked id has a label");
(id, score, label)
})
.collect()
}
pub fn combine_leg_first_appearance(earlier: HitLabel, later: HitLabel) -> HitLabel {
let joined = earlier.source.union(later.source);
if joined == earlier.source {
return earlier;
}
let (held, added) = (earlier.signals, later.signals);
HitLabel {
rank: earlier.rank,
signals: SearchSignals {
vector_similarity: held.vector_similarity.or(added.vector_similarity),
keyword_score: held.keyword_score.or(added.keyword_score),
},
source: joined,
title: earlier.title.or(later.title),
snippet: earlier.snippet.or(later.snippet),
}
}
pub fn combine_best_ranked_evidence(earlier: HitLabel, later: HitLabel) -> HitLabel {
let (rank, signals) = if later.rank < earlier.rank {
(later.rank, later.signals)
} else {
(earlier.rank, earlier.signals)
};
HitLabel {
rank,
signals,
source: earlier.source.union(later.source),
title: earlier.title.or(later.title),
snippet: earlier.snippet.or(later.snippet),
}
}