use crate::errors::SpartError;
use crate::geometry::{
BoundingVolume, BoundingVolumeFromPoint, Cube, DistanceMetric, HasMinDistance, Point2D,
Point3D, Rectangle,
};
use crate::rtree_common::{
KnnCandidate, compute_group_mbr as common_compute_group_mbr,
delete_entry as common_delete_entry, search_node as common_search_node,
};
use ordered_float::OrderedFloat;
#[cfg(feature = "serde")]
use serde::{Deserialize, Serialize};
use std::cmp::Ordering;
use std::collections::BinaryHeap;
use tracing::{debug, info};
const EPSILON: f64 = 1e-10;
#[cfg(feature = "serde")]
pub trait RTreeObject: std::fmt::Debug + Clone {
type B: BoundingVolume
+ std::fmt::Debug
+ Clone
+ serde::Serialize
+ for<'de> serde::Deserialize<'de>;
fn mbr(&self) -> Self::B;
}
#[cfg(not(feature = "serde"))]
pub trait RTreeObject: std::fmt::Debug + Clone {
type B: BoundingVolume + std::fmt::Debug + Clone;
fn mbr(&self) -> Self::B;
}
#[derive(Debug, Clone)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub enum RTreeEntry<T: RTreeObject> {
Leaf { mbr: T::B, object: T },
Node { mbr: T::B, child: Box<RTreeNode<T>> },
}
impl<T: RTreeObject> RTreeEntry<T> {
pub fn mbr(&self) -> &T::B {
match self {
RTreeEntry::Leaf { mbr, .. } => mbr,
RTreeEntry::Node { mbr, .. } => mbr,
}
}
}
#[derive(Debug, Clone)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub struct RTreeNode<T: RTreeObject> {
pub entries: Vec<RTreeEntry<T>>,
pub is_leaf: bool,
}
#[derive(Debug, Clone)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub struct RTree<T: RTreeObject> {
root: RTreeNode<T>,
max_entries: usize,
min_entries: usize,
}
impl<T: RTreeObject> crate::rtree_common::EntryAccess for RTreeEntry<T> {
type BV = T::B;
type Node = RTreeNode<T>;
type Obj = T;
fn mbr(&self) -> &Self::BV {
RTreeEntry::mbr(self)
}
fn as_leaf_obj(&self) -> Option<&Self::Obj> {
match self {
RTreeEntry::Leaf { object, .. } => Some(object),
_ => None,
}
}
fn child(&self) -> Option<&<Self as crate::rtree_common::EntryAccess>::Node> {
match self {
RTreeEntry::Node { child, .. } => Some(child),
_ => None,
}
}
fn child_mut(&mut self) -> Option<&mut <Self as crate::rtree_common::EntryAccess>::Node> {
match self {
RTreeEntry::Node { child, .. } => Some(child),
_ => None,
}
}
fn set_mbr(&mut self, new_mbr: Self::BV) {
if let RTreeEntry::Node { mbr, .. } = self {
*mbr = new_mbr;
}
}
fn into_child(self) -> Option<Box<<Self as crate::rtree_common::EntryAccess>::Node>>
where
Self: Sized,
{
match self {
RTreeEntry::Node { child, .. } => Some(child),
_ => None,
}
}
}
impl<T: RTreeObject> crate::rtree_common::NodeAccess for RTreeNode<T> {
type Entry = RTreeEntry<T>;
fn is_leaf(&self) -> bool {
self.is_leaf
}
fn entries(&self) -> &Vec<Self::Entry> {
&self.entries
}
fn entries_mut(&mut self) -> &mut Vec<Self::Entry> {
&mut self.entries
}
}
impl<T: RTreeObject> RTree<T> {
pub fn new(max_entries: usize) -> Result<Self, SpartError> {
if max_entries < 2 {
return Err(SpartError::InvalidCapacity {
capacity: max_entries,
});
}
info!("Creating new RTree with max_entries: {}", max_entries);
Ok(RTree {
root: RTreeNode {
entries: Vec::new(),
is_leaf: true,
},
max_entries,
min_entries: (max_entries as f64 * 0.4).ceil() as usize,
})
}
pub fn insert(&mut self, object: T) {
info!("Inserting object into RTree: {:?}", object);
let entry = RTreeEntry::Leaf {
mbr: object.mbr(),
object,
};
insert_entry_node(&mut self.root, entry);
if self.root.entries.len() > self.max_entries {
info!("Root has exceeded max_entries; splitting root");
self.split_root();
}
}
fn split_root(&mut self) {
info!("Splitting root node");
let old_entries = std::mem::take(&mut self.root.entries);
let (group1, group2) = split_entries(old_entries, self.max_entries);
let child1 = RTreeNode {
entries: group1,
is_leaf: self.root.is_leaf,
};
let child2 = RTreeNode {
entries: group2,
is_leaf: self.root.is_leaf,
};
let mbr1 = common_compute_group_mbr(&child1.entries)
.unwrap_or_else(|| unreachable!("non-empty group must have MBR"));
let mbr2 = common_compute_group_mbr(&child2.entries)
.unwrap_or_else(|| unreachable!("non-empty group must have MBR"));
self.root.is_leaf = false;
self.root.entries.push(RTreeEntry::Node {
mbr: mbr1,
child: Box::new(child1),
});
self.root.entries.push(RTreeEntry::Node {
mbr: mbr2,
child: Box::new(child2),
});
}
pub fn range_search_bbox(&self, query: &T::B) -> Vec<&T> {
info!("Performing range search with query: {:?}", query);
let mut result = Vec::new();
common_search_node(&self.root, query, &mut result);
result
}
pub fn insert_bulk(&mut self, objects: Vec<T>) {
if objects.is_empty() {
return;
}
let mut entries: Vec<RTreeEntry<T>> = objects
.into_iter()
.map(|obj| RTreeEntry::Leaf {
mbr: obj.mbr(),
object: obj,
})
.collect();
while entries.len() > self.max_entries {
let mut new_level_entries = Vec::new();
let chunks = entries.chunks(self.max_entries);
for chunk in chunks {
let child_node = RTreeNode {
entries: chunk.to_vec(),
is_leaf: self.root.is_leaf,
};
if let Some(mbr) = common_compute_group_mbr(&child_node.entries) {
new_level_entries.push(RTreeEntry::Node {
mbr,
child: Box::new(child_node),
});
}
}
entries = new_level_entries;
self.root.is_leaf = false;
}
self.root.entries.extend(entries);
}
}
fn insert_entry_node<T: RTreeObject>(node: &mut RTreeNode<T>, entry: RTreeEntry<T>) {
if node.is_leaf {
debug!("Inserting entry into leaf node");
node.entries.push(entry);
} else {
let mut best_index: Option<usize> = None;
let mut best_enlargement = f64::INFINITY;
for (i, child_entry) in node.entries.iter().enumerate() {
if let RTreeEntry::Node { mbr, .. } = child_entry {
let enlargement = mbr.enlargement(entry.mbr());
if enlargement < best_enlargement {
best_enlargement = enlargement;
best_index = Some(i);
} else if (enlargement - best_enlargement).abs() < f64::EPSILON {
if let Some(current_best) = best_index {
if mbr.area() < node.entries[current_best].mbr().area() {
best_index = Some(i);
}
}
}
}
}
if let Some(best_index) = best_index {
if let RTreeEntry::Node { mbr, child } = &mut node.entries[best_index] {
*mbr = mbr.union(entry.mbr());
insert_entry_node(child, entry);
if let Some(new_mbr) = common_compute_group_mbr(&child.entries) {
*mbr = new_mbr;
}
}
} else {
node.entries.push(entry);
}
}
}
fn split_entries<T: RTreeObject>(
entries: Vec<RTreeEntry<T>>,
_max_entries: usize,
) -> (Vec<RTreeEntry<T>>, Vec<RTreeEntry<T>>) {
let mut entries = entries;
if entries.len() < 2 {
return (entries, Vec::new());
}
let seed1 = entries.remove(0);
let seed2 = entries.remove(0);
let mut group1 = vec![seed1];
let mut group2 = vec![seed2];
for entry in entries {
let mbr1 = common_compute_group_mbr(&group1)
.unwrap_or_else(|| unreachable!("non-empty group must have MBR"));
let mbr2 = common_compute_group_mbr(&group2)
.unwrap_or_else(|| unreachable!("non-empty group must have MBR"));
let enlargement1 = mbr1.enlargement(entry.mbr());
let enlargement2 = mbr2.enlargement(entry.mbr());
if enlargement1 < enlargement2 {
group1.push(entry);
} else {
group2.push(entry);
}
}
(group1, group2)
}
impl<T: RTreeObject> RTree<T>
where
T: PartialEq,
{
pub fn delete(&mut self, object: &T) -> bool {
info!("Attempting to delete object: {:?}", object);
let object_mbr = object.mbr();
let mut reinsert_list = Vec::new();
let deleted = common_delete_entry(
&mut self.root,
object,
&object_mbr,
self.min_entries,
&mut reinsert_list,
);
if deleted {
for entry in reinsert_list {
self.insert_entry(entry);
}
if !self.root.is_leaf && self.root.entries.len() == 1 {
if let Some(RTreeEntry::Node { child, .. }) = self.root.entries.pop() {
self.root = *child;
}
}
}
deleted
}
fn insert_entry(&mut self, entry: RTreeEntry<T>) {
insert_entry_node(&mut self.root, entry);
if self.root.entries.len() > self.max_entries {
self.split_root();
}
}
}
impl<T: std::fmt::Debug + Clone> RTreeObject for Point2D<T> {
type B = Rectangle;
fn mbr(&self) -> Self::B {
Rectangle {
x: self.x,
y: self.y,
width: EPSILON,
height: EPSILON,
}
}
}
impl<T: std::fmt::Debug + Clone> RTreeObject for Point3D<T> {
type B = Cube;
fn mbr(&self) -> Self::B {
Cube {
x: self.x,
y: self.y,
z: self.z,
width: EPSILON,
height: EPSILON,
depth: EPSILON,
}
}
}
impl Rectangle {
pub fn min_distance<T>(&self, point: &Point2D<T>) -> f64 {
let dx = if point.x < self.x {
self.x - point.x
} else if point.x > self.x + self.width {
point.x - (self.x + self.width)
} else {
0.0
};
let dy = if point.y < self.y {
self.y - point.y
} else if point.y > self.y + self.height {
point.y - (self.y + self.height)
} else {
0.0
};
(dx * dx + dy * dy).sqrt()
}
}
impl Cube {
pub fn min_distance<T>(&self, point: &Point3D<T>) -> f64 {
let dx = if point.x < self.x {
self.x - point.x
} else if point.x > self.x + self.width {
point.x - (self.x + self.width)
} else {
0.0
};
let dy = if point.y < self.y {
self.y - point.y
} else if point.y > self.y + self.height {
point.y - (self.y + self.height)
} else {
0.0
};
let dz = if point.z < self.z {
self.z - point.z
} else if point.z > self.z + self.depth {
point.z - (self.z + self.depth)
} else {
0.0
};
(dx * dx + dy * dy + dz * dz).sqrt()
}
}
impl<T: std::fmt::Debug + Clone> RTree<Point2D<T>> {
pub fn knn_search<M: DistanceMetric<Point2D<T>>>(
&self,
query: &Point2D<T>,
k: usize,
) -> Vec<&Point2D<T>> {
if k == 0 {
return Vec::new();
}
let mut heap: BinaryHeap<crate::rtree_common::KnnCandidate<RTreeEntry<Point2D<T>>>> =
BinaryHeap::new();
for entry in &self.root.entries {
let dist_sq = entry.mbr().min_distance(query).powi(2);
heap.push(KnnCandidate {
dist: dist_sq,
entry,
});
}
type OrdDist = OrderedFloat<f64>;
#[inline]
#[allow(non_snake_case)]
fn OrdDist(x: f64) -> OrderedFloat<f64> {
OrderedFloat(x)
}
struct HeapItem<'a, P> {
key: OrdDist,
idx: usize,
obj: &'a P,
}
impl<P> PartialEq for HeapItem<'_, P> {
fn eq(&self, other: &Self) -> bool {
self.key == other.key && self.idx == other.idx
}
}
impl<P> Eq for HeapItem<'_, P> {}
impl<P> Ord for HeapItem<'_, P> {
fn cmp(&self, other: &Self) -> Ordering {
match self.key.cmp(&other.key) {
Ordering::Equal => self.idx.cmp(&other.idx),
ord => ord,
}
}
}
impl<P> PartialOrd for HeapItem<'_, P> {
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
Some(self.cmp(other))
}
}
let mut results: BinaryHeap<HeapItem<Point2D<T>>> = BinaryHeap::new();
let mut counter: usize = 0;
while let Some(KnnCandidate { dist, entry }) = heap.pop() {
if results.len() >= k {
if let Some(worst_result) = results.peek() {
if dist > worst_result.key.0 {
break;
}
}
}
match entry {
RTreeEntry::Leaf { object, .. } => {
let d_sq = M::distance_sq(query, object);
if results.len() < k {
counter += 1;
results.push(HeapItem {
key: OrdDist(d_sq),
idx: counter,
obj: object,
});
} else if let Some(peek) = results.peek() {
if d_sq < peek.key.0 {
results.pop();
counter += 1;
results.push(HeapItem {
key: OrdDist(d_sq),
idx: counter,
obj: object,
});
}
}
}
RTreeEntry::Node { child, .. } => {
for child_entry in &child.entries {
let d_sq = child_entry.mbr().min_distance(query).powi(2);
if results.len() < k {
heap.push(KnnCandidate {
dist: d_sq,
entry: child_entry,
});
} else if let Some(peek) = results.peek() {
if d_sq < peek.key.0 {
heap.push(KnnCandidate {
dist: d_sq,
entry: child_entry,
});
}
}
}
}
}
}
let mut sorted_results = results.into_vec();
sorted_results.sort_by(|a, b| a.key.partial_cmp(&b.key).unwrap_or(Ordering::Equal));
sorted_results.into_iter().map(|r| r.obj).collect()
}
}
impl<T: std::fmt::Debug + Clone> RTree<Point3D<T>> {
pub fn knn_search<M: DistanceMetric<Point3D<T>>>(
&self,
query: &Point3D<T>,
k: usize,
) -> Vec<&Point3D<T>> {
if k == 0 {
return Vec::new();
}
let mut heap: BinaryHeap<crate::rtree_common::KnnCandidate<RTreeEntry<Point3D<T>>>> =
BinaryHeap::new();
for entry in &self.root.entries {
let dist_sq = entry.mbr().min_distance(query).powi(2);
heap.push(KnnCandidate {
dist: dist_sq,
entry,
});
}
type OrdDist = OrderedFloat<f64>;
#[inline]
#[allow(non_snake_case)]
fn OrdDist(x: f64) -> OrderedFloat<f64> {
OrderedFloat(x)
}
struct HeapItem<'a, P> {
key: OrdDist,
idx: usize,
obj: &'a P,
}
impl<P> PartialEq for HeapItem<'_, P> {
fn eq(&self, other: &Self) -> bool {
self.key == other.key && self.idx == other.idx
}
}
impl<P> Eq for HeapItem<'_, P> {}
impl<P> Ord for HeapItem<'_, P> {
fn cmp(&self, other: &Self) -> Ordering {
match self.key.cmp(&other.key) {
Ordering::Equal => self.idx.cmp(&other.idx),
ord => ord,
}
}
}
impl<P> PartialOrd for HeapItem<'_, P> {
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
Some(self.cmp(other))
}
}
let mut results: BinaryHeap<HeapItem<Point3D<T>>> = BinaryHeap::new();
let mut counter: usize = 0;
while let Some(KnnCandidate { dist, entry }) = heap.pop() {
if results.len() >= k {
if let Some(worst_result) = results.peek() {
if dist > worst_result.key.0 {
break;
}
}
}
match entry {
RTreeEntry::Leaf { object, .. } => {
let d_sq = M::distance_sq(query, object);
if results.len() < k {
counter += 1;
results.push(HeapItem {
key: OrdDist(d_sq),
idx: counter,
obj: object,
});
} else if let Some(peek) = results.peek() {
if d_sq < peek.key.0 {
results.pop();
counter += 1;
results.push(HeapItem {
key: OrdDist(d_sq),
idx: counter,
obj: object,
});
}
}
}
RTreeEntry::Node { child, .. } => {
for child_entry in &child.entries {
let d_sq = child_entry.mbr().min_distance(query).powi(2);
if results.len() < k {
heap.push(KnnCandidate {
dist: d_sq,
entry: child_entry,
});
} else if let Some(peek) = results.peek() {
if d_sq < peek.key.0 {
heap.push(KnnCandidate {
dist: d_sq,
entry: child_entry,
});
}
}
}
}
}
}
let mut sorted_results = results.into_vec();
sorted_results.sort_by(|a, b| a.key.partial_cmp(&b.key).unwrap_or(Ordering::Equal));
sorted_results.into_iter().map(|r| r.obj).collect()
}
}
impl<T> RTree<T>
where
T: RTreeObject + PartialEq + std::fmt::Debug,
T::B: BoundingVolumeFromPoint<T> + HasMinDistance<T> + Clone,
{
pub fn range_search<M: DistanceMetric<T>>(&self, query: &T, radius: f64) -> Vec<&T> {
if radius < 0.0 {
return Vec::new();
}
let query_volume = T::B::from_point_radius(query, radius);
let candidates = self.range_search_bbox(&query_volume);
candidates
.into_iter()
.filter(|object| M::distance_sq(query, object) <= radius * radius)
.collect()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::geometry::EuclideanDistance;
#[test]
fn test_range_search_radius_zero_2d() {
let mut tree: RTree<Point2D<&str>> = RTree::new(4).unwrap();
let target = Point2D::new(5.0, 5.0, Some("T"));
tree.insert(target.clone());
tree.insert(Point2D::new(5.0, 6.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_range_search_bbox_filters_results() {
let mut tree: RTree<Point2D<&str>> = RTree::new(4).unwrap();
let inside = Point2D::new(1.0, 1.0, Some("I"));
let outside = Point2D::new(20.0, 20.0, Some("O"));
tree.insert(inside.clone());
tree.insert(outside);
let query = Rectangle {
x: 0.0,
y: 0.0,
width: 5.0,
height: 5.0,
};
let results = tree.range_search_bbox(&query);
assert_eq!(results.len(), 1);
assert_eq!(*results[0], inside);
}
#[test]
fn test_delete_removes_point_3d() {
let mut tree: RTree<Point3D<&str>> = RTree::new(4).unwrap();
let a = Point3D::new(1.0, 1.0, 1.0, Some("A"));
let b = Point3D::new(2.0, 2.0, 2.0, Some("B"));
tree.insert(a.clone());
tree.insert(b.clone());
assert!(tree.delete(&a));
let removed = tree.range_search::<EuclideanDistance>(&a, 0.0);
let remaining = tree.range_search::<EuclideanDistance>(&b, 0.0);
assert!(removed.is_empty());
assert_eq!(remaining.len(), 1);
assert_eq!(*remaining[0], b);
}
#[test]
fn test_delete_underflow() {
let mut tree: RTree<Point2D<i32>> = RTree::new(4).unwrap();
let points: Vec<_> = (0..10)
.map(|i| Point2D::new(i as f64, i as f64, Some(i)))
.collect();
for p in &points {
tree.insert(p.clone());
}
assert!(tree.delete(&points[0]));
assert!(tree.delete(&points[1]));
assert!(tree.delete(&points[2]));
let all_points = tree.range_search_bbox(&crate::geometry::Rectangle {
x: -1.0,
y: -1.0,
width: 12.0,
height: 12.0,
});
assert_eq!(all_points.len(), 7);
for i in 3..10 {
assert!(tree.delete(&points[i]));
}
let all_points_after_all_deleted = tree.range_search_bbox(&crate::geometry::Rectangle {
x: -1.0,
y: -1.0,
width: 12.0,
height: 12.0,
});
assert!(all_points_after_all_deleted.is_empty());
}
#[test]
fn test_empty_tree_queries() {
let mut tree: RTree<Point2D<&str>> = RTree::new(4).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 mut tree: RTree<Point2D<&str>> = RTree::new(4).unwrap();
let points = vec![
Point2D::new(1.0, 1.0, Some("A")),
Point2D::new(2.0, 2.0, Some("B")),
Point2D::new(3.0, 3.0, Some("C")),
];
let num_points = points.len();
tree.insert_bulk(points.clone());
let target = Point2D::new(1.5, 1.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_duplicates_delete_one() {
let mut tree: RTree<Point2D<&str>> = RTree::new(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_negative_radius_empty() {
let mut tree: RTree<Point2D<&str>> = RTree::new(4).unwrap();
let target = Point2D::new(5.0, 5.0, Some("T"));
tree.insert(target.clone());
let results = tree.range_search::<EuclideanDistance>(&target, -1.0);
assert!(results.is_empty());
}
}