use super::BvhBuildError;
use super::source::BvhSource;
use crate::{Aabb, Vec3};
use std::ops::Range;
use super::build;
use build::{checked_index, checked_u32};
#[cfg(test)]
#[path = "bvh_tests.rs"]
mod tests;
#[path = "bvh_query.rs"]
mod query;
#[cfg(test)]
use query::point_box_distance_squared;
pub(super) const LEAF_SIZE: usize = 4;
pub(super) const MORTON_LEVELS: u32 = (1 << 21) - 1;
const PARALLEL_SORT_THRESHOLD: usize = 1 << 16;
pub(super) const COUNT_SHIFT: u32 = 29;
pub(super) const INDEX_MASK: u32 = (1 << COUNT_SHIFT) - 1;
pub(super) const MAX_BRANCH_LEFT: u32 = INDEX_MASK - 1;
#[repr(C, align(16))]
#[derive(Clone, Copy, PartialEq, Debug, Default, bytemuck::Pod, bytemuck::Zeroable)]
pub struct BvhNode {
pub min_left: [f32; 4],
pub max_radius: [f32; 4],
}
impl BvhNode {
#[inline]
pub(super) fn leaf(bounds: Aabb, first: u32, count: u32, max_radius: f32) -> Self {
let metadata = first | (count << COUNT_SHIFT);
Self {
min_left: [
bounds.min.x,
bounds.min.y,
bounds.min.z,
f32::from_bits(metadata),
],
max_radius: [bounds.max.x, bounds.max.y, bounds.max.z, max_radius],
}
}
#[inline]
pub(super) fn branch(bounds: Aabb, left: u32, max_radius: f32) -> Self {
Self {
min_left: [
bounds.min.x,
bounds.min.y,
bounds.min.z,
f32::from_bits(left),
],
max_radius: [bounds.max.x, bounds.max.y, bounds.max.z, max_radius],
}
}
#[must_use]
#[inline]
pub fn bounds(self) -> Aabb {
Aabb::new(
Vec3::from_array([self.min_left[0], self.min_left[1], self.min_left[2]]),
Vec3::from_array([self.max_radius[0], self.max_radius[1], self.max_radius[2]]),
)
}
#[must_use]
#[inline]
pub fn maximum_radius(self) -> f32 {
self.max_radius[3]
}
#[must_use]
#[inline]
pub fn is_leaf(self) -> bool {
self.min_left[3].to_bits() >> COUNT_SHIFT != 0
}
#[must_use]
#[inline]
pub fn children(self) -> Option<(u32, u32)> {
if self.is_leaf() {
return None;
}
let left = self.min_left[3].to_bits() & INDEX_MASK;
Some((left, left + 1))
}
#[must_use]
#[inline]
pub fn primitive_range(self) -> Option<Range<u32>> {
if !self.is_leaf() {
return None;
}
let metadata = self.min_left[3].to_bits();
let first = metadata & INDEX_MASK;
let count = metadata >> COUNT_SHIFT;
Some(first..first + count)
}
}
#[derive(Clone, PartialEq, Debug, Default)]
pub struct Bvh {
pub nodes: Vec<BvhNode>,
pub primitive_indices: Vec<u32>,
pub escape: Vec<u32>,
}
#[derive(Clone, Debug, Default)]
pub struct BvhBuildScratch {
entries: Vec<MortonEntry>,
radix: Vec<MortonEntry>,
stack: Vec<u32>,
}
impl Bvh {
pub const ESCAPE_END: u32 = u32::MAX;
pub fn build<S: BvhSource + ?Sized>(source: &S) -> Result<Self, BvhBuildError> {
let mut hierarchy = Self::default();
hierarchy.rebuild(source, &mut BvhBuildScratch::default())?;
Ok(hierarchy)
}
pub fn rebuild<S: BvhSource + ?Sized>(
&mut self,
source: &S,
scratch: &mut BvhBuildScratch,
) -> Result<(), BvhBuildError> {
self.nodes.clear();
self.primitive_indices.clear();
self.escape.clear();
scratch.entries.clear();
let count = checked_u32("BVH source primitive", source.len(), u32::MAX)?;
let scene_bounds = build::scene_bounds(source, count);
if scene_bounds.is_empty() {
return Ok(());
}
build::extend_keys(&mut scratch.entries, source, count, scene_bounds);
checked_index("BVH primitive table", scratch.entries.len(), INDEX_MASK)?;
super::radix::sort_by_code(
&mut scratch.entries,
&mut scratch.radix,
PARALLEL_SORT_THRESHOLD,
);
let node_capacity = scratch.entries.len().checked_mul(2).ok_or(BvhBuildError {
resource: "BVH node capacity",
index: u64::MAX,
maximum: INDEX_MASK,
})?;
checked_index("BVH node table", node_capacity, INDEX_MASK)?;
self.nodes.reserve(node_capacity);
self.primitive_indices.reserve(scratch.entries.len());
self.nodes.push(BvhNode::default());
build::build_range(&scratch.entries, source, 0, scratch.entries.len(), 0, self)?;
self.thread_escapes(&mut scratch.stack);
Ok(())
}
pub fn refit<S: BvhSource + ?Sized>(&mut self, source: &S) -> Result<(), BvhBuildError> {
for node_index in 0..self.nodes.len() {
let Some(node) = self.nodes.get(node_index).copied() else {
continue;
};
let Some(range) = node.primitive_range() else {
continue;
};
let first = range.start;
let count = range.end - range.start;
let mut aggregate = Aabb::EMPTY;
let mut maximum_radius = 0.0f32;
for offset in range {
let Some(&source_row) = self.primitive_indices.get(offset as usize) else {
return Err(refit_error("BVH primitive table", u64::from(offset)));
};
if source_row as usize >= source.len() {
return Err(refit_error("BVH source primitive", u64::from(source_row)));
}
let bound = source.bound(source_row);
if !build::valid_bound(&bound) {
return Err(refit_error("BVH source bounds", u64::from(source_row)));
}
aggregate = aggregate.union(&bound);
maximum_radius = maximum_radius.max(bound.half_extents().max_element());
}
self.nodes[node_index] = BvhNode::leaf(aggregate, first, count, maximum_radius);
}
for node_index in (0..self.nodes.len()).rev() {
let Some(node) = self.nodes.get(node_index).copied() else {
continue;
};
let Some((left, right)) = node.children() else {
continue;
};
let Some(left_node) = self.nodes.get(left as usize).copied() else {
return Err(refit_error("BVH left child", u64::from(left)));
};
let Some(right_node) = self.nodes.get(right as usize).copied() else {
return Err(refit_error("BVH right child", u64::from(right)));
};
self.nodes[node_index] = BvhNode::branch(
left_node.bounds().union(&right_node.bounds()),
left,
left_node.maximum_radius().max(right_node.maximum_radius()),
);
}
Ok(())
}
fn thread_escapes(&mut self, stack: &mut Vec<u32>) {
self.escape.resize(self.nodes.len(), Self::ESCAPE_END);
stack.clear();
if self.nodes.is_empty() {
return;
}
stack.push(0);
while let Some(index) = stack.pop() {
let Some(node) = self.nodes.get(index as usize).copied() else {
continue;
};
let Some((left, right)) = node.children() else {
continue;
};
let Some(&after) = self.escape.get(index as usize) else {
continue;
};
if let Some(slot) = self.escape.get_mut(left as usize) {
*slot = right;
}
if let Some(slot) = self.escape.get_mut(right as usize) {
*slot = after;
}
stack.push(right);
stack.push(left);
}
}
}
#[inline]
fn refit_error(resource: &'static str, index: u64) -> BvhBuildError {
BvhBuildError {
resource,
index,
maximum: u32::MAX,
}
}
#[derive(Clone, Copy, Debug)]
pub(super) struct MortonEntry {
pub(super) code: u64,
pub(super) source: u32,
}
impl MortonEntry {
pub(super) const PLACEHOLDER: Self = Self { code: 0, source: 0 };
}