use std::hash::Hash;
use async_trait::async_trait;
use khive_score::DeterministicScore;
use crate::error::Result;
use khive_fusion::{fuse, FusionStrategy};
use super::config::{HybridConfig, Query};
#[async_trait]
pub trait VectorSearch: Send + Sync {
type Id: Eq + Hash + Clone + Ord + Send + Sync;
async fn vector_search(
&self,
embedding: &[f32],
top_k: usize,
) -> Result<Vec<(Self::Id, DeterministicScore)>>;
}
#[async_trait]
pub trait KeywordSearch: Send + Sync {
type Id: Eq + Hash + Clone + Ord + Send + Sync;
async fn keyword_search(
&self,
text: &str,
top_k: usize,
) -> Result<Vec<(Self::Id, DeterministicScore)>>;
}
#[async_trait]
pub trait HybridSearcher: VectorSearch + KeywordSearch<Id = <Self as VectorSearch>::Id> {
async fn hybrid_search(
&self,
query: &Query,
config: &HybridConfig,
) -> Result<Vec<(<Self as VectorSearch>::Id, DeterministicScore)>>;
}
#[async_trait]
pub trait Reranker<Id: Send + Sync + 'static>: Send + Sync {
async fn rerank(
&self,
query: &str,
results: Vec<(Id, DeterministicScore)>,
top_k: usize,
) -> Result<Vec<(Id, DeterministicScore)>>;
}
pub fn fuse_search_results<Id: Eq + Hash + Clone + Ord>(
sources: Vec<Vec<(Id, DeterministicScore)>>,
config: &HybridConfig,
) -> Vec<(Id, DeterministicScore)> {
if sources.is_empty() {
return Vec::new();
}
if sources.len() == 1 {
let mut results = sources.into_iter().next().unwrap();
if let Some(min_score) = config.min_score {
results.retain(|(_, score)| *score >= min_score);
}
results.truncate(config.top_k);
return results;
}
let strategy = match &config.fusion_strategy {
FusionStrategy::Weighted { .. } => {
debug_assert_eq!(
sources.len(),
2,
"Weighted fusion expects exactly 2 sources"
);
let (v, k) = config.normalized_weights();
FusionStrategy::weighted(vec![v, k])
}
other => other.clone(),
};
let mut fused = fuse(sources, &strategy, config.top_k);
if let Some(min_score) = config.min_score {
fused.retain(|(_, score)| *score >= min_score);
}
fused
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_fuse_empty_sources() {
let sources: Vec<Vec<(String, DeterministicScore)>> = vec![];
let config = HybridConfig::default();
let results = fuse_search_results(sources, &config);
assert!(results.is_empty());
}
#[test]
fn test_fuse_single_source() {
let sources = vec![vec![
("a".to_string(), DeterministicScore::from_f64(0.9)),
("b".to_string(), DeterministicScore::from_f64(0.8)),
]];
let config = HybridConfig::new(10);
let results = fuse_search_results(sources, &config);
assert_eq!(results.len(), 2);
assert_eq!(results[0].0, "a");
}
#[test]
fn test_fuse_multiple_sources_rrf() {
let source1 = vec![
("a".to_string(), DeterministicScore::from_f64(0.9)),
("b".to_string(), DeterministicScore::from_f64(0.8)),
];
let source2 = vec![
("b".to_string(), DeterministicScore::from_f64(0.95)),
("c".to_string(), DeterministicScore::from_f64(0.7)),
];
let config = HybridConfig::new(10);
let results = fuse_search_results(vec![source1, source2], &config);
assert_eq!(results.len(), 3);
assert_eq!(results[0].0, "b");
}
#[test]
fn test_fuse_with_min_score() {
let sources = vec![vec![
("a".to_string(), DeterministicScore::from_f64(0.9)),
("b".to_string(), DeterministicScore::from_f64(0.1)),
]];
let config = HybridConfig::new(10).with_min_score(DeterministicScore::from_f64(0.5));
let results = fuse_search_results(sources, &config);
assert!(!results.is_empty());
}
#[test]
fn test_fuse_top_k_limit() {
let sources = vec![vec![
("a".to_string(), DeterministicScore::from_f64(0.9)),
("b".to_string(), DeterministicScore::from_f64(0.8)),
("c".to_string(), DeterministicScore::from_f64(0.7)),
("d".to_string(), DeterministicScore::from_f64(0.6)),
("e".to_string(), DeterministicScore::from_f64(0.5)),
]];
let config = HybridConfig::new(3);
let results = fuse_search_results(sources, &config);
assert_eq!(results.len(), 3);
}
}