1use crate::point::min_inline;
2use crate::{
3 node::{ParentNode, RTreeNode},
4 object::Distance,
5};
6use crate::{Envelope, PointDistance, RTreeObject};
7
8#[cfg(doc)]
9use crate::RTree;
10
11use alloc::collections::BinaryHeap;
12#[cfg(not(test))]
13use alloc::{vec, vec::Vec};
14use core::mem::replace;
15use heapless::binary_heap as static_heap;
16use num_traits::Bounded;
17
18struct RTreeNodeDistanceWrapper<'a, T>
19where
20 T: PointDistance + 'a,
21{
22 node: &'a RTreeNode<T>,
23 distance: Distance<T>,
24}
25
26impl<T> PartialEq for RTreeNodeDistanceWrapper<'_, T>
27where
28 T: PointDistance,
29{
30 fn eq(&self, other: &Self) -> bool {
31 self.distance == other.distance
32 }
33}
34
35impl<T> PartialOrd for RTreeNodeDistanceWrapper<'_, T>
36where
37 T: PointDistance,
38{
39 fn partial_cmp(&self, other: &Self) -> Option<::core::cmp::Ordering> {
40 Some(self.cmp(other))
41 }
42}
43
44impl<T> Eq for RTreeNodeDistanceWrapper<'_, T> where T: PointDistance {}
45
46impl<T> Ord for RTreeNodeDistanceWrapper<'_, T>
47where
48 T: PointDistance,
49{
50 fn cmp(&self, other: &Self) -> ::core::cmp::Ordering {
51 other.distance.partial_cmp(&self.distance).unwrap()
53 }
54}
55
56impl<'a, T> NearestNeighborDistance2Iterator<'a, T>
57where
58 T: PointDistance,
59{
60 pub(crate) fn new(
61 root: &'a ParentNode<T>,
62 query_point: <T::Envelope as Envelope>::Point,
63 ) -> Self {
64 let mut result = NearestNeighborDistance2Iterator {
65 nodes: SmallHeap::new(),
66 query_point,
67 };
68 result.extend_heap(&root.children);
69 result
70 }
71
72 fn extend_heap(&mut self, children: &'a [RTreeNode<T>]) {
73 let &mut NearestNeighborDistance2Iterator {
74 ref mut nodes,
75 ref query_point,
76 } = self;
77 nodes.extend(children.iter().map(|child: &RTreeNode<T>| {
78 let distance = match child {
79 RTreeNode::Parent(ref data) => data.envelope.distance_2(query_point),
80 RTreeNode::Leaf(ref t) => t.distance_2(query_point),
81 };
82
83 RTreeNodeDistanceWrapper {
84 node: child,
85 distance,
86 }
87 }));
88 }
89}
90
91impl<'a, T> Iterator for NearestNeighborDistance2Iterator<'a, T>
92where
93 T: PointDistance,
94{
95 type Item = (&'a T, Distance<T>);
96
97 fn next(&mut self) -> Option<Self::Item> {
98 while let Some(current) = self.nodes.pop() {
99 match current {
100 RTreeNodeDistanceWrapper {
101 node: RTreeNode::Parent(ref data),
102 ..
103 } => {
104 self.extend_heap(&data.children);
105 }
106 RTreeNodeDistanceWrapper {
107 node: RTreeNode::Leaf(ref t),
108 distance,
109 } => {
110 return Some((t, distance));
111 }
112 }
113 }
114 None
115 }
116}
117
118pub struct NearestNeighborDistance2Iterator<'a, T>
120where
121 T: PointDistance + 'a,
122{
123 nodes: SmallHeap<RTreeNodeDistanceWrapper<'a, T>>,
124 query_point: <T::Envelope as Envelope>::Point,
125}
126
127impl<'a, T> NearestNeighborIterator<'a, T>
128where
129 T: PointDistance,
130{
131 pub(crate) fn new(
132 root: &'a ParentNode<T>,
133 query_point: <T::Envelope as Envelope>::Point,
134 ) -> Self {
135 NearestNeighborIterator {
136 iter: NearestNeighborDistance2Iterator::new(root, query_point),
137 }
138 }
139}
140
141impl<'a, T> Iterator for NearestNeighborIterator<'a, T>
142where
143 T: PointDistance,
144{
145 type Item = &'a T;
146
147 fn next(&mut self) -> Option<Self::Item> {
148 self.iter.next().map(|(t, _distance)| t)
149 }
150}
151
152pub struct NearestNeighborIterator<'a, T>
154where
155 T: PointDistance + 'a,
156{
157 iter: NearestNeighborDistance2Iterator<'a, T>,
158}
159
160enum SmallHeap<T: Ord> {
161 Stack(static_heap::BinaryHeap<T, static_heap::Max, 32>),
162 Heap(BinaryHeap<T>),
163}
164
165impl<T: Ord> SmallHeap<T> {
166 pub fn new() -> Self {
167 Self::Stack(static_heap::BinaryHeap::new())
168 }
169
170 pub fn pop(&mut self) -> Option<T> {
171 match self {
172 SmallHeap::Stack(heap) => heap.pop(),
173 SmallHeap::Heap(heap) => heap.pop(),
174 }
175 }
176
177 pub fn push(&mut self, item: T) {
178 match self {
179 SmallHeap::Stack(heap) => {
180 if let Err(item) = heap.push(item) {
181 let capacity = heap.len() + 1;
182 let new_heap = self.spill(capacity);
183 new_heap.push(item);
184 }
185 }
186 SmallHeap::Heap(heap) => heap.push(item),
187 }
188 }
189
190 pub fn extend<I>(&mut self, iter: I)
191 where
192 I: ExactSizeIterator<Item = T>,
193 {
194 match self {
195 SmallHeap::Stack(heap) => {
196 if heap.capacity() >= heap.len() + iter.len() {
197 for item in iter {
198 if heap.push(item).is_err() {
199 unreachable!();
200 }
201 }
202 } else {
203 let capacity = heap.len() + iter.len();
204 let new_heap = self.spill(capacity);
205 new_heap.extend(iter);
206 }
207 }
208 SmallHeap::Heap(heap) => heap.extend(iter),
209 }
210 }
211
212 #[cold]
213 fn spill(&mut self, capacity: usize) -> &mut BinaryHeap<T> {
214 let new_heap = BinaryHeap::with_capacity(capacity);
215 let old_heap = replace(self, SmallHeap::Heap(new_heap));
216
217 let new_heap = match self {
218 SmallHeap::Heap(new_heap) => new_heap,
219 SmallHeap::Stack(_) => unreachable!(),
220 };
221 let old_heap = match old_heap {
222 SmallHeap::Stack(old_heap) => old_heap,
223 SmallHeap::Heap(_) => unreachable!(),
224 };
225
226 new_heap.extend(old_heap.into_vec());
227
228 new_heap
229 }
230}
231
232pub fn nearest_neighbor_with_distance_2<T>(
233 node: &ParentNode<T>,
234 query_point: <T::Envelope as Envelope>::Point,
235) -> Option<(&T, Distance<T>)>
236where
237 T: PointDistance,
238{
239 fn extend_heap<'a, T>(
240 nodes: &mut SmallHeap<RTreeNodeDistanceWrapper<'a, T>>,
241 node: &'a ParentNode<T>,
242 query_point: &<T::Envelope as Envelope>::Point,
243 min_max_distance: &mut Distance<T>,
244 ) where
245 T: PointDistance + 'a,
246 {
247 for child in &node.children {
248 let distance_if_less_or_equal = match child {
249 RTreeNode::Parent(ref data) => {
250 let distance = data.envelope.distance_2(query_point);
251 if distance <= *min_max_distance {
252 Some(distance)
253 } else {
254 None
255 }
256 }
257 RTreeNode::Leaf(ref t) => {
258 t.distance_2_if_less_or_equal(query_point, *min_max_distance)
259 }
260 };
261 if let Some(distance) = distance_if_less_or_equal {
262 *min_max_distance = min_inline(
263 *min_max_distance,
264 child.envelope().min_max_dist_2(query_point),
265 );
266 nodes.push(RTreeNodeDistanceWrapper {
267 node: child,
268 distance,
269 });
270 }
271 }
272 }
273
274 let mut smallest_min_max: Distance<T> = Bounded::max_value();
276 let mut nodes = SmallHeap::new();
277 extend_heap(&mut nodes, node, &query_point, &mut smallest_min_max);
278 while let Some(current) = nodes.pop() {
279 match current {
280 RTreeNodeDistanceWrapper {
281 node: RTreeNode::Parent(ref data),
282 ..
283 } => {
284 extend_heap(&mut nodes, data, &query_point, &mut smallest_min_max);
285 }
286 RTreeNodeDistanceWrapper {
287 node: RTreeNode::Leaf(ref t),
288 distance,
289 } => {
290 return Some((t, distance));
291 }
292 }
293 }
294 None
295}
296
297pub fn nearest_neighbors_with_distance_2<T>(
298 node: &ParentNode<T>,
299 query_point: <T::Envelope as Envelope>::Point,
300) -> Option<(Vec<&T>, Distance<T>)>
301where
302 T: PointDistance,
303{
304 let mut nearest_neighbors = NearestNeighborDistance2Iterator::new(node, query_point);
305
306 let (first, first_distance_2) = nearest_neighbors.next()?;
307
308 let mut result = vec![first];
310
311 result.extend(
315 nearest_neighbors
316 .take_while(|(_, next_distance_2)| next_distance_2 == &first_distance_2)
317 .map(|(next, _)| next),
318 );
319
320 Some((result, first_distance_2))
321}
322
323#[cfg(test)]
324mod test {
325 use crate::object::PointDistance;
326 use crate::rtree::RTree;
327 use crate::test_utilities::*;
328
329 #[test]
330 fn test_nearest_neighbor_empty() {
331 let tree: RTree<[f32; 2]> = RTree::new();
332 assert!(tree.nearest_neighbor([0.0, 213.0]).is_none());
333 }
334
335 #[test]
336 fn test_nearest_neighbor() {
337 let points = create_random_points(1000, SEED_1);
338 let tree = RTree::bulk_load(points.clone());
339
340 let sample_points = create_random_points(100, SEED_2);
341 for sample_point in sample_points {
342 let mut nearest = None;
343 let mut closest_dist = f64::INFINITY;
344 for point in &points {
345 let delta = [point[0] - sample_point[0], point[1] - sample_point[1]];
346 let new_dist = delta[0] * delta[0] + delta[1] * delta[1];
347 if new_dist < closest_dist {
348 closest_dist = new_dist;
349 nearest = Some(point);
350 }
351 }
352 assert_eq!(nearest, tree.nearest_neighbor(sample_point));
353 }
354 }
355
356 #[test]
357 fn test_nearest_neighbors_empty() {
358 let tree: RTree<[f32; 2]> = RTree::new();
359 assert!(tree.nearest_neighbors(&[0.0, 213.0]).is_empty());
360 }
361
362 #[test]
363 fn test_nearest_neighbors() {
364 let points = create_random_points(1000, SEED_1);
365 let tree = RTree::bulk_load(points);
366
367 let sample_points = create_random_points(50, SEED_2);
368 for sample_point in &sample_points {
369 let nearest_neighbors = tree.nearest_neighbors(sample_point);
370 let mut distance = -1.0;
371 for nn in &nearest_neighbors {
372 if distance < 0.0 {
373 distance = sample_point.distance_2(nn);
374 } else {
375 let new_distance = sample_point.distance_2(nn);
376 assert_eq!(new_distance, distance);
377 }
378 }
379 }
380 }
381
382 #[test]
383 fn test_nearest_neighbor_iterator() {
384 let mut points = create_random_points(1000, SEED_1);
385 let tree = RTree::bulk_load(points.clone());
386
387 let sample_points = create_random_points(50, SEED_2);
388 for sample_point in sample_points {
389 points.sort_by(|r, l| {
390 r.distance_2(&sample_point)
391 .partial_cmp(&l.distance_2(&sample_point))
392 .unwrap()
393 });
394 let collected: Vec<_> = tree.nearest_neighbor_iter(sample_point).cloned().collect();
395 assert_eq!(points, collected);
396 }
397 }
398
399 #[test]
400 fn test_nearest_neighbor_iterator_with_distance_2() {
401 let points = create_random_points(1000, SEED_2);
402 let tree = RTree::bulk_load(points);
403
404 let sample_points = create_random_points(50, SEED_1);
405 for sample_point in sample_points {
406 let mut last_distance = 0.0;
407 for (point, distance) in tree.nearest_neighbor_iter_with_distance_2(sample_point) {
408 assert_eq!(point.distance_2(&sample_point), distance);
409 assert!(last_distance < distance);
410 last_distance = distance;
411 }
412 }
413 }
414}