use std::collections::HashMap;
use std::marker::PhantomData;
use crate::bbox::{compute_bounding_box_over_indices, Interval};
use crate::data_source::DataSource;
use crate::dim::Dim;
use crate::filter::PointFilter;
use crate::metric::{Distance, L2};
use crate::node::Node;
use crate::params::SearchParams;
use crate::result_set::{
KeepInsertionOrder, KnnResultSet, RadiusResultSet, ResultItem, ResultSet, RknnResultSet,
TieBreak,
};
use crate::scalar::{DistanceValue, IndexType, Scalar};
use crate::search::{find_neighbors as search_find_neighbors, SearchCtx};
pub(crate) fn first0bit(n: usize) -> usize {
let mut num = n;
let mut pos = 0usize;
while num & 1 == 1 {
num >>= 1;
pos += 1;
}
pos
}
struct Slot<T: Scalar, D: Dim, Idx: IndexType> {
vind: Vec<Idx>,
nodes: Vec<Node<T>>,
root_bbox: D::Array<Interval<T>>,
}
impl<T: Scalar, D: Dim, Idx: IndexType> Slot<T, D, Idx> {
fn empty(dim: D) -> Self {
Slot {
vind: Vec::new(),
nodes: Vec::new(),
root_bbox: dim.filled(Interval::default()),
}
}
}
struct TombstoneFilter<'a> {
tree_index: &'a [i32],
}
impl<'a, Idx: IndexType> PointFilter<Idx> for TombstoneFilter<'a> {
#[inline]
fn is_active(&self, idx: Idx) -> bool {
self.tree_index[idx.to_usize()] != -1
}
}
pub struct DynamicKdTreeBuilder<T, D, DS, M = L2, Idx = u32, TB = KeepInsertionOrder>
where
T: Scalar,
D: Dim,
DS: DataSource<T>,
M: Distance<T>,
Idx: IndexType,
TB: TieBreak,
{
dim: D,
dataset: DS,
metric: M,
leaf_max_size: usize,
maximum_point_count: usize,
_marker: PhantomData<(T, Idx, TB)>,
}
impl<T: Scalar, D: Dim, DS: DataSource<T>> DynamicKdTreeBuilder<T, D, DS>
where
L2: Distance<T>,
{
pub fn new(dim: D, dataset: DS) -> Self {
Self {
dim,
dataset,
metric: L2,
leaf_max_size: 10,
maximum_point_count: 1_000_000_000,
_marker: PhantomData,
}
}
}
impl<T, D, DS, M, Idx, TB> DynamicKdTreeBuilder<T, D, DS, M, Idx, TB>
where
T: Scalar,
D: Dim,
DS: DataSource<T>,
M: Distance<T>,
Idx: IndexType,
TB: TieBreak,
{
pub fn with_metric<M2: Distance<T>>(
self,
metric: M2,
) -> DynamicKdTreeBuilder<T, D, DS, M2, Idx, TB> {
DynamicKdTreeBuilder {
dim: self.dim,
dataset: self.dataset,
metric,
leaf_max_size: self.leaf_max_size,
maximum_point_count: self.maximum_point_count,
_marker: PhantomData,
}
}
pub fn leaf_max_size(mut self, n: usize) -> Self {
assert!(n > 0, "leaf_max_size: n must be > 0");
self.leaf_max_size = n;
self
}
pub fn maximum_point_count(mut self, n: usize) -> Self {
assert!(n >= 1, "maximum_point_count: n must be >= 1");
self.maximum_point_count = n;
self
}
pub fn index_type<Idx2: IndexType>(self) -> DynamicKdTreeBuilder<T, D, DS, M, Idx2, TB> {
DynamicKdTreeBuilder {
dim: self.dim,
dataset: self.dataset,
metric: self.metric,
leaf_max_size: self.leaf_max_size,
maximum_point_count: self.maximum_point_count,
_marker: PhantomData,
}
}
pub fn tie_break<TB2: TieBreak>(self) -> DynamicKdTreeBuilder<T, D, DS, M, Idx, TB2> {
DynamicKdTreeBuilder {
dim: self.dim,
dataset: self.dataset,
metric: self.metric,
leaf_max_size: self.leaf_max_size,
maximum_point_count: self.maximum_point_count,
_marker: PhantomData,
}
}
pub fn build(self) -> DynamicKdTree<T, D, DS, M, Idx, TB> {
let DynamicKdTreeBuilder {
dim,
dataset,
metric,
leaf_max_size,
maximum_point_count,
..
} = self;
let tree_count = ((maximum_point_count as f64).log2() as usize) + 1;
let slots: Vec<Slot<T, D, Idx>> = (0..tree_count).map(|_| Slot::empty(dim)).collect();
let mut tree = DynamicKdTree {
dataset,
metric,
dim,
leaf_max_size,
slots,
tree_index: Vec::new(),
removed: HashMap::new(),
point_count: 0,
_marker: PhantomData,
};
let n = tree.dataset.point_count();
if n > 0 {
tree.add_points(0, n - 1);
}
tree
}
}
pub struct DynamicKdTree<T, D, DS, M = L2, Idx = u32, TB = KeepInsertionOrder>
where
T: Scalar,
D: Dim,
DS: DataSource<T>,
M: Distance<T>,
Idx: IndexType,
TB: TieBreak,
{
dataset: DS,
metric: M,
dim: D,
leaf_max_size: usize,
slots: Vec<Slot<T, D, Idx>>,
tree_index: Vec<i32>,
removed: HashMap<usize, i32>,
point_count: usize,
_marker: PhantomData<TB>,
}
impl<T, D, DS, M, Idx, TB> DynamicKdTree<T, D, DS, M, Idx, TB>
where
T: Scalar,
D: Dim,
DS: DataSource<T>,
M: Distance<T>,
Idx: IndexType,
TB: TieBreak,
{
pub fn tree_count(&self) -> usize {
self.slots.len()
}
pub fn active_count(&self) -> usize {
self.tree_index.iter().filter(|&&v| v != -1).count()
}
pub fn point_indices_of_slot(&self, slot: usize) -> &[Idx] {
&self.slots[slot].vind
}
pub fn tree_index(&self) -> &[i32] {
&self.tree_index
}
pub fn removed_len(&self) -> usize {
self.removed.len()
}
pub fn dataset(&self) -> &DS {
&self.dataset
}
pub fn dataset_mut(&mut self) -> &mut DS {
&mut self.dataset
}
pub fn add_points(&mut self, start: usize, end_inclusive: usize) {
let max_index = self.add_points_bookkeeping(start, end_inclusive);
self.rebuild_slots_sequential(max_index);
}
fn add_points_bookkeeping(&mut self, start: usize, end_inclusive: usize) -> usize {
let mut max_index: usize = 0;
for idx in start..=end_inclusive {
if let Some(&slot) = self.removed.get(&idx) {
self.tree_index[idx] = slot;
self.removed.remove(&idx);
continue;
}
assert_eq!(
idx, self.point_count,
"add_points: index {idx} is a genuinely-new point but does not equal \
point_count ({}) -- nanoflann's addPoints contract requires brand-new \
indices to be a contiguous append starting exactly at point_count \
(treeIndex_ is indexed by the running pointCount_ counter, not by idx); \
see DynamicKdTree::add_points's doc comment",
self.point_count
);
let pos = first0bit(self.point_count);
assert!(
pos < self.slots.len(),
"add_points: point_count ({}) has outgrown this forest's capacity ({} slots) \
-- construct the forest with a larger DynamicKdTreeBuilder::maximum_point_count",
self.point_count,
self.slots.len()
);
if pos > max_index {
max_index = pos;
}
if self.tree_index.len() <= self.point_count {
self.tree_index.resize(self.point_count + 1, -1);
}
self.tree_index[self.point_count] = pos as i32;
for i in 0..pos {
let mut entries = std::mem::take(&mut self.slots[i].vind);
for &e in entries.iter() {
self.slots[pos].vind.push(e);
let e_usize = e.to_usize();
if self.tree_index[e_usize] != -1 {
self.tree_index[e_usize] = pos as i32;
} else {
self.removed.insert(e_usize, pos as i32);
}
}
entries.clear();
self.slots[i].vind = entries;
}
self.slots[pos].vind.push(Idx::from_usize(idx));
self.point_count += 1;
}
max_index
}
fn rebuild_slots_sequential(&mut self, max_index: usize) {
let dim_n = self.dim.dim();
let leaf_max_size = self.leaf_max_size;
for i in 0..=max_index {
self.slots[i].nodes.clear();
if self.slots[i].vind.is_empty() {
continue;
}
let mut bbox = self.dim.filled(Interval::default());
compute_bounding_box_over_indices(
&self.dataset,
dim_n,
&self.slots[i].vind,
bbox.as_mut(),
);
let slot = &mut self.slots[i];
{
let mut builder = crate::build::SubtreeBuilder {
ds: &self.dataset,
dim: dim_n,
leaf_max_size,
base: 0,
vind: &mut slot.vind,
arena: &mut slot.nodes,
};
builder.build(bbox.as_mut());
}
slot.root_bbox = bbox;
}
}
pub fn remove_point(&mut self, idx: usize) -> bool {
if idx >= self.point_count {
return false;
}
if self.tree_index[idx] == -1 {
return false;
}
self.removed.insert(idx, self.tree_index[idx]);
self.tree_index[idx] = -1;
true
}
pub fn find_neighbors<R: ResultSet<M::DistanceType, Idx>>(
&self,
result: &mut R,
query: &[T],
params: &SearchParams,
) -> bool {
let filter = TombstoneFilter {
tree_index: &self.tree_index,
};
let mut scratch = self.dim.filled(M::DistanceType::ZERO);
for slot in &self.slots {
let ctx = SearchCtx {
ds: &self.dataset,
metric: &self.metric,
dim: self.dim,
nodes: &slot.nodes,
vind: &slot.vind,
root_bbox: slot.root_bbox.as_ref(),
};
let _ = search_find_neighbors(&ctx, result, query, params, &filter, scratch.as_mut());
}
result.full()
}
pub fn knn_search(
&self,
query: &[T],
out_indices: &mut [Idx],
out_dists: &mut [M::DistanceType],
) -> usize {
self.knn_search_with(query, out_indices, out_dists, &SearchParams::default())
}
pub fn knn_search_with(
&self,
query: &[T],
out_indices: &mut [Idx],
out_dists: &mut [M::DistanceType],
params: &SearchParams,
) -> usize {
assert_eq!(
out_indices.len(),
out_dists.len(),
"knn_search: out_indices/out_dists length mismatch"
);
let mut rs = KnnResultSet::<M::DistanceType, Idx, TB>::new(out_indices, out_dists);
self.find_neighbors(&mut rs, query, params);
rs.size()
}
pub fn rknn_search(
&self,
query: &[T],
radius: M::DistanceType,
out_indices: &mut [Idx],
out_dists: &mut [M::DistanceType],
) -> usize {
self.rknn_search_with(
query,
radius,
out_indices,
out_dists,
&SearchParams::default(),
)
}
pub fn rknn_search_with(
&self,
query: &[T],
radius: M::DistanceType,
out_indices: &mut [Idx],
out_dists: &mut [M::DistanceType],
params: &SearchParams,
) -> usize {
assert_eq!(
out_indices.len(),
out_dists.len(),
"rknn_search: out_indices/out_dists length mismatch"
);
let mut rs = RknnResultSet::<M::DistanceType, Idx, TB>::new(out_indices, out_dists, radius);
self.find_neighbors(&mut rs, query, params);
rs.size()
}
pub fn radius_search(
&self,
query: &[T],
radius: M::DistanceType,
out: &mut Vec<ResultItem<Idx, M::DistanceType>>,
) -> usize {
self.radius_search_with(query, radius, out, &SearchParams::default())
}
pub fn radius_search_with(
&self,
query: &[T],
radius: M::DistanceType,
out: &mut Vec<ResultItem<Idx, M::DistanceType>>,
params: &SearchParams,
) -> usize {
let mut rs = RadiusResultSet::new(radius, out);
self.find_neighbors(&mut rs, query, params);
rs.size()
}
}
#[cfg(feature = "parallel")]
impl<T, D, DS, M, Idx, TB> DynamicKdTree<T, D, DS, M, Idx, TB>
where
T: Scalar,
D: Dim,
DS: DataSource<T> + Sync,
M: Distance<T>,
Idx: IndexType,
TB: TieBreak,
{
pub fn add_points_parallel(&mut self, start: usize, end_inclusive: usize) {
let max_index = self.add_points_bookkeeping(start, end_inclusive);
let dim = self.dim;
let dim_n = dim.dim();
let leaf_max_size = self.leaf_max_size;
let dataset = &self.dataset;
let slots = &mut self.slots[0..=max_index];
rayon::scope(|scope| {
for slot in slots.iter_mut() {
scope.spawn(move |_| {
slot.nodes.clear();
if slot.vind.is_empty() {
return;
}
let mut bbox = dim.filled(Interval::default());
compute_bounding_box_over_indices(dataset, dim_n, &slot.vind, bbox.as_mut());
slot.nodes = crate::build_parallel::build_tree_parallel(
dataset,
dim_n,
leaf_max_size,
&mut slot.vind,
bbox.as_mut(),
);
slot.root_bbox = bbox;
});
}
});
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::dim::ConstDim;
struct Ungated<const N: usize>(Vec<[f64; N]>);
impl<const N: usize> DataSource<f64> for Ungated<N> {
fn point_count(&self) -> usize {
0
}
fn point_component(&self, idx: usize, dim: usize) -> f64 {
self.0[idx][dim]
}
}
#[test]
fn first0bit_matches_hand_derived_vector() {
assert_eq!(first0bit(0), 0);
assert_eq!(first0bit(1), 1);
assert_eq!(first0bit(2), 0);
assert_eq!(first0bit(3), 2);
assert_eq!(first0bit(4), 0);
assert_eq!(first0bit(5), 1);
assert_eq!(first0bit(6), 0);
assert_eq!(first0bit(7), 3);
assert_eq!(first0bit(8), 0);
assert_eq!(first0bit(15), 4);
}
#[test]
fn add_4_points_one_batch_matches_first0bit_slot_occupancy() {
let ds = Ungated(vec![[10.0], [20.0], [30.0], [40.0]]);
let mut tree = DynamicKdTreeBuilder::new(ConstDim::<1>, ds)
.maximum_point_count(1000)
.build();
tree.add_points(0, 3);
assert!(
tree.point_indices_of_slot(0).is_empty(),
"slot 0 must be empty"
);
assert!(
tree.point_indices_of_slot(1).is_empty(),
"slot 1 must be empty"
);
assert_eq!(
tree.point_indices_of_slot(2),
&[2u32, 0, 1, 3],
"slot 2 must hold the EXACT merge-append order, not just the right membership"
);
assert_eq!(tree.tree_index(), &[2, 2, 2, 2]);
let mut union: Vec<u32> = (0..tree.tree_count())
.flat_map(|s| tree.point_indices_of_slot(s).iter().copied())
.collect();
union.sort_unstable();
assert_eq!(union, vec![0, 1, 2, 3]);
}
#[test]
fn merge_migrates_a_tombstones_recorded_slot() {
let ds = Ungated(vec![[10.0], [20.0], [30.0], [40.0]]);
let mut tree = DynamicKdTreeBuilder::new(ConstDim::<1>, ds)
.maximum_point_count(1000)
.build();
tree.add_points(0, 1);
assert!(tree.remove_point(0));
tree.add_points(2, 2);
tree.add_points(3, 3);
assert_eq!(
tree.tree_index()[0],
-1,
"removed point must stay -1 through a merge"
);
assert_eq!(tree.removed_len(), 1);
assert_eq!(tree.tree_index()[1], 2);
assert_eq!(tree.tree_index()[2], 2);
assert_eq!(tree.tree_index()[3], 2);
assert_eq!(
tree.point_indices_of_slot(2),
&[2u32, 0, 1, 3],
"slot 2 must hold the EXACT merge-append order"
);
let mut slot2_sorted: Vec<u32> = tree.point_indices_of_slot(2).to_vec();
slot2_sorted.sort_unstable();
assert_eq!(slot2_sorted, vec![0, 1, 2, 3]);
}
#[test]
fn permuting_rebuild_with_leaf_max_size_one_matches_hand_derived_sort() {
let ds = Ungated(vec![
[70.0, 100.0], [10.0, 100.0], [60.0, 100.0], [20.0, 100.0], [50.0, 100.0], [30.0, 100.0], [80.0, 100.0], [40.0, 100.0], ]);
let mut tree = DynamicKdTreeBuilder::new(ConstDim::<2>, ds)
.leaf_max_size(1)
.maximum_point_count(1000)
.build();
tree.add_points(0, 7);
assert!(tree.point_indices_of_slot(0).is_empty());
assert!(tree.point_indices_of_slot(1).is_empty());
assert!(tree.point_indices_of_slot(2).is_empty());
assert_eq!(
tree.point_indices_of_slot(3),
&[1u32, 3, 5, 7, 4, 2, 0, 6],
"leaf_max_size=1 must fully sort slot 3 by dim0 -- a genuinely PERMUTING rebuild, \
not just a membership-preserving pass-through"
);
for &v in tree.tree_index() {
assert_eq!(v, 3, "all 8 points must live in slot 3");
}
assert_eq!(tree.active_count(), 8);
}
#[test]
fn reactivation_restores_tree_index_from_removed_map_without_growing_slots() {
let ds = Ungated(vec![[10.0], [20.0], [30.0], [40.0]]);
let mut tree = DynamicKdTreeBuilder::new(ConstDim::<1>, ds)
.maximum_point_count(1000)
.build();
tree.add_points(0, 3); let slot_before_removal = tree.tree_index()[1];
assert_ne!(slot_before_removal, -1);
assert!(tree.remove_point(1));
assert_eq!(tree.tree_index()[1], -1);
assert_eq!(tree.removed_len(), 1);
let physical_count_before: usize = (0..tree.tree_count())
.map(|s| tree.point_indices_of_slot(s).len())
.sum();
tree.add_points(1, 1);
assert_eq!(
tree.tree_index()[1],
slot_before_removal,
"reactivation must restore the ORIGINAL (removed-from) slot value"
);
assert_eq!(
tree.removed_len(),
0,
"removed map must be empty after reactivation"
);
let physical_count_after: usize = (0..tree.tree_count())
.map(|s| tree.point_indices_of_slot(s).len())
.sum();
assert_eq!(
physical_count_before, physical_count_after,
"reactivation must not grow any slot's physical point list (no duplicate)"
);
}
#[test]
fn remove_point_returns_true_once_then_false_and_active_count_drops() {
let ds = Ungated(vec![[10.0], [20.0], [30.0]]);
let mut tree = DynamicKdTreeBuilder::new(ConstDim::<1>, ds)
.maximum_point_count(1000)
.build();
tree.add_points(0, 2);
assert_eq!(tree.active_count(), 3);
assert!(tree.remove_point(1), "first removal must return true");
assert_eq!(tree.active_count(), 2);
assert!(!tree.remove_point(1), "repeat removal must return false");
assert_eq!(tree.active_count(), 2, "active_count must not drop again");
assert!(
!tree.remove_point(999),
"out-of-range removal must return false"
);
assert!(
!tree.remove_point(2usize.wrapping_add(1_000_000)),
"wildly out-of-range must return false too"
);
}
#[test]
fn reactivation_after_a_migrating_merge_lands_in_the_migrated_slot() {
let ds = Ungated(vec![[10.0], [20.0], [30.0], [40.0]]);
let mut tree = DynamicKdTreeBuilder::new(ConstDim::<1>, ds)
.maximum_point_count(1000)
.build();
tree.add_points(0, 1);
assert!(tree.remove_point(0));
tree.add_points(2, 2);
tree.add_points(3, 3);
assert_eq!(tree.tree_index()[0], -1);
let slot2_len_before = tree.point_indices_of_slot(2).len();
tree.add_points(0, 0);
assert_eq!(
tree.tree_index()[0],
2,
"must reactivate into the MIGRATED slot, not the original one"
);
assert_eq!(tree.removed_len(), 0);
assert_eq!(tree.point_indices_of_slot(2).len(), slot2_len_before);
assert!(
tree.point_indices_of_slot(1).is_empty(),
"slot 1 (the stale original) must stay empty"
);
}
#[test]
fn identical_op_sequences_produce_identical_bookkeeping() {
let pts: Vec<[f64; 2]> = (0..20)
.map(|i| [i as f64 * 1.3, (i * i) as f64 * 0.7])
.collect();
let run = |pts: &[[f64; 2]]| -> (Vec<Vec<u32>>, Vec<i32>, usize) {
let ds = Ungated(pts.to_vec());
let mut tree = DynamicKdTreeBuilder::new(ConstDim::<2>, ds)
.maximum_point_count(1000)
.build();
tree.add_points(0, 7);
assert!(tree.remove_point(2));
assert!(tree.remove_point(5));
tree.add_points(8, 13);
tree.add_points(2, 2); tree.add_points(14, 19);
let slots: Vec<Vec<u32>> = (0..tree.tree_count())
.map(|s| tree.point_indices_of_slot(s).to_vec())
.collect();
(slots, tree.tree_index().to_vec(), tree.removed_len())
};
let (slots_a, ti_a, removed_a) = run(&pts);
let (slots_b, ti_b, removed_b) = run(&pts);
assert_eq!(
slots_a, slots_b,
"slot vind must be element-wise identical across identical runs"
);
assert_eq!(ti_a, ti_b);
assert_eq!(removed_a, removed_b);
}
#[test]
fn ctor_auto_adds_existing_points_immediately_queryable_state() {
let pts: Vec<[f64; 2]> = (0..5).map(|i| [i as f64, i as f64 * 2.0]).collect();
let tree = DynamicKdTreeBuilder::new(ConstDim::<2>, pts.as_slice())
.maximum_point_count(1000)
.build();
assert_eq!(
tree.active_count(),
5,
"ctor must auto-add every point already in the dataset"
);
let mut union: Vec<u32> = (0..tree.tree_count())
.flat_map(|s| tree.point_indices_of_slot(s).iter().copied())
.collect();
union.sort_unstable();
assert_eq!(union, vec![0, 1, 2, 3, 4]);
for &v in tree.tree_index() {
assert_ne!(v, -1);
}
}
use std::cell::Cell;
use std::rc::Rc;
struct GrowableDataSource {
data: [[f64; 2]; 4],
count: Rc<Cell<usize>>,
}
impl DataSource<f64> for GrowableDataSource {
fn point_count(&self) -> usize {
self.count.get()
}
fn point_component(&self, idx: usize, dim: usize) -> f64 {
self.data[idx][dim]
}
}
#[test]
fn empty_dataset_yields_empty_forest_then_add_points_works_after_growth() {
let count = Rc::new(Cell::new(0));
let ds = GrowableDataSource {
data: [[0.0, 0.0], [1.0, 1.0], [2.0, 2.0], [3.0, 3.0]],
count: count.clone(),
};
let mut tree = DynamicKdTreeBuilder::new(ConstDim::<2>, ds)
.maximum_point_count(1000)
.build();
assert_eq!(
tree.active_count(),
0,
"empty dataset at build time must yield an empty forest"
);
for s in 0..tree.tree_count() {
assert!(tree.point_indices_of_slot(s).is_empty());
}
count.set(4);
tree.add_points(0, 3);
assert_eq!(tree.active_count(), 4);
let mut union: Vec<u32> = (0..tree.tree_count())
.flat_map(|s| tree.point_indices_of_slot(s).iter().copied())
.collect();
union.sort_unstable();
assert_eq!(union, vec![0, 1, 2, 3]);
}
#[test]
fn drain_and_refill_all_live_and_bookkeeping_consistent() {
let pts: Vec<[f64; 1]> = (0..8).map(|i| [i as f64 * 3.3]).collect();
let ds = Ungated(pts);
let mut tree = DynamicKdTreeBuilder::new(ConstDim::<1>, ds)
.maximum_point_count(1000)
.build();
tree.add_points(0, 7);
assert_eq!(tree.active_count(), 8);
for i in 0..8 {
assert!(tree.remove_point(i));
}
assert_eq!(tree.active_count(), 0);
assert_eq!(tree.removed_len(), 8);
tree.add_points(0, 7);
assert_eq!(tree.active_count(), 8);
assert_eq!(tree.removed_len(), 0);
for &v in tree.tree_index() {
assert_ne!(v, -1);
}
let mut union: Vec<u32> = (0..tree.tree_count())
.flat_map(|s| tree.point_indices_of_slot(s).iter().copied())
.collect();
union.sort_unstable();
assert_eq!(union, (0u32..8).collect::<Vec<u32>>());
}
#[cfg(feature = "parallel")]
#[test]
fn add_points_parallel_matches_sequential_slot_contents() {
let pts: Vec<[f64; 2]> = (0..64)
.map(|i| [i as f64 * 1.7, (i * i) as f64 * 0.3])
.collect();
let run_seq = || -> (Vec<Vec<u32>>, Vec<i32>) {
let ds = Ungated(pts.clone());
let mut tree = DynamicKdTreeBuilder::new(ConstDim::<2>, ds)
.maximum_point_count(1000)
.build();
tree.add_points(0, 31);
assert!(tree.remove_point(3));
assert!(tree.remove_point(9));
tree.add_points(32, 47);
tree.add_points(3, 3); tree.add_points(48, 63);
let slots: Vec<Vec<u32>> = (0..tree.tree_count())
.map(|s| tree.point_indices_of_slot(s).to_vec())
.collect();
(slots, tree.tree_index().to_vec())
};
let run_par = || -> (Vec<Vec<u32>>, Vec<i32>) {
let ds = Ungated(pts.clone());
let mut tree = DynamicKdTreeBuilder::new(ConstDim::<2>, ds)
.maximum_point_count(1000)
.build();
tree.add_points_parallel(0, 31);
assert!(tree.remove_point(3));
assert!(tree.remove_point(9));
tree.add_points_parallel(32, 47);
tree.add_points_parallel(3, 3); tree.add_points_parallel(48, 63);
let slots: Vec<Vec<u32>> = (0..tree.tree_count())
.map(|s| tree.point_indices_of_slot(s).to_vec())
.collect();
(slots, tree.tree_index().to_vec())
};
let (seq_slots, seq_ti) = run_seq();
let (par_slots, par_ti) = run_par();
assert_eq!(
seq_slots, par_slots,
"parallel rebuild must produce byte-identical slot vind to the sequential path"
);
assert_eq!(seq_ti, par_ti);
}
#[test]
#[should_panic(expected = "does not equal point_count")]
fn add_points_panics_on_misaligned_genuinely_new_index() {
let ds = Ungated(vec![[10.0], [20.0], [30.0]]);
let mut tree = DynamicKdTreeBuilder::new(ConstDim::<1>, ds)
.maximum_point_count(1000)
.build();
tree.add_points(1, 1);
}
#[test]
#[should_panic(expected = "outgrown this forest's capacity")]
fn add_points_panics_when_point_count_outgrows_maximum_point_count() {
let ds = Ungated(vec![[10.0], [20.0]]);
let mut tree = DynamicKdTreeBuilder::new(ConstDim::<1>, ds)
.maximum_point_count(1)
.build();
assert_eq!(tree.tree_count(), 1);
tree.add_points(0, 1);
}
#[test]
fn maximum_point_count_zero_panics() {
let result = std::panic::catch_unwind(|| {
DynamicKdTreeBuilder::new(ConstDim::<1>, Ungated::<1>(vec![])).maximum_point_count(0)
});
assert!(
result.is_err(),
"maximum_point_count(0) must panic, not silently accept an unusable capacity"
);
}
#[test]
fn reactivation_only_call_is_exempt_from_the_contiguity_assert() {
let ds = Ungated(vec![[10.0], [20.0], [30.0], [40.0]]);
let mut tree = DynamicKdTreeBuilder::new(ConstDim::<1>, ds)
.maximum_point_count(1000)
.build();
tree.add_points(0, 3);
assert!(tree.remove_point(1));
tree.add_points(1, 1); assert_eq!(tree.removed_len(), 0);
}
struct Lcg(u64);
impl Lcg {
fn next_f64(&mut self) -> f64 {
self.0 = self
.0
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
((self.0 >> 11) as f64) / ((1u64 << 53) as f64)
}
}
fn brute_force_knn_live<const N: usize>(
pts: &[[f64; N]],
live: &[bool],
query: &[f64; N],
k: usize,
) -> (Vec<u32>, Vec<f64>) {
let metric = L2;
let mut indices = vec![0u32; k];
let mut dists = vec![0.0f64; k];
let count;
{
let mut rs = KnnResultSet::<f64, u32>::new(&mut indices, &mut dists);
for (i, &is_live) in live.iter().enumerate() {
if !is_live {
continue;
}
let d = metric.eval(query.as_slice(), &pts, i, ConstDim::<N>);
rs.add_point(d, i as u32);
}
count = rs.size();
}
indices.truncate(count);
dists.truncate(count);
(indices, dists)
}
fn brute_force_rknn_live<const N: usize>(
pts: &[[f64; N]],
live: &[bool],
query: &[f64; N],
radius: f64,
k: usize,
) -> (Vec<u32>, Vec<f64>) {
let metric = L2;
let mut scored: Vec<(f64, u32)> = (0..pts.len())
.filter(|&i| live[i])
.map(|i| {
(
metric.eval(query.as_slice(), &pts, i, ConstDim::<N>),
i as u32,
)
})
.filter(|&(d, _)| d < radius)
.collect();
scored.sort_by(|a, b| a.0.partial_cmp(&b.0).unwrap().then(a.1.cmp(&b.1)));
scored.truncate(k);
let indices = scored.iter().map(|&(_, i)| i).collect();
let dists = scored.iter().map(|&(d, _)| d).collect();
(indices, dists)
}
fn brute_force_radius_live<const N: usize>(
pts: &[[f64; N]],
live: &[bool],
query: &[f64; N],
radius: f64,
) -> Vec<(u32, f64)> {
let metric = L2;
let mut out: Vec<(u32, f64)> = (0..pts.len())
.filter(|&i| live[i])
.filter_map(|i| {
let d = metric.eval(query.as_slice(), &pts, i, ConstDim::<N>);
if d < radius {
Some((i as u32, d))
} else {
None
}
})
.collect();
out.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap().then(a.0.cmp(&b.0)));
out
}
#[test]
fn removed_points_never_returned_by_knn_or_radius() {
let pts: Vec<[f64; 2]> = (0..20).map(|i| [i as f64, (i * 2) as f64]).collect();
let pts_slice: &[[f64; 2]] = pts.as_slice();
let mut tree = DynamicKdTreeBuilder::new(ConstDim::<2>, pts_slice)
.maximum_point_count(1000)
.build();
let removed_indices = [1usize, 3, 5, 7, 9, 11, 13];
for &i in &removed_indices {
assert!(tree.remove_point(i));
}
assert_eq!(tree.active_count(), 13);
let mut idx = [0u32; 20];
let mut dist = [0.0f64; 20];
let found = tree.knn_search(&[0.0, 0.0], &mut idx, &mut dist);
assert_eq!(
found, 13,
"knn(k=20) over 13 live points must return exactly 13"
);
for &i in &idx[..found] {
assert!(
!removed_indices.contains(&(i as usize)),
"removed point {i} leaked into knn results"
);
}
let mut radius_out = Vec::new();
let radius_found = tree.radius_search(&[0.0, 0.0], 1_000_000.0, &mut radius_out);
assert_eq!(
radius_found, 13,
"radius over everything must return exactly the 13 live points"
);
for item in &radius_out {
assert!(
!removed_indices.contains(&(item.index as usize)),
"removed point leaked into radius results"
);
}
}
#[test]
fn filtered_brute_force_equality_100_points_dim3_random_removals() {
let mut rng = Lcg(0xF0A_15E7u64);
let n = 100;
let pts: Vec<[f64; 3]> = (0..n)
.map(|_| {
[
rng.next_f64() * 200.0 - 100.0,
rng.next_f64() * 200.0 - 100.0,
rng.next_f64() * 200.0 - 100.0,
]
})
.collect();
let pts_slice: &[[f64; 3]] = pts.as_slice();
let mut tree = DynamicKdTreeBuilder::new(ConstDim::<3>, pts_slice)
.maximum_point_count(1000)
.build();
let mut live = vec![true; n];
for (i, live_i) in live.iter_mut().enumerate() {
if rng.next_f64() < 0.3 {
assert!(tree.remove_point(i));
*live_i = false;
}
}
for _case in 0..15 {
let query = [
rng.next_f64() * 200.0 - 100.0,
rng.next_f64() * 200.0 - 100.0,
rng.next_f64() * 200.0 - 100.0,
];
let k = 10;
let mut idx = vec![0u32; k];
let mut dist = vec![0.0f64; k];
let found = tree.knn_search(&query, &mut idx, &mut dist);
let (want_idx, want_dist) = brute_force_knn_live(pts_slice, &live, &query, k);
assert_eq!(found, want_idx.len());
assert_eq!(&idx[..found], want_idx.as_slice());
assert_eq!(&dist[..found], want_dist.as_slice());
let radius = 5000.0;
let mut ridx = vec![0u32; k];
let mut rdist = vec![0.0f64; k];
let rfound = tree.rknn_search(&query, radius, &mut ridx, &mut rdist);
let (want_ridx, want_rdist) =
brute_force_rknn_live(pts_slice, &live, &query, radius, k);
assert_eq!(rfound, want_ridx.len());
assert_eq!(&ridx[..rfound], want_ridx.as_slice());
assert_eq!(&rdist[..rfound], want_rdist.as_slice());
let mut radius_out = Vec::new();
tree.radius_search(&query, radius, &mut radius_out);
let want_radius = brute_force_radius_live(pts_slice, &live, &query, radius);
let got_radius: Vec<(u32, f64)> = radius_out
.iter()
.map(|it| (it.index, it.distance))
.collect();
assert_eq!(got_radius, want_radius);
}
}
#[test]
fn reactivated_point_is_returned_again_with_same_distance_bits() {
let pts: Vec<[f64; 2]> = vec![[0.0, 0.0], [3.0, 4.0], [10.0, 10.0], [-5.0, -5.0]];
let pts_slice: &[[f64; 2]] = pts.as_slice();
let mut tree = DynamicKdTreeBuilder::new(ConstDim::<2>, pts_slice)
.maximum_point_count(1000)
.build();
let query = [0.0, 0.0];
let expected_dist = L2.eval(query.as_slice(), &pts_slice, 1, ConstDim::<2>);
assert!(tree.remove_point(1));
let mut idx = [0u32; 4];
let mut dist = [0.0f64; 4];
let found = tree.knn_search(&query, &mut idx, &mut dist);
assert!(
!idx[..found].contains(&1),
"removed point must not be found by query"
);
tree.add_points(1, 1);
let mut idx2 = [0u32; 4];
let mut dist2 = [0.0f64; 4];
let found2 = tree.knn_search(&query, &mut idx2, &mut dist2);
let pos = idx2[..found2]
.iter()
.position(|&i| i == 1)
.expect("reactivated point must be found again");
assert_eq!(
dist2[pos], expected_dist,
"reactivated point's distance must be bit-identical to its pre-removal distance"
);
}
#[test]
fn empty_forest_quirk_true_for_radius_false_for_knn_zero_for_radius_search_wrapper() {
let pts: Vec<[f64; 2]> = Vec::new();
let pts_slice: &[[f64; 2]] = pts.as_slice();
let tree = DynamicKdTreeBuilder::new(ConstDim::<2>, pts_slice)
.maximum_point_count(1000)
.build();
assert_eq!(tree.active_count(), 0);
let mut idx = [0u32; 3];
let mut dist = [0.0f64; 3];
let mut knn_rs = KnnResultSet::<f64, u32>::new(&mut idx, &mut dist);
let full = tree.find_neighbors(&mut knn_rs, &[0.0, 0.0], &SearchParams::default());
assert!(
!full,
"KNN find_neighbors on an empty forest must return false"
);
assert_eq!(knn_rs.size(), 0);
let mut items = Vec::new();
let mut radius_rs = RadiusResultSet::new(100.0, &mut items);
let radius_full =
tree.find_neighbors(&mut radius_rs, &[0.0, 0.0], &SearchParams::default());
assert!(
radius_full,
"RadiusResultSet::full() is hardwired true -- the quirk"
);
assert_eq!(items.len(), 0);
let mut out = Vec::new();
let count = tree.radius_search(&[0.0, 0.0], 100.0, &mut out);
assert_eq!(
count, 0,
"the additive radius_search wrapper returns the COUNT, hiding the quirk"
);
}
#[test]
fn multi_slot_correctness_spanning_two_slots_then_merged() {
let pts: Vec<[f64; 2]> = vec![[0.0, 0.0], [5.0, 5.0], [-3.0, 2.0], [8.0, -1.0]];
let pts_slice: &[[f64; 2]] = pts.as_slice();
let ds = Ungated(pts.clone());
let mut tree = DynamicKdTreeBuilder::new(ConstDim::<2>, ds)
.maximum_point_count(1000)
.build();
tree.add_points(0, 2);
assert!(
!tree.point_indices_of_slot(0).is_empty(),
"slot 0 must be occupied"
);
assert!(
!tree.point_indices_of_slot(1).is_empty(),
"slot 1 must be occupied"
);
assert!(tree.point_indices_of_slot(2).is_empty());
let query = [1.0, 1.0];
let k = 3;
let mut idx = vec![0u32; k];
let mut dist = vec![0.0f64; k];
let found = tree.knn_search(&query, &mut idx, &mut dist);
let (want_idx, want_dist) =
brute_force_knn_live(pts_slice, &[true, true, true, false], &query, k);
assert_eq!(found, 3);
assert_eq!(idx, want_idx.as_slice());
assert_eq!(dist, want_dist.as_slice());
tree.add_points(3, 3);
assert!(tree.point_indices_of_slot(0).is_empty());
assert!(tree.point_indices_of_slot(1).is_empty());
assert!(
!tree.point_indices_of_slot(2).is_empty(),
"everything merged into slot 2"
);
let k2 = 4;
let mut idx2 = vec![0u32; k2];
let mut dist2 = vec![0.0f64; k2];
let found2 = tree.knn_search(&query, &mut idx2, &mut dist2);
let (want_idx2, want_dist2) = brute_force_knn_live(pts_slice, &[true; 4], &query, k2);
assert_eq!(found2, 4);
assert_eq!(idx2, want_idx2.as_slice());
assert_eq!(dist2, want_dist2.as_slice());
}
#[test]
fn eps_plumbing_zero_matches_brute_force_ten_returns_valid_member() {
let pts: Vec<[f64; 2]> = vec![[0.0, 0.0], [5.0, 5.0], [-3.0, 2.0]];
let pts_slice: &[[f64; 2]] = pts.as_slice();
let tree = DynamicKdTreeBuilder::new(ConstDim::<2>, pts_slice)
.maximum_point_count(1000)
.build();
assert!(!tree.point_indices_of_slot(0).is_empty());
assert!(!tree.point_indices_of_slot(1).is_empty());
let query = [1.0, 1.0];
let k = 1;
let params_zero = SearchParams {
eps: 0.0,
sorted: true,
};
let mut idx0 = [0u32; 1];
let mut dist0 = [0.0f64; 1];
tree.knn_search_with(&query, &mut idx0, &mut dist0, ¶ms_zero);
let (want_idx, want_dist) = brute_force_knn_live(pts_slice, &[true; 3], &query, k);
assert_eq!(
idx0.as_slice(),
want_idx.as_slice(),
"eps=0 must match brute force exactly"
);
assert_eq!(dist0.as_slice(), want_dist.as_slice());
let params_eps = SearchParams {
eps: 10.0,
sorted: true,
};
let mut idx_eps = [0u32; 1];
let mut dist_eps = [0.0f64; 1];
let found = tree.knn_search_with(&query, &mut idx_eps, &mut dist_eps, ¶ms_eps);
assert_eq!(found, 1);
assert!(
(idx_eps[0] as usize) < pts.len(),
"approximate result must still be a member of the live set"
);
}
#[test]
fn sorted_radius_search_ascending_across_slots() {
let pts: Vec<[f64; 2]> = vec![[0.0, 0.0], [5.0, 5.0], [-3.0, 2.0], [8.0, -1.0]];
let ds = Ungated(pts.clone());
let mut tree = DynamicKdTreeBuilder::new(ConstDim::<2>, ds)
.maximum_point_count(1000)
.build();
tree.add_points(0, 2); assert!(!tree.point_indices_of_slot(0).is_empty());
assert!(!tree.point_indices_of_slot(1).is_empty());
let mut out = Vec::new();
let params = SearchParams {
eps: 0.0,
sorted: true,
};
tree.radius_search_with(&[1.0, 1.0], 1000.0, &mut out, ¶ms);
let dists: Vec<f64> = out.iter().map(|it| it.distance).collect();
let mut sorted = dists.clone();
sorted.sort_by(|a, b| a.partial_cmp(b).unwrap());
assert_eq!(
dists, sorted,
"sorted=true must yield ascending distances across multiple slot passes"
);
assert!(
out.len() >= 2,
"test must actually exercise multiple slots' worth of results"
);
}
#[test]
fn dynamic_over_owned_rows_matches_fresh_static_tree_over_live_rows() {
use crate::data_source::{FlatSlice, OwnedRows};
use crate::tree::KdTreeBuilder;
const DIM: usize = 2;
let batch1: [f64; 6] = [0.0, 0.0, 10.0, 0.0, 0.0, 10.0];
let batch2: [f64; 4] = [5.0, 5.0, -5.0, 3.0];
let batch3: [f64; 4] = [2.0, -7.0, 9.0, 9.0];
let mut tree =
DynamicKdTreeBuilder::new(ConstDim::<DIM>, OwnedRows::<f64>::with_capacity(DIM, 0))
.maximum_point_count(1000)
.build();
tree.dataset_mut().push_rows(&batch1);
tree.add_points(0, 2);
tree.dataset_mut().push_rows(&batch2);
tree.add_points(3, 4);
tree.dataset_mut().push_rows(&batch3);
tree.add_points(5, 6);
assert!(tree.remove_point(6), "point 6 must be live before removal");
assert_eq!(tree.active_count(), 6);
let query = [1.0, 1.0];
let k = 3;
let mut idx = [0u32; 3];
let mut dist = [0.0f64; 3];
let found = tree.knn_search(&query, &mut idx, &mut dist);
assert_eq!(found, k);
let live_rows = &tree.dataset().as_slice()[..6 * DIM];
let fresh = KdTreeBuilder::new(ConstDim::<DIM>, FlatSlice::new(live_rows, DIM)).build();
let mut want_idx = [0u32; 3];
let mut want_dist = [0.0f64; 3];
let want_found = fresh.knn_search(&query, &mut want_idx, &mut want_dist);
assert_eq!(want_found, k);
let mut got: Vec<(u32, f64)> = idx.iter().copied().zip(dist.iter().copied()).collect();
let mut want: Vec<(u32, f64)> = want_idx
.iter()
.copied()
.zip(want_dist.iter().copied())
.collect();
got.sort_by_key(|&(i, _)| i);
want.sort_by_key(|&(i, _)| i);
assert_eq!(
got, want,
"dynamic-over-OwnedRows knn must match a fresh static tree over the same live rows"
);
}
}