Skip to main content

lance_index/vector/hnsw/
builder.rs

1// SPDX-License-Identifier: Apache-2.0
2// SPDX-FileCopyrightText: Copyright The Lance Authors
3
4//! Builder of Hnsw Graph.
5
6use arrow::array::{AsArray, ListBuilder, UInt32Builder};
7use arrow::compute::concat_batches;
8use arrow::datatypes::{DataType, UInt32Type};
9use arrow_array::{ArrayRef, Float32Array, ListArray, RecordBatch, UInt64Array};
10use crossbeam_queue::ArrayQueue;
11use itertools::Itertools;
12use lance_core::deepsize::DeepSizeOf;
13use lance_core::utils::row_addr_remap::RowAddrRemap;
14
15use lance_core::utils::tokio::get_num_compute_intensive_cpus;
16use lance_linalg::distance::DistanceType;
17use rayon::prelude::*;
18use std::cmp::min;
19use std::collections::{BinaryHeap, HashMap, VecDeque};
20use std::fmt::Debug;
21use std::iter;
22use std::sync::Arc;
23use std::sync::RwLock;
24use std::sync::atomic::{AtomicUsize, Ordering};
25use tracing::instrument;
26
27use lance_core::{Error, Result};
28use rand::{Rng, SeedableRng, rngs::SmallRng};
29use serde::{Deserialize, Serialize};
30
31use super::super::graph::beam_search;
32use super::{
33    HNSW_TYPE, HnswMetadata, VECTOR_ID_COL, VECTOR_ID_FIELD, select_neighbors_heuristic_owned,
34};
35use crate::metrics::MetricsCollector;
36use crate::prefilter::PreFilter;
37use crate::vector::flat::storage::{FlatBinStorage, FlatFloatStorage};
38use crate::vector::graph::builder::GraphBuilderNode;
39use crate::vector::graph::{
40    BorrowingGraph, DISTS_FIELD, Graph, NEIGHBORS_COL, NEIGHBORS_FIELD, OrderedFloat, OrderedNode,
41    VisitedGenerator,
42};
43use crate::vector::graph::{
44    Visited, beam_search_acorn, beam_search_borrowed, greedy_search, greedy_search_borrowed,
45};
46use crate::vector::storage::{DistCalculator, VectorStore};
47use crate::vector::v3::subindex::IvfSubIndex;
48use crate::vector::{ApproxMode, Query, VECTOR_RESULT_SCHEMA};
49
50pub const HNSW_METADATA_KEY: &str = "lance:hnsw";
51
52/// Fixed seed for HNSW node-level assignment.
53///
54/// A constant seed makes graph construction reproducible (same data + params =>
55/// same graph), which keeps index builds deterministic and tests stable. Recall
56/// is statistically unaffected — the level distribution is identical, only the
57/// random draws become fixed. Shared by the offline ([`HNSWBuilder`]) and online
58/// ([`super::online::OnlineHnswBuilder`]) builders so both produce comparable graphs.
59pub(crate) const HNSW_LEVEL_RNG_SEED: u64 = 42;
60
61/// Parameters of building HNSW index
62#[derive(Debug, Clone, Serialize, Deserialize, DeepSizeOf)]
63pub struct HnswBuildParams {
64    /// max level ofm
65    pub max_level: u16,
66
67    /// number of connections to establish while inserting new element
68    pub m: usize,
69
70    /// size of the dynamic list for the candidates
71    pub ef_construction: usize,
72
73    /// number of vectors ahead to prefetch while building the graph
74    pub prefetch_distance: Option<usize>,
75}
76
77impl From<&HnswBuildParams> for crate::pb::HnswParameters {
78    fn from(params: &HnswBuildParams) -> Self {
79        Self {
80            max_connections: params.m as u32,
81            construction_ef: params.ef_construction as u32,
82            max_level: params.max_level as u32,
83        }
84    }
85}
86
87impl Default for HnswBuildParams {
88    fn default() -> Self {
89        Self {
90            max_level: 7,
91            m: 20,
92            ef_construction: 150,
93            prefetch_distance: Some(2),
94        }
95    }
96}
97
98impl HnswBuildParams {
99    /// The maximum level of the graph.
100    /// The default value is `8`.
101    pub fn max_level(mut self, max_level: u16) -> Self {
102        self.max_level = max_level;
103        self
104    }
105
106    /// The number of connections to establish while inserting new element
107    /// The default value is `30`.
108    pub fn num_edges(mut self, m: usize) -> Self {
109        self.m = m;
110        self
111    }
112
113    /// Number of candidates to be considered when searching for the nearest neighbors
114    /// during the construction of the graph.
115    ///
116    /// The default value is `100`.
117    pub fn ef_construction(mut self, ef_construction: usize) -> Self {
118        self.ef_construction = ef_construction;
119        self
120    }
121
122    /// Build the HNSW index from the given data.
123    ///
124    /// # Parameters
125    /// - `data`: A FixedSizeList to build the HNSW.
126    /// - `distance_type`: The distance type to use.
127    pub async fn build(self, data: ArrayRef, distance_type: DistanceType) -> Result<HNSW> {
128        let vectors = data.as_fixed_size_list().clone();
129        match (vectors.value_type(), distance_type) {
130            (DataType::UInt8, DistanceType::Hamming) => {
131                let vec_store = Arc::new(FlatBinStorage::new(vectors, distance_type));
132                HNSW::index_vectors(vec_store.as_ref(), self)
133            }
134            (DataType::UInt8, _) => Err(Error::invalid_input(format!(
135                "HNSW only supports hamming distance for UInt8 vectors, got {}",
136                distance_type
137            ))),
138            (_, DistanceType::Hamming) => Err(Error::invalid_input(format!(
139                "HNSW hamming distance only supports UInt8 vectors, got {}",
140                vectors.value_type()
141            ))),
142            _ => {
143                let vec_store = Arc::new(FlatFloatStorage::new(vectors, distance_type));
144                HNSW::index_vectors(vec_store.as_ref(), self)
145            }
146        }
147    }
148}
149
150/// Build a HNSW graph.
151///
152/// Currently, the HNSW graph is fully built in memory.
153///
154/// During the build, the graph is built layer by layer.
155///
156/// Each node in the graph has a global ID which is the index on the base layer.
157#[derive(Clone, DeepSizeOf)]
158pub struct HNSW {
159    inner: Arc<HnswCore>,
160}
161
162struct HnswCore {
163    params: HnswBuildParams,
164    graph: HnswGraph,
165    level_count: Vec<usize>,
166    entry_point: u32,
167    visited_generator_queue: Arc<ArrayQueue<VisitedGenerator>>,
168}
169
170impl DeepSizeOf for HnswCore {
171    fn deep_size_of_children(&self, context: &mut lance_core::deepsize::Context) -> usize {
172        self.params.deep_size_of_children(context)
173            + self.graph.deep_size_of_children(context)
174            + self.level_count.deep_size_of_children(context)
175        // Skipping the visited_generator_queue
176    }
177}
178
179impl HnswCore {
180    fn max_level(&self) -> u16 {
181        self.params.max_level
182    }
183
184    fn num_nodes(&self, level: usize) -> usize {
185        self.level_count[level]
186    }
187}
188
189impl Debug for HNSW {
190    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
191        write!(f, "HNSW(max_layers: {})", self.inner.max_level() as usize,)
192    }
193}
194
195impl HNSW {
196    /// Construct an HNSW from its constituent parts. Used by the online
197    /// builder when finalizing.
198    pub(crate) fn from_parts(
199        params: HnswBuildParams,
200        nodes: Vec<GraphBuilderNode>,
201        level_count: Vec<usize>,
202        entry_point: u32,
203    ) -> Self {
204        let queue_size = get_num_compute_intensive_cpus().max(1) * 2;
205        let visited_generator_queue = Arc::new(ArrayQueue::new(queue_size));
206        for _ in 0..queue_size {
207            let _ = visited_generator_queue.push(VisitedGenerator::new(0));
208        }
209        Self {
210            inner: Arc::new(HnswCore {
211                params,
212                graph: HnswGraph::Built(Arc::new(nodes)),
213                level_count,
214                entry_point,
215                visited_generator_queue,
216            }),
217        }
218    }
219
220    pub fn empty() -> Self {
221        Self {
222            inner: Arc::new(HnswCore {
223                params: HnswBuildParams::default(),
224                graph: HnswGraph::Built(Arc::new(Vec::new())),
225                level_count: Vec::new(),
226                entry_point: 0,
227                visited_generator_queue: Arc::new(ArrayQueue::new(1)),
228            }),
229        }
230    }
231
232    pub fn len(&self) -> usize {
233        match &self.inner.graph {
234            HnswGraph::Built(nodes) => nodes.len(),
235            // `level_count[0]` is the bottom-level (== total) node count.
236            HnswGraph::Loaded(graph) => graph.level_count[0],
237        }
238    }
239
240    pub fn is_empty(&self) -> bool {
241        self.len() == 0
242    }
243
244    pub fn max_level(&self) -> u16 {
245        self.inner.max_level()
246    }
247
248    pub fn num_nodes(&self, level: usize) -> usize {
249        self.inner.num_nodes(level)
250    }
251
252    /// Returns the in-memory builder nodes, if this graph was freshly built.
253    ///
254    /// A disk-loaded graph is Arrow-backed and has no `GraphBuilderNode`s,
255    /// so this returns `None` for it.
256    pub fn nodes(&self) -> Option<Arc<Vec<GraphBuilderNode>>> {
257        match &self.inner.graph {
258            HnswGraph::Built(nodes) => Some(nodes.clone()),
259            HnswGraph::Loaded(_) => None,
260        }
261    }
262
263    #[allow(clippy::too_many_arguments)]
264    pub fn search_inner(
265        &self,
266        query: ArrayRef,
267        k: usize,
268        params: &HnswQueryParams,
269        bitset: Option<Visited>,
270        visited_generator: &mut VisitedGenerator,
271        storage: &impl VectorStore,
272        prefetch_distance: Option<usize>,
273    ) -> Result<Vec<OrderedNode>> {
274        let dist_calc = storage.dist_calculator(query, params.dist_q_c);
275        let entry = self.inner.entry_point;
276        let ep = OrderedNode::new(entry, dist_calc.distance(entry).into());
277
278        // The level descent + bottom beam search are identical across
279        // graph backends; only the view types differ. `run_search` is
280        // generic over those view types so the loop is single-sourced:
281        // each backend supplies a per-level view closure and a
282        // bottom-level view.
283        let result = match &self.inner.graph {
284            HnswGraph::Built(nodes) => {
285                let nodes = nodes.as_slice();
286                self.run_search(
287                    ep,
288                    k,
289                    params,
290                    bitset.as_ref(),
291                    visited_generator,
292                    storage.len(),
293                    prefetch_distance,
294                    &dist_calc,
295                    |level| ImmutableHnswLevelView::new(level, nodes),
296                    ImmutableHnswBottomView::new(nodes),
297                )
298            }
299            HnswGraph::Loaded(graph) => {
300                let graph = graph.as_ref();
301                self.run_search(
302                    ep,
303                    k,
304                    params,
305                    bitset.as_ref(),
306                    visited_generator,
307                    storage.len(),
308                    prefetch_distance,
309                    &dist_calc,
310                    |level| LoadedHnswLevelView::new(level, graph),
311                    LoadedHnswBottomView::new(graph),
312                )
313            }
314        };
315        Ok(result)
316    }
317
318    /// Drives the shared HNSW query path over backend-specific graph
319    /// views: a per-level view produced by `make_level` and a
320    /// bottom-level view `bottom`. The views borrow their backing store
321    /// and are created, used, and dropped entirely within this call;
322    /// only the owned result escapes. Monomorphizing over `L`/`B` is the
323    /// single seam that lets the in-memory and disk-loaded backends
324    /// share one search loop.
325    #[allow(clippy::too_many_arguments)]
326    fn run_search<L, B>(
327        &self,
328        ep: OrderedNode,
329        k: usize,
330        params: &HnswQueryParams,
331        bitset: Option<&Visited>,
332        visited_generator: &mut VisitedGenerator,
333        storage_len: usize,
334        prefetch_distance: Option<usize>,
335        dist_calc: &impl DistCalculator,
336        make_level: impl Fn(u16) -> L,
337        bottom: B,
338    ) -> Vec<OrderedNode>
339    where
340        L: BorrowingGraph,
341        B: BorrowingGraph,
342    {
343        let mut ep = ep;
344        for level in (0..self.max_level()).rev() {
345            let cur_level = make_level(level);
346            ep = greedy_search_borrowed(
347                &cur_level,
348                ep,
349                dist_calc,
350                self.inner.params.prefetch_distance,
351            );
352        }
353        let mut visited = visited_generator.generate(storage_len);
354        beam_search_borrowed(
355            &bottom,
356            &ep,
357            params,
358            dist_calc,
359            bitset,
360            prefetch_distance,
361            &mut visited,
362        )
363        .into_iter()
364        .take(k)
365        .collect::<Vec<OrderedNode>>()
366    }
367
368    #[instrument(level = "debug", skip(self, query, bitset, storage))]
369    pub fn search_basic(
370        &self,
371        query: ArrayRef,
372        k: usize,
373        params: &HnswQueryParams,
374        bitset: Option<Visited>,
375        storage: &impl VectorStore,
376    ) -> Result<Vec<OrderedNode>> {
377        let mut visited_generator = self
378            .inner
379            .visited_generator_queue
380            .pop()
381            .unwrap_or_else(|| VisitedGenerator::new(storage.len()));
382        let result = self.search_inner(
383            query,
384            k,
385            params,
386            bitset,
387            &mut visited_generator,
388            storage,
389            Some(2),
390        );
391
392        match self.inner.visited_generator_queue.push(visited_generator) {
393            Ok(_) => {}
394            Err(_) => {
395                log::warn!("visited_generator_queue is full");
396            }
397        }
398
399        result
400    }
401
402    /// Like [Self::search_basic] but the bottom level runs
403    /// [beam_search_acorn], which only scores mask-passing nodes.
404    pub fn search_acorn(
405        &self,
406        query: ArrayRef,
407        k: usize,
408        params: &HnswQueryParams,
409        bitset: &Visited,
410        storage: &impl VectorStore,
411    ) -> Result<Vec<OrderedNode>> {
412        let mut visited_generator = self
413            .inner
414            .visited_generator_queue
415            .pop()
416            .unwrap_or_else(|| VisitedGenerator::new(storage.len()));
417        let mut expanded_generator = self
418            .inner
419            .visited_generator_queue
420            .pop()
421            .unwrap_or_else(|| VisitedGenerator::new(storage.len()));
422
423        let result = self.search_acorn_inner(
424            query,
425            k,
426            params,
427            bitset,
428            &mut visited_generator,
429            &mut expanded_generator,
430            storage,
431            Some(2),
432        );
433
434        // if the queue is full, we just don't push it back, so ignore the error here
435        let _ = self.inner.visited_generator_queue.push(visited_generator);
436        let _ = self.inner.visited_generator_queue.push(expanded_generator);
437        result
438    }
439
440    #[allow(clippy::too_many_arguments)]
441    fn search_acorn_inner(
442        &self,
443        query: ArrayRef,
444        k: usize,
445        params: &HnswQueryParams,
446        bitset: &Visited,
447        visited_generator: &mut VisitedGenerator,
448        expanded_generator: &mut VisitedGenerator,
449        storage: &impl VectorStore,
450        prefetch_distance: Option<usize>,
451    ) -> Result<Vec<OrderedNode>> {
452        let dist_calc = storage.dist_calculator(query, params.dist_q_c);
453        let entry = self.inner.entry_point;
454        let ep = OrderedNode::new(entry, dist_calc.distance(entry).into());
455
456        let result = match &self.inner.graph {
457            HnswGraph::Built(nodes) => {
458                let nodes = nodes.as_slice();
459                self.run_search_acorn(
460                    ep,
461                    params,
462                    bitset,
463                    visited_generator,
464                    expanded_generator,
465                    storage.len(),
466                    prefetch_distance,
467                    &dist_calc,
468                    |level| ImmutableHnswLevelView::new(level, nodes),
469                    ImmutableHnswBottomView::new(nodes),
470                )
471            }
472            HnswGraph::Loaded(graph) => {
473                let graph = graph.as_ref();
474                self.run_search_acorn(
475                    ep,
476                    params,
477                    bitset,
478                    visited_generator,
479                    expanded_generator,
480                    storage.len(),
481                    prefetch_distance,
482                    &dist_calc,
483                    |level| LoadedHnswLevelView::new(level, graph),
484                    LoadedHnswBottomView::new(graph),
485                )
486            }
487        };
488        Ok(result.into_iter().take(k).collect())
489    }
490
491    /// [Self::run_search] for the ACORN traversal: same level descent, but
492    /// the bottom level runs [beam_search_acorn].
493    #[allow(clippy::too_many_arguments)]
494    fn run_search_acorn<L, B>(
495        &self,
496        ep: OrderedNode,
497        params: &HnswQueryParams,
498        bitset: &Visited,
499        visited_generator: &mut VisitedGenerator,
500        expanded_generator: &mut VisitedGenerator,
501        storage_len: usize,
502        prefetch_distance: Option<usize>,
503        dist_calc: &impl DistCalculator,
504        make_level: impl Fn(u16) -> L,
505        bottom: B,
506    ) -> Vec<OrderedNode>
507    where
508        L: BorrowingGraph,
509        B: BorrowingGraph,
510    {
511        let mut ep = ep;
512        for level in (0..self.max_level()).rev() {
513            let cur_level = make_level(level);
514            ep = greedy_search_borrowed(
515                &cur_level,
516                ep,
517                dist_calc,
518                self.inner.params.prefetch_distance,
519            );
520        }
521        let mut visited = visited_generator.generate(storage_len);
522        let mut expanded = expanded_generator.generate(storage_len);
523        beam_search_acorn(
524            &bottom,
525            &ep,
526            params,
527            dist_calc,
528            bitset,
529            prefetch_distance,
530            &mut visited,
531            &mut expanded,
532        )
533    }
534
535    #[instrument(level = "debug", skip(self, storage, query, prefilter_bitset))]
536    fn flat_search(
537        &self,
538        storage: &impl VectorStore,
539        query: ArrayRef,
540        k: usize,
541        prefilter_bitset: Visited,
542        params: &HnswQueryParams,
543    ) -> Vec<OrderedNode> {
544        let lower_bound: OrderedFloat = params.lower_bound.unwrap_or(f32::MIN).into();
545        let upper_bound: OrderedFloat = params.upper_bound.unwrap_or(f32::MAX).into();
546
547        let dist_calc = storage.dist_calculator(query, params.dist_q_c);
548        let mut heap = BinaryHeap::<OrderedNode>::with_capacity(k);
549
550        match self.inner.params.prefetch_distance {
551            Some(ahead) if ahead > 0 => {
552                let mut ids_iter = prefilter_bitset.iter_ones().map(|i| i as u32);
553                let mut buffer = VecDeque::with_capacity(ahead + 1);
554                for _ in 0..=ahead {
555                    if let Some(id) = ids_iter.next() {
556                        buffer.push_back(id);
557                    } else {
558                        break;
559                    }
560                }
561
562                while let Some(node_id) = buffer.pop_front() {
563                    if let Some(&prefetch_id) = buffer.get(ahead - 1) {
564                        dist_calc.prefetch(prefetch_id);
565                    }
566                    if let Some(next) = ids_iter.next() {
567                        buffer.push_back(next);
568                    }
569
570                    let dist: OrderedFloat = dist_calc.distance(node_id).into();
571                    if dist <= lower_bound || dist > upper_bound {
572                        continue;
573                    }
574                    if heap.len() < k {
575                        heap.push((dist, node_id).into());
576                    } else if dist < heap.peek().unwrap().dist {
577                        heap.pop();
578                        heap.push((dist, node_id).into());
579                    }
580                }
581            }
582            _ => {
583                for node_id in prefilter_bitset.iter_ones().map(|i| i as u32) {
584                    let dist: OrderedFloat = dist_calc.distance(node_id).into();
585                    if dist <= lower_bound || dist > upper_bound {
586                        continue;
587                    }
588                    if heap.len() < k {
589                        heap.push((dist, node_id).into());
590                    } else if dist < heap.peek().unwrap().dist {
591                        heap.pop();
592                        heap.push((dist, node_id).into());
593                    }
594                }
595            }
596        };
597        heap.into_sorted_vec()
598    }
599
600    /// Returns the metadata of this [`HNSW`].
601    pub fn metadata(&self) -> HnswMetadata {
602        // calculate the offsets of each level,
603        // start from 0
604        let level_offsets = self
605            .inner
606            .level_count
607            .iter()
608            .chain(iter::once(&0))
609            .scan(0, |state, x| {
610                let start = *state;
611                *state += *x;
612                Some(start)
613            })
614            .collect();
615
616        HnswMetadata {
617            entry_point: self.inner.entry_point,
618            params: self.inner.params.clone(),
619            level_offsets,
620        }
621    }
622}
623
624struct HnswBuilder {
625    params: HnswBuildParams,
626
627    nodes: Arc<Vec<RwLock<GraphBuilderNode>>>,
628    level_count: Vec<AtomicUsize>,
629
630    entry_point: u32,
631
632    visited_generator_queue: Arc<ArrayQueue<VisitedGenerator>>,
633}
634
635impl DeepSizeOf for HnswBuilder {
636    fn deep_size_of_children(&self, context: &mut lance_core::deepsize::Context) -> usize {
637        self.params.deep_size_of_children(context)
638            + self.nodes.deep_size_of_children(context)
639            + self.level_count.deep_size_of_children(context)
640        // Skipping the visited_generator_queue
641    }
642}
643
644impl HnswBuilder {
645    fn finish(self) -> HNSW {
646        let nodes = match Arc::try_unwrap(self.nodes) {
647            Ok(nodes) => nodes
648                .into_iter()
649                .map(|node| node.into_inner().expect("builder lock poisoned"))
650                .collect(),
651            Err(nodes) => nodes
652                .iter()
653                .map(|node| node.read().expect("builder lock poisoned").clone())
654                .collect(),
655        };
656
657        let level_count = self
658            .level_count
659            .into_iter()
660            .map(|count| count.load(Ordering::Relaxed))
661            .collect();
662
663        HNSW {
664            inner: Arc::new(HnswCore {
665                params: self.params,
666                graph: HnswGraph::Built(Arc::new(nodes)),
667                level_count,
668                entry_point: self.entry_point,
669                visited_generator_queue: self.visited_generator_queue,
670            }),
671        }
672    }
673
674    /// Create a new [`HNSWBuilder`] with prepared params and in memory vector storage.
675    pub fn with_params(params: HnswBuildParams, storage: &impl VectorStore) -> Self {
676        let len = storage.len();
677        let max_level = params.max_level;
678
679        let level_count = (0..max_level)
680            .map(|_| AtomicUsize::new(0))
681            .collect::<Vec<_>>();
682
683        let visited_generator_queue = Arc::new(ArrayQueue::new(get_num_compute_intensive_cpus()));
684        for _ in 0..get_num_compute_intensive_cpus() {
685            visited_generator_queue
686                .push(VisitedGenerator::new(0))
687                .unwrap();
688        }
689        let mut builder = Self {
690            params,
691            nodes: Arc::new(Vec::new()),
692            level_count,
693            entry_point: 0,
694            visited_generator_queue,
695        };
696
697        if storage.is_empty() {
698            return builder;
699        }
700
701        let mut nodes = Vec::with_capacity(len);
702        {
703            if len > 0 {
704                nodes.push(RwLock::new(GraphBuilderNode::new(0, max_level as usize)));
705            }
706            let mut level_rng = SmallRng::seed_from_u64(HNSW_LEVEL_RNG_SEED);
707            for i in 1..len {
708                nodes.push(RwLock::new(GraphBuilderNode::new(
709                    i as u32,
710                    builder.random_level(&mut level_rng) as usize + 1,
711                )));
712            }
713        }
714        builder.nodes = Arc::new(nodes);
715
716        builder
717    }
718
719    /// New node's level
720    ///
721    /// See paper `Algorithm 1`
722    fn random_level<R: Rng + ?Sized>(&self, rng: &mut R) -> u16 {
723        let ml = 1.0 / (self.params.m as f32).ln();
724        min(
725            (-rng.random::<f32>().ln() * ml) as u16,
726            self.params.max_level - 1,
727        )
728    }
729
730    /// Insert one node.
731    fn insert(
732        &self,
733        node: u32,
734        visited_generator: &mut VisitedGenerator,
735        storage: &impl VectorStore,
736    ) {
737        let nodes = &self.nodes;
738        let target_level = nodes[node as usize].read().unwrap().level_neighbors.len() as u16 - 1;
739        let dist_calc = storage.dist_calculator_from_id(node);
740        let mut ep = OrderedNode::new(
741            self.entry_point,
742            dist_calc.distance(self.entry_point).into(),
743        );
744
745        //
746        // Search for entry point in paper.
747        // ```
748        //   for l_c in (L..l+1) {
749        //     W = Search-Layer(q, ep, ef=1, l_c)
750        //    ep = Select-Neighbors(W, 1)
751        //  }
752        // ```
753        for level in (target_level + 1..self.params.max_level).rev() {
754            let cur_level = HnswLevelView::new(level, nodes);
755            ep = greedy_search(&cur_level, ep, &dist_calc, self.params.prefetch_distance);
756        }
757
758        let mut pruned_neighbors_per_level: Vec<Vec<_>> =
759            vec![Vec::new(); (target_level + 1) as usize];
760        {
761            let mut current_node = nodes[node as usize].write().unwrap();
762            for level in (0..=target_level).rev() {
763                self.level_count[level as usize].fetch_add(1, Ordering::Relaxed);
764
765                let neighbors = self.search_level(&ep, level, &dist_calc, nodes, visited_generator);
766                for neighbor in &neighbors {
767                    current_node.add_neighbor(neighbor.id, neighbor.dist, level);
768                }
769                self.prune(storage, &mut current_node, level);
770                pruned_neighbors_per_level[level as usize]
771                    .clone_from(&current_node.level_neighbors_ranked[level as usize]);
772
773                ep = neighbors[0].clone();
774            }
775        }
776        for (level, pruned_neighbors) in pruned_neighbors_per_level.iter().enumerate() {
777            for unpruned_edge in pruned_neighbors {
778                let level = level as u16;
779                let m_max = match level {
780                    0 => self.params.m * 2,
781                    _ => self.params.m,
782                };
783                if unpruned_edge.dist
784                    < nodes[unpruned_edge.id as usize]
785                        .read()
786                        .unwrap()
787                        .cutoff(level, m_max)
788                {
789                    let mut chosen_node = nodes[unpruned_edge.id as usize].write().unwrap();
790                    chosen_node.add_neighbor(node, unpruned_edge.dist, level);
791                    self.prune(storage, &mut chosen_node, level);
792                }
793            }
794        }
795    }
796
797    fn search_level(
798        &self,
799        ep: &OrderedNode,
800        level: u16,
801        dist_calc: &impl DistCalculator,
802        nodes: &[RwLock<GraphBuilderNode>],
803        visited_generator: &mut VisitedGenerator,
804    ) -> Vec<OrderedNode> {
805        let cur_level = HnswLevelView::new(level, nodes);
806        let mut visited = visited_generator.generate(nodes.len());
807        beam_search(
808            &cur_level,
809            ep,
810            &HnswQueryParams {
811                ef: self.params.ef_construction,
812                lower_bound: None,
813                upper_bound: None,
814                dist_q_c: 0.0,
815                use_acorn: false,
816            },
817            dist_calc,
818            None,
819            self.params.prefetch_distance,
820            &mut visited,
821        )
822    }
823
824    fn prune(&self, storage: &impl VectorStore, builder_node: &mut GraphBuilderNode, level: u16) {
825        let m_max = match level {
826            0 => self.params.m * 2,
827            _ => self.params.m,
828        };
829
830        let neighbors_ranked = &mut builder_node.level_neighbors_ranked[level as usize];
831        if neighbors_ranked.len() <= m_max {
832            builder_node.update_from_ranked_neighbors(level);
833            return;
834        }
835
836        let level_neighbors = std::mem::take(neighbors_ranked);
837        *neighbors_ranked = select_neighbors_heuristic_owned(storage, level_neighbors, m_max);
838        builder_node.update_from_ranked_neighbors(level);
839    }
840}
841
842// View of a level in HNSW graph.
843// This is used to iterate over neighbors in a specific level.
844pub(crate) struct HnswLevelView<'a> {
845    level: u16,
846    nodes: &'a [RwLock<GraphBuilderNode>],
847}
848
849impl<'a> HnswLevelView<'a> {
850    pub fn new(level: u16, nodes: &'a [RwLock<GraphBuilderNode>]) -> Self {
851        Self { level, nodes }
852    }
853}
854
855impl Graph for HnswLevelView<'_> {
856    fn len(&self) -> usize {
857        self.nodes.len()
858    }
859
860    fn neighbors(&self, key: u32) -> Arc<Vec<u32>> {
861        let node = &self.nodes[key as usize];
862        node.read().unwrap().level_neighbors[self.level as usize].clone()
863    }
864}
865
866pub(crate) struct ImmutableHnswLevelView<'a> {
867    level: u16,
868    nodes: &'a [GraphBuilderNode],
869}
870
871impl<'a> ImmutableHnswLevelView<'a> {
872    pub fn new(level: u16, nodes: &'a [GraphBuilderNode]) -> Self {
873        Self { level, nodes }
874    }
875}
876
877impl Graph for ImmutableHnswLevelView<'_> {
878    fn len(&self) -> usize {
879        self.nodes.len()
880    }
881
882    fn neighbors(&self, key: u32) -> Arc<Vec<u32>> {
883        self.nodes[key as usize].level_neighbors[self.level as usize].clone()
884    }
885}
886
887impl BorrowingGraph for ImmutableHnswLevelView<'_> {
888    fn len(&self) -> usize {
889        self.nodes.len()
890    }
891
892    fn neighbors(&self, key: u32) -> &[u32] {
893        self.nodes[key as usize].level_neighbors[self.level as usize].as_slice()
894    }
895}
896
897pub(crate) struct ImmutableHnswBottomView<'a> {
898    nodes: &'a [GraphBuilderNode],
899}
900
901impl<'a> ImmutableHnswBottomView<'a> {
902    pub fn new(nodes: &'a [GraphBuilderNode]) -> Self {
903        Self { nodes }
904    }
905}
906
907impl Graph for ImmutableHnswBottomView<'_> {
908    fn len(&self) -> usize {
909        self.nodes.len()
910    }
911
912    fn neighbors(&self, key: u32) -> Arc<Vec<u32>> {
913        self.nodes[key as usize].bottom_neighbors.clone()
914    }
915}
916
917impl BorrowingGraph for ImmutableHnswBottomView<'_> {
918    fn len(&self) -> usize {
919        self.nodes.len()
920    }
921
922    fn neighbors(&self, key: u32) -> &[u32] {
923        self.nodes[key as usize].bottom_neighbors.as_slice()
924    }
925}
926
927/// Per-level node-id -> row-index lookup for a disk-loaded HNSW graph.
928enum LevelLookup {
929    /// `row == node id`. Used only for level 0, where [`HNSW::to_batch`]
930    /// writes every node once in ascending `__vector_id` (== node id) order,
931    /// so the level-0 slice is exactly `[0, N)` with `row == id`.
932    Dense,
933    /// Upper level: an explicit `node_id -> row` map built from the level's
934    /// `__vector_id` column.
935    ///
936    /// We do *not* assume the column is sorted or that the slice is aligned
937    /// to a true level boundary: `level_offsets`/`level_count` omit the
938    /// entry-point node (it is written at every level by `to_batch` but only
939    /// counted at level 0), so upper-level slices can be off-by-one and
940    /// non-monotonic. Keying by the `__vector_id` value -- exactly what the
941    /// old per-node `load` did -- preserves behavior bit-for-bit. Upper
942    /// levels shrink geometrically, so this map stays tiny.
943    Sparse(HashMap<u32, u32>),
944}
945
946/// A search-only HNSW graph backed directly by the Arrow buffers of the
947/// on-disk `RecordBatch`.
948///
949/// Loading performs no per-node reconstruction: neighbor adjacency is served
950/// as `&[u32]` slices straight out of the `__neighbors` `ListArray` value
951/// buffer (zero copy). The full `batch` is retained so [`HNSW::to_batch`] is a
952/// near-free passthrough -- required, because the IVF partition cache
953/// re-serializes loaded indices through `to_batch()`
954/// (`lance/src/index/vector/ivf/partition_serde.rs`) -- and so a future
955/// zero-copy `CacheCodec` (#6745) can write/read it through
956/// `lance_arrow::ipc` without rebuilding the graph.
957struct LoadedHnswGraph {
958    /// The full loaded batch (all levels concatenated, level 0 first),
959    /// retained verbatim for `to_batch()` and #6745.
960    batch: RecordBatch,
961    /// Per-level `__neighbors` `List<UInt32>`, zero-copy slices of `batch`.
962    level_neighbors: Vec<ListArray>,
963    /// Per-level node-id -> row lookup (see [`LevelLookup`]).
964    level_lookup: Vec<LevelLookup>,
965    /// Number of nodes present at each level (`level_count[0]` == total).
966    level_count: Vec<usize>,
967}
968
969impl DeepSizeOf for LoadedHnswGraph {
970    fn deep_size_of_children(&self, _context: &mut lance_core::deepsize::Context) -> usize {
971        // `level_neighbors` are zero-copy views into `batch`, so counting
972        // `batch` alone avoids double counting (mirrors
973        // `vector/flat/storage.rs`). The upper-level `level_lookup` maps are
974        // sized to the geometrically-shrinking node counts above level 0 --
975        // negligible next to the batch and not separately accounted here.
976        self.batch.get_array_memory_size()
977    }
978}
979
980impl LoadedHnswGraph {
981    /// Borrow the neighbor ids of `key` at `level` directly from the Arrow
982    /// `ListArray` value buffer -- no allocation, no copy.
983    #[inline]
984    fn neighbors_at(&self, level: usize, key: u32) -> &[u32] {
985        let row = match &self.level_lookup[level] {
986            LevelLookup::Dense => key as usize,
987            LevelLookup::Sparse(id_to_row) => match id_to_row.get(&key) {
988                Some(&row) => row as usize,
989                // The node is absent at this level -- e.g. an empty upper
990                // level the search descends through, or a node that only
991                // exists at lower levels. Mirror the old representation
992                // (`level_neighbors[level]` defaulted to empty): no
993                // neighbors here, so greedy search stays put and descends.
994                None => return &[],
995            },
996        };
997        let list = &self.level_neighbors[level];
998        let offsets = list.value_offsets();
999        let start = offsets[row] as usize;
1000        let end = offsets[row + 1] as usize;
1001        // The `__neighbors` list child is `UInt32` per `HNSW::schema()`.
1002        // Validity bitmap is ignored on purpose: `to_batch` never writes null
1003        // neighbor lists, matching the previous `.unwrap()`-based load.
1004        let values = list.values().as_primitive::<UInt32Type>();
1005        &values.values()[start..end]
1006    }
1007}
1008
1009/// Per-level search view over a disk-loaded [`LoadedHnswGraph`].
1010pub(crate) struct LoadedHnswLevelView<'a> {
1011    level: usize,
1012    graph: &'a LoadedHnswGraph,
1013}
1014
1015impl<'a> LoadedHnswLevelView<'a> {
1016    fn new(level: u16, graph: &'a LoadedHnswGraph) -> Self {
1017        Self {
1018            level: level as usize,
1019            graph,
1020        }
1021    }
1022}
1023
1024impl Graph for LoadedHnswLevelView<'_> {
1025    fn len(&self) -> usize {
1026        // Mirrors `ImmutableHnswLevelView::len` (total node count).
1027        self.graph.level_count[0]
1028    }
1029
1030    fn neighbors(&self, key: u32) -> Arc<Vec<u32>> {
1031        // Non-hot fallback: HNSW search goes through `BorrowingGraph`. Kept
1032        // only so the `Graph` trait / legacy `greedy_search` need no
1033        // special-casing for loaded graphs.
1034        Arc::new(self.graph.neighbors_at(self.level, key).to_vec())
1035    }
1036}
1037
1038impl BorrowingGraph for LoadedHnswLevelView<'_> {
1039    fn len(&self) -> usize {
1040        self.graph.level_count[0]
1041    }
1042
1043    fn neighbors(&self, key: u32) -> &[u32] {
1044        self.graph.neighbors_at(self.level, key)
1045    }
1046}
1047
1048/// Bottom-level (level 0) search view over a disk-loaded [`LoadedHnswGraph`].
1049pub(crate) struct LoadedHnswBottomView<'a> {
1050    graph: &'a LoadedHnswGraph,
1051}
1052
1053impl<'a> LoadedHnswBottomView<'a> {
1054    fn new(graph: &'a LoadedHnswGraph) -> Self {
1055        Self { graph }
1056    }
1057}
1058
1059impl Graph for LoadedHnswBottomView<'_> {
1060    fn len(&self) -> usize {
1061        self.graph.level_count[0]
1062    }
1063
1064    fn neighbors(&self, key: u32) -> Arc<Vec<u32>> {
1065        Arc::new(self.graph.neighbors_at(0, key).to_vec())
1066    }
1067}
1068
1069impl BorrowingGraph for LoadedHnswBottomView<'_> {
1070    fn len(&self) -> usize {
1071        self.graph.level_count[0]
1072    }
1073
1074    fn neighbors(&self, key: u32) -> &[u32] {
1075        self.graph.neighbors_at(0, key)
1076    }
1077}
1078
1079/// The graph backing an [`HNSW`]: either built in memory or disk-loaded.
1080enum HnswGraph {
1081    /// Built in memory by the (online) builder / `index_vectors` /
1082    /// `from_parts`. Mutable-shaped `GraphBuilderNode`s; `to_batch()`
1083    /// re-encodes from these (it needs the per-node ranked distances).
1084    Built(Arc<Vec<GraphBuilderNode>>),
1085    /// Loaded from disk, Arrow-backed, search-only.
1086    Loaded(Arc<LoadedHnswGraph>),
1087}
1088
1089impl DeepSizeOf for HnswGraph {
1090    fn deep_size_of_children(&self, context: &mut lance_core::deepsize::Context) -> usize {
1091        match self {
1092            Self::Built(nodes) => nodes.deep_size_of_children(context),
1093            Self::Loaded(graph) => graph.deep_size_of_children(context),
1094        }
1095    }
1096}
1097
1098#[derive(Debug, Clone, Copy)]
1099pub struct HnswQueryParams {
1100    pub ef: usize,
1101    pub lower_bound: Option<f32>,
1102    pub upper_bound: Option<f32>,
1103    pub dist_q_c: f32,
1104    pub use_acorn: bool,
1105}
1106
1107impl From<&Query> for HnswQueryParams {
1108    fn from(query: &Query) -> Self {
1109        let k = query.k * query.refine_factor.unwrap_or(1) as usize;
1110        Self {
1111            ef: query.ef.unwrap_or(k + k / 2),
1112            lower_bound: query.lower_bound,
1113            upper_bound: query.upper_bound,
1114            dist_q_c: query.dist_q_c,
1115            use_acorn: query.approx_mode == ApproxMode::Fast,
1116        }
1117    }
1118}
1119
1120impl IvfSubIndex for HNSW {
1121    type BuildParams = HnswBuildParams;
1122    type QueryParams = HnswQueryParams;
1123
1124    fn load(data: RecordBatch) -> Result<Self>
1125    where
1126        Self: Sized,
1127    {
1128        if data.num_rows() == 0 {
1129            return Ok(Self::empty());
1130        }
1131
1132        let hnsw_metadata = data
1133            .schema_ref()
1134            .metadata()
1135            .get(HNSW_METADATA_KEY)
1136            .ok_or(Error::index(format!("{} not found", HNSW_METADATA_KEY)))?;
1137        let hnsw_metadata: HnswMetadata = serde_json::from_str(hnsw_metadata).map_err(|e| {
1138            Error::index(format!(
1139                "Failed to decode HNSW metadata: {}, json: {}",
1140                e, hnsw_metadata
1141            ))
1142        })?;
1143
1144        // Slice the concatenated batch into one (zero-copy) view per level.
1145        let level_batches: Vec<RecordBatch> = hnsw_metadata
1146            .level_offsets
1147            .iter()
1148            .tuple_windows()
1149            .map(|(start, end)| data.slice(*start, end - start))
1150            .collect();
1151
1152        let level_count = level_batches
1153            .iter()
1154            .map(|b| b.num_rows())
1155            .collect::<Vec<_>>();
1156
1157        // No per-node reconstruction: keep the Arrow adjacency buffers as-is
1158        // and only build the tiny per-upper-level id->row lookups. The
1159        // `__distance` column is never materialized here -- search doesn't
1160        // need it, and `to_batch()` returns the retained `data` verbatim.
1161        let mut level_neighbors = Vec::with_capacity(level_batches.len());
1162        let mut level_lookup = Vec::with_capacity(level_batches.len());
1163        for (level, batch) in level_batches.iter().enumerate() {
1164            // `.clone()` on an Arrow array bumps a refcount; buffers stay
1165            // shared with `data` (zero copy).
1166            let neighbors = batch[NEIGHBORS_COL].as_list::<i32>().clone();
1167            let ids = batch[VECTOR_ID_COL].as_primitive::<UInt32Type>();
1168            if level == 0 {
1169                // `to_batch` writes every node at level 0 exactly once in
1170                // ascending `__vector_id` (== node id) order, so the level-0
1171                // slice is exactly `[0, N)` and the row index *is* the node
1172                // id. The `Dense` lookup below depends on this: in a release
1173                // build a violated invariant would silently make search read
1174                // the wrong neighbor list, so enforce it at load time (not via
1175                // `debug_assert!`) and reject a malformed or version-
1176                // incompatible batch.
1177                if let Some((row, id)) = ids
1178                    .values()
1179                    .iter()
1180                    .enumerate()
1181                    .find(|&(row, id)| *id != row as u32)
1182                {
1183                    return Err(Error::index(format!(
1184                        "HNSW level-0 __vector_id must equal the row index, but \
1185                         row {row} has __vector_id {id}; the on-disk batch is \
1186                         malformed or was written by an incompatible version"
1187                    )));
1188                }
1189                level_lookup.push(LevelLookup::Dense);
1190            } else {
1191                // Upper levels: explicit id -> row map. No ordering/alignment
1192                // assumption (see `LevelLookup::Sparse`). On the rare
1193                // duplicate id (a misaligned slice can repeat one across a
1194                // level boundary) the last wins, matching the old load's
1195                // `nodes[id].level_neighbors[level] = ...` last-write.
1196                let id_to_row: HashMap<u32, u32> = ids
1197                    .values()
1198                    .iter()
1199                    .enumerate()
1200                    .map(|(row, id)| (*id, row as u32))
1201                    .collect();
1202                level_lookup.push(LevelLookup::Sparse(id_to_row));
1203            }
1204            level_neighbors.push(neighbors);
1205        }
1206
1207        // `entry_point` is read from untrusted metadata and indexes the `Dense`
1208        // level-0 lookup directly; an out-of-range value would read past the
1209        // level-0 neighbor buffer during search. Validate it under the same
1210        // persisted-format invariant as the level-0 ids above.
1211        let num_nodes = level_count[0];
1212        if hnsw_metadata.entry_point as usize >= num_nodes {
1213            return Err(Error::index(format!(
1214                "HNSW entry_point {} is out of range for a graph with {num_nodes} \
1215                 nodes; the on-disk batch is malformed or was written by an \
1216                 incompatible version",
1217                hnsw_metadata.entry_point
1218            )));
1219        }
1220
1221        let visited_generator_queue =
1222            Arc::new(ArrayQueue::new(get_num_compute_intensive_cpus() * 2));
1223        for _ in 0..get_num_compute_intensive_cpus() * 2 {
1224            visited_generator_queue
1225                .push(VisitedGenerator::new(0))
1226                .unwrap();
1227        }
1228
1229        let graph = LoadedHnswGraph {
1230            batch: data,
1231            level_neighbors,
1232            level_lookup,
1233            level_count: level_count.clone(),
1234        };
1235        let inner = HnswCore {
1236            params: hnsw_metadata.params,
1237            graph: HnswGraph::Loaded(Arc::new(graph)),
1238            level_count,
1239            entry_point: hnsw_metadata.entry_point,
1240            visited_generator_queue,
1241        };
1242
1243        Ok(Self {
1244            inner: Arc::new(inner),
1245        })
1246    }
1247
1248    fn name() -> &'static str {
1249        HNSW_TYPE
1250    }
1251
1252    fn metadata_key() -> &'static str {
1253        "lance:hnsw"
1254    }
1255
1256    /// Return the schema of the sub index
1257    fn schema() -> arrow_schema::SchemaRef {
1258        arrow_schema::Schema::new(vec![
1259            VECTOR_ID_FIELD.clone(),
1260            NEIGHBORS_FIELD.clone(),
1261            DISTS_FIELD.clone(),
1262        ])
1263        .into()
1264    }
1265
1266    #[instrument(level = "debug", skip(self, query, storage, prefilter, _metrics))]
1267    fn search(
1268        &self,
1269        query: ArrayRef,
1270        k: usize,
1271        params: Self::QueryParams,
1272        storage: &impl VectorStore,
1273        prefilter: Arc<dyn PreFilter>,
1274        _metrics: &dyn MetricsCollector,
1275    ) -> Result<RecordBatch> {
1276        if params.ef < k {
1277            return Err(Error::index(
1278                "ef must be greater than or equal to k".to_string(),
1279            ));
1280        }
1281
1282        let schema = VECTOR_RESULT_SCHEMA.clone();
1283        if self.is_empty() {
1284            return Ok(RecordBatch::new_empty(schema));
1285        }
1286
1287        let mut prefilter_generator = self
1288            .inner
1289            .visited_generator_queue
1290            .pop()
1291            .unwrap_or_else(|| VisitedGenerator::new(storage.len()));
1292        let results = if prefilter.is_empty() {
1293            self.search_basic(query, k, &params, None, storage)?
1294        } else {
1295            // the bitset must be moved into a callee on every path so its
1296            // borrow of `prefilter_generator` ends before the push below
1297            let indices = prefilter.filter_row_ids(Box::new(storage.row_ids()));
1298            let mut prefilter_bitset = prefilter_generator.generate(storage.len());
1299            for index in indices {
1300                prefilter_bitset.insert(index as u32);
1301            }
1302            let remained = prefilter_bitset.count_ones();
1303            if remained == storage.len() {
1304                // mask passes every row: same as unfiltered
1305                drop(prefilter_bitset);
1306                self.search_basic(query, k, &params, None, storage)?
1307            } else if remained < self.len() * 10 / 100 {
1308                // few matching rows: brute force is cheaper and exact
1309                self.flat_search(storage, query, k, prefilter_bitset, &params)
1310            } else if params.use_acorn {
1311                let acorn_results =
1312                    self.search_acorn(query.clone(), k, &params, &prefilter_bitset, storage)?;
1313                // under-delivery means the budget ran out on a fragmented
1314                // mask, except range-bounded queries which return short
1315                // legitimately
1316                let bounded = params.lower_bound.is_some() || params.upper_bound.is_some();
1317                if !bounded && acorn_results.len() < k.min(remained) {
1318                    self.search_basic(query, k, &params, Some(prefilter_bitset), storage)?
1319                } else {
1320                    drop(prefilter_bitset);
1321                    acorn_results
1322                }
1323            } else {
1324                self.search_basic(query, k, &params, Some(prefilter_bitset), storage)?
1325            }
1326        };
1327        // if the queue is full, we just don't push it back, so ignore the error here
1328        let _ = self.inner.visited_generator_queue.push(prefilter_generator);
1329
1330        // need to unique by row ids in case of searching multivector
1331        let (row_ids, dists): (Vec<_>, Vec<_>) = results
1332            .into_iter()
1333            .map(|r| (storage.row_id(r.id), r.dist.0))
1334            .unique_by(|r| r.0)
1335            .unzip();
1336        let row_ids = Arc::new(UInt64Array::from(row_ids));
1337        let distances = Arc::new(Float32Array::from(dists));
1338
1339        Ok(RecordBatch::try_new(schema, vec![distances, row_ids])?)
1340    }
1341
1342    /// Given a vector storage, containing all the data for the IVF partition, build the sub index.
1343    fn index_vectors(storage: &impl VectorStore, params: Self::BuildParams) -> Result<Self>
1344    where
1345        Self: Sized,
1346    {
1347        let builder = HnswBuilder::with_params(params, storage);
1348
1349        log::debug!(
1350            "Building HNSW graph: num={}, max_levels={}, m={}, ef_construction={}, distance_type:{}",
1351            storage.len(),
1352            builder.params.max_level,
1353            builder.params.m,
1354            builder.params.ef_construction,
1355            storage.distance_type(),
1356        );
1357
1358        if storage.is_empty() {
1359            return Ok(builder.finish());
1360        }
1361
1362        let len = storage.len();
1363        builder.level_count[0].fetch_add(1, Ordering::Relaxed);
1364        (1..len).into_par_iter().for_each_init(
1365            || VisitedGenerator::new(len),
1366            |visited_generator, node| {
1367                builder.insert(node as u32, visited_generator, storage);
1368            },
1369        );
1370
1371        assert_eq!(builder.level_count[0].load(Ordering::Relaxed), len);
1372        Ok(builder.finish())
1373    }
1374
1375    fn remap(
1376        &self,
1377        _mapping: &RowAddrRemap, // we don't need the mapping here because we rebuild the graph from remapped storage
1378        store: &impl VectorStore,
1379    ) -> Result<Self> {
1380        // We can't simply remap the row ids in the graph because the vectors are changed,
1381        // so the graph needs to be rebuilt.
1382        Self::index_vectors(store, self.inner.params.clone())
1383    }
1384
1385    /// Encode the sub index into a record batch
1386    fn to_batch(&self) -> Result<RecordBatch> {
1387        let nodes = match &self.inner.graph {
1388            HnswGraph::Built(nodes) => nodes,
1389            HnswGraph::Loaded(graph) => {
1390                // A loaded graph is already Arrow-backed: return the retained
1391                // batch verbatim, re-stamped with up-to-date HNSW metadata.
1392                // The IVF partition cache re-serializes loaded indices through
1393                // here (`ivf/partition_serde.rs`), so this must round-trip.
1394                //
1395                // Merge into (not replace) the existing schema metadata: a
1396                // disk-loaded batch inherits other keys from the index file
1397                // schema (e.g. `INDEX_METADATA_SCHEMA_KEY`, `IVF_METADATA_KEY`),
1398                // and `RecordBatch::with_schema` requires the new metadata to be
1399                // a superset of the current one. Dropping those keys here makes
1400                // the new schema a non-superset and fails the round-trip with
1401                // "target schema is not superset of current schema".
1402                let metadata = serde_json::to_string(&self.metadata())?;
1403                let mut schema_metadata = graph.batch.schema_ref().metadata().clone();
1404                schema_metadata.insert(HNSW_METADATA_KEY.to_string(), metadata);
1405                let schema = graph
1406                    .batch
1407                    .schema()
1408                    .as_ref()
1409                    .clone()
1410                    .with_metadata(schema_metadata);
1411                return Ok(graph.batch.clone().with_schema(Arc::new(schema))?);
1412            }
1413        };
1414
1415        let mut vector_id_builder = UInt32Builder::with_capacity(self.len());
1416        let mut neighbors_builder = ListBuilder::with_capacity(UInt32Builder::new(), self.len());
1417        let mut distances_builder =
1418            ListBuilder::with_capacity(arrow_array::builder::Float32Builder::new(), self.len());
1419        let mut batches = Vec::with_capacity(self.max_level() as usize);
1420        for level in 0..self.max_level() {
1421            let level = level as usize;
1422            for (id, node) in nodes.iter().enumerate() {
1423                if level >= node.level_neighbors.len() {
1424                    continue;
1425                }
1426                let neighbors = node.level_neighbors[level].iter().map(|n| Some(*n));
1427                let distances = node.level_neighbors_ranked[level]
1428                    .iter()
1429                    .map(|n| Some(n.dist.0));
1430                vector_id_builder.append_value(id as u32);
1431                neighbors_builder.append_value(neighbors);
1432                distances_builder.append_value(distances);
1433            }
1434
1435            let batch = RecordBatch::try_new(
1436                Self::schema(),
1437                vec![
1438                    Arc::new(vector_id_builder.finish()),
1439                    Arc::new(neighbors_builder.finish()),
1440                    Arc::new(distances_builder.finish()),
1441                ],
1442            )?;
1443            batches.push(batch);
1444        }
1445
1446        let metadata = self.metadata();
1447        let metadata = serde_json::to_string(&metadata)?;
1448        let schema = Self::schema()
1449            .as_ref()
1450            .clone()
1451            .with_metadata(HashMap::from_iter(vec![(
1452                HNSW_METADATA_KEY.to_string(),
1453                metadata,
1454            )]));
1455        let batch = concat_batches(&Self::schema(), batches.iter())?;
1456        let batch = batch.with_schema(Arc::new(schema))?;
1457        Ok(batch)
1458    }
1459}
1460
1461#[cfg(test)]
1462mod tests {
1463    use std::sync::Arc;
1464
1465    use arrow_array::{ArrayRef, FixedSizeListArray, RecordBatch, UInt8Array, UInt32Array};
1466    use arrow_schema::Schema;
1467    use lance_arrow::FixedSizeListArrayExt;
1468    use lance_core::deepsize::DeepSizeOf;
1469    use lance_file::versions::v1::{
1470        reader::FileReader as V1FileReader,
1471        writer::{FileWriter as V1FileWriter, FileWriterOptions as V1FileWriterOptions},
1472    };
1473    use lance_io::object_store::ObjectStore;
1474    use lance_linalg::distance::DistanceType;
1475    use lance_table::format::SelfDescribingFileReader;
1476    use lance_table::io::manifest::ManifestDescribing;
1477    use lance_testing::datagen::generate_random_array;
1478    use object_store::path::Path;
1479    use rstest::rstest;
1480
1481    use super::HnswGraph;
1482    use crate::vector::storage::{DistCalculator, VectorStore};
1483    use crate::vector::v3::subindex::IvfSubIndex;
1484    use crate::vector::{
1485        flat::storage::{FlatBinStorage, FlatFloatStorage},
1486        graph::{DISTS_FIELD, NEIGHBORS_FIELD, VisitedGenerator},
1487        hnsw::{
1488            HNSW, VECTOR_ID_FIELD,
1489            builder::{HnswBuildParams, HnswQueryParams},
1490        },
1491    };
1492
1493    #[tokio::test]
1494    async fn test_builder_write_load() {
1495        const DIM: usize = 32;
1496        const TOTAL: usize = 2048;
1497        const NUM_EDGES: usize = 20;
1498        let data = generate_random_array(TOTAL * DIM);
1499        let fsl = FixedSizeListArray::try_new_from_values(data, DIM as i32).unwrap();
1500        let store = Arc::new(FlatFloatStorage::new(fsl.clone(), DistanceType::L2));
1501        let builder = HNSW::index_vectors(
1502            store.as_ref(),
1503            HnswBuildParams::default()
1504                .num_edges(NUM_EDGES)
1505                .ef_construction(50),
1506        )
1507        .unwrap();
1508
1509        let object_store = ObjectStore::memory();
1510        let path = Path::from("test_builder_write_load");
1511        let writer = object_store.create(&path).await.unwrap();
1512        let schema = Schema::new(vec![
1513            VECTOR_ID_FIELD.clone(),
1514            NEIGHBORS_FIELD.clone(),
1515            DISTS_FIELD.clone(),
1516        ]);
1517        let schema = lance_core::datatypes::Schema::try_from(&schema).unwrap();
1518        let mut writer = V1FileWriter::<ManifestDescribing>::with_object_writer(
1519            writer,
1520            schema,
1521            &V1FileWriterOptions::default(),
1522        )
1523        .unwrap();
1524        let batch = builder.to_batch().unwrap();
1525        let metadata = batch.schema_ref().metadata().clone();
1526        writer.write(&[batch]).await.unwrap();
1527        writer.finish_with_metadata(&metadata).await.unwrap();
1528
1529        let reader = V1FileReader::try_new_self_described(&object_store, &path, None)
1530            .await
1531            .unwrap();
1532        let batch = reader
1533            .read_range(0..reader.len(), reader.schema())
1534            .await
1535            .unwrap();
1536        let loaded_hnsw = HNSW::load(batch).unwrap();
1537
1538        let query = fsl.value(0);
1539        let k = 10;
1540        let params = HnswQueryParams {
1541            ef: 50,
1542            lower_bound: None,
1543            upper_bound: None,
1544            dist_q_c: 0.0,
1545            use_acorn: false,
1546        };
1547        let builder_results = builder
1548            .search_basic(query.clone(), k, &params, None, store.as_ref())
1549            .unwrap();
1550        let loaded_results = loaded_hnsw
1551            .search_basic(query, k, &params, None, store.as_ref())
1552            .unwrap();
1553        assert_eq!(builder_results, loaded_results);
1554    }
1555
1556    #[tokio::test]
1557    async fn test_builder_write_load_binary_hamming() {
1558        const DIM: usize = 8;
1559        const TOTAL: usize = 256;
1560        const NUM_EDGES: usize = 20;
1561        let data = UInt8Array::from_iter_values((0..TOTAL * DIM).map(|v| (v % 16) as u8));
1562        let fsl = FixedSizeListArray::try_new_from_values(data, DIM as i32).unwrap();
1563        let store = Arc::new(FlatBinStorage::new(fsl.clone(), DistanceType::Hamming));
1564        let builder = HnswBuildParams::default()
1565            .num_edges(NUM_EDGES)
1566            .ef_construction(50)
1567            .build(Arc::new(fsl.clone()), DistanceType::Hamming)
1568            .await
1569            .unwrap();
1570
1571        let object_store = ObjectStore::memory();
1572        let path = Path::from("test_builder_write_load_binary_hamming");
1573        let writer = object_store.create(&path).await.unwrap();
1574        let schema = Schema::new(vec![
1575            VECTOR_ID_FIELD.clone(),
1576            NEIGHBORS_FIELD.clone(),
1577            DISTS_FIELD.clone(),
1578        ]);
1579        let schema = lance_core::datatypes::Schema::try_from(&schema).unwrap();
1580        let mut writer = V1FileWriter::<ManifestDescribing>::with_object_writer(
1581            writer,
1582            schema,
1583            &V1FileWriterOptions::default(),
1584        )
1585        .unwrap();
1586        let batch = builder.to_batch().unwrap();
1587        let metadata = batch.schema_ref().metadata().clone();
1588        writer.write(&[batch]).await.unwrap();
1589        writer.finish_with_metadata(&metadata).await.unwrap();
1590
1591        let reader = V1FileReader::try_new_self_described(&object_store, &path, None)
1592            .await
1593            .unwrap();
1594        let batch = reader
1595            .read_range(0..reader.len(), reader.schema())
1596            .await
1597            .unwrap();
1598        let loaded_hnsw = HNSW::load(batch).unwrap();
1599
1600        let query = fsl.value(0);
1601        let k = 10;
1602        let params = HnswQueryParams {
1603            ef: 50,
1604            lower_bound: None,
1605            upper_bound: None,
1606            dist_q_c: 0.0,
1607            use_acorn: false,
1608        };
1609        let builder_results = builder
1610            .search_basic(query.clone(), k, &params, None, store.as_ref())
1611            .unwrap();
1612        let loaded_results = loaded_hnsw
1613            .search_basic(query, k, &params, None, store.as_ref())
1614            .unwrap();
1615        assert_eq!(builder_results, loaded_results);
1616    }
1617
1618    /// Brute-force top-`k` node ids by distance -- recall ground truth.
1619    fn brute_force_topk(store: &FlatFloatStorage, query: ArrayRef, k: usize) -> Vec<u32> {
1620        let dist_calc = store.dist_calculator(query, 0.0);
1621        let mut all: Vec<(f32, u32)> = (0..store.len() as u32)
1622            .map(|id| (dist_calc.distance(id), id))
1623            .collect();
1624        all.sort_by(|a, b| a.0.partial_cmp(&b.0).unwrap());
1625        all.into_iter().take(k).map(|(_, id)| id).collect()
1626    }
1627
1628    /// The Arrow-backed loaded graph must search bit-identically to the
1629    /// in-memory build, across distance types and graph sizes (single node,
1630    /// pair, and a multi-level graph exercising the sparse upper-level
1631    /// id->row lookup).
1632    #[rstest]
1633    #[case::l2_single(DistanceType::L2, 1)]
1634    #[case::l2_pair(DistanceType::L2, 2)]
1635    #[case::l2_multi_level(DistanceType::L2, 2048)]
1636    #[case::dot_multi_level(DistanceType::Dot, 2048)]
1637    #[tokio::test]
1638    async fn test_loaded_search_parity_and_recall(
1639        #[case] distance_type: DistanceType,
1640        #[case] total: usize,
1641    ) {
1642        const DIM: usize = 32;
1643        let fsl =
1644            FixedSizeListArray::try_new_from_values(generate_random_array(total * DIM), DIM as i32)
1645                .unwrap();
1646        let store = Arc::new(FlatFloatStorage::new(fsl.clone(), distance_type));
1647        let builder = HNSW::index_vectors(
1648            store.as_ref(),
1649            HnswBuildParams::default().num_edges(20).ef_construction(50),
1650        )
1651        .unwrap();
1652        assert!(!matches!(builder.inner.graph, HnswGraph::Loaded(_)));
1653
1654        let loaded = HNSW::load(builder.to_batch().unwrap()).unwrap();
1655        assert!(matches!(loaded.inner.graph, HnswGraph::Loaded(_)));
1656        assert_eq!(loaded.len(), total);
1657
1658        let k = total.min(10);
1659        let params = HnswQueryParams {
1660            ef: 50,
1661            lower_bound: None,
1662            upper_bound: None,
1663            dist_q_c: 0.0,
1664            use_acorn: false,
1665        };
1666        let query = fsl.value(0);
1667
1668        let builder_results = builder
1669            .search_basic(query.clone(), k, &params, None, store.as_ref())
1670            .unwrap();
1671        let loaded_results = loaded
1672            .search_basic(query.clone(), k, &params, None, store.as_ref())
1673            .unwrap();
1674        assert_eq!(builder_results, loaded_results);
1675
1676        // Recall vs brute-force ground truth (project rule: >= 0.5).
1677        let truth: std::collections::HashSet<u32> = brute_force_topk(store.as_ref(), query, k)
1678            .into_iter()
1679            .collect();
1680        let hits = loaded_results
1681            .iter()
1682            .filter(|n| truth.contains(&n.id))
1683            .count();
1684        let recall = hits as f32 / k as f32;
1685        assert!(recall >= 0.5, "recall {recall} below 0.5 (k={k})");
1686    }
1687
1688    /// Brute-force top-`k` restricted to mask-passing ids.
1689    fn brute_force_topk_masked(
1690        store: &FlatFloatStorage,
1691        query: ArrayRef,
1692        k: usize,
1693        passes: impl Fn(u32) -> bool,
1694    ) -> Vec<u32> {
1695        let dist_calc = store.dist_calculator(query, 0.0);
1696        let mut matching: Vec<(f32, u32)> = (0..store.len() as u32)
1697            .filter(|id| passes(*id))
1698            .map(|id| (dist_calc.distance(id), id))
1699            .collect();
1700        matching.sort_by(|a, b| a.0.partial_cmp(&b.0).unwrap());
1701        matching.into_iter().take(k).map(|(_, id)| id).collect()
1702    }
1703
1704    /// ACORN returns only mask-passing nodes, searches built and loaded
1705    /// graphs identically, and holds recall vs brute force over the mask.
1706    #[tokio::test]
1707    async fn test_acorn_filtered_search() {
1708        const DIM: usize = 32;
1709        const TOTAL: usize = 2048;
1710        let fsl =
1711            FixedSizeListArray::try_new_from_values(generate_random_array(TOTAL * DIM), DIM as i32)
1712                .unwrap();
1713        let store = Arc::new(FlatFloatStorage::new(fsl.clone(), DistanceType::L2));
1714        let builder = HNSW::index_vectors(
1715            store.as_ref(),
1716            HnswBuildParams::default().num_edges(20).ef_construction(50),
1717        )
1718        .unwrap();
1719        let loaded = HNSW::load(builder.to_batch().unwrap()).unwrap();
1720
1721        let mut mask_generator = VisitedGenerator::new(TOTAL);
1722        let k = 10;
1723        let params = HnswQueryParams {
1724            ef: 50,
1725            lower_bound: None,
1726            upper_bound: None,
1727            dist_q_c: 0.0,
1728            use_acorn: false,
1729        };
1730        let query = fsl.value(0);
1731        let truth: std::collections::HashSet<u32> =
1732            brute_force_topk_masked(store.as_ref(), query.clone(), k, |id| id % 2 == 0)
1733                .into_iter()
1734                .collect();
1735
1736        let mut all_results = vec![];
1737        for hnsw in [&builder, &loaded] {
1738            let mut bitset = mask_generator.generate(TOTAL);
1739            for id in (0..TOTAL as u32).step_by(2) {
1740                bitset.insert(id);
1741            }
1742            let results = hnsw
1743                .search_acorn(query.clone(), k, &params, &bitset, store.as_ref())
1744                .unwrap();
1745            assert_eq!(results.len(), k);
1746            assert!(results.iter().all(|node| node.id % 2 == 0));
1747            assert!(results.windows(2).all(|w| w[0].dist <= w[1].dist));
1748            let hits = results.iter().filter(|n| truth.contains(&n.id)).count();
1749            let recall = hits as f32 / k as f32;
1750            assert!(recall >= 0.5, "recall {recall} below 0.5 (k={k})");
1751            all_results.push(results);
1752        }
1753        assert_eq!(all_results[0], all_results[1]);
1754
1755        // default ef (k + k/2) and a deletion-style mask (all but a few rows)
1756        let default_ef_params = HnswQueryParams {
1757            ef: k + k / 2,
1758            ..params
1759        };
1760        for excluded_stride in [2, 400] {
1761            let passes = |id: u32| id % excluded_stride != 1;
1762            let mut bitset = mask_generator.generate(TOTAL);
1763            for id in (0..TOTAL as u32).filter(|id| passes(*id)) {
1764                bitset.insert(id);
1765            }
1766            let truth: std::collections::HashSet<u32> =
1767                brute_force_topk_masked(store.as_ref(), query.clone(), k, passes)
1768                    .into_iter()
1769                    .collect();
1770            let results = builder
1771                .search_acorn(
1772                    query.clone(),
1773                    k,
1774                    &default_ef_params,
1775                    &bitset,
1776                    store.as_ref(),
1777                )
1778                .unwrap();
1779            assert_eq!(results.len(), k);
1780            assert!(results.iter().all(|node| passes(node.id)));
1781            let hits = results.iter().filter(|n| truth.contains(&n.id)).count();
1782            let recall = hits as f32 / k as f32;
1783            assert!(recall >= 0.5, "recall {recall} below 0.5 (k={k})");
1784        }
1785    }
1786
1787    /// Dispatch: dense prefilters take the graph traversal, sparse ones the
1788    /// exact flat scan, and both return only mask-passing row ids.
1789    #[tokio::test]
1790    async fn test_subindex_prefilter_dispatch() {
1791        use arrow_array::cast::AsArray;
1792        use async_trait::async_trait;
1793        use lance_core::Result;
1794        use lance_select::{RowAddrMask, RowAddrTreeMap};
1795
1796        use crate::metrics::NoOpMetricsCollector;
1797        use crate::prefilter::PreFilter;
1798
1799        struct MaskPreFilter {
1800            mask: Arc<RowAddrMask>,
1801        }
1802
1803        #[async_trait]
1804        impl PreFilter for MaskPreFilter {
1805            async fn wait_for_ready(&self) -> Result<()> {
1806                Ok(())
1807            }
1808            fn is_empty(&self) -> bool {
1809                false
1810            }
1811            fn mask(&self) -> Arc<RowAddrMask> {
1812                self.mask.clone()
1813            }
1814            fn filter_row_ids<'a>(
1815                &self,
1816                row_ids: Box<dyn Iterator<Item = &'a u64> + 'a>,
1817            ) -> Vec<u64> {
1818                self.mask.selected_indices(row_ids)
1819            }
1820        }
1821
1822        const DIM: usize = 32;
1823        const TOTAL: usize = 2048;
1824        let fsl =
1825            FixedSizeListArray::try_new_from_values(generate_random_array(TOTAL * DIM), DIM as i32)
1826                .unwrap();
1827        let store = Arc::new(FlatFloatStorage::new(fsl.clone(), DistanceType::L2));
1828        let hnsw = HNSW::index_vectors(
1829            store.as_ref(),
1830            HnswBuildParams::default().num_edges(20).ef_construction(50),
1831        )
1832        .unwrap();
1833
1834        let k = 10;
1835        let query_key = fsl.value(0);
1836
1837        let search_row_ids = |allowed: Vec<u64>, use_acorn: bool| {
1838            let params = HnswQueryParams {
1839                ef: 50,
1840                lower_bound: None,
1841                upper_bound: None,
1842                dist_q_c: 0.0,
1843                use_acorn,
1844            };
1845            let filter = Arc::new(MaskPreFilter {
1846                mask: Arc::new(RowAddrMask::from_allowed(RowAddrTreeMap::from_iter(
1847                    allowed,
1848                ))),
1849            });
1850            let batch = hnsw
1851                .search(
1852                    query_key.clone(),
1853                    k,
1854                    params,
1855                    store.as_ref(),
1856                    filter,
1857                    &NoOpMetricsCollector,
1858                )
1859                .unwrap();
1860            batch[lance_core::ROW_ID]
1861                .as_primitive::<arrow_array::types::UInt64Type>()
1862                .values()
1863                .to_vec()
1864        };
1865
1866        // Dense mask (50% of rows), in both modes.
1867        let dense: Vec<u64> = (0..TOTAL as u64).step_by(2).collect();
1868        for use_acorn in [false, true] {
1869            let row_ids = search_row_ids(dense.clone(), use_acorn);
1870            assert_eq!(row_ids.len(), k);
1871            assert!(row_ids.iter().all(|id| id % 2 == 0));
1872        }
1873
1874        // All-pass mask: shortcuts to the unfiltered path.
1875        let all: Vec<u64> = (0..TOTAL as u64).collect();
1876        let unfiltered = hnsw
1877            .search_basic(
1878                query_key.clone(),
1879                k,
1880                &HnswQueryParams {
1881                    ef: 50,
1882                    lower_bound: None,
1883                    upper_bound: None,
1884                    dist_q_c: 0.0,
1885                    use_acorn: false,
1886                },
1887                None,
1888                store.as_ref(),
1889            )
1890            .unwrap();
1891        let row_ids = search_row_ids(all, true);
1892        assert_eq!(
1893            row_ids,
1894            unfiltered.iter().map(|n| n.id as u64).collect::<Vec<_>>()
1895        );
1896
1897        // Sparse mask (< 10% of rows): the flat scan, which is exact.
1898        let sparse: Vec<u64> = (0..TOTAL as u64).step_by(25).collect();
1899        let row_ids = search_row_ids(sparse.clone(), true);
1900        assert_eq!(row_ids.len(), k);
1901        let truth = brute_force_topk_masked(store.as_ref(), query_key.clone(), k, |id| {
1902            sparse.contains(&(id as u64))
1903        });
1904        let mut got: Vec<u32> = row_ids.iter().map(|id| *id as u32).collect();
1905        got.sort_unstable();
1906        let mut expected = truth;
1907        expected.sort_unstable();
1908        assert_eq!(got, expected);
1909    }
1910
1911    /// Regression guard for the `level_offsets` misalignment (issue #6746).
1912    /// `to_batch` writes the entry-point node at *every* level, but
1913    /// `level_count` only counts it at level 0, so the serialized batch has
1914    /// strictly more rows than `sum(level_count)` and the upper-level
1915    /// `level_offsets` slices are off-by-one / non-monotonic. The Arrow-backed
1916    /// loaded graph must still search bit-identically to the in-memory build:
1917    /// it keys upper levels by `__vector_id` value via the `Sparse` map
1918    /// (last-write-wins), never `row == id`. A naive `row == id`
1919    /// reimplementation would pass the small cases but break here.
1920    #[tokio::test]
1921    async fn test_loaded_level_offsets_misalignment_invariant() {
1922        use arrow::array::AsArray;
1923        use arrow::datatypes::UInt32Type;
1924
1925        const DIM: usize = 32;
1926        const TOTAL: usize = 2048;
1927        let fsl =
1928            FixedSizeListArray::try_new_from_values(generate_random_array(TOTAL * DIM), DIM as i32)
1929                .unwrap();
1930        let store = Arc::new(FlatFloatStorage::new(fsl.clone(), DistanceType::L2));
1931        let builder = HNSW::index_vectors(
1932            store.as_ref(),
1933            HnswBuildParams::default().num_edges(20).ef_construction(50),
1934        )
1935        .unwrap();
1936
1937        // The scenario only exists on a multi-level graph.
1938        assert!(
1939            builder.max_level() >= 2,
1940            "expected a multi-level graph (got max_level {})",
1941            builder.max_level()
1942        );
1943
1944        let batch = builder.to_batch().unwrap();
1945        let md = builder.metadata();
1946        let total_counted = *md.level_offsets.last().unwrap();
1947
1948        // The exact misalignment: more serialized rows than `level_count` sums
1949        // to, because the entry-point node is written at every level yet
1950        // counted only at level 0.
1951        assert!(
1952            batch.num_rows() > total_counted,
1953            "expected serialized rows ({}) to exceed sum(level_count) ({}) -- \
1954             entry point should be written at every level",
1955            batch.num_rows(),
1956            total_counted,
1957        );
1958
1959        // Level-0 slice must still be exactly `[0, N)` with
1960        // `__vector_id == row` -- the precondition for `LevelLookup::Dense`.
1961        let n = md.level_offsets[1];
1962        assert_eq!(n, TOTAL);
1963        let level0 = batch.slice(0, n);
1964        let ids = level0.column(0).as_primitive::<UInt32Type>();
1965        assert!(
1966            ids.values()
1967                .iter()
1968                .enumerate()
1969                .all(|(row, id)| *id == row as u32),
1970            "level-0 __vector_id must equal the row index",
1971        );
1972
1973        // Despite the surplus rows and off-by-one upper slices, the loaded
1974        // graph searches bit-identically to the in-memory build (old `load`
1975        // semantics preserved via the `Sparse` last-write-wins map).
1976        let loaded = HNSW::load(batch).unwrap();
1977        assert!(matches!(loaded.inner.graph, HnswGraph::Loaded(_)));
1978        let params = HnswQueryParams {
1979            ef: 50,
1980            lower_bound: None,
1981            upper_bound: None,
1982            dist_q_c: 0.0,
1983            use_acorn: false,
1984        };
1985        let query = fsl.value(0);
1986        let builder_results = builder
1987            .search_basic(query.clone(), 10, &params, None, store.as_ref())
1988            .unwrap();
1989        let loaded_results = loaded
1990            .search_basic(query, 10, &params, None, store.as_ref())
1991            .unwrap();
1992        assert_eq!(builder_results, loaded_results);
1993    }
1994
1995    /// `load()` must reject a batch whose level-0 `__vector_id` no longer
1996    /// matches the row index. The `LevelLookup::Dense` fast path relies on
1997    /// `row == id`, and the old `debug_assert!` was compiled out of release
1998    /// builds -- so a corrupt batch must fail at the `load()` boundary instead
1999    /// of silently searching the wrong neighbor lists.
2000    #[tokio::test]
2001    async fn test_load_rejects_misaligned_level0_id() {
2002        use arrow::array::AsArray;
2003        use arrow::datatypes::UInt32Type;
2004
2005        const DIM: usize = 16;
2006        const TOTAL: usize = 256;
2007        let fsl =
2008            FixedSizeListArray::try_new_from_values(generate_random_array(TOTAL * DIM), DIM as i32)
2009                .unwrap();
2010        let store = Arc::new(FlatFloatStorage::new(fsl, DistanceType::L2));
2011        let builder = HNSW::index_vectors(
2012            store.as_ref(),
2013            HnswBuildParams::default().num_edges(20).ef_construction(50),
2014        )
2015        .unwrap();
2016
2017        let batch = builder.to_batch().unwrap();
2018        // Row 0 is always a level-0 node; break its `__vector_id == row`
2019        // invariant while preserving the (metadata-bearing) schema.
2020        let mut ids = batch
2021            .column(0)
2022            .as_primitive::<UInt32Type>()
2023            .values()
2024            .to_vec();
2025        ids[0] = ids.len() as u32;
2026        let mut columns = batch.columns().to_vec();
2027        columns[0] = Arc::new(UInt32Array::from(ids));
2028        let corrupted = RecordBatch::try_new(batch.schema(), columns).unwrap();
2029
2030        assert!(
2031            HNSW::load(corrupted).is_err(),
2032            "load() must reject a misaligned level-0 __vector_id"
2033        );
2034    }
2035
2036    /// `load()` must reject metadata whose `entry_point` is out of range for
2037    /// the node count: it indexes the `Dense` level-0 lookup directly, so an
2038    /// out-of-range value would read past the level-0 neighbor buffer at search
2039    /// time.
2040    #[tokio::test]
2041    async fn test_load_rejects_out_of_range_entry_point() {
2042        use super::{HNSW_METADATA_KEY, HnswMetadata};
2043
2044        const DIM: usize = 16;
2045        const TOTAL: usize = 256;
2046        let fsl =
2047            FixedSizeListArray::try_new_from_values(generate_random_array(TOTAL * DIM), DIM as i32)
2048                .unwrap();
2049        let store = Arc::new(FlatFloatStorage::new(fsl, DistanceType::L2));
2050        let builder = HNSW::index_vectors(
2051            store.as_ref(),
2052            HnswBuildParams::default().num_edges(20).ef_construction(50),
2053        )
2054        .unwrap();
2055
2056        let batch = builder.to_batch().unwrap();
2057        let mut metadata = batch.schema_ref().metadata().clone();
2058        let mut md: HnswMetadata =
2059            serde_json::from_str(metadata.get(HNSW_METADATA_KEY).unwrap()).unwrap();
2060        // Valid entry points are `[0, N)`; `level_offsets[1]` == N is one past.
2061        let n = md.level_offsets[1];
2062        md.entry_point = n as u32;
2063        metadata.insert(
2064            HNSW_METADATA_KEY.to_string(),
2065            serde_json::to_string(&md).unwrap(),
2066        );
2067        // Rebuild the batch under the rewritten metadata. `with_schema` would
2068        // reject this: it requires the new metadata to be a superset, but we
2069        // are changing an existing key's value, not adding one.
2070        let schema = batch.schema().as_ref().clone().with_metadata(metadata);
2071        let corrupted = RecordBatch::try_new(Arc::new(schema), batch.columns().to_vec()).unwrap();
2072
2073        assert!(
2074            HNSW::load(corrupted).is_err(),
2075            "load() must reject an out-of-range entry_point"
2076        );
2077    }
2078
2079    /// An empty index round-trips: 0-row `to_batch` -> `load` -> empty graph.
2080    #[tokio::test]
2081    async fn test_loaded_empty_index() {
2082        const DIM: usize = 16;
2083        let fsl =
2084            FixedSizeListArray::try_new_from_values(generate_random_array(0), DIM as i32).unwrap();
2085        let store = Arc::new(FlatFloatStorage::new(fsl, DistanceType::L2));
2086        let builder = HNSW::index_vectors(store.as_ref(), HnswBuildParams::default()).unwrap();
2087        assert!(builder.is_empty());
2088
2089        let batch = builder.to_batch().unwrap();
2090        assert_eq!(batch.num_rows(), 0);
2091
2092        let loaded = HNSW::load(batch).unwrap();
2093        assert!(loaded.is_empty());
2094        assert_eq!(loaded.len(), 0);
2095        // A 0-row load short-circuits to the empty (Built) graph.
2096        assert!(!matches!(loaded.inner.graph, HnswGraph::Loaded(_)));
2097        assert_eq!(loaded.to_batch().unwrap().num_rows(), 0);
2098    }
2099
2100    /// build -> `to_batch` (b1) -> `load` -> `to_batch` (b2) must satisfy
2101    /// `b1 == b2`, and the round-tripped batch must reload and search
2102    /// identically. This is exactly the IVF partition-cache path:
2103    /// `ivf/partition_serde.rs` calls `to_batch()` on a *loaded* index.
2104    #[tokio::test]
2105    async fn test_to_batch_roundtrip_loaded() {
2106        const DIM: usize = 24;
2107        const TOTAL: usize = 1500;
2108        let fsl =
2109            FixedSizeListArray::try_new_from_values(generate_random_array(TOTAL * DIM), DIM as i32)
2110                .unwrap();
2111        let store = Arc::new(FlatFloatStorage::new(fsl.clone(), DistanceType::L2));
2112        let builder = HNSW::index_vectors(
2113            store.as_ref(),
2114            HnswBuildParams::default().num_edges(16).ef_construction(50),
2115        )
2116        .unwrap();
2117
2118        let b1 = builder.to_batch().unwrap();
2119        let loaded = HNSW::load(b1.clone()).unwrap();
2120        assert!(matches!(loaded.inner.graph, HnswGraph::Loaded(_)));
2121        let b2 = loaded.to_batch().unwrap();
2122        assert_eq!(b1, b2);
2123
2124        let reloaded = HNSW::load(b2).unwrap();
2125        let params = HnswQueryParams {
2126            ef: 50,
2127            lower_bound: None,
2128            upper_bound: None,
2129            dist_q_c: 0.0,
2130            use_acorn: false,
2131        };
2132        let query = fsl.value(7);
2133        let a = builder
2134            .search_basic(query.clone(), 10, &params, None, store.as_ref())
2135            .unwrap();
2136        let b = reloaded
2137            .search_basic(query, 10, &params, None, store.as_ref())
2138            .unwrap();
2139        assert_eq!(a, b);
2140    }
2141
2142    /// Regression for the IVF partition-cache round-trip: a disk-loaded batch
2143    /// inherits extra schema metadata keys from the index file (e.g.
2144    /// `INDEX_METADATA_SCHEMA_KEY`, `IVF_METADATA_KEY`). `to_batch()` on the
2145    /// loaded graph must *merge* the HNSW key into that metadata rather than
2146    /// replacing it -- otherwise the new schema is not a superset of the
2147    /// current one and `RecordBatch::with_schema` fails with "target schema is
2148    /// not superset of current schema".
2149    #[tokio::test]
2150    async fn test_to_batch_loaded_preserves_extra_schema_metadata() {
2151        use super::HNSW_METADATA_KEY;
2152
2153        const DIM: usize = 24;
2154        const TOTAL: usize = 512;
2155        let fsl =
2156            FixedSizeListArray::try_new_from_values(generate_random_array(TOTAL * DIM), DIM as i32)
2157                .unwrap();
2158        let store = Arc::new(FlatFloatStorage::new(fsl, DistanceType::L2));
2159        let builder = HNSW::index_vectors(
2160            store.as_ref(),
2161            HnswBuildParams::default().num_edges(16).ef_construction(50),
2162        )
2163        .unwrap();
2164
2165        // Simulate the disk-load path (`ivf/v2.rs::load_partition_entry`):
2166        // the batch reaching `HNSW::load` carries the index file's schema
2167        // metadata in addition to the HNSW key.
2168        let built_batch = builder.to_batch().unwrap();
2169        let mut metadata = built_batch.schema_ref().metadata().clone();
2170        metadata.insert(
2171            "lance:index_metadata".to_string(),
2172            "{\"distance_type\":\"l2\"}".to_string(),
2173        );
2174        metadata.insert("lance:ivf".to_string(), "42".to_string());
2175        let schema = built_batch
2176            .schema()
2177            .as_ref()
2178            .clone()
2179            .with_metadata(metadata);
2180        let batch_with_extra =
2181            RecordBatch::try_new(Arc::new(schema), built_batch.columns().to_vec()).unwrap();
2182
2183        let loaded = HNSW::load(batch_with_extra).unwrap();
2184        assert!(matches!(loaded.inner.graph, HnswGraph::Loaded(_)));
2185
2186        // Before the fix this fails: the re-stamped schema dropped the extra
2187        // keys, so `with_schema`'s superset check rejected the round-trip.
2188        let out = loaded.to_batch().unwrap();
2189        let out_metadata = out.schema_ref().metadata();
2190        assert!(out_metadata.contains_key(HNSW_METADATA_KEY));
2191        assert_eq!(
2192            out_metadata.get("lance:index_metadata").map(String::as_str),
2193            Some("{\"distance_type\":\"l2\"}"),
2194        );
2195        assert_eq!(
2196            out_metadata.get("lance:ivf").map(String::as_str),
2197            Some("42")
2198        );
2199
2200        // The HNSW key must still decode to valid metadata after the merge.
2201        let reloaded = HNSW::load(out).unwrap();
2202        assert_eq!(reloaded.len(), loaded.len());
2203    }
2204
2205    /// The loaded graph shares the Arrow batch and reconstructs no per-node
2206    /// `Vec<GraphBuilderNode>` / `Vec<OrderedNode>`, so it is strictly
2207    /// lighter than the in-memory build representation.
2208    #[tokio::test]
2209    async fn test_loaded_graph_is_arrow_backed() {
2210        const DIM: usize = 32;
2211        const TOTAL: usize = 2048;
2212        let fsl =
2213            FixedSizeListArray::try_new_from_values(generate_random_array(TOTAL * DIM), DIM as i32)
2214                .unwrap();
2215        let store = Arc::new(FlatFloatStorage::new(fsl, DistanceType::L2));
2216        let builder = HNSW::index_vectors(
2217            store.as_ref(),
2218            HnswBuildParams::default().num_edges(20).ef_construction(50),
2219        )
2220        .unwrap();
2221        assert!(!matches!(builder.inner.graph, HnswGraph::Loaded(_)));
2222
2223        let loaded = HNSW::load(builder.to_batch().unwrap()).unwrap();
2224        assert!(matches!(loaded.inner.graph, HnswGraph::Loaded(_)));
2225        assert!(
2226            loaded.deep_size_of() < builder.deep_size_of(),
2227            "loaded graph ({}) should be lighter than built ({})",
2228            loaded.deep_size_of(),
2229            builder.deep_size_of(),
2230        );
2231    }
2232}