1use 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#[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}