Skip to main content

proof_engine/svogi/
octree.rs

1use glam::{Vec3, Vec4, UVec3, IVec3};
2use std::collections::VecDeque;
3
4/// Axis-aligned bounding box.
5#[derive(Debug, Clone, Copy)]
6pub struct Aabb {
7    pub min: Vec3,
8    pub max: Vec3,
9}
10
11impl Aabb {
12    pub fn new(min: Vec3, max: Vec3) -> Self {
13        Self { min, max }
14    }
15
16    pub fn center(&self) -> Vec3 {
17        (self.min + self.max) * 0.5
18    }
19
20    pub fn size(&self) -> Vec3 {
21        self.max - self.min
22    }
23
24    pub fn half_size(&self) -> Vec3 {
25        self.size() * 0.5
26    }
27
28    pub fn contains(&self, point: Vec3) -> bool {
29        point.x >= self.min.x && point.x <= self.max.x
30            && point.y >= self.min.y && point.y <= self.max.y
31            && point.z >= self.min.z && point.z <= self.max.z
32    }
33
34    pub fn intersects(&self, other: &Aabb) -> bool {
35        self.min.x <= other.max.x && self.max.x >= other.min.x
36            && self.min.y <= other.max.y && self.max.y >= other.min.y
37            && self.min.z <= other.max.z && self.max.z >= other.min.z
38    }
39
40    /// Return the child octant bounding box (0..7).
41    pub fn subdivide_octant(&self, octant: u8) -> Aabb {
42        let c = self.center();
43        let min = Vec3::new(
44            if octant & 1 == 0 { self.min.x } else { c.x },
45            if octant & 2 == 0 { self.min.y } else { c.y },
46            if octant & 4 == 0 { self.min.z } else { c.z },
47        );
48        let max = Vec3::new(
49            if octant & 1 == 0 { c.x } else { self.max.x },
50            if octant & 2 == 0 { c.y } else { self.max.y },
51            if octant & 4 == 0 { c.z } else { self.max.z },
52        );
53        Aabb { min, max }
54    }
55
56    /// Intersect a ray with this AABB, returning (t_enter, t_exit). None if miss.
57    pub fn intersect_ray(&self, origin: Vec3, inv_dir: Vec3) -> Option<(f32, f32)> {
58        let t1 = (self.min - origin) * inv_dir;
59        let t2 = (self.max - origin) * inv_dir;
60        let t_min = t1.min(t2);
61        let t_max = t1.max(t2);
62        let t_enter = t_min.x.max(t_min.y).max(t_min.z);
63        let t_exit = t_max.x.min(t_max.y).min(t_max.z);
64        if t_enter <= t_exit && t_exit >= 0.0 {
65            Some((t_enter.max(0.0), t_exit))
66        } else {
67            None
68        }
69    }
70}
71
72/// Data stored per voxel.
73#[derive(Debug, Clone, Copy)]
74pub struct VoxelData {
75    pub radiance: Vec4,
76    pub normal: Vec3,
77    pub opacity: f32,
78    pub sh_coeffs: [f32; 9],
79}
80
81impl Default for VoxelData {
82    fn default() -> Self {
83        Self {
84            radiance: Vec4::ZERO,
85            normal: Vec3::ZERO,
86            opacity: 0.0,
87            sh_coeffs: [0.0; 9],
88        }
89    }
90}
91
92impl VoxelData {
93    pub fn is_empty(&self) -> bool {
94        self.opacity <= 0.0
95    }
96
97    pub fn average(datas: &[&VoxelData]) -> VoxelData {
98        if datas.is_empty() {
99            return VoxelData::default();
100        }
101        let n = datas.len() as f32;
102        let mut result = VoxelData::default();
103        for d in datas {
104            result.radiance += d.radiance;
105            result.normal += d.normal;
106            result.opacity += d.opacity;
107            for i in 0..9 {
108                result.sh_coeffs[i] += d.sh_coeffs[i];
109            }
110        }
111        result.radiance /= n;
112        result.normal = if result.normal.length_squared() > 1e-8 {
113            result.normal.normalize()
114        } else {
115            Vec3::ZERO
116        };
117        result.opacity /= n;
118        for i in 0..9 {
119            result.sh_coeffs[i] /= n;
120        }
121        result
122    }
123}
124
125/// A single node in the sparse voxel octree.
126#[derive(Debug, Clone)]
127pub struct OctreeNode {
128    pub children: [Option<u32>; 8],
129    pub data: VoxelData,
130    pub level: u8,
131    pub morton_code: u64,
132}
133
134impl OctreeNode {
135    pub fn new_leaf(data: VoxelData, level: u8, morton_code: u64) -> Self {
136        Self {
137            children: [None; 8],
138            data,
139            level,
140            morton_code,
141        }
142    }
143
144    pub fn new_internal(level: u8, morton_code: u64) -> Self {
145        Self {
146            children: [None; 8],
147            data: VoxelData::default(),
148            level,
149            morton_code,
150        }
151    }
152
153    pub fn is_leaf(&self) -> bool {
154        self.children.iter().all(|c| c.is_none())
155    }
156
157    pub fn has_children(&self) -> bool {
158        self.children.iter().any(|c| c.is_some())
159    }
160}
161
162/// Morton code encoding: interleave bits of x, y, z into a 64-bit code.
163pub fn morton_encode_3d(x: u32, y: u32, z: u32) -> u64 {
164    fn split_by_3(mut v: u32) -> u64 {
165        let mut v = v as u64 & 0x1fffff; // 21 bits
166        v = (v | (v << 32)) & 0x1f00000000ffff;
167        v = (v | (v << 16)) & 0x1f0000ff0000ff;
168        v = (v | (v << 8))  & 0x100f00f00f00f00f;
169        v = (v | (v << 4))  & 0x10c30c30c30c30c3;
170        v = (v | (v << 2))  & 0x1249249249249249;
171        v
172    }
173    split_by_3(x) | (split_by_3(y) << 1) | (split_by_3(z) << 2)
174}
175
176/// Morton code decoding: extract x, y, z from a 64-bit code.
177pub fn morton_decode_3d(code: u64) -> (u32, u32, u32) {
178    fn compact_by_3(mut v: u64) -> u32 {
179        v &= 0x1249249249249249;
180        v = (v | (v >> 2))  & 0x10c30c30c30c30c3;
181        v = (v | (v >> 4))  & 0x100f00f00f00f00f;
182        v = (v | (v >> 8))  & 0x1f0000ff0000ff;
183        v = (v | (v >> 16)) & 0x1f00000000ffff;
184        v = (v | (v >> 32)) & 0x1fffff;
185        v as u32
186    }
187    (compact_by_3(code), compact_by_3(code >> 1), compact_by_3(code >> 2))
188}
189
190/// Dense 3D voxel grid.
191#[derive(Debug, Clone)]
192pub struct VoxelGrid {
193    pub data: Vec<VoxelData>,
194    pub resolution: UVec3,
195}
196
197impl VoxelGrid {
198    pub fn new(resolution: UVec3) -> Self {
199        let count = (resolution.x * resolution.y * resolution.z) as usize;
200        Self {
201            data: vec![VoxelData::default(); count],
202            resolution,
203        }
204    }
205
206    pub fn index(&self, x: u32, y: u32, z: u32) -> usize {
207        (z * self.resolution.y * self.resolution.x + y * self.resolution.x + x) as usize
208    }
209
210    pub fn get(&self, x: u32, y: u32, z: u32) -> &VoxelData {
211        &self.data[self.index(x, y, z)]
212    }
213
214    pub fn get_mut(&mut self, x: u32, y: u32, z: u32) -> &mut VoxelData {
215        let idx = self.index(x, y, z);
216        &mut self.data[idx]
217    }
218
219    pub fn set(&mut self, x: u32, y: u32, z: u32, data: VoxelData) {
220        let idx = self.index(x, y, z);
221        self.data[idx] = data;
222    }
223
224    pub fn in_bounds(&self, x: i32, y: i32, z: i32) -> bool {
225        x >= 0 && y >= 0 && z >= 0
226            && (x as u32) < self.resolution.x
227            && (y as u32) < self.resolution.y
228            && (z as u32) < self.resolution.z
229    }
230
231    pub fn filled_count(&self) -> usize {
232        self.data.iter().filter(|v| !v.is_empty()).count()
233    }
234}
235
236/// The sparse voxel octree.
237#[derive(Debug, Clone)]
238pub struct SparseVoxelOctree {
239    pub nodes: Vec<OctreeNode>,
240    pub root: u32,
241    pub max_depth: u8,
242    pub world_bounds: Aabb,
243}
244
245impl SparseVoxelOctree {
246    pub fn new(world_bounds: Aabb, max_depth: u8) -> Self {
247        let root_node = OctreeNode::new_internal(0, 0);
248        Self {
249            nodes: vec![root_node],
250            root: 0,
251            max_depth,
252            world_bounds,
253        }
254    }
255
256    /// Determine which octant a point falls into within the given bounds.
257    fn octant_for_point(bounds: &Aabb, point: Vec3) -> u8 {
258        let c = bounds.center();
259        let mut octant = 0u8;
260        if point.x >= c.x { octant |= 1; }
261        if point.y >= c.y { octant |= 2; }
262        if point.z >= c.z { octant |= 4; }
263        octant
264    }
265
266    /// Insert a voxel at the given world position, creating intermediate nodes as needed.
267    pub fn insert(&mut self, position: Vec3, data: VoxelData) {
268        if !self.world_bounds.contains(position) {
269            return;
270        }
271        let max_depth = self.max_depth;
272        self.insert_recursive(self.root, self.world_bounds, position, data, 0, max_depth);
273    }
274
275    fn insert_recursive(
276        &mut self,
277        node_idx: u32,
278        bounds: Aabb,
279        position: Vec3,
280        data: VoxelData,
281        depth: u8,
282        max_depth: u8,
283    ) {
284        if depth >= max_depth {
285            self.nodes[node_idx as usize].data = data;
286            return;
287        }
288
289        let octant = Self::octant_for_point(&bounds, position);
290        let child_bounds = bounds.subdivide_octant(octant);
291
292        let child_idx = if let Some(idx) = self.nodes[node_idx as usize].children[octant as usize] {
293            idx
294        } else {
295            let idx = self.nodes.len() as u32;
296            let morton = morton_encode_3d(
297                ((position.x - self.world_bounds.min.x) / self.world_bounds.size().x * ((1u32 << max_depth) as f32)) as u32,
298                ((position.y - self.world_bounds.min.y) / self.world_bounds.size().y * ((1u32 << max_depth) as f32)) as u32,
299                ((position.z - self.world_bounds.min.z) / self.world_bounds.size().z * ((1u32 << max_depth) as f32)) as u32,
300            );
301            let new_node = OctreeNode::new_internal(depth + 1, morton);
302            self.nodes.push(new_node);
303            self.nodes[node_idx as usize].children[octant as usize] = Some(idx);
304            idx
305        };
306
307        self.insert_recursive(child_idx, child_bounds, position, data, depth + 1, max_depth);
308    }
309
310    /// Look up voxel data at a world position, at the given LOD level.
311    /// Level 0 = root (coarsest), max_depth = leaf (finest).
312    pub fn lookup(&self, position: Vec3, level: u8) -> Option<&VoxelData> {
313        if !self.world_bounds.contains(position) {
314            return None;
315        }
316        self.lookup_recursive(self.root, self.world_bounds, position, 0, level)
317    }
318
319    fn lookup_recursive(
320        &self,
321        node_idx: u32,
322        bounds: Aabb,
323        position: Vec3,
324        depth: u8,
325        target_level: u8,
326    ) -> Option<&VoxelData> {
327        let node = &self.nodes[node_idx as usize];
328
329        if depth >= target_level || node.is_leaf() {
330            if node.data.is_empty() && node.is_leaf() && depth < target_level {
331                return None;
332            }
333            return Some(&node.data);
334        }
335
336        let octant = Self::octant_for_point(&bounds, position);
337        if let Some(child_idx) = node.children[octant as usize] {
338            let child_bounds = bounds.subdivide_octant(octant);
339            self.lookup_recursive(child_idx, child_bounds, position, depth + 1, target_level)
340        } else {
341            // No child at this octant; return this node's data if non-empty
342            if !node.data.is_empty() {
343                Some(&node.data)
344            } else {
345                None
346            }
347        }
348    }
349
350    /// Remove a voxel at the given position (clear leaf data).
351    pub fn remove(&mut self, position: Vec3) {
352        if !self.world_bounds.contains(position) {
353            return;
354        }
355        let max_depth = self.max_depth;
356        self.remove_recursive(self.root, self.world_bounds, position, 0, max_depth);
357    }
358
359    fn remove_recursive(
360        &mut self,
361        node_idx: u32,
362        bounds: Aabb,
363        position: Vec3,
364        depth: u8,
365        max_depth: u8,
366    ) -> bool {
367        if depth >= max_depth {
368            self.nodes[node_idx as usize].data = VoxelData::default();
369            return true; // node is now empty
370        }
371
372        let octant = Self::octant_for_point(&bounds, position);
373        let child_idx = match self.nodes[node_idx as usize].children[octant as usize] {
374            Some(idx) => idx,
375            None => return false,
376        };
377
378        let child_bounds = bounds.subdivide_octant(octant);
379        let child_empty = self.remove_recursive(child_idx, child_bounds, position, depth + 1, max_depth);
380
381        if child_empty {
382            let child_node = &self.nodes[child_idx as usize];
383            if child_node.is_leaf() && child_node.data.is_empty() {
384                self.nodes[node_idx as usize].children[octant as usize] = None;
385            }
386        }
387
388        // Prune: check if all children are None
389        self.nodes[node_idx as usize].children.iter().all(|c| c.is_none())
390            && self.nodes[node_idx as usize].data.is_empty()
391    }
392
393    /// Traverse a ray through the octree, collecting hits.
394    pub fn traverse_ray(&self, origin: Vec3, direction: Vec3, max_t: f32) -> Vec<(Vec3, VoxelData, f32)> {
395        let mut results = Vec::new();
396        let dir_safe = Vec3::new(
397            if direction.x.abs() < 1e-8 { 1e-8 } else { direction.x },
398            if direction.y.abs() < 1e-8 { 1e-8 } else { direction.y },
399            if direction.z.abs() < 1e-8 { 1e-8 } else { direction.z },
400        );
401        let inv_dir = Vec3::new(1.0 / dir_safe.x, 1.0 / dir_safe.y, 1.0 / dir_safe.z);
402
403        self.traverse_ray_recursive(
404            self.root,
405            self.world_bounds,
406            origin,
407            direction,
408            inv_dir,
409            max_t,
410            &mut results,
411        );
412
413        results.sort_by(|a, b| a.2.partial_cmp(&b.2).unwrap_or(std::cmp::Ordering::Equal));
414        results
415    }
416
417    fn traverse_ray_recursive(
418        &self,
419        node_idx: u32,
420        bounds: Aabb,
421        origin: Vec3,
422        direction: Vec3,
423        inv_dir: Vec3,
424        max_t: f32,
425        results: &mut Vec<(Vec3, VoxelData, f32)>,
426    ) {
427        let hit = bounds.intersect_ray(origin, inv_dir);
428        let (t_enter, t_exit) = match hit {
429            Some((te, tx)) if te <= max_t => (te, tx),
430            _ => return,
431        };
432
433        let node = &self.nodes[node_idx as usize];
434
435        if node.is_leaf() {
436            if !node.data.is_empty() {
437                let hit_point = origin + direction * t_enter;
438                results.push((hit_point, node.data, t_enter));
439            }
440            return;
441        }
442
443        // Traverse children sorted by t_enter for front-to-back ordering
444        let mut child_order: Vec<(u8, f32)> = Vec::new();
445        for octant in 0..8u8 {
446            if let Some(child_idx) = node.children[octant as usize] {
447                let child_bounds = bounds.subdivide_octant(octant);
448                if let Some((te, _)) = child_bounds.intersect_ray(origin, inv_dir) {
449                    if te <= max_t {
450                        child_order.push((octant, te));
451                    }
452                }
453            }
454        }
455        child_order.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal));
456
457        for (octant, _) in child_order {
458            let child_idx = self.nodes[node_idx as usize].children[octant as usize].unwrap();
459            let child_bounds = bounds.subdivide_octant(octant);
460            self.traverse_ray_recursive(child_idx, child_bounds, origin, direction, inv_dir, max_t, results);
461        }
462    }
463
464    /// Average child voxels into parent for LOD (mipmap). Process the given level.
465    pub fn mipmap_level(&mut self, level: u8) {
466        // Collect nodes at the given level that have children
467        let indices: Vec<u32> = (0..self.nodes.len() as u32)
468            .filter(|&i| {
469                let n = &self.nodes[i as usize];
470                n.level == level && n.has_children()
471            })
472            .collect();
473
474        for idx in indices {
475            let children_indices: Vec<u32> = self.nodes[idx as usize]
476                .children
477                .iter()
478                .filter_map(|c| *c)
479                .collect();
480
481            if children_indices.is_empty() {
482                continue;
483            }
484
485            let child_datas: Vec<&VoxelData> = children_indices
486                .iter()
487                .map(|&ci| &self.nodes[ci as usize].data)
488                .filter(|d| !d.is_empty())
489                .collect();
490
491            if !child_datas.is_empty() {
492                self.nodes[idx as usize].data = VoxelData::average(&child_datas);
493            }
494        }
495    }
496
497    /// Build the full mipmap chain from leaves up to root.
498    pub fn build_mipmaps(&mut self) {
499        if self.max_depth == 0 {
500            return;
501        }
502        for level in (0..self.max_depth).rev() {
503            self.mipmap_level(level);
504        }
505    }
506
507    /// Build a sparse voxel octree from a dense grid.
508    pub fn build_from_voxel_grid(grid: &VoxelGrid, world_bounds: Aabb) -> SparseVoxelOctree {
509        let max_dim = grid.resolution.x.max(grid.resolution.y).max(grid.resolution.z);
510        let max_depth = (max_dim as f32).log2().ceil() as u8;
511        let max_depth = max_depth.max(1);
512
513        let mut octree = SparseVoxelOctree::new(world_bounds, max_depth);
514        let voxel_size = world_bounds.size() / Vec3::new(
515            grid.resolution.x as f32,
516            grid.resolution.y as f32,
517            grid.resolution.z as f32,
518        );
519
520        for z in 0..grid.resolution.z {
521            for y in 0..grid.resolution.y {
522                for x in 0..grid.resolution.x {
523                    let voxel = grid.get(x, y, z);
524                    if !voxel.is_empty() {
525                        let pos = world_bounds.min + Vec3::new(
526                            (x as f32 + 0.5) * voxel_size.x,
527                            (y as f32 + 0.5) * voxel_size.y,
528                            (z as f32 + 0.5) * voxel_size.z,
529                        );
530                        octree.insert(pos, *voxel);
531                    }
532                }
533            }
534        }
535
536        octree.build_mipmaps();
537        octree
538    }
539
540    /// Total memory usage in bytes (approximate).
541    pub fn memory_usage(&self) -> usize {
542        self.nodes.len() * std::mem::size_of::<OctreeNode>()
543    }
544
545    /// Total number of nodes.
546    pub fn node_count(&self) -> usize {
547        self.nodes.len()
548    }
549
550    /// Count of leaf nodes (no children).
551    pub fn leaf_count(&self) -> usize {
552        self.nodes.iter().filter(|n| n.is_leaf()).count()
553    }
554
555    /// Maximum depth actually present.
556    pub fn depth(&self) -> u8 {
557        self.nodes.iter().map(|n| n.level).max().unwrap_or(0)
558    }
559
560    /// Iterate all occupied leaf voxels.
561    pub fn iter_leaves(&self) -> OctreeIterator<'_> {
562        OctreeIterator {
563            octree: self,
564            stack: vec![(self.root, self.world_bounds)],
565        }
566    }
567
568    /// Sample the octree at a world position with fractional LOD.
569    /// Interpolates between two integer LOD levels.
570    pub fn sample_lod(&self, position: Vec3, lod: f32) -> Option<VoxelData> {
571        let lod_low = lod.floor() as u8;
572        let lod_high = lod.ceil() as u8;
573        let frac = lod - lod.floor();
574
575        let data_low = self.lookup(position, lod_low);
576        let data_high = self.lookup(position, lod_high);
577
578        match (data_low, data_high) {
579            (Some(a), Some(b)) => {
580                let mut result = VoxelData::default();
581                result.radiance = a.radiance * (1.0 - frac) + b.radiance * frac;
582                result.normal = if (a.normal * (1.0 - frac) + b.normal * frac).length_squared() > 1e-8 {
583                    (a.normal * (1.0 - frac) + b.normal * frac).normalize()
584                } else {
585                    Vec3::ZERO
586                };
587                result.opacity = a.opacity * (1.0 - frac) + b.opacity * frac;
588                for i in 0..9 {
589                    result.sh_coeffs[i] = a.sh_coeffs[i] * (1.0 - frac) + b.sh_coeffs[i] * frac;
590                }
591                Some(result)
592            }
593            (Some(a), None) => Some(*a),
594            (None, Some(b)) => Some(*b),
595            (None, None) => None,
596        }
597    }
598}
599
600/// Iterator over all occupied leaf voxels.
601pub struct OctreeIterator<'a> {
602    octree: &'a SparseVoxelOctree,
603    stack: Vec<(u32, Aabb)>,
604}
605
606impl<'a> Iterator for OctreeIterator<'a> {
607    type Item = (Vec3, &'a VoxelData);
608
609    fn next(&mut self) -> Option<Self::Item> {
610        while let Some((node_idx, bounds)) = self.stack.pop() {
611            let node = &self.octree.nodes[node_idx as usize];
612
613            if node.is_leaf() {
614                if !node.data.is_empty() {
615                    return Some((bounds.center(), &node.data));
616                }
617                continue;
618            }
619
620            for octant in (0..8u8).rev() {
621                if let Some(child_idx) = node.children[octant as usize] {
622                    self.stack.push((child_idx, bounds.subdivide_octant(octant)));
623                }
624            }
625        }
626        None
627    }
628}
629
630#[cfg(test)]
631mod tests {
632    use super::*;
633
634    #[test]
635    fn test_aabb_basics() {
636        let aabb = Aabb::new(Vec3::ZERO, Vec3::splat(10.0));
637        assert_eq!(aabb.center(), Vec3::splat(5.0));
638        assert_eq!(aabb.size(), Vec3::splat(10.0));
639        assert!(aabb.contains(Vec3::splat(5.0)));
640        assert!(!aabb.contains(Vec3::splat(11.0)));
641    }
642
643    #[test]
644    fn test_aabb_intersects() {
645        let a = Aabb::new(Vec3::ZERO, Vec3::splat(5.0));
646        let b = Aabb::new(Vec3::splat(3.0), Vec3::splat(8.0));
647        let c = Aabb::new(Vec3::splat(6.0), Vec3::splat(10.0));
648        assert!(a.intersects(&b));
649        assert!(!a.intersects(&c));
650    }
651
652    #[test]
653    fn test_morton_roundtrip() {
654        for x in 0..16u32 {
655            for y in 0..16u32 {
656                for z in 0..16u32 {
657                    let code = morton_encode_3d(x, y, z);
658                    let (dx, dy, dz) = morton_decode_3d(code);
659                    assert_eq!((x, y, z), (dx, dy, dz), "Morton roundtrip failed for ({x},{y},{z})");
660                }
661            }
662        }
663    }
664
665    #[test]
666    fn test_morton_ordering() {
667        // Morton codes should maintain Z-curve ordering
668        let c1 = morton_encode_3d(0, 0, 0);
669        let c2 = morton_encode_3d(1, 0, 0);
670        let c3 = morton_encode_3d(0, 1, 0);
671        assert!(c1 < c2);
672        assert!(c1 < c3);
673    }
674
675    #[test]
676    fn test_insert_and_lookup() {
677        let bounds = Aabb::new(Vec3::ZERO, Vec3::splat(8.0));
678        let mut octree = SparseVoxelOctree::new(bounds, 3);
679
680        let data = VoxelData {
681            radiance: Vec4::new(1.0, 0.0, 0.0, 1.0),
682            normal: Vec3::Y,
683            opacity: 1.0,
684            sh_coeffs: [0.0; 9],
685        };
686
687        octree.insert(Vec3::new(1.0, 1.0, 1.0), data);
688
689        let result = octree.lookup(Vec3::new(1.0, 1.0, 1.0), 3);
690        assert!(result.is_some());
691        let r = result.unwrap();
692        assert!((r.radiance.x - 1.0).abs() < 1e-5);
693        assert!((r.opacity - 1.0).abs() < 1e-5);
694    }
695
696    #[test]
697    fn test_insert_multiple_and_lookup() {
698        let bounds = Aabb::new(Vec3::ZERO, Vec3::splat(8.0));
699        let mut octree = SparseVoxelOctree::new(bounds, 3);
700
701        let data_red = VoxelData {
702            radiance: Vec4::new(1.0, 0.0, 0.0, 1.0),
703            normal: Vec3::Y,
704            opacity: 1.0,
705            sh_coeffs: [0.0; 9],
706        };
707        let data_blue = VoxelData {
708            radiance: Vec4::new(0.0, 0.0, 1.0, 1.0),
709            normal: Vec3::X,
710            opacity: 0.5,
711            sh_coeffs: [0.0; 9],
712        };
713
714        octree.insert(Vec3::new(1.0, 1.0, 1.0), data_red);
715        octree.insert(Vec3::new(6.0, 6.0, 6.0), data_blue);
716
717        let r1 = octree.lookup(Vec3::new(1.0, 1.0, 1.0), 3).unwrap();
718        assert!((r1.radiance.x - 1.0).abs() < 1e-5);
719
720        let r2 = octree.lookup(Vec3::new(6.0, 6.0, 6.0), 3).unwrap();
721        assert!((r2.radiance.z - 1.0).abs() < 1e-5);
722    }
723
724    #[test]
725    fn test_remove() {
726        let bounds = Aabb::new(Vec3::ZERO, Vec3::splat(8.0));
727        let mut octree = SparseVoxelOctree::new(bounds, 3);
728
729        let data = VoxelData {
730            radiance: Vec4::ONE,
731            normal: Vec3::Y,
732            opacity: 1.0,
733            sh_coeffs: [0.0; 9],
734        };
735
736        octree.insert(Vec3::new(2.0, 2.0, 2.0), data);
737        assert!(octree.lookup(Vec3::new(2.0, 2.0, 2.0), 3).is_some());
738
739        octree.remove(Vec3::new(2.0, 2.0, 2.0));
740        // After remove, the data should be empty
741        let r = octree.lookup(Vec3::new(2.0, 2.0, 2.0), 3);
742        match r {
743            Some(d) => assert!(d.is_empty()),
744            None => {} // also acceptable
745        }
746    }
747
748    #[test]
749    fn test_ray_traversal() {
750        let bounds = Aabb::new(Vec3::ZERO, Vec3::splat(8.0));
751        let mut octree = SparseVoxelOctree::new(bounds, 3);
752
753        let data = VoxelData {
754            radiance: Vec4::new(1.0, 0.0, 0.0, 1.0),
755            normal: Vec3::Z,
756            opacity: 1.0,
757            sh_coeffs: [0.0; 9],
758        };
759
760        // Place a voxel at roughly (4, 4, 4)
761        octree.insert(Vec3::new(4.0, 4.0, 4.0), data);
762
763        // Shoot ray along Z towards it
764        let hits = octree.traverse_ray(
765            Vec3::new(4.0, 4.0, 0.0),
766            Vec3::new(0.0, 0.0, 1.0),
767            20.0,
768        );
769        assert!(!hits.is_empty(), "Ray should hit the voxel");
770        assert!((hits[0].1.radiance.x - 1.0).abs() < 1e-5);
771    }
772
773    #[test]
774    fn test_ray_miss() {
775        let bounds = Aabb::new(Vec3::ZERO, Vec3::splat(8.0));
776        let mut octree = SparseVoxelOctree::new(bounds, 3);
777
778        let data = VoxelData {
779            radiance: Vec4::ONE,
780            normal: Vec3::Y,
781            opacity: 1.0,
782            sh_coeffs: [0.0; 9],
783        };
784        octree.insert(Vec3::new(1.0, 1.0, 1.0), data);
785
786        // Shoot ray that misses entirely
787        let hits = octree.traverse_ray(
788            Vec3::new(7.0, 7.0, 0.0),
789            Vec3::new(0.0, 0.0, 1.0),
790            20.0,
791        );
792        // The ray goes through a different octant, should not hit
793        // (it may or may not depending on resolution, but at least check it runs)
794        let _ = hits;
795    }
796
797    #[test]
798    fn test_mipmap_averaging() {
799        let bounds = Aabb::new(Vec3::ZERO, Vec3::splat(8.0));
800        let mut octree = SparseVoxelOctree::new(bounds, 2);
801
802        let data = VoxelData {
803            radiance: Vec4::new(1.0, 0.0, 0.0, 1.0),
804            normal: Vec3::Y,
805            opacity: 1.0,
806            sh_coeffs: [0.0; 9],
807        };
808
809        octree.insert(Vec3::new(1.0, 1.0, 1.0), data);
810        octree.insert(Vec3::new(3.0, 1.0, 1.0), data);
811
812        octree.build_mipmaps();
813
814        // Parent should have averaged radiance
815        let parent = octree.lookup(Vec3::new(2.0, 1.0, 1.0), 1);
816        assert!(parent.is_some());
817    }
818
819    #[test]
820    fn test_build_from_grid() {
821        let mut grid = VoxelGrid::new(UVec3::new(4, 4, 4));
822        grid.set(1, 1, 1, VoxelData {
823            radiance: Vec4::new(0.5, 0.5, 0.0, 1.0),
824            normal: Vec3::Y,
825            opacity: 1.0,
826            sh_coeffs: [0.0; 9],
827        });
828
829        let bounds = Aabb::new(Vec3::ZERO, Vec3::splat(4.0));
830        let octree = SparseVoxelOctree::build_from_voxel_grid(&grid, bounds);
831
832        assert!(octree.node_count() > 1);
833        assert!(octree.leaf_count() > 0);
834    }
835
836    #[test]
837    fn test_iterator() {
838        let bounds = Aabb::new(Vec3::ZERO, Vec3::splat(8.0));
839        let mut octree = SparseVoxelOctree::new(bounds, 3);
840
841        let data = VoxelData {
842            radiance: Vec4::ONE,
843            normal: Vec3::Y,
844            opacity: 1.0,
845            sh_coeffs: [0.0; 9],
846        };
847
848        octree.insert(Vec3::new(1.0, 1.0, 1.0), data);
849        octree.insert(Vec3::new(6.0, 6.0, 6.0), data);
850
851        let leaves: Vec<_> = octree.iter_leaves().collect();
852        assert_eq!(leaves.len(), 2);
853    }
854
855    #[test]
856    fn test_voxel_grid() {
857        let mut grid = VoxelGrid::new(UVec3::new(4, 4, 4));
858        assert_eq!(grid.filled_count(), 0);
859
860        grid.set(0, 0, 0, VoxelData {
861            radiance: Vec4::ONE,
862            normal: Vec3::Y,
863            opacity: 1.0,
864            sh_coeffs: [0.0; 9],
865        });
866        assert_eq!(grid.filled_count(), 1);
867        assert!(grid.in_bounds(3, 3, 3));
868        assert!(!grid.in_bounds(4, 0, 0));
869    }
870
871    #[test]
872    fn test_aabb_ray_intersection() {
873        let aabb = Aabb::new(Vec3::ZERO, Vec3::splat(2.0));
874        // Axis-aligned ray through the middle: enters at z = 0 (t = 5) and
875        // leaves at z = 2 (t = 7). Infinite reciprocals are fine here.
876        let hit = aabb.intersect_ray(Vec3::new(1.0, 1.0, -5.0), Vec3::new(0.0, 0.0, 1.0).recip());
877        let (t0, t1) = hit.expect("axis ray hits");
878        assert!((t0 - 5.0).abs() < 1e-5 && (t1 - 7.0).abs() < 1e-5);
879        // Diagonal through the box corner to corner: t from 1 to 3.
880        let (t0, t1) = aabb
881            .intersect_ray(Vec3::splat(-1.0), Vec3::splat(1.0).recip())
882            .expect("diagonal ray hits");
883        assert!((t0 - 1.0).abs() < 1e-5 && (t1 - 3.0).abs() < 1e-5);
884        // The old test's ray, from (1, 1, -2) along (1, 1, 1), is already at
885        // x = 3 when it reaches z = 0, so it really misses the box.
886        assert!(aabb.intersect_ray(Vec3::new(1.0, 1.0, -2.0), Vec3::splat(1.0).recip()).is_none());
887    }
888
889    #[test]
890    fn test_memory_usage() {
891        let bounds = Aabb::new(Vec3::ZERO, Vec3::splat(8.0));
892        let octree = SparseVoxelOctree::new(bounds, 3);
893        assert!(octree.memory_usage() > 0);
894    }
895
896    #[test]
897    fn test_subdivide_octant_covers_full_volume() {
898        let aabb = Aabb::new(Vec3::ZERO, Vec3::splat(8.0));
899        for octant in 0..8u8 {
900            let child = aabb.subdivide_octant(octant);
901            assert!(child.size().x > 0.0);
902            assert!(child.size().y > 0.0);
903            assert!(child.size().z > 0.0);
904            // Child should be half the size
905            assert!((child.size().x - 4.0).abs() < 1e-5);
906        }
907    }
908
909    #[test]
910    fn test_sample_lod() {
911        let bounds = Aabb::new(Vec3::ZERO, Vec3::splat(8.0));
912        let mut octree = SparseVoxelOctree::new(bounds, 3);
913        let data = VoxelData {
914            radiance: Vec4::new(1.0, 0.0, 0.0, 1.0),
915            normal: Vec3::Y,
916            opacity: 1.0,
917            sh_coeffs: [0.0; 9],
918        };
919        octree.insert(Vec3::new(1.0, 1.0, 1.0), data);
920        octree.build_mipmaps();
921
922        let sampled = octree.sample_lod(Vec3::new(1.0, 1.0, 1.0), 2.5);
923        assert!(sampled.is_some());
924    }
925
926    #[test]
927    fn test_lookup_out_of_bounds_returns_none() {
928        let bounds = Aabb::new(Vec3::ZERO, Vec3::splat(8.0));
929        let octree = SparseVoxelOctree::new(bounds, 3);
930        assert!(octree.lookup(Vec3::splat(100.0), 3).is_none());
931    }
932
933    #[test]
934    fn test_insert_out_of_bounds_noop() {
935        let bounds = Aabb::new(Vec3::ZERO, Vec3::splat(8.0));
936        let mut octree = SparseVoxelOctree::new(bounds, 3);
937        let initial_count = octree.node_count();
938        octree.insert(Vec3::splat(-5.0), VoxelData {
939            radiance: Vec4::ONE,
940            normal: Vec3::Y,
941            opacity: 1.0,
942            sh_coeffs: [0.0; 9],
943        });
944        assert_eq!(octree.node_count(), initial_count);
945    }
946
947    #[test]
948    fn test_depth_tracking() {
949        let bounds = Aabb::new(Vec3::ZERO, Vec3::splat(16.0));
950        let mut octree = SparseVoxelOctree::new(bounds, 4);
951        octree.insert(Vec3::splat(1.0), VoxelData {
952            radiance: Vec4::ONE,
953            normal: Vec3::Y,
954            opacity: 1.0,
955            sh_coeffs: [0.0; 9],
956        });
957        assert!(octree.depth() >= 1);
958        assert!(octree.depth() <= 4);
959    }
960
961    #[test]
962    fn test_voxel_data_average() {
963        let a = VoxelData {
964            radiance: Vec4::new(1.0, 0.0, 0.0, 1.0),
965            normal: Vec3::Y,
966            opacity: 1.0,
967            sh_coeffs: [1.0; 9],
968        };
969        let b = VoxelData {
970            radiance: Vec4::new(0.0, 1.0, 0.0, 1.0),
971            normal: Vec3::Y,
972            opacity: 0.5,
973            sh_coeffs: [0.0; 9],
974        };
975        let avg = VoxelData::average(&[&a, &b]);
976        assert!((avg.radiance.x - 0.5).abs() < 1e-5);
977        assert!((avg.radiance.y - 0.5).abs() < 1e-5);
978        assert!((avg.opacity - 0.75).abs() < 1e-5);
979    }
980
981    #[test]
982    fn test_morton_large_values() {
983        let code = morton_encode_3d(1000, 2000, 3000);
984        let (x, y, z) = morton_decode_3d(code);
985        assert_eq!((x, y, z), (1000, 2000, 3000));
986    }
987
988    #[test]
989    fn test_grid_index_consistency() {
990        let grid = VoxelGrid::new(UVec3::new(8, 8, 8));
991        let idx1 = grid.index(0, 0, 0);
992        let idx2 = grid.index(7, 7, 7);
993        assert_eq!(idx1, 0);
994        assert_eq!(idx2, 511);
995    }
996}