1use crate::graph::{Graph, Metric, CAT_NAV_PENDING};
28use crate::keys;
29use crate::vecquant::{self, AffineQuery};
30use crate::{Error, Result};
31
32pub const NAV_R: usize = 32;
35pub const NAV_L_BUILD: usize = 64;
37pub const NAV_ALPHA: f32 = 1.2;
44pub const NAV_FOLD_CHECKPOINT_EVERY: u64 = 25_000;
51
52pub const NAV_R_SLACK: usize = NAV_R * 2;
57
58pub(crate) fn sortable(d: f32) -> u64 {
67 let b = d.to_bits();
68 (if d >= 0.0 { b | 0x8000_0000 } else { !b }) as u64
69}
70
71struct NavRow {
72 norm: f32,
73 code_off: usize,
74 code_len: usize,
75 neighbors: Vec<u64>,
76}
77
78fn decode_row(v: &[u8], code_len: usize) -> Result<NavRow> {
79 let count_at = 4usize.checked_add(code_len).ok_or(Error::Corrupt {
80 page_no: 0,
81 why: "navigation row code length overflows",
82 })?;
83 let header_end = count_at.checked_add(2).ok_or(Error::Corrupt {
84 page_no: 0,
85 why: "navigation row header length overflows",
86 })?;
87 let count_bytes = v.get(count_at..header_end).ok_or(Error::Corrupt {
88 page_no: 0,
89 why: "navigation row is shorter than its code and count",
90 })?;
91 let norm = f32::from_le_bytes(v.get(..4).ok_or(Error::Corrupt {
92 page_no: 0,
93 why: "navigation row has no norm",
94 })?.try_into().unwrap());
95 let n = u16::from_le_bytes(count_bytes.try_into().unwrap()) as usize;
96 if n > NAV_R_SLACK {
97 return Err(Error::Corrupt { page_no: 0, why: "navigation row has too many neighbours" });
98 }
99 let expected = n.checked_mul(8).and_then(|bytes| header_end.checked_add(bytes))
100 .ok_or(Error::Corrupt { page_no: 0, why: "navigation row length overflows" })?;
101 if v.len() != expected {
102 return Err(Error::Corrupt { page_no: 0, why: "navigation row length is invalid" });
103 }
104 let mut neighbors = Vec::with_capacity(n);
105 for i in 0..n {
106 let o = header_end + i * 8;
107 neighbors.push(u64::from_le_bytes(v[o..o + 8].try_into().unwrap()));
108 }
109 Ok(NavRow { norm, code_off: 4, code_len, neighbors })
110}
111
112fn encode_row(norm: f32, code: &[u8], neighbors: &[u64]) -> Vec<u8> {
113 let mut v = Vec::with_capacity(4 + code.len() + 2 + neighbors.len() * 8);
114 v.extend_from_slice(&norm.to_le_bytes());
115 v.extend_from_slice(code);
116 v.extend_from_slice(&(neighbors.len() as u16).to_le_bytes());
117 for n in neighbors { v.extend_from_slice(&n.to_le_bytes()); }
118 v
119}
120
121impl Graph {
122 fn nav_pending_ids(&self, field: u64, out: &mut Vec<u64>) -> Result<()> {
125 let prefix = keys::catalog_field(CAT_NAV_PENDING, field);
126 self.store_ref().scan(&prefix)?.for_each_ref(|key, _| {
127 if !key.starts_with(&prefix) { return false; }
128 if key.len() == 25 { out.push(keys::u64_at(key, 17)); }
129 true
130 })
131 }
132
133 fn nav_code_len(&self, field: u64) -> usize {
134 vecquant::code_len(vecquant::pad_dim(self.vec_dim(field) as usize),
135 self.vec_bits_pub(field))
136 }
137
138 pub fn nav_walk_diag(&self, field: u64, q: &[f32], ef: usize)
142 -> Result<(Vec<u64>, std::collections::HashSet<u64>)>
143 {
144 let medoid = match self.vec_meta(field) {
145 Some(m) if m.medoid != 0 => m.medoid,
146 _ => return Ok((Vec::new(), Default::default())),
147 };
148 let Some(enc) = self.encoder(field) else { return Ok((Vec::new(), Default::default())) };
149 let aq = enc.affine_query(q);
150 let mut visited = std::collections::HashSet::new();
151 let kept = self.nav_beam_inner(field, &aq, Metric::L2, ef, medoid, Some(&mut visited))?;
152 Ok((kept.into_iter().map(|(_, id)| id).collect(), visited))
153 }
154
155 fn nav_beam(&self, field: u64, aq: &AffineQuery, metric: Metric, ef: usize, start: u64)
158 -> Result<Vec<(f32, u64)>>
159 {
160 self.nav_beam_inner(field, aq, metric, ef, start, None)
161 }
162
163 fn nav_beam_inner(&self, field: u64, aq: &AffineQuery, metric: Metric, ef: usize, start: u64,
164 mut diag: Option<&mut std::collections::HashSet<u64>>)
165 -> Result<Vec<(f32, u64)>>
166 {
167 let code_len = self.nav_code_len(field);
168 let est = |row: &NavRow, v: &[u8]| -> f32 {
169 let dot = vecquant::dot_est_affine2(row.norm, &v[row.code_off..row.code_off + row.code_len], aq);
170 match metric {
171 Metric::L2 | Metric::L1 => row.norm * row.norm - 2.0 * dot,
172 Metric::Dot => -dot,
173 Metric::Cosine => if row.norm > 0.0 { -dot / row.norm } else { 0.0 },
174 }
175 };
176 let mut visited: std::collections::HashSet<u64> = std::collections::HashSet::new();
177 let mut frontier2: std::collections::BinaryHeap<std::cmp::Reverse<(u64, u64)>> =
179 std::collections::BinaryHeap::new();
180 let mut kept: std::collections::BinaryHeap<(u64, u64)> = std::collections::BinaryHeap::new();
181 let seed = match self.store_ref().get(&keys::nav_key(field, start))? {
182 Some(v) => v,
183 None => return Ok(Vec::new()),
184 };
185 let row = decode_row(&seed, code_len)?;
186 let d0 = est(&row, &seed);
187 visited.insert(start);
188 frontier2.push(std::cmp::Reverse((sortable(d0), start)));
189 kept.push((sortable(d0), start));
190 let mut dists: std::collections::HashMap<u64, f32> = std::collections::HashMap::new();
191 dists.insert(start, d0);
192
193 while let Some(std::cmp::Reverse((ds, id))) = frontier2.pop() {
194 if kept.len() >= ef {
196 if let Some(&(worst, _)) = kept.peek() {
197 if ds > worst { break; }
198 }
199 }
200 let Some(v) = self.store_ref().get(&keys::nav_key(field, id))? else { continue };
201 let row = decode_row(&v, code_len)?;
202 for &nb in &row.neighbors {
203 if !visited.insert(nb) { continue; }
204 if let Some(d) = diag.as_deref_mut() { d.insert(nb); }
205 let Some(nv) = self.store_ref().get(&keys::nav_key(field, nb))? else { continue };
206 let nrow = decode_row(&nv, code_len)?;
207 let nd = est(&nrow, &nv);
208 let nds = sortable(nd);
209 let admit = kept.len() < ef || kept.peek().map(|&(w, _)| nds < w).unwrap_or(true);
210 if admit {
211 kept.push((nds, nb));
212 if kept.len() > ef { kept.pop(); }
213 frontier2.push(std::cmp::Reverse((nds, nb)));
214 dists.insert(nb, nd);
215 }
216 }
217 }
218 let mut out: Vec<(f32, u64)> = kept.into_iter()
219 .map(|(_, id)| (dists.get(&id).copied().unwrap_or(f32::MAX), id))
220 .collect();
221 out.sort_by(|a, b| a.0.total_cmp(&b.0));
222 Ok(out)
223 }
224
225 fn nav_insert(&mut self, field: u64, id: u64, medoid: u64, refresh: bool) -> Result<()> {
229 let code_len = self.nav_code_len(field);
230 let Some(vrow) = self.store_ref().get(&keys::vcode_key(field, id))? else { return Ok(()) };
231 if vrow.len() != 4usize.saturating_add(code_len) {
232 return Err(Error::Corrupt { page_no: 0, why: "vector code row length is invalid" });
233 }
234 let norm = f32::from_le_bytes(vrow[0..4].try_into().unwrap());
235 let code = vrow[4..].to_vec();
236 let Some(v) = self.get_vec(field, id)? else { return Ok(()) };
238 let Some(enc) = self.encoder(field) else { return Ok(()) };
239 let aq = enc.affine_query(&v);
240
241 let mut cands = if id == medoid && !refresh {
242 Vec::new()
243 } else {
244 self.nav_beam(field, &aq, Metric::L2, NAV_L_BUILD, medoid)?
245 };
246 cands.retain(|&(_, candidate)| candidate != id);
249 let mut vcache: std::collections::HashMap<u64, Vec<f32>> =
252 std::collections::HashMap::new();
253 vcache.insert(id, v.clone());
254 for c in cands.iter_mut() {
255 if let Some(d) = self.pair_dist_cached(field, &mut vcache, id, c.1)? { c.0 = d; }
256 }
257 let pruned = self.robust_prune(field, &mut vcache, id, &cands, code_len)?;
258 self.store().put(&keys::nav_key(field, id), &encode_row(norm, &code, &pruned))?;
259 for &nb in &pruned {
261 let Some(nv) = self.store_ref().get(&keys::nav_key(field, nb))? else { continue };
262 let mut nrow = decode_row(&nv, code_len)?;
263 if nrow.neighbors.contains(&id) { continue; }
264 nrow.neighbors.push(id);
265 let links = if nrow.neighbors.len() > NAV_R_SLACK {
266 let mut cds = Vec::with_capacity(nrow.neighbors.len());
267 for &x in &nrow.neighbors {
268 cds.push((self.pair_dist_cached(field, &mut vcache, nb, x)?.unwrap_or(f32::MAX), x));
269 }
270 self.robust_prune(field, &mut vcache, nb, &cds, code_len)?
271 } else { nrow.neighbors.clone() };
272 let ncode = nv[4..4 + code_len].to_vec();
273 self.store().put(&keys::nav_key(field, nb), &encode_row(nrow.norm, &ncode, &links))?;
274 }
275 Ok(())
276 }
277
278 fn pair_dist_cached(&self, field: u64, cache: &mut std::collections::HashMap<u64, Vec<f32>>,
283 a: u64, b: u64) -> Result<Option<f32>> {
284 for id in [a, b] {
285 if !cache.contains_key(&id) {
286 let Some(vector) = self.get_vec(field, id)? else { return Ok(None) };
287 cache.insert(id, vector);
288 }
289 }
290 let (Some(va), Some(vb)) = (cache.get(&a), cache.get(&b)) else { return Ok(None) };
291 Ok(Some(va.iter().zip(vb).map(|(x, y)| (x - y) * (x - y)).sum()))
292 }
293
294 fn robust_prune(&self, field: u64, vcache: &mut std::collections::HashMap<u64, Vec<f32>>,
298 _node: u64, cands: &[(f32, u64)], code_len: usize)
299 -> Result<Vec<u64>>
300 {
301 let mut sorted: Vec<(f32, u64)> = cands.to_vec();
302 sorted.sort_by(|a, b| a.0.total_cmp(&b.0));
303 sorted.dedup_by_key(|c| c.1);
304 let mut kept: Vec<(f32, u64)> = Vec::new();
305 let _ = code_len;
306 'cand: for &(d, c) in &sorted {
307 if kept.len() >= NAV_R { break; }
308 for &(_, k) in &kept {
309 let Some(dkc) = self.pair_dist_cached(field, vcache, k, c)? else { continue };
310 if NAV_ALPHA * NAV_ALPHA * dkc < d { continue 'cand; }
311 }
312 kept.push((d, c));
313 }
314 Ok(kept.into_iter().map(|(_, id)| id).collect())
315 }
316
317 pub fn fold_nav(&mut self, field: u64) -> Result<u64> {
320 let Some(mut meta) = self.vec_meta(field) else { return Ok(0) };
321 let mut stragglers: Vec<u64> = Vec::new();
322 self.nav_pending_ids(field, &mut stragglers)?;
323 stragglers.sort_unstable();
324 stragglers.dedup();
325 let mut pending = stragglers.clone();
326 if let Some(next) = meta.watermark.checked_add(1) {
329 let from = keys::vcode_key(field, next);
330 let it = self.store_ref().scan(&from)?;
331 it.for_each_ref(|key, _| {
332 if key.len() != 17 || key[0] != keys::TAG_VCODE
333 || keys::u64_at(key, 1) != field { return false; }
334 pending.push(keys::u64_at(key, 9));
335 true
336 })?;
337 }
338 pending.sort_unstable();
339 pending.dedup();
340 let n = pending.len() as u64;
341 let mut since = 0u64;
342 for id in pending {
343 let m = match meta.medoid {
344 0 => {
345 meta.medoid = id;
346 self.set_vec_meta(field, meta)?;
347 id
348 }
349 m => m,
350 };
351 let refresh = stragglers.binary_search(&id).is_ok();
352 self.nav_insert(field, id, m, refresh)?;
353 if refresh {
354 self.store().delete(
355 &keys::catalog_field_item(CAT_NAV_PENDING, field, id))?;
356 }
357 meta.watermark = meta.watermark.max(id);
358 self.set_vec_meta(field, meta)?;
359 since += 1;
360 if since >= self.nav_fold_every {
361 self.commit()?;
362 self.checkpoint()?; since = 0;
364 }
365 }
366 self.commit()?;
367 Ok(n)
368 }
369
370 pub fn nearest_nav(&self, field: u64, q: &[f32], k: usize, metric: Metric, oversample: usize)
374 -> Result<Vec<(u64, f32)>>
375 {
376 let meta = match self.vec_meta(field) {
377 Some(m) if m.medoid != 0 => m,
378 _ => return self.nearest(field, q, k, metric, oversample),
379 };
380 let Some(enc) = self.encoder(field) else {
381 return self.nearest(field, q, k, metric, oversample);
382 };
383 let aq = enc.affine_query(q);
384 let ef = (k * oversample).max(k).max(64);
385 let mut cands: Vec<u64> = self.nav_beam(field, &aq, metric, ef, meta.medoid)?
386 .into_iter().map(|(_, id)| id).collect();
387 let mut stragglers = Vec::new();
388 self.nav_pending_ids(field, &mut stragglers)?;
389 for id in stragglers {
390 if !cands.contains(&id) { cands.push(id); }
393 }
394 if let Some(next) = meta.watermark.checked_add(1) {
396 let head_from = keys::vcode_key(field, next);
397 let it = self.store_ref().scan(&head_from)?;
398 it.for_each_ref(|key, _| {
399 if key.len() != 17 || key[0] != keys::TAG_VCODE
400 || keys::u64_at(key, 1) != field { return false; }
401 cands.push(keys::u64_at(key, 9));
402 true
403 })?;
404 }
405 self.rescore(field, &cands, q, metric, k)
406 }
407}
408
409#[cfg(test)]
410mod tests {
411 use super::*;
412 use crate::store::{Config, Store};
413
414 #[test]
418 fn sortable_orders_floats_across_the_sign_boundary() {
419 let seq = [-3.5f32, -1.0, -0.25, 0.0, 0.25, 1.0, 3.5];
420 for w in seq.windows(2) {
421 assert!(sortable(w[0]) < sortable(w[1]),
422 "{} must sort below {}", w[0], w[1]);
423 }
424 }
425
426 #[test]
427 fn malformed_navigation_rows_are_errors_not_missing_nodes() {
428 assert!(matches!(decode_row(&[0; 3], 8), Err(Error::Corrupt { .. })));
429
430 let mut trailing = encode_row(1.0, &[0; 8], &[7]);
431 trailing.push(0);
432 assert!(matches!(decode_row(&trailing, 8), Err(Error::Corrupt { .. })));
433 }
434
435 #[test]
440 fn robust_prune_alpha_acts_on_true_distances() {
441 let d = tempfile::TempDir::new().unwrap();
442 let g = Graph::new(Store::create(d.path(), Config::default()).unwrap()).unwrap();
443 let mut vc: std::collections::HashMap<u64, Vec<f32>> = std::collections::HashMap::new();
444 vc.insert(10, vec![1.0, 0.0]); vc.insert(20, vec![0.68375, 1.10567]); vc.insert(30, vec![1.05, 0.0]); let cands = vec![
448 (1.0f32, 10u64), (1.69, 20), (1.1025, 30), ];
452 let kept = g.robust_prune(0, &mut vc, 0, &cands, 0).unwrap();
453 assert_eq!(kept, vec![10, 20], "alpha must apply to true distances, not squared");
456 }
457}