Skip to main content

uqa_storage/key_value/
vector_index.rs

1//
2// Unified Query Algebra
3//
4// Copyright (c) 2023-2026 Cognica, Inc.
5//
6
7//! Vector-index adapter over an ordered key/value store.
8
9use std::collections::BTreeMap;
10
11use super::codec::{
12    blob_to_vector, other_error, read_str, read_u64, usize_to_u64, validate_vector_ordinal_count,
13    vector_doc_prefix, vector_field_prefix, vector_key, vector_to_blob,
14};
15use super::{
16    cosine_similarity, validate_vector_values, Arc, DocId, KeyValueBatch, KeyValueStore, Payload,
17    PostingEntry, PostingList, StorageBackendResult, VectorIndex,
18};
19
20/// Brute-force vector index implemented over [`KeyValueStore`].
21#[derive(Clone)]
22pub struct KeyValueVectorIndex {
23    store: Arc<dyn KeyValueStore>,
24    table: String,
25    field: String,
26    dimensions: u32,
27}
28
29impl KeyValueVectorIndex {
30    pub fn new(
31        store: Arc<dyn KeyValueStore>,
32        table: impl Into<String>,
33        field: impl Into<String>,
34        dimensions: u32,
35    ) -> Self {
36        Self {
37            store,
38            table: table.into(),
39            field: field.into(),
40            dimensions,
41        }
42    }
43
44    pub(super) fn load_all_with_ordinals(
45        &self,
46    ) -> StorageBackendResult<Vec<(DocId, u32, Vec<f32>)>> {
47        let mut vectors = Vec::new();
48        let mut current_doc = None;
49        let mut expected_ordinal = 0_u32;
50        for (key, value) in self
51            .store
52            .scan_prefix(&vector_field_prefix(&self.table, &self.field)?)?
53        {
54            let mut offset = 1;
55            let _table = read_str(&key, &mut offset)?;
56            let _field = read_str(&key, &mut offset)?;
57            let doc_id = read_u64(&key, &mut offset)?;
58            let ordinal = read_u64(&key, &mut offset)?;
59            let ordinal = u32::try_from(ordinal)
60                .map_err(|_| other_error("persisted vector ordinal exceeds u32 index format"))?;
61            if offset != key.len() {
62                return Err(other_error("persisted vector key has trailing bytes"));
63            }
64            if current_doc != Some(doc_id) {
65                current_doc = Some(doc_id);
66                expected_ordinal = 0;
67            }
68            if ordinal != expected_ordinal {
69                return Err(other_error(format!(
70                    "invalid persisted vector ordinal sequence for document {doc_id}: expected {expected_ordinal}, found {ordinal}"
71                )));
72            }
73            expected_ordinal = expected_ordinal
74                .checked_add(1)
75                .ok_or_else(|| other_error("persisted vector ordinal sequence overflow"))?;
76            let vector = blob_to_vector(&value)?;
77            self.validate_dimensions(&vector)?;
78            vectors.push((doc_id, ordinal, vector));
79        }
80        Ok(vectors)
81    }
82
83    pub(super) fn load_all(&self) -> StorageBackendResult<Vec<(DocId, Vec<f32>)>> {
84        Ok(self
85            .load_all_with_ordinals()?
86            .into_iter()
87            .map(|(doc_id, _, vector)| (doc_id, vector))
88            .collect())
89    }
90
91    pub(super) fn load_by_document(&self) -> StorageBackendResult<BTreeMap<DocId, Vec<Vec<f32>>>> {
92        let mut grouped = BTreeMap::<DocId, Vec<Vec<f32>>>::new();
93        for (doc_id, ordinal, vector) in self.load_all_with_ordinals()? {
94            let vectors = grouped.entry(doc_id).or_default();
95            if usize::try_from(ordinal).ok() != Some(vectors.len()) {
96                return Err(other_error(format!(
97                    "invalid canonical vector ordinal for document {doc_id}: expected {}, found {ordinal}",
98                    vectors.len()
99                )));
100            }
101            vectors.push(vector);
102        }
103        Ok(grouped)
104    }
105
106    pub(super) fn stage_replace(
107        &self,
108        batch: &mut dyn KeyValueBatch,
109        doc_id: DocId,
110        vectors: &[Vec<f32>],
111    ) -> StorageBackendResult<()> {
112        for vector in vectors {
113            self.validate_dimensions(vector)?;
114        }
115        validate_vector_ordinal_count(usize_to_u64(vectors.len(), "vector count")?)?;
116        batch.delete_prefix(&vector_doc_prefix(&self.table, &self.field, doc_id)?)?;
117        for (ordinal, vector) in vectors.iter().enumerate() {
118            let ordinal = u32::try_from(ordinal)
119                .map_err(|_| other_error("vector ordinal exceeds u32 index format"))?;
120            batch.put(
121                &vector_key(&self.table, &self.field, doc_id, ordinal)?,
122                &vector_to_blob(vector)?,
123            )?;
124        }
125        Ok(())
126    }
127
128    pub(super) fn stage_clear(&self, batch: &mut dyn KeyValueBatch) -> StorageBackendResult<()> {
129        batch.delete_prefix(&vector_field_prefix(&self.table, &self.field)?)
130    }
131
132    fn validate_dimensions(&self, vector: &[f32]) -> StorageBackendResult<()> {
133        validate_vector_values(self.dimensions, vector).map_err(|error| {
134            other_error(format!(
135                "invalid vector for {}.{}: {error}",
136                self.table, self.field
137            ))
138        })
139    }
140}
141
142impl VectorIndex for KeyValueVectorIndex {
143    fn dimensions(&self) -> u32 {
144        self.dimensions
145    }
146
147    fn index_kind(&self) -> &'static str {
148        "keyvalue-bruteforce"
149    }
150
151    fn add(&mut self, doc_id: DocId, vector: Vec<f32>) -> StorageBackendResult<()> {
152        self.add_many(doc_id, vec![vector])
153    }
154
155    fn add_many(&mut self, doc_id: DocId, vectors: Vec<Vec<f32>>) -> StorageBackendResult<()> {
156        let mut batch = self.store.batch();
157        self.stage_replace(batch.as_mut(), doc_id, &vectors)?;
158        batch.commit()
159    }
160
161    fn delete(&mut self, doc_id: DocId) -> StorageBackendResult<()> {
162        let mut batch = self.store.batch();
163        batch.delete_prefix(&vector_doc_prefix(&self.table, &self.field, doc_id)?)?;
164        batch.commit()
165    }
166
167    fn clear(&mut self) -> StorageBackendResult<()> {
168        let mut batch = self.store.batch();
169        self.stage_clear(batch.as_mut())?;
170        batch.commit()
171    }
172
173    fn search_knn(&self, query: &[f32], k: usize) -> StorageBackendResult<PostingList> {
174        self.validate_dimensions(query)?;
175        if k == 0 {
176            return Ok(PostingList::new());
177        }
178        let entries = self.load_all()?;
179        let mut best_by_doc = BTreeMap::<DocId, f32>::new();
180        for (doc_id, vector) in &entries {
181            let sim = cosine_similarity(query, vector);
182            best_by_doc
183                .entry(*doc_id)
184                .and_modify(|best| {
185                    if sim > *best {
186                        *best = sim;
187                    }
188                })
189                .or_insert(sim);
190        }
191        let mut scored = best_by_doc.into_iter().collect::<Vec<_>>();
192        scored.sort_by(|a, b| b.1.total_cmp(&a.1).then_with(|| a.0.cmp(&b.0)));
193        scored.truncate(k);
194        scored.sort_by_key(|(doc_id, _)| *doc_id);
195        Ok(PostingList::from_sorted_unchecked(
196            scored
197                .into_iter()
198                .map(|(doc_id, sim)| PostingEntry::new(doc_id, Payload::with_score(f64::from(sim))))
199                .collect(),
200        ))
201    }
202
203    fn search_threshold(&self, query: &[f32], threshold: f32) -> StorageBackendResult<PostingList> {
204        self.validate_dimensions(query)?;
205        if !threshold.is_finite() {
206            return Err(other_error(format!(
207                "vector similarity threshold must be finite, got {threshold}"
208            )));
209        }
210        let mut best_by_doc = BTreeMap::<DocId, f32>::new();
211        for (doc_id, vector) in self.load_all()? {
212            let sim = cosine_similarity(query, &vector);
213            if sim >= threshold {
214                best_by_doc
215                    .entry(doc_id)
216                    .and_modify(|best| {
217                        if sim > *best {
218                            *best = sim;
219                        }
220                    })
221                    .or_insert(sim);
222            }
223        }
224        let mut entries = best_by_doc
225            .into_iter()
226            .map(|(doc_id, sim)| PostingEntry::new(doc_id, Payload::with_score(f64::from(sim))))
227            .collect::<Vec<_>>();
228        entries.sort_by_key(|entry| entry.doc_id);
229        Ok(PostingList::from_sorted_unchecked(entries))
230    }
231
232    fn count(&self) -> StorageBackendResult<usize> {
233        Ok(self
234            .store
235            .scan_prefix(&vector_field_prefix(&self.table, &self.field)?)?
236            .len())
237    }
238
239    fn snapshot(&self) -> StorageBackendResult<Arc<dyn VectorIndex>> {
240        Ok(Arc::new(self.clone()))
241    }
242}