harden: NaN-safe sort + dimension-mismatch guard in VectorIndex
This commit is contained in:
+7
-1
@@ -13,11 +13,17 @@ impl VectorIndex {
|
|||||||
|
|
||||||
/// (entry_index, cosine_similarity) descending.
|
/// (entry_index, cosine_similarity) descending.
|
||||||
pub fn top_k(&self, query: &[f32], k: usize) -> Vec<(usize, f32)> {
|
pub fn top_k(&self, query: &[f32], k: usize) -> Vec<(usize, f32)> {
|
||||||
|
// dimension mismatch → no results (caller bug / corrupt cache); never silently mispair
|
||||||
|
let dim = self.vectors.first().map(|v| v.len()).unwrap_or(0);
|
||||||
|
if !self.vectors.is_empty() && query.len() != dim {
|
||||||
|
return Vec::new();
|
||||||
|
}
|
||||||
|
|
||||||
let qn = dot(query, query).sqrt().max(1e-8);
|
let qn = dot(query, query).sqrt().max(1e-8);
|
||||||
let mut scored: Vec<(usize, f32)> = self.vectors.iter().enumerate()
|
let mut scored: Vec<(usize, f32)> = self.vectors.iter().enumerate()
|
||||||
.map(|(i, v)| (i, dot(v, query) / (self.norms[i] * qn)))
|
.map(|(i, v)| (i, dot(v, query) / (self.norms[i] * qn)))
|
||||||
.collect();
|
.collect();
|
||||||
scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap());
|
scored.sort_by(|a, b| b.1.total_cmp(&a.1));
|
||||||
scored.truncate(k);
|
scored.truncate(k);
|
||||||
scored
|
scored
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -12,3 +12,14 @@ fn cosine_topk_ranks_nearest_first() {
|
|||||||
assert_eq!(hits[1].0, 2);
|
assert_eq!(hits[1].0, 2);
|
||||||
assert!(hits[0].1 > hits[1].1);
|
assert!(hits[0].1 > hits[1].1);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn nan_score_does_not_panic_and_dim_mismatch_returns_empty() {
|
||||||
|
// NaN in a stored vector must not panic the sort.
|
||||||
|
let vi = VectorIndex::from_vectors(vec![vec![1.0, 0.0], vec![f32::NAN, 1.0]]);
|
||||||
|
let hits = vi.top_k(&[1.0, 0.0], 2); // must not panic
|
||||||
|
assert!(hits.iter().any(|&(i, _)| i == 0)); // the valid vector still found
|
||||||
|
// Dimension mismatch → empty, never silently mispaired.
|
||||||
|
let vi2 = VectorIndex::from_vectors(vec![vec![1.0, 0.0, 0.0]]);
|
||||||
|
assert!(vi2.top_k(&[1.0, 0.0], 1).is_empty());
|
||||||
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user