diff --git a/src/vector.rs b/src/vector.rs index 7abb060..cf29f96 100644 --- a/src/vector.rs +++ b/src/vector.rs @@ -13,11 +13,17 @@ impl VectorIndex { /// (entry_index, cosine_similarity) descending. 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 mut scored: Vec<(usize, f32)> = self.vectors.iter().enumerate() .map(|(i, v)| (i, dot(v, query) / (self.norms[i] * qn))) .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 } diff --git a/tests/vector_tests.rs b/tests/vector_tests.rs index 67756fa..368e748 100644 --- a/tests/vector_tests.rs +++ b/tests/vector_tests.rs @@ -12,3 +12,14 @@ fn cosine_topk_ranks_nearest_first() { assert_eq!(hits[1].0, 2); 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()); +}