use crate::errors::SpartError;
use crate::geometry::{DistanceMetric, HeapItem, Point2D, Rectangle};
use ordered_float::OrderedFloat;
#[cfg(feature = "serde")]
use serde::{Deserialize, Serialize};
use std::collections::BinaryHeap;
use tracing::{debug, info};
#[derive(Debug, Clone)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub struct Quadtree<T: Clone + PartialEq> {
boundary: Rectangle,
points: Vec<Point2D<T>>,
capacity: usize,
divided: bool,
northeast: Option<Box<Quadtree<T>>>,
northwest: Option<Box<Quadtree<T>>>,
southeast: Option<Box<Quadtree<T>>>,
southwest: Option<Box<Quadtree<T>>>,
}
impl<T: Clone + PartialEq + std::fmt::Debug> Quadtree<T> {
pub fn new(boundary: &Rectangle, capacity: usize) -> Result<Self, SpartError> {
if capacity == 0 {
return Err(SpartError::InvalidCapacity { capacity });
}
info!(
"Creating new Quadtree with boundary: {:?} and capacity: {}",
boundary, capacity
);
Ok(Quadtree {
boundary: boundary.clone(),
points: Vec::new(),
capacity,
divided: false,
northeast: None,
northwest: None,
southeast: None,
southwest: None,
})
}
fn subdivide(&mut self) {
info!("Subdividing Quadtree at boundary: {:?}", self.boundary);
let x = self.boundary.x;
let y = self.boundary.y;
let w = self.boundary.width / 2.0;
let h = self.boundary.height / 2.0;
self.northeast = Some(Box::new({
let child = Quadtree::new(
&Rectangle {
x: x + w,
y,
width: w,
height: h,
},
self.capacity,
);
match child {
Ok(c) => c,
Err(_) => unreachable!("capacity validated at construction"),
}
}));
self.northwest = Some(Box::new({
let child = Quadtree::new(
&Rectangle {
x,
y,
width: w,
height: h,
},
self.capacity,
);
match child {
Ok(c) => c,
Err(_) => unreachable!("capacity validated at construction"),
}
}));
self.southeast = Some(Box::new({
let child = Quadtree::new(
&Rectangle {
x: x + w,
y: y + h,
width: w,
height: h,
},
self.capacity,
);
match child {
Ok(c) => c,
Err(_) => unreachable!("capacity validated at construction"),
}
}));
self.southwest = Some(Box::new({
let child = Quadtree::new(
&Rectangle {
x,
y: y + h,
width: w,
height: h,
},
self.capacity,
);
match child {
Ok(c) => c,
Err(_) => unreachable!("capacity validated at construction"),
}
}));
self.divided = true;
let old_points = std::mem::take(&mut self.points);
for point in old_points {
let inserted = self.insert(point);
if !inserted {
debug!("Failed to reinsert point during subdivision");
}
}
}
pub fn insert(&mut self, point: Point2D<T>) -> bool {
if !self.boundary.contains(&point) {
return false;
}
if !self.divided {
if self.points.len() < self.capacity {
self.points.push(point);
return true;
}
self.subdivide();
}
if self
.northwest
.as_mut()
.is_some_and(|c| c.insert(point.clone()))
{
return true;
}
if self
.northeast
.as_mut()
.is_some_and(|c| c.insert(point.clone()))
{
return true;
}
if self
.southwest
.as_mut()
.is_some_and(|c| c.insert(point.clone()))
{
return true;
}
if self
.southeast
.as_mut()
.is_some_and(|c| c.insert(point.clone()))
{
return true;
}
unreachable!("A point within the parent boundary should always fit in a child boundary.");
}
pub fn insert_bulk(&mut self, points: &[Point2D<T>]) {
if points.is_empty() {
return;
}
let points_within_boundary: Vec<Point2D<T>> = points
.iter()
.filter(|p| self.boundary.contains(p))
.cloned()
.collect();
if points_within_boundary.is_empty() {
return;
}
if !self.divided && self.points.len() + points_within_boundary.len() <= self.capacity {
self.points.extend(points_within_boundary);
return;
}
if !self.divided {
self.subdivide();
}
let mut points_to_insert = points_within_boundary;
if self.divided {
let mut children_points: [Vec<Point2D<T>>; 4] = [vec![], vec![], vec![], vec![]];
for point in points_to_insert.drain(..) {
if self
.northeast
.as_ref()
.map(|c| c.boundary.contains(&point))
.unwrap_or(false)
{
children_points[0].push(point);
} else if self
.northwest
.as_ref()
.map(|c| c.boundary.contains(&point))
.unwrap_or(false)
{
children_points[1].push(point);
} else if self
.southeast
.as_ref()
.map(|c| c.boundary.contains(&point))
.unwrap_or(false)
{
children_points[2].push(point);
} else if self
.southwest
.as_ref()
.map(|c| c.boundary.contains(&point))
.unwrap_or(false)
{
children_points[3].push(point);
}
}
if !children_points[0].is_empty() {
if let Some(c) = self.northeast.as_mut() {
c.insert_bulk(&children_points[0]);
}
}
if !children_points[1].is_empty() {
if let Some(c) = self.northwest.as_mut() {
c.insert_bulk(&children_points[1]);
}
}
if !children_points[2].is_empty() {
if let Some(c) = self.southeast.as_mut() {
c.insert_bulk(&children_points[2]);
}
}
if !children_points[3].is_empty() {
if let Some(c) = self.southwest.as_mut() {
c.insert_bulk(&children_points[3]);
}
}
}
}
fn children_mut(&mut self) -> Vec<&mut Quadtree<T>> {
let mut children = Vec::with_capacity(4);
if let Some(ref mut child) = self.northeast {
children.push(child.as_mut());
}
if let Some(ref mut child) = self.northwest {
children.push(child.as_mut());
}
if let Some(ref mut child) = self.southeast {
children.push(child.as_mut());
}
if let Some(ref mut child) = self.southwest {
children.push(child.as_mut());
}
children
}
fn children(&self) -> Vec<&Quadtree<T>> {
let mut children = Vec::with_capacity(4);
if let Some(ref child) = self.northeast {
children.push(child.as_ref());
}
if let Some(ref child) = self.northwest {
children.push(child.as_ref());
}
if let Some(ref child) = self.southeast {
children.push(child.as_ref());
}
if let Some(ref child) = self.southwest {
children.push(child.as_ref());
}
children
}
fn min_distance_sq(&self, target: &Point2D<T>) -> f64 {
let mut dx = 0.0;
if target.x < self.boundary.x {
dx = self.boundary.x - target.x;
} else if target.x > self.boundary.x + self.boundary.width {
dx = target.x - (self.boundary.x + self.boundary.width);
}
let mut dy = 0.0;
if target.y < self.boundary.y {
dy = self.boundary.y - target.y;
} else if target.y > self.boundary.y + self.boundary.height {
dy = target.y - (self.boundary.y + self.boundary.height);
}
dx * dx + dy * dy
}
pub fn knn_search<M: DistanceMetric<Point2D<T>>>(
&self,
target: &Point2D<T>,
k: usize,
) -> Vec<Point2D<T>> {
if k == 0 {
return Vec::new();
}
let mut heap: BinaryHeap<HeapItem<T>> = BinaryHeap::new();
self.knn_search_helper::<M>(target, k, &mut heap);
heap.into_sorted_vec()
.into_iter()
.filter_map(|item| item.point_2d)
.collect()
}
fn knn_search_helper<M: DistanceMetric<Point2D<T>>>(
&self,
target: &Point2D<T>,
k: usize,
heap: &mut BinaryHeap<HeapItem<T>>,
) {
for point in &self.points {
let dist_sq = M::distance_sq(point, target);
let item = HeapItem {
neg_distance: OrderedFloat(-dist_sq),
point_2d: Some(point.clone()),
point_3d: None,
};
heap.push(item);
if heap.len() > k {
heap.pop();
}
}
if self.divided {
for child in self.children() {
if heap.len() == k {
if let Some(top) = heap.peek() {
let current_farthest = -top.neg_distance.into_inner();
if child.min_distance_sq(target) > current_farthest {
continue;
}
}
}
child.knn_search_helper::<M>(target, k, heap);
}
}
}
pub fn range_search<M: DistanceMetric<Point2D<T>>>(
&self,
center: &Point2D<T>,
radius: f64,
) -> Vec<Point2D<T>> {
if radius < 0.0 {
return Vec::new();
}
let mut found = Vec::new();
let radius_sq = radius * radius;
if self.min_distance_sq(center) > radius_sq {
return found;
}
for point in &self.points {
if M::distance_sq(point, center) <= radius_sq {
found.push(point.clone());
}
}
if self.divided {
for child in self.children() {
found.extend(child.range_search::<M>(center, radius));
}
}
found
}
pub fn delete(&mut self, point: &Point2D<T>) -> bool {
if !self.boundary.contains(point) {
return false;
}
let mut deleted = false;
if self.divided {
for child in self.children_mut() {
if child.delete(point) {
deleted = true;
break;
}
}
self.try_merge();
return deleted;
}
if let Some(pos) = self.points.iter().position(|p| p == point) {
self.points.remove(pos);
info!("Deleting point {:?} from Quadtree", point);
true
} else {
false
}
}
fn try_merge(&mut self) {
if !self.divided {
return;
}
for child in self.children_mut() {
child.try_merge();
}
let children = self.children();
if children.iter().all(|child| !child.divided) {
let total_points: usize = children.iter().map(|child| child.points.len()).sum();
if total_points <= self.capacity {
let mut merged_points = Vec::with_capacity(total_points);
if let Some(child) = self.northeast.take() {
merged_points.extend(child.points);
}
if let Some(child) = self.northwest.take() {
merged_points.extend(child.points);
}
if let Some(child) = self.southeast.take() {
merged_points.extend(child.points);
}
if let Some(child) = self.southwest.take() {
merged_points.extend(child.points);
}
info!(
"Merging children into parent node at boundary {:?} with {} points",
self.boundary,
merged_points.len()
);
self.points.extend(merged_points);
self.divided = false;
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::geometry::EuclideanDistance;
#[test]
fn test_insert_rejects_outside_boundary() {
let boundary = Rectangle {
x: 0.0,
y: 0.0,
width: 10.0,
height: 10.0,
};
let mut tree: Quadtree<&str> = Quadtree::new(&boundary, 2).unwrap();
let outside = Point2D::new(20.0, 20.0, Some("O"));
assert!(!tree.insert(outside));
}
#[test]
fn test_insert_accepts_boundary_points() {
let boundary = Rectangle {
x: 0.0,
y: 0.0,
width: 10.0,
height: 10.0,
};
let mut tree: Quadtree<&str> = Quadtree::new(&boundary, 1).unwrap();
let edge = Point2D::new(10.0, 10.0, Some("E"));
assert!(tree.insert(edge));
}
#[test]
fn test_range_search_zero_radius_returns_exact_match() {
let boundary = Rectangle {
x: 0.0,
y: 0.0,
width: 100.0,
height: 100.0,
};
let mut tree: Quadtree<&str> = Quadtree::new(&boundary, 2).unwrap();
let target = Point2D::new(25.0, 25.0, Some("T"));
tree.insert(target.clone());
tree.insert(Point2D::new(26.0, 25.0, Some("N")));
let results = tree.range_search::<EuclideanDistance>(&target, 0.0);
assert_eq!(results.len(), 1);
assert_eq!(results[0], target);
}
#[test]
fn test_delete_existing_point() {
let boundary = Rectangle {
x: 0.0,
y: 0.0,
width: 100.0,
height: 100.0,
};
let mut tree: Quadtree<&str> = Quadtree::new(&boundary, 2).unwrap();
let p1 = Point2D::new(10.0, 10.0, Some("A"));
let p2 = Point2D::new(20.0, 20.0, Some("B"));
tree.insert(p1.clone());
tree.insert(p2);
assert!(tree.delete(&p1));
let results = tree.knn_search::<EuclideanDistance>(&p1, 1);
assert_ne!(results[0], p1);
assert!(!tree.delete(&p1));
}
#[test]
fn test_empty_tree_queries() {
let boundary = Rectangle {
x: 0.0,
y: 0.0,
width: 100.0,
height: 100.0,
};
let mut tree: Quadtree<&str> = Quadtree::new(&boundary, 2).unwrap();
let target = Point2D::new(5.0, 5.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 boundary = Rectangle {
x: 0.0,
y: 0.0,
width: 100.0,
height: 100.0,
};
let mut tree: Quadtree<&str> = Quadtree::new(&boundary, 4).unwrap();
let points = vec![
Point2D::new(10.0, 10.0, Some("A")),
Point2D::new(20.0, 20.0, Some("B")),
Point2D::new(30.0, 30.0, Some("C")),
];
let num_points = points.len();
tree.insert_bulk(&points);
let target = Point2D::new(15.0, 15.0, 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_duplicates_delete_one() {
let boundary = Rectangle {
x: 0.0,
y: 0.0,
width: 100.0,
height: 100.0,
};
let mut tree: Quadtree<&str> = Quadtree::new(&boundary, 4).unwrap();
let p1 = Point2D::new(10.0, 10.0, Some("A"));
let p2 = p1.clone();
tree.insert(p1.clone());
tree.insert(p2.clone());
let results = tree.knn_search::<EuclideanDistance>(&p1, 2);
assert_eq!(results.len(), 2);
assert!(tree.delete(&p1));
let results_after_delete = tree.knn_search::<EuclideanDistance>(&p1, 2);
assert_eq!(results_after_delete.len(), 1);
}
#[test]
fn test_range_search_includes_boundary_point() {
let boundary = Rectangle {
x: 0.0,
y: 0.0,
width: 100.0,
height: 100.0,
};
let mut tree: Quadtree<&str> = Quadtree::new(&boundary, 4).unwrap();
let center = Point2D::new(50.0, 50.0, Some("C"));
let boundary_point = Point2D::new(60.0, 50.0, Some("B"));
tree.insert(center.clone());
tree.insert(boundary_point.clone());
let results = tree.range_search::<EuclideanDistance>(¢er, 10.0);
assert!(results.contains(&boundary_point));
assert!(results.contains(¢er));
}
#[test]
fn test_bulk_insert_empty_noop() {
let boundary = Rectangle {
x: 0.0,
y: 0.0,
width: 100.0,
height: 100.0,
};
let mut tree: Quadtree<i32> = Quadtree::new(&boundary, 4).unwrap();
let empty: Vec<Point2D<i32>> = Vec::new();
tree.insert_bulk(&empty);
let target = Point2D::new(10.0, 10.0, None::<i32>);
let results = tree.knn_search::<EuclideanDistance>(&target, 1);
assert!(results.is_empty());
}
#[test]
fn test_zero_capacity_rejected() {
let boundary = Rectangle {
x: 0.0,
y: 0.0,
width: 100.0,
height: 100.0,
};
let result = Quadtree::<i32>::new(&boundary, 0);
assert!(result.is_err());
}
#[test]
fn test_range_search_negative_radius_empty() {
let boundary = Rectangle {
x: 0.0,
y: 0.0,
width: 100.0,
height: 100.0,
};
let mut tree: Quadtree<&str> = Quadtree::new(&boundary, 2).unwrap();
let target = Point2D::new(10.0, 10.0, Some("T"));
tree.insert(target.clone());
let results = tree.range_search::<EuclideanDistance>(&target, -1.0);
assert!(results.is_empty());
}
}