1use crate::topology::point::PointId;
4use crate::topology::sieve::{InMemorySieve, Sieve};
5use std::collections::{HashMap, HashSet};
6
7#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
9pub struct TreeCell<const D: usize> {
10 pub level: u8,
12 pub coords: [u32; D],
14}
15
16impl<const D: usize> TreeCell<D> {
17 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 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#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
54pub struct ForestVertex<const D: usize> {
55 pub coords: [u32; D],
56}
57
58#[derive(Debug, Clone)]
60pub struct ForestMeshView<const D: usize> {
61 pub sieve: InMemorySieve<PointId, ()>,
63 pub cell_points: HashMap<TreeCell<D>, PointId>,
65 pub vertex_points: HashMap<ForestVertex<D>, PointId>,
67 pub max_level: u8,
69}
70
71#[derive(Debug, Clone)]
73pub struct Forest<const D: usize> {
74 leaves: HashSet<TreeCell<D>>,
75}
76
77pub type QuadForest = Forest<2>;
79pub 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 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 pub fn leaves(&self) -> impl Iterator<Item = &TreeCell<D>> {
101 self.leaves.iter()
102 }
103
104 pub fn leaf_count(&self) -> usize {
106 self.leaves.len()
107 }
108
109 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 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 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}