Skip to main content

mesh_sieve/
forest.rs

1//! Quad/oct-tree AMR forest representation with Sieve mappings.
2
3use crate::topology::point::PointId;
4use crate::topology::sieve::{InMemorySieve, Sieve};
5use std::collections::{HashMap, HashSet};
6
7/// A cell in a quadtree/octree forest.
8#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
9pub struct TreeCell<const D: usize> {
10    /// Refinement level (0 is root).
11    pub level: u8,
12    /// Integer coordinates at the given level.
13    pub coords: [u32; D],
14}
15
16impl<const D: usize> TreeCell<D> {
17    /// Returns the parent cell, or `None` for the root.
18    pub fn parent(&self) -> Option<Self> {
19        if self.level == 0 {
20            None
21        } else {
22            let mut coords = self.coords;
23            for coord in &mut coords {
24                *coord /= 2;
25            }
26            Some(Self {
27                level: self.level - 1,
28                coords,
29            })
30        }
31    }
32
33    /// Returns the `2^D` children of this cell.
34    pub fn children(&self) -> Vec<Self> {
35        let count = 1usize << D;
36        let mut children = Vec::with_capacity(count);
37        for idx in 0..count {
38            let mut coords = [0u32; D];
39            for axis in 0..D {
40                let bit = (idx >> axis) & 1;
41                coords[axis] = self.coords[axis] * 2 + bit as u32;
42            }
43            children.push(Self {
44                level: self.level + 1,
45                coords,
46            });
47        }
48        children
49    }
50}
51
52/// A vertex on the conforming mesh view (coordinates on the max-level grid).
53#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
54pub struct ForestVertex<const D: usize> {
55    pub coords: [u32; D],
56}
57
58/// A conforming mesh view extracted from the forest.
59#[derive(Debug, Clone)]
60pub struct ForestMeshView<const D: usize> {
61    /// Sieve topology with cell → vertex arrows.
62    pub sieve: InMemorySieve<PointId, ()>,
63    /// Mapping from forest leaf cells to sieve cell points.
64    pub cell_points: HashMap<TreeCell<D>, PointId>,
65    /// Mapping from vertex coordinates to sieve vertex points.
66    pub vertex_points: HashMap<ForestVertex<D>, PointId>,
67    /// Maximum refinement level represented by the view.
68    pub max_level: u8,
69}
70
71/// Forest representation for quadtrees (`D = 2`) or octrees (`D = 3`).
72#[derive(Debug, Clone)]
73pub struct Forest<const D: usize> {
74    leaves: HashSet<TreeCell<D>>,
75}
76
77/// A quadtree forest (`D = 2`).
78pub type QuadForest = Forest<2>;
79/// An octree forest (`D = 3`).
80pub type OctForest = Forest<3>;
81
82impl<const D: usize> Default for Forest<D> {
83    fn default() -> Self {
84        Self::new()
85    }
86}
87
88impl<const D: usize> Forest<D> {
89    /// Create a new forest with a single root cell.
90    pub fn new() -> Self {
91        let mut leaves = HashSet::new();
92        leaves.insert(TreeCell {
93            level: 0,
94            coords: [0; D],
95        });
96        Self { leaves }
97    }
98
99    /// Return an iterator over leaf cells.
100    pub fn leaves(&self) -> impl Iterator<Item = &TreeCell<D>> {
101        self.leaves.iter()
102    }
103
104    /// Return the number of leaf cells.
105    pub fn leaf_count(&self) -> usize {
106        self.leaves.len()
107    }
108
109    /// Refine all leaf cells whose indicator exceeds the threshold.
110    pub fn refine_by_indicator<F>(&mut self, indicator: F, threshold: f64) -> usize
111    where
112        F: Fn(&TreeCell<D>) -> f64,
113    {
114        let to_refine: Vec<_> = self
115            .leaves
116            .iter()
117            .copied()
118            .filter(|cell| indicator(cell) > threshold)
119            .collect();
120        self.refine_cells(&to_refine)
121    }
122
123    /// Coarsen all leaf siblings whose indicators are below the threshold.
124    pub fn coarsen_by_indicator<F>(&mut self, indicator: F, threshold: f64) -> usize
125    where
126        F: Fn(&TreeCell<D>) -> f64,
127    {
128        let mut parent_to_children: HashMap<TreeCell<D>, Vec<TreeCell<D>>> = HashMap::new();
129        for leaf in &self.leaves {
130            if let Some(parent) = leaf.parent() {
131                parent_to_children.entry(parent).or_default().push(*leaf);
132            }
133        }
134
135        let mut to_coarsen = Vec::new();
136        let sibling_count = 1 << D;
137        for (parent, children) in parent_to_children {
138            if children.len() == sibling_count
139                && children.iter().all(|child| indicator(child) < threshold)
140            {
141                to_coarsen.push((parent, children));
142            }
143        }
144
145        let mut coarsened = 0;
146        for (parent, children) in to_coarsen {
147            let mut removed = 0;
148            for child in children {
149                if self.leaves.remove(&child) {
150                    removed += 1;
151                }
152            }
153            if removed == sibling_count {
154                self.leaves.insert(parent);
155                coarsened += 1;
156            }
157        }
158        coarsened
159    }
160
161    /// Build a conforming mesh view by refining until neighboring cells match in level.
162    pub fn conforming_view(&self) -> ForestMeshView<D> {
163        let mut balanced = self.clone();
164        balanced.balance();
165        balanced.build_view()
166    }
167
168    fn refine_cells(&mut self, cells: &[TreeCell<D>]) -> usize {
169        let mut refined = 0;
170        for cell in cells {
171            if self.leaves.remove(cell) {
172                for child in cell.children() {
173                    self.leaves.insert(child);
174                }
175                refined += 1;
176            }
177        }
178        refined
179    }
180
181    fn max_level(&self) -> u8 {
182        self.leaves.iter().map(|cell| cell.level).max().unwrap_or(0)
183    }
184
185    fn balance(&mut self) {
186        loop {
187            let leaves: Vec<_> = self.leaves.iter().copied().collect();
188            let max_level = leaves.iter().map(|cell| cell.level).max().unwrap_or(0);
189            let mut to_refine = HashSet::new();
190            for (i, cell) in leaves.iter().enumerate() {
191                for other in leaves.iter().skip(i + 1) {
192                    if are_face_neighbors(cell, other, max_level) {
193                        if cell.level < other.level {
194                            to_refine.insert(*cell);
195                        } else if other.level < cell.level {
196                            to_refine.insert(*other);
197                        }
198                    }
199                }
200            }
201
202            if to_refine.is_empty() {
203                break;
204            }
205
206            let cells: Vec<_> = to_refine.into_iter().collect();
207            self.refine_cells(&cells);
208        }
209    }
210
211    fn build_view(&self) -> ForestMeshView<D> {
212        let max_level = self.max_level();
213        let leaves: Vec<_> = self.leaves.iter().copied().collect();
214        let mut sieve = InMemorySieve::<PointId, ()>::default();
215        let mut cell_points = HashMap::new();
216        let mut vertex_points = HashMap::new();
217        let mut next_id = 1u64;
218
219        for cell in &leaves {
220            let point = PointId::new(next_id).expect("cell point id");
221            next_id += 1;
222            cell_points.insert(*cell, point);
223        }
224
225        for cell in &leaves {
226            let cell_point = cell_points[cell];
227            for vertex in cell_vertices(cell, max_level) {
228                let vertex_point = vertex_points.entry(vertex).or_insert_with(|| {
229                    let point = PointId::new(next_id).expect("vertex point id");
230                    next_id += 1;
231                    point
232                });
233                sieve.add_arrow(cell_point, *vertex_point, ());
234            }
235        }
236
237        ForestMeshView {
238            sieve,
239            cell_points,
240            vertex_points,
241            max_level,
242        }
243    }
244}
245
246fn cell_bounds<const D: usize>(cell: &TreeCell<D>, max_level: u8) -> [(u32, u32); D] {
247    let scale = 1u32 << (max_level - cell.level);
248    let mut bounds = [(0u32, 0u32); D];
249    for axis in 0..D {
250        let start = cell.coords[axis] * scale;
251        bounds[axis] = (start, start + scale);
252    }
253    bounds
254}
255
256fn are_face_neighbors<const D: usize>(a: &TreeCell<D>, b: &TreeCell<D>, max_level: u8) -> bool {
257    let a_bounds = cell_bounds(a, max_level);
258    let b_bounds = cell_bounds(b, max_level);
259    let mut touching_axis = None;
260    for axis in 0..D {
261        let (a0, a1) = a_bounds[axis];
262        let (b0, b1) = b_bounds[axis];
263        if a1 == b0 || b1 == a0 {
264            if touching_axis.is_some() {
265                return false;
266            }
267            touching_axis = Some(axis);
268        } else if a0 >= b1 || b0 >= a1 {
269            return false;
270        }
271    }
272    if let Some(axis) = touching_axis {
273        for other_axis in 0..D {
274            if other_axis == axis {
275                continue;
276            }
277            let (a0, a1) = a_bounds[other_axis];
278            let (b0, b1) = b_bounds[other_axis];
279            if a0 >= b1 || b0 >= a1 {
280                return false;
281            }
282        }
283        true
284    } else {
285        false
286    }
287}
288
289fn cell_vertices<const D: usize>(cell: &TreeCell<D>, max_level: u8) -> Vec<ForestVertex<D>> {
290    let bounds = cell_bounds(cell, max_level);
291    let mut vertices = Vec::with_capacity(1 << D);
292    for idx in 0..(1 << D) {
293        let mut coords = [0u32; D];
294        for axis in 0..D {
295            let bit = (idx >> axis) & 1;
296            coords[axis] = if bit == 0 {
297                bounds[axis].0
298            } else {
299                bounds[axis].1
300            };
301        }
302        vertices.push(ForestVertex { coords });
303    }
304    vertices
305}
306
307#[cfg(test)]
308mod tests {
309    use super::*;
310
311    #[test]
312    fn forest_refine_and_coarsen_by_indicator() {
313        let mut forest = QuadForest::new();
314        assert_eq!(forest.leaf_count(), 1);
315
316        let refined = forest.refine_by_indicator(
317            |cell| {
318                if cell.level == 0 { 1.0 } else { 0.0 }
319            },
320            0.5,
321        );
322        assert_eq!(refined, 1);
323        assert_eq!(forest.leaf_count(), 4);
324
325        forest.refine_by_indicator(
326            |cell| {
327                if cell.level == 1 && cell.coords == [0, 0] {
328                    1.0
329                } else {
330                    0.0
331                }
332            },
333            0.5,
334        );
335        assert_eq!(forest.leaf_count(), 7);
336
337        let coarsened = forest.coarsen_by_indicator(|_| 0.0, 0.1);
338        assert!(coarsened > 0);
339        assert_eq!(forest.leaf_count(), 4);
340
341        let coarsened_again = forest.coarsen_by_indicator(|_| 0.0, 0.1);
342        assert!(coarsened_again > 0);
343        assert_eq!(forest.leaf_count(), 1);
344    }
345
346    #[test]
347    fn forest_conforming_view_has_consistent_topology() {
348        let mut forest = QuadForest::new();
349        forest.refine_by_indicator(|cell| if cell.level == 0 { 1.0 } else { 0.0 }, 0.0);
350        forest.refine_by_indicator(
351            |cell| {
352                if cell.level == 1 && cell.coords == [0, 0] {
353                    1.0
354                } else {
355                    0.0
356                }
357            },
358            0.5,
359        );
360
361        let view = forest.conforming_view();
362        assert_eq!(view.max_level, 2);
363        assert_eq!(view.cell_points.len(), 16);
364
365        for cell_point in view.cell_points.values() {
366            let cone: Vec<_> = view.sieve.cone_points(*cell_point).collect();
367            assert_eq!(cone.len(), 4);
368        }
369    }
370}