1use std::collections::HashSet;
2
3use bincode_trait_derive::{Decode, Encode};
4use derive_more::{From, Into};
5use itertools::Itertools;
6use serde::{Deserialize, Serialize};
7use symbolica::{
8 atom::{Atom, AtomCore},
9 parse,
10};
11use typed_index_collections::TiVec;
12
13#[derive(
16 Debug,
17 From,
18 Into,
19 Copy,
20 Clone,
21 Serialize,
22 Deserialize,
23 Encode,
24 Decode,
25 PartialEq,
26 Eq,
27 PartialOrd,
28 Ord,
29 Hash,
30)]
31pub struct NodeId(usize);
32
33impl NodeId {
34 pub const fn root() -> Self {
35 NodeId(0)
36 }
37}
38
39#[derive(Debug, Clone, Serialize, Deserialize, Encode, Decode, Eq, PartialEq)]
40pub struct TreeNode<T> {
41 pub data: T,
42 pub node_id: NodeId,
43 pub children: Vec<NodeId>,
44 pub parent: Option<NodeId>,
45}
46
47#[allow(dead_code)]
48fn determine_shifted_id(removed_ids_sorted: &[NodeId], original_id: NodeId) -> NodeId {
49 let shift = removed_ids_sorted
50 .iter()
51 .filter(|&&removed_id| removed_id < original_id)
52 .count();
53 NodeId(original_id.0 - shift)
54}
55
56impl<T> TreeNode<T> {
57 #[allow(dead_code)]
58 fn update_node_ids(&mut self, removed_ids: &[NodeId]) {
59 self.node_id = determine_shifted_id(removed_ids, self.node_id);
60
61 self.children
62 .retain(|child_id| !removed_ids.contains(child_id));
63 self.children.iter_mut().for_each(|child_id| {
64 *child_id = determine_shifted_id(removed_ids, *child_id);
65 });
66
67 if let Some(parent_id) = self.parent {
68 if removed_ids.contains(&parent_id) {
69 self.parent = None;
70 } else {
71 let shifted_parent_id = determine_shifted_id(removed_ids, parent_id);
72 self.parent = Some(shifted_parent_id);
73 }
74 }
75 }
76}
77
78#[derive(Debug, Clone, Serialize, Deserialize, Encode, Decode, Eq, PartialEq)]
79pub struct Tree<T> {
80 nodes: TiVec<NodeId, TreeNode<T>>,
81}
82
83impl<T> Tree<T> {
84 pub(crate) fn from_root(data: T) -> Self {
85 let node_id = NodeId(0);
86 let root_node = TreeNode {
87 data,
88 node_id,
89 children: Vec::new(),
90 parent: None,
91 };
92 Tree {
93 nodes: vec![root_node].into(),
94 }
95 }
96
97 pub(crate) fn insert_node(&mut self, parent_id: NodeId, data: T) {
98 let node_id = NodeId(self.nodes.len());
99 let new_node = TreeNode {
100 data,
101 node_id,
102 children: Vec::new(),
103 parent: Some(parent_id),
104 };
105 self.nodes.push(new_node);
106 self.nodes[parent_id].children.push(node_id);
107 }
108
109 pub(crate) fn get_node(&self, node_id: NodeId) -> &TreeNode<T> {
110 &self.nodes[node_id]
111 }
112
113 #[allow(dead_code)]
114 pub(crate) fn get_data_node_mut(&mut self, node_id: NodeId) -> &mut T {
115 &mut self.nodes[node_id].data
116 }
117
118 pub(crate) fn apply_mut_closure(&mut self, node_id: NodeId, closure: impl Fn(&mut T)) {
119 closure(&mut self.nodes[node_id].data);
120 }
121
122 pub(crate) fn get_bottom_layer(&self) -> Vec<NodeId> {
123 self.nodes
124 .iter()
125 .filter(|node| node.children.is_empty())
126 .map(|node| node.node_id)
127 .collect()
128 }
129
130 pub(crate) fn map<G>(self, f: impl Fn(T) -> G) -> Tree<G> {
131 Tree {
132 nodes: self
133 .nodes
134 .into_iter()
135 .map(|node| TreeNode {
136 data: f(node.data),
137 node_id: node.node_id,
138 children: node.children,
139 parent: node.parent,
140 })
141 .collect(),
142 }
143 }
144
145 pub(crate) fn map_mut(&mut self, f: impl Fn(&mut T)) {
146 for node in &mut self.nodes {
147 f(&mut node.data);
148 }
149 }
150
151 #[allow(dead_code)]
152 pub(crate) fn get_num_nodes(&self) -> usize {
153 self.nodes.len()
154 }
155
156 pub(crate) fn iter_nodes(&self) -> impl Iterator<Item = &TreeNode<T>> {
157 self.nodes.iter()
158 }
159
160 fn path_to_root(&self, mut node_id: NodeId) -> Vec<NodeId> {
161 let mut path = Vec::new();
162 loop {
163 path.push(node_id);
164 if let Some(parent_id) = self.nodes[node_id].parent {
165 node_id = parent_id;
166 } else {
167 break;
168 }
169 }
170 path.reverse();
171 path
172 }
173
174 #[allow(dead_code)]
175 fn obtain_subtree_node_ids_impl(&self, node_id: NodeId, collected_ids: &mut Vec<NodeId>) {
176 for &child_id in &self.nodes[node_id].children {
177 collected_ids.push(child_id);
178 self.obtain_subtree_node_ids_impl(child_id, collected_ids);
179 }
180 }
181
182 #[allow(dead_code)]
183 fn obtain_subtree_node_ids(&self, node_id: NodeId) -> Vec<NodeId> {
184 let mut collected_ids = vec![node_id];
185 self.obtain_subtree_node_ids_impl(node_id, &mut collected_ids);
186 collected_ids
187 }
188
189 #[allow(dead_code)]
190 pub(crate) fn remove_node(&mut self, node_id: NodeId) {
191 let mut subtree = self.obtain_subtree_node_ids(node_id);
192 subtree.sort();
193
194 for id in subtree.iter().rev() {
196 self.nodes.remove(*id);
197 }
198
199 for node in &mut self.nodes {
200 node.update_node_ids(&subtree);
202 }
203 }
204
205 #[allow(dead_code)]
206 pub(crate) fn filter_mut(&mut self, predicate: impl Fn(&T) -> bool) {
207 let mut nodes_to_remove = HashSet::new();
208
209 for node in &self.nodes {
210 if !predicate(&node.data) {
211 nodes_to_remove.extend(self.obtain_subtree_node_ids(node.node_id));
212 }
213 }
214
215 let nodes_to_remove: Vec<NodeId> = nodes_to_remove.into_iter().sorted().collect();
216
217 for id in nodes_to_remove.iter().rev() {
218 self.nodes.remove(*id);
219 }
220
221 for node in &mut self.nodes {
222 node.update_node_ids(&nodes_to_remove);
224 }
225 }
226
227 pub(crate) fn keep_branches_with_value_count_mut(&mut self, value: &T, n: usize)
228 where
229 T: Eq,
230 {
231 if self.nodes.is_empty() {
232 return;
233 }
234
235 let leaves = self.get_bottom_layer();
236 let mut nodes_to_keep = HashSet::new();
237 let mut has_match = false;
238
239 for leaf in leaves {
240 let path = self.path_to_root(leaf);
241 let count = path
242 .iter()
243 .filter(|&&node_id| self.nodes[node_id].data == *value)
244 .count();
245
246 if count == n {
247 has_match = true;
248 nodes_to_keep.extend(path);
249 }
250 }
251
252 if !has_match {
253 self.nodes.clear();
254 return;
255 }
256
257 let nodes_to_remove: Vec<NodeId> = self
258 .nodes
259 .iter()
260 .filter(|node| !nodes_to_keep.contains(&node.node_id))
261 .map(|node| node.node_id)
262 .sorted()
263 .collect();
264
265 for id in nodes_to_remove.iter().rev() {
266 self.nodes.remove(*id);
267 }
268
269 for node in &mut self.nodes {
270 node.update_node_ids(&nodes_to_remove);
271 }
272 }
273
274 pub(crate) fn max_value_count_on_branch(&self, value: &T) -> usize
275 where
276 T: Eq,
277 {
278 if self.nodes.is_empty() {
279 return 0;
280 }
281
282 self.get_bottom_layer()
283 .into_iter()
284 .map(|leaf| {
285 self.path_to_root(leaf)
286 .into_iter()
287 .filter(|&node_id| self.nodes[node_id].data == *value)
288 .count()
289 })
290 .max()
291 .unwrap_or(0)
292 }
293}
294
295impl<T> Tree<T>
296where
297 Atom: From<T>,
298 T: Copy,
299{
300 fn to_atom_inv_impl(&self, cur_node: NodeId) -> Atom {
301 let node = &self.nodes[cur_node];
302 let inv_data_esurface = (Atom::num(1) / Atom::from(node.data))
303 .replace(parse!("η_inf^-1"))
304 .with(Atom::num(0));
305
306 let child_sum = node
307 .children
308 .iter()
309 .map(|&child| self.to_atom_inv_impl(child))
310 .reduce(|acc, x| acc + x)
311 .unwrap_or(Atom::num(1));
312
313 inv_data_esurface * child_sum
314 }
315
316 pub(crate) fn to_atom_inv(&self) -> Atom {
317 if self.nodes.is_empty() {
318 return Atom::num(0);
319 }
320 self.to_atom_inv_impl(NodeId::root())
321 }
322}
323
324#[cfg(test)]
325mod tests {
326
327 use crate::cff::tree::{NodeId, Tree};
328
329 #[test]
330 fn test_remove_node() {
331 let mut tree = Tree::from_root(0);
332 tree.insert_node(NodeId(0), 1);
333 tree.insert_node(NodeId(0), 2);
334 tree.insert_node(NodeId(2), 3);
335 tree.insert_node(NodeId(2), 4);
336 tree.insert_node(NodeId(4), 5);
337 tree.insert_node(NodeId(4), 6);
338 tree.insert_node(NodeId(4), 7);
339 tree.insert_node(NodeId(7), 8);
340 tree.insert_node(NodeId(7), 9);
341
342 tree.remove_node(NodeId(4));
343
344 let mut expected_tree = Tree::from_root(0);
345 expected_tree.insert_node(NodeId(0), 1);
346 expected_tree.insert_node(NodeId(0), 2);
347 expected_tree.insert_node(NodeId(2), 3);
348
349 assert_eq!(tree, expected_tree);
350
351 let mut tree = Tree::from_root(0);
352 tree.insert_node(NodeId(0), 1);
353 tree.insert_node(NodeId(1), 2);
354 tree.insert_node(NodeId(1), 3);
355 tree.insert_node(NodeId(1), 4);
356 tree.insert_node(NodeId(4), 5);
357 tree.insert_node(NodeId(4), 6);
358 tree.insert_node(NodeId(0), 7);
359 tree.insert_node(NodeId(7), 8);
360 tree.insert_node(NodeId(7), 9);
361
362 tree.remove_node(NodeId(1));
363
364 let mut expected_tree = Tree::from_root(0);
365 expected_tree.insert_node(NodeId(0), 7);
366 expected_tree.insert_node(NodeId(1), 8);
367 expected_tree.insert_node(NodeId(1), 9);
368
369 assert_eq!(tree, expected_tree);
370 }
371
372 #[test]
373 fn test_keep_branches_with_value_count_mut() {
374 let mut tree = Tree::from_root(0);
375 tree.insert_node(NodeId(0), 1);
376 tree.insert_node(NodeId(1), 2);
377 tree.insert_node(NodeId(2), 1);
378 tree.insert_node(NodeId(0), 1);
379 tree.insert_node(NodeId(4), 3);
380 tree.insert_node(NodeId(5), 4);
381
382 tree.keep_branches_with_value_count_mut(&1, 1);
383
384 let mut expected_tree = Tree::from_root(0);
385 expected_tree.insert_node(NodeId(0), 1);
386 expected_tree.insert_node(NodeId(1), 3);
387 expected_tree.insert_node(NodeId(2), 4);
388
389 assert_eq!(tree, expected_tree);
390 }
391
392 #[test]
393 fn test_keep_branches_with_value_count_mut_no_match() {
394 let mut tree = Tree::from_root(0);
395 tree.insert_node(NodeId(0), 1);
396 tree.insert_node(NodeId(1), 2);
397 tree.insert_node(NodeId(2), 1);
398
399 tree.keep_branches_with_value_count_mut(&1, 3);
400
401 assert_eq!(tree.get_num_nodes(), 0);
402 }
403
404 #[test]
405 fn test_keep_branches_with_value_count_mut_zero_occurrences() {
406 let mut tree = Tree::from_root(0);
407 tree.insert_node(NodeId(0), 2);
408 tree.insert_node(NodeId(1), 3);
409 tree.insert_node(NodeId(0), 1);
410 tree.insert_node(NodeId(3), 4);
411
412 tree.keep_branches_with_value_count_mut(&1, 0);
413
414 let mut expected_tree = Tree::from_root(0);
415 expected_tree.insert_node(NodeId(0), 2);
416 expected_tree.insert_node(NodeId(1), 3);
417
418 assert_eq!(tree, expected_tree);
419 }
420
421 #[test]
422 fn test_keep_branches_with_value_count_mut_keeps_shared_prefix() {
423 let mut tree = Tree::from_root(0);
424 tree.insert_node(NodeId(0), 9);
425 tree.insert_node(NodeId(1), 1);
426 tree.insert_node(NodeId(2), 5);
427 tree.insert_node(NodeId(1), 2);
428 tree.insert_node(NodeId(4), 6);
429
430 tree.keep_branches_with_value_count_mut(&1, 1);
431
432 let mut expected_tree = Tree::from_root(0);
433 expected_tree.insert_node(NodeId(0), 9);
434 expected_tree.insert_node(NodeId(1), 1);
435 expected_tree.insert_node(NodeId(2), 5);
436
437 assert_eq!(tree, expected_tree);
438 }
439
440 #[test]
441 fn test_max_value_count_on_branch() {
442 let mut tree = Tree::from_root(1);
443 tree.insert_node(NodeId(0), 1);
444 tree.insert_node(NodeId(1), 2);
445 tree.insert_node(NodeId(2), 1);
446 tree.insert_node(NodeId(0), 3);
447 tree.insert_node(NodeId(4), 1);
448
449 assert_eq!(tree.max_value_count_on_branch(&1), 3);
450 }
451
452 #[test]
453 fn test_max_value_count_on_branch_value_absent() {
454 let mut tree = Tree::from_root(0);
455 tree.insert_node(NodeId(0), 2);
456 tree.insert_node(NodeId(1), 3);
457 tree.insert_node(NodeId(0), 4);
458
459 assert_eq!(tree.max_value_count_on_branch(&1), 0);
460 }
461
462 #[test]
463 fn test_max_value_count_on_branch_single_branch() {
464 let mut tree = Tree::from_root(1);
465 tree.insert_node(NodeId(0), 1);
466
467 assert_eq!(tree.max_value_count_on_branch(&1), 2);
468 }
469}