use std::cmp::Ordering;
#[cfg(feature = "serde")]
use serde::{Deserialize, Serialize};
use tracing::info;
use crate::{errors::SpartError, geometry::DistanceMetric, knn::KnnHeap};
pub trait KdPoint: Clone + PartialEq + std::fmt::Debug {
fn dims(&self) -> usize;
fn coord(&self, axis: usize) -> Result<f64, SpartError>;
}
impl<T> KdPoint for crate::geometry::Point2D<T>
where
T: std::fmt::Debug + Clone + PartialEq,
{
fn dims(&self) -> usize {
2
}
fn coord(&self, axis: usize) -> Result<f64, SpartError> {
match axis {
0 => Ok(self.x),
1 => Ok(self.y),
_ => Err(SpartError::InvalidDimension {
requested: axis,
available: 2,
}),
}
}
}
impl<T> KdPoint for crate::geometry::Point3D<T>
where
T: std::fmt::Debug + Clone + PartialEq,
{
fn dims(&self) -> usize {
3
}
fn coord(&self, axis: usize) -> Result<f64, SpartError> {
match axis {
0 => Ok(self.x),
1 => Ok(self.y),
2 => Ok(self.z),
_ => Err(SpartError::InvalidDimension {
requested: axis,
available: 3,
}),
}
}
}
const REBUILD_RATIO_NUM: usize = 2;
const REBUILD_RATIO_DEN: usize = 3;
const REBUILD_MIN_SIZE: usize = 8;
#[derive(Debug, Clone)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
struct KdNode<P: KdPoint> {
point: P,
left: Option<Box<KdNode<P>>>,
right: Option<Box<KdNode<P>>>,
size: usize,
}
impl<P: KdPoint> KdNode<P> {
fn new(point: P) -> Self {
KdNode {
point,
left: None,
right: None,
size: 1,
}
}
fn child_size(child: &Option<Box<KdNode<P>>>) -> usize {
child.as_ref().map_or(0, |node| node.size)
}
fn update_size(&mut self) {
self.size = 1 + Self::child_size(&self.left) + Self::child_size(&self.right);
}
fn is_unbalanced(&self) -> bool {
if self.size < REBUILD_MIN_SIZE {
return false;
}
let heaviest = Self::child_size(&self.left).max(Self::child_size(&self.right));
heaviest * REBUILD_RATIO_DEN > self.size * REBUILD_RATIO_NUM
}
}
#[derive(Debug, Clone)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub struct KdTree<P: KdPoint> {
root: Option<Box<KdNode<P>>>,
k: Option<usize>,
}
impl<P: KdPoint> Default for KdTree<P> {
fn default() -> Self {
Self::new()
}
}
impl<P: KdPoint> KdTree<P> {
pub fn new() -> Self {
KdTree {
root: None,
k: None,
}
}
pub fn with_dimension(k: usize) -> Self {
KdTree {
root: None,
k: Some(k),
}
}
pub fn len(&self) -> usize {
self.root.as_ref().map_or(0, |node| node.size)
}
pub fn is_empty(&self) -> bool {
self.root.is_none()
}
pub fn clear(&mut self) {
self.root = None;
}
fn bbox_search_rec<'a>(
node: &'a Option<Box<KdNode<P>>>,
lo: &[f64],
hi: &[f64],
depth: usize,
found: &mut Vec<&'a P>,
) {
let Some(n) = node else {
return;
};
let axes = lo.len();
let inside = (0..axes).all(|axis| {
let c = n.point.coord(axis).unwrap_or(f64::NAN);
lo[axis] <= c && c <= hi[axis]
});
if inside {
found.push(&n.point);
}
let axis = depth % axes;
let node_coord = n.point.coord(axis).unwrap_or(f64::NAN);
if lo[axis] <= node_coord {
Self::bbox_search_rec(&n.left, lo, hi, depth + 1, found);
}
if node_coord <= hi[axis] {
Self::bbox_search_rec(&n.right, lo, hi, depth + 1, found);
}
}
pub fn contains(&self, point: &P) -> bool {
let k = match self.k {
Some(k) => k,
None => return false,
};
Self::contains_rec(&self.root, point, 0, k)
}
fn contains_rec(node: &Option<Box<KdNode<P>>>, point: &P, depth: usize, k: usize) -> bool {
match node {
None => false,
Some(n) => {
if n.point == *point {
return true;
}
let axis = depth % k;
let p_coord = point
.coord(axis)
.unwrap_or_else(|_| unreachable!("axis computed from dims, must be valid"));
let c_coord = n
.point
.coord(axis)
.unwrap_or_else(|_| unreachable!("axis computed from dims, must be valid"));
if p_coord < c_coord {
Self::contains_rec(&n.left, point, depth + 1, k)
} else if p_coord > c_coord {
Self::contains_rec(&n.right, point, depth + 1, k)
} else {
Self::contains_rec(&n.right, point, depth + 1, k)
|| Self::contains_rec(&n.left, point, depth + 1, k)
}
}
}
}
pub fn insert(&mut self, point: P) -> Result<(), SpartError> {
let k = match self.k {
Some(k) => {
if point.dims() != k {
return Err(SpartError::DimensionMismatch {
expected: k,
actual: point.dims(),
});
}
k
}
None => {
let k = point.dims();
self.k = Some(k);
k
}
};
info!("Inserting point: {:?}", point);
self.root = Some(Self::insert_rec(self.root.take(), point.clone(), 0, k));
Self::rebalance_along_path(&mut self.root, &point, k);
Ok(())
}
fn rebalance_along_path(root: &mut Option<Box<KdNode<P>>>, point: &P, k: usize) {
let mut target_depth = None;
let mut current = &*root;
let mut depth = 0;
while let Some(node) = current {
if node.is_unbalanced() {
target_depth = Some(depth);
break;
}
current = if Self::goes_left(point, &node.point, depth % k) {
&node.left
} else {
&node.right
};
depth += 1;
}
let Some(target_depth) = target_depth else {
return;
};
let mut current = root;
for depth in 0..target_depth {
let Some(node) = current.as_mut() else {
return;
};
current = if Self::goes_left(point, &node.point, depth % k) {
&mut node.left
} else {
&mut node.right
};
}
let mut points = Vec::new();
Self::collect_points(current, &mut points);
*current = Self::insert_bulk_rec(&mut points, target_depth, k);
}
fn goes_left(point: &P, pivot: &P, axis: usize) -> bool {
let p = point.coord(axis).unwrap_or(f64::NAN);
let c = pivot.coord(axis).unwrap_or(f64::NAN);
p < c
}
pub fn insert_bulk(&mut self, mut points: Vec<P>) -> Result<(), SpartError> {
if points.is_empty() {
return Ok(());
}
let k = self.k.unwrap_or_else(|| points[0].dims());
for p in &points {
if p.dims() != k {
return Err(SpartError::DimensionMismatch {
expected: k,
actual: p.dims(),
});
}
}
self.k = Some(k);
if self.root.is_some() {
let mut existing = Vec::new();
Self::collect_points(&self.root, &mut existing);
points.extend(existing);
}
self.root = Self::insert_bulk_rec(&mut points[..], 0, k);
Ok(())
}
fn collect_points(node: &Option<Box<KdNode<P>>>, result: &mut Vec<P>) {
if let Some(n) = node {
result.push(n.point.clone());
Self::collect_points(&n.left, result);
Self::collect_points(&n.right, result);
}
}
fn insert_bulk_rec(points: &mut [P], depth: usize, k: usize) -> Option<Box<KdNode<P>>> {
if points.is_empty() {
return None;
}
let axis = depth % k;
let median_idx = points.len() / 2;
points.select_nth_unstable_by(median_idx, |a, b| {
let ac = a.coord(axis).unwrap_or(f64::NAN);
let bc = b.coord(axis).unwrap_or(f64::NAN);
ac.partial_cmp(&bc).unwrap_or(Ordering::Equal)
});
let mut node = KdNode::new(points[median_idx].clone());
let (left_slice, right_slice) = points.split_at_mut(median_idx);
let right_slice = &mut right_slice[1..];
node.left = Self::insert_bulk_rec(left_slice, depth + 1, k);
node.right = Self::insert_bulk_rec(right_slice, depth + 1, k);
node.update_size();
Some(Box::new(node))
}
fn insert_rec(
node: Option<Box<KdNode<P>>>,
point: P,
depth: usize,
k: usize,
) -> Box<KdNode<P>> {
if let Some(mut current) = node {
if Self::goes_left(&point, ¤t.point, depth % k) {
current.left = Some(Self::insert_rec(current.left.take(), point, depth + 1, k));
} else {
current.right = Some(Self::insert_rec(current.right.take(), point, depth + 1, k));
}
current.update_size();
current
} else {
Box::new(KdNode::new(point))
}
}
pub fn knn_search<'a, M: DistanceMetric<P>>(
&'a self,
target: &P,
k_neighbors: usize,
) -> Vec<&'a P> {
if k_neighbors == 0 {
return Vec::new();
}
let k = match self.k {
Some(k) => k,
None => return Vec::new(),
};
if target.dims() != k {
return Vec::new();
}
info!(
"Performing k\u{2011}NN search for target {:?} with k={}",
target, k_neighbors
);
let mut heap = KnnHeap::new(k_neighbors);
Self::knn_search_rec::<M>(&self.root, target, 0, k, &mut heap);
heap.into_sorted_vec()
}
fn knn_search_rec<'a, M: DistanceMetric<P>>(
node: &'a Option<Box<KdNode<P>>>,
target: &P,
depth: usize,
k: usize,
heap: &mut KnnHeap<&'a P>,
) {
let Some(n) = node else {
return;
};
heap.offer(M::distance_sq(target, &n.point), &n.point);
let axis = depth % k;
let target_coord = target.coord(axis).unwrap_or(f64::NAN);
let node_coord = n.point.coord(axis).unwrap_or(f64::NAN);
let (near, far) = if target_coord < node_coord {
(&n.left, &n.right)
} else {
(&n.right, &n.left)
};
Self::knn_search_rec::<M>(near, target, depth + 1, k, heap);
let plane_distance = target_coord - node_coord;
if plane_distance * plane_distance < heap.worst() {
Self::knn_search_rec::<M>(far, target, depth + 1, k, heap);
}
}
pub fn range_search<'a, M: DistanceMetric<P>>(&'a self, center: &P, radius: f64) -> Vec<&'a P> {
info!("Finding points within radius {} of {:?}", radius, center);
if radius < 0.0 {
return Vec::new();
}
let k = match self.k {
Some(k) => k,
None => return Vec::new(),
};
if center.dims() != k {
return Vec::new();
}
let mut found = Vec::new();
let radius_sq = radius * radius;
Self::range_search_rec::<M>(&self.root, center, radius_sq, 0, radius, &mut found);
found
}
fn range_search_rec<'a, M: DistanceMetric<P>>(
node: &'a Option<Box<KdNode<P>>>,
center: &P,
radius_sq: f64,
depth: usize,
radius: f64,
found: &mut Vec<&'a P>,
) {
if let Some(n) = node {
let dist_sq = M::distance_sq(center, &n.point);
if dist_sq <= radius_sq {
found.push(&n.point);
}
let axis = depth % center.dims();
let center_coord = center
.coord(axis)
.unwrap_or_else(|_| unreachable!("axis computed from dims, must be valid"));
let node_coord = n
.point
.coord(axis)
.unwrap_or_else(|_| unreachable!("axis computed from dims, must be valid"));
if center_coord - radius <= node_coord {
Self::range_search_rec::<M>(&n.left, center, radius_sq, depth + 1, radius, found);
}
if center_coord + radius >= node_coord {
Self::range_search_rec::<M>(&n.right, center, radius_sq, depth + 1, radius, found);
}
}
}
pub fn delete(&mut self, point: &P) -> bool {
if self.root.is_none() {
return false;
}
info!("Attempting to delete point: {:?}", point);
let k = match self.k {
Some(k) => k,
None => return false,
};
let (new_root, deleted) = Self::delete_rec(self.root.take(), point, 0, k);
self.root = new_root;
if deleted {
Self::rebalance_along_path(&mut self.root, point, k);
}
deleted
}
fn delete_rec(
node: Option<Box<KdNode<P>>>,
point: &P,
depth: usize,
k: usize,
) -> (Option<Box<KdNode<P>>>, bool) {
match node {
None => (None, false),
Some(mut current) => {
let axis = depth % k;
if current.point == *point {
if let Some(right_subtree) = current.right.take() {
let successor = Self::find_min(&right_subtree, axis, depth + 1, k).clone();
let (new_right, _) =
Self::delete_rec(Some(right_subtree), &successor, depth + 1, k);
current.point = successor;
current.right = new_right;
current.update_size();
(Some(current), true)
} else if let Some(left_subtree) = current.left.take() {
let successor = Self::find_min(&left_subtree, axis, depth + 1, k).clone();
let (mut new_left, _) =
Self::delete_rec(Some(left_subtree), &successor, depth + 1, k);
current.point = successor;
current.right = new_left.take();
current.left = None;
current.update_size();
(Some(current), true)
} else {
(None, true)
}
} else {
let p_coord = point
.coord(axis)
.unwrap_or_else(|_| unreachable!("axis computed from dims, must be valid"));
let c_coord = current
.point
.coord(axis)
.unwrap_or_else(|_| unreachable!("axis computed from dims, must be valid"));
if p_coord < c_coord {
let (new_left, deleted) =
Self::delete_rec(current.left.take(), point, depth + 1, k);
current.left = new_left;
current.update_size();
(Some(current), deleted)
} else if p_coord > c_coord {
let (new_right, deleted) =
Self::delete_rec(current.right.take(), point, depth + 1, k);
current.right = new_right;
current.update_size();
(Some(current), deleted)
} else {
let (new_right, deleted_right) =
Self::delete_rec(current.right.take(), point, depth + 1, k);
current.right = new_right;
if deleted_right {
current.update_size();
(Some(current), true)
} else {
let (new_left, deleted_left) =
Self::delete_rec(current.left.take(), point, depth + 1, k);
current.left = new_left;
current.update_size();
(Some(current), deleted_left)
}
}
}
}
}
}
fn find_min(node: &KdNode<P>, d: usize, depth: usize, k: usize) -> &P {
let axis = depth % k;
let mut min = &node.point;
if axis == d {
if let Some(ref left) = node.left {
let left_min = Self::find_min(left, d, depth + 1, k);
let left_c = left_min
.coord(d)
.unwrap_or_else(|_| unreachable!("axis computed from dims, must be valid"));
let min_c = min
.coord(d)
.unwrap_or_else(|_| unreachable!("axis computed from dims, must be valid"));
if left_c < min_c {
min = left_min;
}
}
} else {
if let Some(ref left) = node.left {
let left_min = Self::find_min(left, d, depth + 1, k);
let left_c = left_min
.coord(d)
.unwrap_or_else(|_| unreachable!("axis computed from dims, must be valid"));
let min_c = min
.coord(d)
.unwrap_or_else(|_| unreachable!("axis computed from dims, must be valid"));
if left_c < min_c {
min = left_min;
}
}
if let Some(ref right) = node.right {
let right_min = Self::find_min(right, d, depth + 1, k);
let right_c = right_min
.coord(d)
.unwrap_or_else(|_| unreachable!("axis computed from dims, must be valid"));
let min_c = min
.coord(d)
.unwrap_or_else(|_| unreachable!("axis computed from dims, must be valid"));
if right_c < min_c {
min = right_min;
}
}
}
min
}
}
impl<T: std::fmt::Debug + Clone + PartialEq> KdTree<crate::geometry::Point2D<T>> {
pub fn range_search_bbox(
&self,
query: &crate::geometry::Rectangle,
) -> Vec<&crate::geometry::Point2D<T>> {
let lo = [query.x, query.y];
let hi = [query.x + query.width, query.y + query.height];
if self.k.is_some_and(|k| k != lo.len()) {
return Vec::new();
}
let mut found = Vec::new();
Self::bbox_search_rec(&self.root, &lo, &hi, 0, &mut found);
found
}
}
impl<T: std::fmt::Debug + Clone + PartialEq> KdTree<crate::geometry::Point3D<T>> {
pub fn range_search_bbox(
&self,
query: &crate::geometry::Cube,
) -> Vec<&crate::geometry::Point3D<T>> {
let lo = [query.x, query.y, query.z];
let hi = [
query.x + query.width,
query.y + query.height,
query.z + query.depth,
];
if self.k.is_some_and(|k| k != lo.len()) {
return Vec::new();
}
let mut found = Vec::new();
Self::bbox_search_rec(&self.root, &lo, &hi, 0, &mut found);
found
}
}
macro_rules! impl_kdtree_spatial_index {
($point:ident, $volume:ident) => {
impl<T: std::fmt::Debug + Clone + PartialEq> crate::index::SpatialIndex
for KdTree<crate::geometry::$point<T>>
{
type Item = crate::geometry::$point<T>;
type Volume = crate::geometry::$volume;
fn len(&self) -> usize {
KdTree::len(self)
}
fn clear(&mut self) {
KdTree::clear(self);
}
fn contains(&self, item: &Self::Item) -> bool {
KdTree::contains(self, item)
}
fn insert(&mut self, item: Self::Item) -> Result<bool, SpartError> {
KdTree::insert(self, item).map(|()| true)
}
fn insert_bulk(&mut self, items: Vec<Self::Item>) -> Result<usize, SpartError> {
let count = items.len();
KdTree::insert_bulk(self, items).map(|()| count)
}
fn delete(&mut self, item: &Self::Item) -> bool {
KdTree::delete(self, item)
}
fn knn_search<M: DistanceMetric<Self::Item>>(
&self,
query: &Self::Item,
k: usize,
) -> Vec<&Self::Item> {
KdTree::knn_search::<M>(self, query, k)
}
fn range_search<M: DistanceMetric<Self::Item>>(
&self,
query: &Self::Item,
radius: f64,
) -> Vec<&Self::Item> {
KdTree::range_search::<M>(self, query, radius)
}
fn range_search_bbox(&self, query: &Self::Volume) -> Vec<&Self::Item> {
<KdTree<crate::geometry::$point<T>>>::range_search_bbox(self, query)
}
}
};
}
impl_kdtree_spatial_index!(Point2D, Rectangle);
impl_kdtree_spatial_index!(Point3D, Cube);
#[cfg(test)]
mod tests {
use super::*;
use crate::geometry::{EuclideanDistance, Point2D, Point3D};
#[test]
fn test_insert_bulk_consecutive_preserves_points() {
let mut tree: KdTree<Point2D<&str>> = KdTree::new();
let first = vec![
Point2D::new(1.0, 1.0, Some("A")),
Point2D::new(2.0, 2.0, Some("B")),
];
let second = vec![
Point2D::new(3.0, 3.0, Some("C")),
Point2D::new(4.0, 4.0, Some("D")),
];
tree.insert_bulk(first.clone()).unwrap();
tree.insert_bulk(second.clone()).unwrap();
for p in first.into_iter().chain(second) {
assert!(tree.contains(&p));
}
let target = Point2D::new(2.5, 2.5, None::<&str>);
let knn = tree.knn_search::<EuclideanDistance>(&target, 4);
assert_eq!(knn.len(), 4);
}
#[test]
fn test_insert_bulk_dimension_mismatch() {
let mut tree: KdTree<Point2D<()>> = KdTree::with_dimension(3);
let points = vec![Point2D::new(1.0, 2.0, None)];
let result = tree.insert_bulk(points);
assert!(matches!(
result,
Err(SpartError::DimensionMismatch {
expected: 3,
actual: 2
})
));
}
#[test]
fn test_dimension_inference() {
let mut tree: KdTree<Point2D<()>> = KdTree::new();
let p = Point2D::new(1.0, 2.0, None);
tree.insert(p).unwrap();
let p2 = Point2D::new(3.0, 4.0, None);
assert!(tree.insert(p2).is_ok());
}
#[test]
fn test_empty_tree_queries() {
let mut tree: KdTree<Point2D<&str>> = KdTree::new();
let target = Point2D::new(1.0, 2.0, None::<&str>);
let knn_results = tree.knn_search::<EuclideanDistance>(&target, 5);
assert!(knn_results.is_empty());
let range_results = tree.range_search::<EuclideanDistance>(&target, 10.0);
assert!(range_results.is_empty());
assert!(!tree.delete(&target));
}
#[test]
fn test_knn_edge_cases() {
let mut tree: KdTree<Point2D<&str>> = KdTree::new();
let points = vec![
Point2D::new(0.0, 0.0, Some("A")),
Point2D::new(1.0, 1.0, Some("B")),
Point2D::new(2.0, 2.0, Some("C")),
];
let num_points = points.len();
tree.insert_bulk(points).unwrap();
let target = Point2D::new(0.5, 0.5, None::<&str>);
let knn_results = tree.knn_search::<EuclideanDistance>(&target, 0);
assert!(knn_results.is_empty());
let knn_results = tree.knn_search::<EuclideanDistance>(&target, num_points + 5);
assert_eq!(knn_results.len(), num_points);
}
#[test]
fn test_range_search_negative_radius_empty() {
let mut tree: KdTree<Point2D<&str>> = KdTree::new();
let target = Point2D::new(5.0, 5.0, Some("T"));
tree.insert(target.clone()).unwrap();
tree.insert(Point2D::new(5.2, 5.0, Some("N"))).unwrap();
assert!(
tree.range_search::<EuclideanDistance>(&target, -1.0)
.is_empty()
);
}
#[test]
fn test_range_zero_radius_exact_match() {
let mut tree: KdTree<Point2D<&str>> = KdTree::new();
let target = Point2D::new(10.0, 10.0, Some("A"));
tree.insert(target.clone()).unwrap();
tree.insert(Point2D::new(11.0, 11.0, Some("B"))).unwrap();
let results = tree.range_search::<EuclideanDistance>(&target, 0.0);
assert_eq!(results.len(), 1);
assert_eq!(*results[0], target);
}
#[test]
fn test_duplicates_delete_one() {
let mut tree: KdTree<Point2D<&str>> = KdTree::new();
let p1 = Point2D::new(10.0, 10.0, Some("A"));
let p2 = Point2D::new(10.0, 10.0, Some("A"));
tree.insert(p1.clone()).unwrap();
tree.insert(p2.clone()).unwrap();
let target = Point2D::new(10.0, 10.0, None::<&str>);
let results = tree.knn_search::<EuclideanDistance>(&target, 2);
assert_eq!(results.len(), 2);
assert!(tree.delete(&p1));
let results_after_delete = tree.knn_search::<EuclideanDistance>(&target, 2);
assert_eq!(results_after_delete.len(), 1);
}
#[test]
fn test_delete_many() {
let mut tree: KdTree<Point2D<&str>> = KdTree::new();
let points = [
Point2D::new(1.0, 2.0, Some("A")),
Point2D::new(3.0, 4.0, Some("B")),
Point2D::new(-1.0, -2.0, Some("C")),
Point2D::new(1.5, 3.2, Some("D")),
Point2D::new(0.5, 2.0, Some("E")),
Point2D::new(0.25, 2.0, Some("F")),
Point2D::new(0.5, 1.0, Some("G")),
];
for p in points.clone() {
tree.insert(p).unwrap();
}
for p in &points {
assert!(tree.delete(p));
let knn_after = tree.knn_search::<EuclideanDistance>(p, 2);
for pt in &knn_after {
assert_ne!(pt.data, p.data);
}
}
}
#[test]
fn test_delete_same_coords_different_data() {
let mut tree: KdTree<Point2D<&str>> = KdTree::new();
let p1 = Point2D::new(10.0, 10.0, Some("A"));
let p2 = Point2D::new(10.0, 10.0, Some("B"));
let p3 = Point2D::new(10.0, 10.0, Some("C"));
tree.insert(p1.clone()).unwrap();
tree.insert(p2.clone()).unwrap();
tree.insert(p3.clone()).unwrap();
assert!(tree.delete(&p2));
assert!(tree.contains(&p1));
assert!(tree.contains(&p3));
assert!(!tree.contains(&p2));
let tgt = Point2D::new(10.0, 10.0, None::<&str>);
let res = tree.knn_search::<EuclideanDistance>(&tgt, 3);
assert_eq!(res.len(), 2);
for r in res {
assert_ne!(r.data, Some("B"));
}
}
#[test]
fn test_delete_nonexistent_with_equal_axis() {
let mut tree: KdTree<Point2D<&str>> = KdTree::new();
let a = Point2D::new(1.0, 0.0, Some("A"));
let b = Point2D::new(1.0, 1.0, Some("B"));
let c = Point2D::new(1.0, -1.0, Some("C"));
tree.insert(a.clone()).unwrap();
tree.insert(b.clone()).unwrap();
tree.insert(c.clone()).unwrap();
let not_present = Point2D::new(1.0, 2.0, Some("X"));
assert!(!tree.delete(¬_present));
assert!(tree.contains(&a));
assert!(tree.contains(&b));
assert!(tree.contains(&c));
}
#[test]
fn test_delete_root_with_only_left() {
let mut tree: KdTree<Point2D<&str>> = KdTree::new();
let root = Point2D::new(5.0, 5.0, Some("R"));
let l1 = Point2D::new(2.0, 2.0, Some("L1"));
let l2 = Point2D::new(1.0, 1.0, Some("L2"));
tree.insert(root.clone()).unwrap();
tree.insert(l1.clone()).unwrap();
tree.insert(l2.clone()).unwrap();
assert!(tree.delete(&root));
assert!(!tree.contains(&root));
assert!(tree.contains(&l1));
assert!(tree.contains(&l2));
assert!(tree.delete(&l1));
assert!(tree.contains(&l2));
}
#[test]
fn test_delete_all_and_reinsert() {
let mut tree: KdTree<Point2D<&str>> = KdTree::new();
let pts = [
Point2D::new(0.0, 0.0, Some("A")),
Point2D::new(1.0, 1.0, Some("B")),
Point2D::new(-1.0, -1.0, Some("C")),
];
for p in pts.iter().cloned() {
tree.insert(p).unwrap();
}
for p in &pts {
assert!(tree.delete(p));
}
for p in &pts {
assert!(!tree.delete(p));
}
let new_pts = [
Point2D::new(2.0, 2.0, Some("D")),
Point2D::new(3.0, 3.0, Some("E")),
];
for p in new_pts.iter().cloned() {
tree.insert(p).unwrap();
}
let tgt = Point2D::new(2.1, 2.1, None::<&str>);
let res = tree.knn_search::<EuclideanDistance>(&tgt, 2);
assert_eq!(res.len(), 2);
}
#[test]
fn test_delete_many_equal_on_axis() {
let mut tree: KdTree<Point2D<&str>> = KdTree::new();
let pts = [
Point2D::new(0.0, 0.0, Some("A")),
Point2D::new(0.0, 1.0, Some("B")),
Point2D::new(0.0, 2.0, Some("C")),
Point2D::new(0.0, 3.0, Some("D")),
Point2D::new(0.0, -1.0, Some("E")),
];
for p in pts.iter().cloned() {
tree.insert(p).unwrap();
}
for p in &pts {
assert!(tree.delete(p));
assert!(!tree.contains(p));
}
let tgt = Point2D::new(0.0, 0.0, None::<&str>);
let res = tree.knn_search::<EuclideanDistance>(&tgt, 1);
assert!(res.is_empty());
}
#[test]
fn test_insert_bulk_3d_smoke() {
let mut tree: KdTree<Point3D<&str>> = KdTree::new();
let points = vec![
Point3D::new(1.0, 2.0, 3.0, Some("A")),
Point3D::new(4.0, 5.0, 6.0, Some("B")),
];
tree.insert_bulk(points).unwrap();
let target = Point3D::new(2.0, 3.0, 4.0, None::<&str>);
let results = tree.knn_search::<EuclideanDistance>(&target, 1);
assert_eq!(results.len(), 1);
}
#[test]
fn test_insert_bulk_empty_is_ok() {
let mut tree: KdTree<Point2D<i32>> = KdTree::new();
let result = tree.insert_bulk(Vec::new());
assert!(result.is_ok());
}
#[test]
fn test_knn_dimension_mismatch_returns_empty() {
let tree: KdTree<Point2D<&str>> = KdTree::with_dimension(3);
let target = Point2D::new(1.0, 2.0, None::<&str>);
let results = tree.knn_search::<EuclideanDistance>(&target, 1);
assert!(results.is_empty());
}
#[test]
fn test_range_dimension_mismatch_returns_empty() {
let tree: KdTree<Point2D<&str>> = KdTree::with_dimension(3);
let target = Point2D::new(1.0, 2.0, None::<&str>);
let results = tree.range_search::<EuclideanDistance>(&target, 1.0);
assert!(results.is_empty());
}
}