1use core::cmp::Ordering;
11use core::ops::ControlFlow;
12use std::collections::BinaryHeap;
13
14use axiolid_core::{Aabb, Ray3, Scalar};
15
16use crate::{RayHit, SpatialIndex, SpatialItem};
17
18const LEAF_SIZE: usize = 8;
19
20#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
22pub struct SpatialQueryStats {
23 pub visited_nodes: usize,
25 pub tested_items: usize,
27}
28
29#[derive(Debug, Clone, PartialEq)]
31pub struct CandidatePair<K> {
32 pub a: K,
34 pub b: K,
36 pub lower_bound: Scalar,
38}
39
40#[derive(Debug, Clone, PartialEq)]
42pub struct PairCandidates<K> {
43 pub pairs: Vec<CandidatePair<K>>,
45 pub stats: SpatialQueryStats,
47}
48
49#[derive(Debug, Clone, PartialEq)]
51pub struct NearestCandidate<K> {
52 pub key: K,
54 pub lower_bound: Scalar,
56 pub stats: SpatialQueryStats,
58}
59
60#[derive(Debug)]
61enum NodeKind {
62 Leaf(Vec<usize>),
63 Branch { left: usize, right: usize },
64}
65
66#[derive(Debug)]
67struct Node {
68 bounds: Aabb,
69 kind: NodeKind,
70}
71
72#[derive(Debug)]
79pub struct Bvh<K> {
80 items: Vec<SpatialItem<K>>,
81 nodes: Vec<Node>,
82 root: Option<usize>,
83 rejected_items: usize,
84}
85
86impl<K> Bvh<K> {
87 pub fn build(items: impl IntoIterator<Item = SpatialItem<K>>) -> Self {
89 let mut rejected_items = 0;
90 let items = items
91 .into_iter()
92 .filter(|item| {
93 let accepted = item.bounds.is_finite() && !item.bounds.is_empty();
94 rejected_items += usize::from(!accepted);
95 accepted
96 })
97 .collect();
98 let mut tree = Self {
99 items,
100 nodes: Vec::new(),
101 root: None,
102 rejected_items,
103 };
104 if !tree.items.is_empty() {
105 let indices = (0..tree.items.len()).collect();
106 tree.root = Some(tree.build_node(indices));
107 }
108 tree
109 }
110
111 pub fn len(&self) -> usize {
113 self.items.len()
114 }
115
116 pub fn is_empty(&self) -> bool {
118 self.items.is_empty()
119 }
120
121 pub fn rejected_items(&self) -> usize {
123 self.rejected_items
124 }
125
126 pub fn item(&self, index: usize) -> Option<&SpatialItem<K>> {
128 self.items.get(index)
129 }
130
131 fn build_node(&mut self, mut indices: Vec<usize>) -> usize {
132 let bounds = union_bounds(indices.iter().map(|&index| self.items[index].bounds));
133 let node_index = self.nodes.len();
134 self.nodes.push(Node {
135 bounds,
136 kind: NodeKind::Leaf(Vec::new()),
137 });
138 if indices.len() <= LEAF_SIZE {
139 self.nodes[node_index].kind = NodeKind::Leaf(indices);
140 return node_index;
141 }
142
143 let extent = bounds.diagonal();
144 let axis = if extent.x >= extent.y && extent.x >= extent.z {
145 0
146 } else if extent.y >= extent.z {
147 1
148 } else {
149 2
150 };
151 indices.sort_unstable_by(|&left, &right| {
152 component(self.items[left].bounds.center(), axis)
153 .total_cmp(&component(self.items[right].bounds.center(), axis))
154 .then_with(|| left.cmp(&right))
155 });
156 let right_indices = indices.split_off(indices.len() / 2);
157 let left = self.build_node(indices);
158 let right = self.build_node(right_indices);
159 self.nodes[node_index].kind = NodeKind::Branch { left, right };
160 node_index
161 }
162}
163
164impl<K: Clone> Bvh<K> {
165 pub fn query_aabb(&self, probe: &Aabb, out: &mut Vec<usize>) {
177 out.clear();
178 let Some(root) = self.root else {
179 return;
180 };
181 let mut stack = vec![root];
182 while let Some(index) = stack.pop() {
183 let node = &self.nodes[index];
184 if !node.bounds.intersects(probe) {
185 continue;
186 }
187 match &node.kind {
188 NodeKind::Leaf(items) => out.extend_from_slice(items),
189 NodeKind::Branch { left, right } => {
190 stack.push(*left);
191 stack.push(*right);
192 }
193 }
194 }
195 }
196
197 pub fn overlap_pairs(&self, min_penetration: Scalar) -> PairCandidates<K> {
198 assert!(
199 min_penetration.is_finite() && min_penetration >= 0.0,
200 "minimum penetration must be finite and non-negative"
201 );
202 self.collect_pairs(|left, right| {
203 penetrates(left, right, min_penetration).then_some(left.gap(right))
204 })
205 }
206
207 pub fn pairs_within_distance(&self, max_distance: Scalar) -> PairCandidates<K> {
213 assert!(
214 max_distance.is_finite() && max_distance >= 0.0,
215 "maximum distance must be finite and non-negative"
216 );
217 self.collect_pairs(|left, right| {
218 let gap = left.gap(right);
219 (gap <= max_distance).then_some(gap)
220 })
221 }
222
223 pub fn nearest_to(
228 &self,
229 query: &Aabb,
230 accept: impl Fn(&K) -> bool,
231 ) -> Option<NearestCandidate<K>> {
232 assert!(
233 query.is_finite() && !query.is_empty(),
234 "nearest-neighbour query bounds must be finite and non-empty"
235 );
236 let root = self.root?;
237 let mut stats = SpatialQueryStats::default();
238 let mut pending = BinaryHeap::new();
239 pending.push(NearestQueueEntry::new(
240 query.gap(&self.nodes[root].bounds),
241 root,
242 ));
243 let mut best = None;
244
245 while let Some(entry) = pending.pop() {
246 stats.visited_nodes += 1;
247 if best.is_some_and(|(distance, _)| entry.distance > distance) {
248 break;
249 }
250 match &self.nodes[entry.node].kind {
251 NodeKind::Leaf(indices) => {
252 for &index in indices {
253 if !accept(&self.items[index].key) {
254 continue;
255 }
256 stats.tested_items += 1;
257 let distance = query.gap(&self.items[index].bounds);
258 if best.is_none_or(|(current, current_index)| {
259 distance < current || (distance == current && index < current_index)
260 }) {
261 best = Some((distance, index));
262 }
263 }
264 }
265 NodeKind::Branch { left, right } => {
266 for child in [*left, *right] {
267 let distance = query.gap(&self.nodes[child].bounds);
268 if best.is_none_or(|(current, _)| distance <= current) {
269 pending.push(NearestQueueEntry::new(distance, child));
270 }
271 }
272 }
273 }
274 }
275
276 best.map(|(lower_bound, index)| NearestCandidate {
277 key: self.items[index].key.clone(),
278 lower_bound,
279 stats,
280 })
281 }
282
283 fn collect_pairs(&self, matches: impl Fn(&Aabb, &Aabb) -> Option<Scalar>) -> PairCandidates<K> {
284 let mut pairs = Vec::new();
285 let mut stats = SpatialQueryStats::default();
286 let Some(root) = self.root else {
287 return PairCandidates { pairs, stats };
288 };
289
290 for index in 0..self.items.len() {
291 let bounds = &self.items[index].bounds;
292 let mut stack = vec![root];
293 let mut matches_for_item = Vec::new();
294 while let Some(node_index) = stack.pop() {
295 stats.visited_nodes += 1;
296 let node = &self.nodes[node_index];
297 if matches(bounds, &node.bounds).is_none() {
298 continue;
299 }
300 match &node.kind {
301 NodeKind::Leaf(indices) => {
302 for &other_index in indices {
303 if other_index <= index {
304 continue;
305 }
306 stats.tested_items += 1;
307 if let Some(lower_bound) =
308 matches(bounds, &self.items[other_index].bounds)
309 {
310 matches_for_item.push((other_index, lower_bound));
311 }
312 }
313 }
314 NodeKind::Branch { left, right } => {
315 stack.push(*right);
316 stack.push(*left);
317 }
318 }
319 }
320 matches_for_item.sort_unstable_by_key(|(other_index, _)| *other_index);
321 pairs.extend(
322 matches_for_item
323 .into_iter()
324 .map(|(other_index, lower_bound)| CandidatePair {
325 a: self.items[index].key.clone(),
326 b: self.items[other_index].key.clone(),
327 lower_bound,
328 }),
329 );
330 }
331 PairCandidates { pairs, stats }
332 }
333}
334
335impl<K> SpatialIndex<K> for Bvh<K>
336where
337 K: core::fmt::Debug + Send + Sync,
338{
339 fn visit_aabb(&self, query: &Aabb, visitor: &mut dyn FnMut(&K) -> ControlFlow<()>) {
340 if query.is_empty() || !query.is_finite() {
341 return;
342 }
343 let Some(root) = self.root else {
344 return;
345 };
346 let mut stack = vec![root];
347 while let Some(node_index) = stack.pop() {
348 let node = &self.nodes[node_index];
349 if !query.intersects(&node.bounds) {
350 continue;
351 }
352 match &node.kind {
353 NodeKind::Leaf(indices) => {
354 for &item_index in indices {
355 let item = &self.items[item_index];
356 if query.intersects(&item.bounds) && visitor(&item.key).is_break() {
357 return;
358 }
359 }
360 }
361 NodeKind::Branch { left, right } => {
362 stack.push(*right);
363 stack.push(*left);
364 }
365 }
366 }
367 }
368
369 fn visit_ray(&self, ray: &Ray3, visitor: &mut dyn FnMut(RayHit<&K>) -> ControlFlow<()>) {
370 let Some(root) = self.root else {
371 return;
372 };
373 let Some(root_distance) = ray_aabb_entry(ray, &self.nodes[root].bounds) else {
374 return;
375 };
376 let mut pending = BinaryHeap::new();
377 pending.push(RayQueueEntry::node(root_distance, root));
378 while let Some(entry) = pending.pop() {
379 match entry.kind {
380 RayQueueKind::Node(node_index) => match &self.nodes[node_index].kind {
381 NodeKind::Leaf(indices) => {
382 for &item_index in indices {
383 if let Some(distance) =
384 ray_aabb_entry(ray, &self.items[item_index].bounds)
385 {
386 pending.push(RayQueueEntry::item(distance, item_index));
387 }
388 }
389 }
390 NodeKind::Branch { left, right } => {
391 for child in [*left, *right] {
392 if let Some(distance) = ray_aabb_entry(ray, &self.nodes[child].bounds) {
393 pending.push(RayQueueEntry::node(distance, child));
394 }
395 }
396 }
397 },
398 RayQueueKind::Item(item_index) => {
399 if visitor(RayHit {
400 key: &self.items[item_index].key,
401 distance: entry.distance,
402 })
403 .is_break()
404 {
405 return;
406 }
407 }
408 }
409 }
410 }
411
412 fn len(&self) -> usize {
413 self.len()
414 }
415}
416
417#[derive(Debug, Clone, Copy)]
418struct NearestQueueEntry {
419 distance: Scalar,
420 node: usize,
421}
422
423impl NearestQueueEntry {
424 const fn new(distance: Scalar, node: usize) -> Self {
425 Self { distance, node }
426 }
427}
428
429impl PartialEq for NearestQueueEntry {
430 fn eq(&self, other: &Self) -> bool {
431 self.distance == other.distance && self.node == other.node
432 }
433}
434impl Eq for NearestQueueEntry {}
435impl PartialOrd for NearestQueueEntry {
436 fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
437 Some(self.cmp(other))
438 }
439}
440impl Ord for NearestQueueEntry {
441 fn cmp(&self, other: &Self) -> Ordering {
442 other
443 .distance
444 .total_cmp(&self.distance)
445 .then_with(|| other.node.cmp(&self.node))
446 }
447}
448
449#[derive(Debug, Clone, Copy)]
450enum RayQueueKind {
451 Node(usize),
452 Item(usize),
453}
454
455#[derive(Debug, Clone, Copy)]
456struct RayQueueEntry {
457 distance: Scalar,
458 kind: RayQueueKind,
459}
460
461impl RayQueueEntry {
462 fn node(distance: Scalar, index: usize) -> Self {
463 Self {
464 distance,
465 kind: RayQueueKind::Node(index),
466 }
467 }
468
469 fn item(distance: Scalar, index: usize) -> Self {
470 Self {
471 distance,
472 kind: RayQueueKind::Item(index),
473 }
474 }
475
476 fn order_key(self) -> (u8, usize) {
477 match self.kind {
478 RayQueueKind::Node(index) => (0, index),
481 RayQueueKind::Item(index) => (1, index),
482 }
483 }
484}
485
486impl PartialEq for RayQueueEntry {
487 fn eq(&self, other: &Self) -> bool {
488 self.distance == other.distance && self.order_key() == other.order_key()
489 }
490}
491impl Eq for RayQueueEntry {}
492impl PartialOrd for RayQueueEntry {
493 fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
494 Some(self.cmp(other))
495 }
496}
497impl Ord for RayQueueEntry {
498 fn cmp(&self, other: &Self) -> Ordering {
499 other
500 .distance
501 .total_cmp(&self.distance)
502 .then_with(|| other.order_key().cmp(&self.order_key()))
503 }
504}
505
506fn union_bounds(bounds: impl IntoIterator<Item = Aabb>) -> Aabb {
507 let mut union = Aabb::empty();
508 for bounds in bounds {
509 union.union(&bounds);
510 }
511 union
512}
513
514fn component(point: axiolid_core::Point3, axis: usize) -> Scalar {
515 match axis {
516 0 => point.x,
517 1 => point.y,
518 _ => point.z,
519 }
520}
521
522fn penetrates(left: &Aabb, right: &Aabb, minimum: Scalar) -> bool {
523 let overlap = left.max.min(right.max) - left.min.max(right.min);
524 overlap.x >= minimum && overlap.y >= minimum && overlap.z >= minimum
525}
526
527fn ray_aabb_entry(ray: &Ray3, bounds: &Aabb) -> Option<Scalar> {
528 if bounds.is_empty()
529 || !bounds.is_finite()
530 || !ray.origin.is_finite()
531 || !ray.direction.is_finite()
532 {
533 return None;
534 }
535 let mut entry = Scalar::NEG_INFINITY;
536 let mut exit = Scalar::INFINITY;
537 for (origin, direction, min, max) in [
538 (ray.origin.x, ray.direction.x, bounds.min.x, bounds.max.x),
539 (ray.origin.y, ray.direction.y, bounds.min.y, bounds.max.y),
540 (ray.origin.z, ray.direction.z, bounds.min.z, bounds.max.z),
541 ] {
542 if direction == 0.0 {
543 if origin < min || origin > max {
544 return None;
545 }
546 continue;
547 }
548 let first = (min - origin) / direction;
549 let second = (max - origin) / direction;
550 entry = entry.max(first.min(second));
551 exit = exit.min(first.max(second));
552 if exit < entry {
553 return None;
554 }
555 }
556 (exit >= 0.0).then_some(entry.max(0.0))
557}