Skip to main content

uqa_storage/
vector_index.rs

1//
2// Unified Query Algebra
3//
4// Copyright (c) 2023-2026 Cognica, Inc.
5//
6
7//! Vector index abstraction and an in-memory brute-force fallback.
8//!
9//! Operators (`KNNOperator`, `VectorSimilarityOperator`,
10//! `QueryPoolVectorScoreOperator`) depend only on this trait. IVF and HNSW
11//! backends slot in by implementing the same surface.
12
13use std::collections::BTreeMap;
14use std::sync::Arc;
15
16use uqa_core::{DocId, Payload, PostingEntry, PostingList};
17
18use crate::{StorageBackendError, StorageBackendResult};
19
20mod config;
21
22pub use config::{HNSWIndexParams, IVFIndexParams, VectorIndexOpenMode, VectorIndexSpec};
23
24pub(crate) fn validate_vector_values(dimensions: u32, vector: &[f32]) -> StorageBackendResult<()> {
25    let dimensions = usize::try_from(dimensions).map_err(|_| {
26        StorageBackendError::Other(format!(
27            "vector dimension {dimensions} exceeds the platform usize range"
28        ))
29    })?;
30    if vector.len() != dimensions {
31        return Err(StorageBackendError::Other(format!(
32            "vector dimension mismatch: expected {dimensions}, got {}",
33            vector.len()
34        )));
35    }
36    if let Some((index, value)) = vector
37        .iter()
38        .copied()
39        .enumerate()
40        .find(|(_, value)| !value.is_finite())
41    {
42        return Err(StorageBackendError::Other(format!(
43            "vector component {index} must be finite, got {value}"
44        )));
45    }
46    Ok(())
47}
48
49fn checked_vector_count(counts: impl IntoIterator<Item = usize>) -> StorageBackendResult<usize> {
50    counts.into_iter().try_fold(0_usize, |total, count| {
51        total
52            .checked_add(count)
53            .ok_or_else(|| StorageBackendError::Other("vector count overflow".into()))
54    })
55}
56
57fn validate_threshold(threshold: f32) -> StorageBackendResult<()> {
58    if threshold.is_finite() {
59        Ok(())
60    } else {
61        Err(StorageBackendError::Other(format!(
62            "vector similarity threshold must be finite, got {threshold}"
63        )))
64    }
65}
66
67pub(crate) fn select_top_k_scored(scored: &mut Vec<(DocId, f32)>, k: usize) {
68    if scored.len() > k {
69        scored.select_nth_unstable_by(k, |a, b| b.1.total_cmp(&a.1).then_with(|| a.0.cmp(&b.0)));
70        scored.truncate(k);
71    }
72}
73
74/// Collapse tensor-vector scores to the best score for each document without
75/// allocating one tree node per candidate. The final sort performed by each
76/// caller restores the posting-list invariant after top-k selection.
77pub(crate) fn deduplicate_scored_by_doc(scored: &mut Vec<(DocId, f32)>) {
78    if scored.len() < 2 {
79        return;
80    }
81    scored.sort_unstable_by_key(|(doc_id, _)| *doc_id);
82    let mut write = 1;
83    for read in 1..scored.len() {
84        let (doc_id, score) = scored[read];
85        if scored[write - 1].0 == doc_id {
86            scored[write - 1].1 = scored[write - 1].1.max(score);
87        } else {
88            scored[write] = (doc_id, score);
89            write += 1;
90        }
91    }
92    scored.truncate(write);
93}
94
95pub(crate) fn vector_norm(vector: &[f32]) -> f32 {
96    let mut squared_norm = 0.0_f32;
97    for value in vector {
98        squared_norm += value * value;
99    }
100    squared_norm.sqrt()
101}
102
103/// Cosine similarity when both vector norms were computed once outside the
104/// candidate loop. This preserves the raw-vector score while avoiding two
105/// norm reductions and two square roots for every candidate.
106pub(crate) fn cosine_similarity_with_norms(a: &[f32], b: &[f32], norm_a: f32, norm_b: f32) -> f32 {
107    if a.len() != b.len() || a.is_empty() || norm_a == 0.0 || norm_b == 0.0 {
108        return 0.0;
109    }
110    let mut dot = 0.0_f32;
111    for (x, y) in a.iter().zip(b) {
112        dot += x * y;
113    }
114    dot / (norm_a * norm_b)
115}
116
117/// Cosine similarity between two equal-length vectors. Returns `0.0` when
118/// either vector has zero norm or the dimensions differ.
119///
120/// Arithmetic stays in `f32` end-to-end so the result is bit-equal to
121/// the reference `NumPy` implementation (`np.dot(q, v) / (||q|| * ||v||)`
122/// over `float32` arrays).
123pub fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
124    if a.len() != b.len() || a.is_empty() {
125        return 0.0;
126    }
127    let mut dot = 0.0f32;
128    let mut norm_a = 0.0f32;
129    let mut norm_b = 0.0f32;
130    for (x, y) in a.iter().zip(b.iter()) {
131        dot += x * y;
132        norm_a += x * x;
133        norm_b += y * y;
134    }
135    if norm_a == 0.0 || norm_b == 0.0 {
136        return 0.0;
137    }
138    dot / (norm_a.sqrt() * norm_b.sqrt())
139}
140
141pub trait VectorIndex: Send + Sync {
142    fn dimensions(&self) -> u32;
143    fn index_kind(&self) -> &'static str {
144        "vector"
145    }
146    fn add(&mut self, doc_id: DocId, vector: Vec<f32>) -> StorageBackendResult<()>;
147    fn add_many(&mut self, doc_id: DocId, vectors: Vec<Vec<f32>>) -> StorageBackendResult<()>;
148    fn delete(&mut self, doc_id: DocId) -> StorageBackendResult<()>;
149    fn clear(&mut self) -> StorageBackendResult<()>;
150    fn search_knn(&self, query: &[f32], k: usize) -> StorageBackendResult<PostingList>;
151    fn search_threshold(&self, query: &[f32], threshold: f32) -> StorageBackendResult<PostingList>;
152    fn count(&self) -> StorageBackendResult<usize>;
153
154    /// Build any auxiliary physical metadata required by this index from its
155    /// current vector contents. Brute-force indexes need no extra work;
156    /// persistent IVF and HNSW implementations use this during explicit
157    /// index creation, while restore paths deliberately skip it.
158    fn initialize(&mut self) -> StorageBackendResult<()> {
159        Ok(())
160    }
161
162    /// Read-only handle suitable for an `ExecutionContext`.
163    fn snapshot(&self) -> StorageBackendResult<Arc<dyn VectorIndex>>;
164
165    /// Independent writable copy used by in-memory engine rollback. The
166    /// default keeps third-party and persistent implementations source
167    /// compatible; only indexes hosted by a memory engine need to support it.
168    fn writable_snapshot(&self) -> StorageBackendResult<Box<dyn VectorIndex>> {
169        Err(StorageBackendError::Other(
170            "writable vector-index snapshots are not supported by this backend".into(),
171        ))
172    }
173}
174
175#[derive(Debug, Clone)]
176pub struct MemoryVectorIndex {
177    dimensions: u32,
178    vectors: BTreeMap<DocId, Vec<Vec<f32>>>,
179}
180
181impl MemoryVectorIndex {
182    pub fn new(dimensions: u32) -> Self {
183        Self {
184            dimensions,
185            vectors: BTreeMap::new(),
186        }
187    }
188
189    pub fn vectors(&self) -> &BTreeMap<DocId, Vec<Vec<f32>>> {
190        &self.vectors
191    }
192}
193
194impl VectorIndex for MemoryVectorIndex {
195    fn dimensions(&self) -> u32 {
196        self.dimensions
197    }
198
199    fn index_kind(&self) -> &'static str {
200        "memory-bruteforce"
201    }
202
203    fn add(&mut self, doc_id: DocId, vector: Vec<f32>) -> StorageBackendResult<()> {
204        self.validate_dimensions(&vector)?;
205        self.vectors.insert(doc_id, vec![vector]);
206        Ok(())
207    }
208
209    fn add_many(&mut self, doc_id: DocId, vectors: Vec<Vec<f32>>) -> StorageBackendResult<()> {
210        for vector in &vectors {
211            self.validate_dimensions(vector)?;
212        }
213        if vectors.is_empty() {
214            self.vectors.remove(&doc_id);
215        } else {
216            self.vectors.insert(doc_id, vectors);
217        }
218        Ok(())
219    }
220
221    fn delete(&mut self, doc_id: DocId) -> StorageBackendResult<()> {
222        self.vectors.remove(&doc_id);
223        Ok(())
224    }
225
226    fn clear(&mut self) -> StorageBackendResult<()> {
227        self.vectors.clear();
228        Ok(())
229    }
230
231    /// Brute-force top-k by cosine similarity, descending.
232    fn search_knn(&self, query: &[f32], k: usize) -> StorageBackendResult<PostingList> {
233        self.validate_dimensions(query)?;
234        if k == 0 || self.vectors.is_empty() {
235            return Ok(PostingList::new());
236        }
237        let mut scored: Vec<(DocId, f32)> = self
238            .vectors
239            .iter()
240            .filter_map(|(&doc_id, vectors)| best_vector_score(query, vectors).map(|s| (doc_id, s)))
241            .collect();
242        select_top_k_scored(&mut scored, k);
243        // The output of `top_k` is re-sorted by doc_id ascending so the
244        // posting list invariant holds; the score lives in the payload.
245        scored.sort_by_key(|(id, _)| *id);
246        let entries = scored
247            .into_iter()
248            .map(|(doc_id, sim)| PostingEntry::new(doc_id, Payload::with_score(f64::from(sim))))
249            .collect::<Vec<_>>();
250        Ok(PostingList::from_sorted_unchecked(entries))
251    }
252
253    /// Brute-force threshold scan: keep all docs with `cosine >= threshold`.
254    fn search_threshold(&self, query: &[f32], threshold: f32) -> StorageBackendResult<PostingList> {
255        self.validate_dimensions(query)?;
256        validate_threshold(threshold)?;
257        let mut entries: Vec<PostingEntry> = self
258            .vectors
259            .iter()
260            .filter_map(|(&doc_id, vectors)| {
261                let sim = best_vector_score(query, vectors)?;
262                if sim >= threshold {
263                    Some(PostingEntry::new(
264                        doc_id,
265                        Payload::with_score(f64::from(sim)),
266                    ))
267                } else {
268                    None
269                }
270            })
271            .collect();
272        // The BTreeMap iteration is already doc_id-ascending, so the
273        // filter preserves the invariant.
274        entries.sort_by_key(|e| e.doc_id);
275        Ok(PostingList::from_sorted_unchecked(entries))
276    }
277
278    fn count(&self) -> StorageBackendResult<usize> {
279        checked_vector_count(self.vectors.values().map(Vec::len))
280    }
281
282    fn snapshot(&self) -> StorageBackendResult<Arc<dyn VectorIndex>> {
283        Ok(Arc::new(self.clone()))
284    }
285
286    fn writable_snapshot(&self) -> StorageBackendResult<Box<dyn VectorIndex>> {
287        Ok(Box::new(self.clone()))
288    }
289}
290
291impl MemoryVectorIndex {
292    fn validate_dimensions(&self, vector: &[f32]) -> StorageBackendResult<()> {
293        validate_vector_values(self.dimensions, vector)
294    }
295}
296
297fn best_vector_score(query: &[f32], vectors: &[Vec<f32>]) -> Option<f32> {
298    vectors
299        .iter()
300        .map(|vector| cosine_similarity(query, vector))
301        .max_by(f32::total_cmp)
302}
303
304#[cfg(test)]
305mod tests {
306    use super::*;
307
308    fn approx_eq(a: f32, b: f32, eps: f32) {
309        assert!((a - b).abs() < eps, "expected {a} ~ {b} within {eps}");
310    }
311
312    #[test]
313    fn cosine_identity_is_one() {
314        let v = vec![1.0, 2.0, 3.0];
315        approx_eq(cosine_similarity(&v, &v), 1.0, 1e-6);
316    }
317
318    #[test]
319    fn vector_count_overflow_is_reported() {
320        let error = checked_vector_count([usize::MAX, 1]).unwrap_err();
321        assert!(error.to_string().contains("vector count overflow"));
322    }
323
324    #[test]
325    fn cosine_orthogonal_is_zero() {
326        let a = vec![1.0, 0.0];
327        let b = vec![0.0, 1.0];
328        approx_eq(cosine_similarity(&a, &b), 0.0, 1e-6);
329    }
330
331    #[test]
332    fn cosine_zero_norm_is_zero() {
333        let a = vec![0.0, 0.0];
334        let b = vec![1.0, 1.0];
335        approx_eq(cosine_similarity(&a, &b), 0.0, 1e-6);
336    }
337
338    #[test]
339    fn knn_orders_by_similarity_descending_then_doc_id() {
340        let mut idx = MemoryVectorIndex::new(2);
341        idx.add(1, vec![1.0, 0.0]).unwrap();
342        idx.add(2, vec![0.5, 0.5]).unwrap();
343        idx.add(3, vec![0.0, 1.0]).unwrap();
344        let pl = idx.search_knn(&[1.0, 0.0], 2).unwrap();
345        let docs: Vec<_> = pl.iter().map(|e| e.doc_id).collect();
346        // posting list is doc_id-sorted but the top-2 should be {1, 2}.
347        assert_eq!(docs, vec![1, 2]);
348        let entry1 = pl.get_entry(1).unwrap();
349        let entry2 = pl.get_entry(2).unwrap();
350        assert!(entry1.payload.score > entry2.payload.score);
351    }
352
353    #[test]
354    fn partial_top_k_keeps_deterministic_doc_id_ties() {
355        let mut scored = vec![(10, 0.5), (3, 0.9), (1, 0.9), (8, 0.7), (2, 0.1)];
356        select_top_k_scored(&mut scored, 2);
357        scored.sort_by_key(|(doc_id, _)| *doc_id);
358        assert_eq!(scored, vec![(1, 0.9), (3, 0.9)]);
359    }
360
361    #[test]
362    fn precomputed_norm_cosine_matches_reference_bits() {
363        let a = [0.25, -3.0, 1.5, 8.0];
364        let b = [2.0, 0.75, -4.0, 0.5];
365        let expected = cosine_similarity(&a, &b);
366        let actual = cosine_similarity_with_norms(&a, &b, vector_norm(&a), vector_norm(&b));
367        assert_eq!(actual.to_bits(), expected.to_bits());
368    }
369
370    #[test]
371    fn score_deduplication_keeps_best_tensor_vector() {
372        let mut scored = vec![(7, 0.3), (2, 0.8), (7, 0.9), (2, 0.4), (9, -0.2)];
373        deduplicate_scored_by_doc(&mut scored);
374        assert_eq!(scored, vec![(2, 0.8), (7, 0.9), (9, -0.2)]);
375    }
376
377    #[test]
378    fn threshold_filters_below_cutoff() {
379        let mut idx = MemoryVectorIndex::new(2);
380        idx.add(1, vec![1.0, 0.0]).unwrap();
381        idx.add(2, vec![0.5, 0.5]).unwrap();
382        idx.add(3, vec![0.0, 1.0]).unwrap();
383        let pl = idx.search_threshold(&[1.0, 0.0], 0.5).unwrap();
384        let docs: Vec<_> = pl.iter().map(|e| e.doc_id).collect();
385        assert_eq!(docs, vec![1, 2]);
386    }
387
388    #[test]
389    fn delete_removes_vector() {
390        let mut idx = MemoryVectorIndex::new(2);
391        idx.add(1, vec![1.0, 0.0]).unwrap();
392        idx.delete(1).unwrap();
393        assert_eq!(idx.count().unwrap(), 0);
394    }
395
396    #[test]
397    fn non_finite_vectors_queries_and_thresholds_are_errors() {
398        let mut idx = MemoryVectorIndex::new(2);
399        assert!(idx.add(1, vec![f32::NAN, 0.0]).is_err());
400        idx.add(1, vec![1.0, 0.0]).unwrap();
401        assert!(idx.search_knn(&[f32::INFINITY, 0.0], 1).is_err());
402        assert!(idx
403            .search_threshold(&[1.0, 0.0], f32::NEG_INFINITY)
404            .is_err());
405    }
406}