pub(crate) fn dense_search<'a, I>(docs: I, query_vec: &[f32], top_k: usize) -> Vec<(String, f32)>
where
I: IntoIterator<Item = (String, &'a [f32])>,
{
let mut ranked: Vec<(String, f32)> = docs
.into_iter()
.map(|(id, vec)| (id, cosine(vec, query_vec)))
.collect();
ranked.sort_by(|a, b| {
b.1.partial_cmp(&a.1)
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| a.0.cmp(&b.0))
});
ranked.truncate(top_k);
ranked
}
fn cosine(a: &[f32], b: &[f32]) -> f32 {
a.iter().zip(b.iter()).map(|(x, y)| x * y).sum()
}
#[cfg(test)]
mod tests {
use super::*;
fn refs(docs: &[(String, Vec<f32>)]) -> Vec<(String, &[f32])> {
docs.iter()
.map(|(id, v)| (id.clone(), v.as_slice()))
.collect()
}
#[test]
fn empty_docs_yield_no_hits() {
let hits = dense_search(Vec::<(String, &[f32])>::new(), &[1.0, 0.0], 5);
assert!(hits.is_empty());
}
#[test]
fn ranks_the_closest_vector_first() {
let docs = vec![
("read".to_string(), vec![1.0, 0.0]),
("write".to_string(), vec![0.0, 1.0]),
];
let hits = dense_search(refs(&docs), &[1.0, 0.0], 5);
assert_eq!(hits.first().map(|(id, _)| id.as_str()), Some("read"));
}
#[test]
fn respects_top_k() {
let docs: Vec<(String, Vec<f32>)> = (0..10)
.map(|i| (format!("doc{i}"), vec![1.0, 0.0]))
.collect();
let hits = dense_search(refs(&docs), &[1.0, 0.0], 3);
assert!(hits.len() <= 3);
}
#[test]
fn tied_scores_break_by_id_with_stable_top_k_membership() {
let docs = vec![
("zeta".to_string(), vec![0.0, 1.0]),
("alpha".to_string(), vec![0.0, 1.0]),
("mid".to_string(), vec![0.0, 1.0]),
];
let hits = dense_search(refs(&docs), &[0.0, 1.0], 2);
assert_eq!(hits.len(), 2);
assert_eq!(hits[0].0, "alpha");
assert_eq!(hits[1].0, "mid");
}
}