use std::collections::HashMap;
use crate::bm25::BM25Index;
use crate::vector_search;
#[allow(clippy::too_many_arguments)]
pub fn hybrid_search(
query_embedding: &[f32],
query_text: &str,
vectors: &[Vec<f32>],
_chunks: &[String],
tombstones: &[u8],
bm25_index: &BM25Index,
vector_weight: f32,
keyword_weight: f32,
k: usize,
) -> Vec<(usize, f32)> {
let vec_scores = {
#[cfg(feature = "parallel")]
{
if vectors.len() > 10_000 {
vector_search::parallel_cosine_batch(
query_embedding, vectors, tombstones, vectors.len(),
)
} else {
vector_search::cosine_similarity_batch(query_embedding, vectors, tombstones)
}
}
#[cfg(not(feature = "parallel"))]
{
vector_search::cosine_similarity_batch(query_embedding, vectors, tombstones)
}
};
let kw_scores = bm25_index.search(query_text, vectors.len());
let vec_normalized = normalize_scores(&vec_scores);
let kw_normalized = normalize_scores(&kw_scores);
let mut merged: HashMap<usize, f32> = HashMap::new();
for (idx, score) in &vec_normalized {
*merged.entry(*idx).or_insert(0.0) += vector_weight * score;
}
for (idx, score) in &kw_normalized {
*merged.entry(*idx).or_insert(0.0) += keyword_weight * score;
}
let mut results: Vec<(usize, f32)> = merged.into_iter().collect();
results.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
results.truncate(k);
results
}
fn normalize_scores(scores: &[(usize, f32)]) -> Vec<(usize, f32)> {
if scores.is_empty() {
return Vec::new();
}
let min = scores
.iter()
.map(|(_, s)| *s)
.fold(f32::INFINITY, f32::min);
let max = scores
.iter()
.map(|(_, s)| *s)
.fold(f32::NEG_INFINITY, f32::max);
let range = max - min;
if range == 0.0 {
return scores.iter().map(|(idx, _)| (*idx, 0.0)).collect();
}
scores
.iter()
.map(|(idx, s)| (*idx, (s - min) / range))
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
fn make_test_data() -> (Vec<Vec<f32>>, Vec<String>, Vec<u8>, BM25Index) {
let vectors = vec![
vec![1.0, 0.0, 0.0], vec![0.0, 1.0, 0.0], vec![0.7, 0.7, 0.0], vec![0.0, 0.0, 1.0], ];
let chunks = vec![
"rust programming language".to_string(),
"python scripting language".to_string(),
"rust and python comparison".to_string(),
"javascript web development".to_string(),
];
let tombstones = vec![0u8; 4];
let bm25 = BM25Index::build(&chunks, &tombstones);
(vectors, chunks, tombstones, bm25)
}
#[test]
fn vector_only_search() {
let (vectors, chunks, tombstones, bm25) = make_test_data();
let query_emb = vec![1.0, 0.0, 0.0];
let results = hybrid_search(
&query_emb,
"nonexistent_xyz",
&vectors,
&chunks,
&tombstones,
&bm25,
1.0, 0.0, 4,
);
assert!(!results.is_empty());
assert_eq!(results[0].0, 0, "doc 0 should be top match for x-direction query");
}
#[test]
fn keyword_only_search() {
let (vectors, chunks, tombstones, bm25) = make_test_data();
let query_emb = vec![0.0, 0.0, 0.0];
let results = hybrid_search(
&query_emb,
"rust programming",
&vectors,
&chunks,
&tombstones,
&bm25,
0.0, 1.0, 4,
);
assert!(!results.is_empty());
assert_eq!(results[0].0, 0);
}
#[test]
fn balanced_merge_ranking() {
let (vectors, chunks, tombstones, bm25) = make_test_data();
let query_emb = vec![0.9, 0.1, 0.0];
let results = hybrid_search(
&query_emb,
"rust",
&vectors,
&chunks,
&tombstones,
&bm25,
0.7,
0.3,
4,
);
assert!(!results.is_empty());
let top_ids: Vec<usize> = results.iter().map(|(idx, _)| *idx).collect();
assert!(
top_ids.contains(&0),
"doc 0 should appear in results"
);
assert!(
top_ids.contains(&2),
"doc 2 should appear in results"
);
}
#[test]
fn empty_results_when_no_data() {
let vectors: Vec<Vec<f32>> = Vec::new();
let chunks: Vec<String> = Vec::new();
let tombstones: Vec<u8> = Vec::new();
let bm25 = BM25Index::build(&chunks, &tombstones);
let results = hybrid_search(
&[],
"anything",
&vectors,
&chunks,
&tombstones,
&bm25,
0.7,
0.3,
10,
);
assert!(results.is_empty());
}
#[test]
fn normalize_scores_empty() {
let result = normalize_scores(&[]);
assert!(result.is_empty());
}
#[test]
fn normalize_scores_single() {
let result = normalize_scores(&[(0, 5.0)]);
assert_eq!(result.len(), 1);
assert_eq!(result[0].1, 0.0);
}
#[test]
fn normalize_scores_range() {
let scores = vec![(0, 2.0), (1, 4.0), (2, 6.0)];
let result = normalize_scores(&scores);
assert_eq!(result.len(), 3);
assert!((result[0].1 - 0.0).abs() < 1e-6); assert!((result[1].1 - 0.5).abs() < 1e-6); assert!((result[2].1 - 1.0).abs() < 1e-6); }
#[test]
fn hybrid_respects_k_limit() {
let (vectors, chunks, tombstones, bm25) = make_test_data();
let query_emb = vec![0.5, 0.5, 0.0];
let results = hybrid_search(
&query_emb,
"language",
&vectors,
&chunks,
&tombstones,
&bm25,
0.5,
0.5,
2,
);
assert!(results.len() <= 2);
}
}