1use derive_more::From;
32use ego_tree::{NodeRef, Tree};
33use itertools::{Itertools, zip_eq};
34use std::collections::VecDeque;
35
36use super::tree_algorithms::TreeIsomorphism;
37
38#[derive(Clone, Debug, From, PartialEq, Eq)]
48pub enum OpenTree<Ty, Op> {
49 Id(Ty),
51
52 #[from]
54 Comp(Tree<Option<Op>>),
55}
56
57impl<Ty, Op> OpenTree<Ty, Op> {
58 pub fn empty(ty: Ty) -> Self {
60 OpenTree::Id(ty)
61 }
62
63 pub fn single(op: Op, arity: usize) -> Self {
65 let mut tree = Tree::new(Some(op));
66 for _ in 0..arity {
67 tree.root_mut().append(None);
68 }
69 tree.into()
70 }
71
72 pub fn graft(subtrees: impl IntoIterator<Item = Self>, op: Op) -> Self {
77 let mut tree = Tree::new(Some(op));
78 for subtree in subtrees {
79 match subtree {
80 OpenTree::Id(_) => tree.root_mut().append(None),
81 OpenTree::Comp(subtree) => tree.root_mut().append_subtree(subtree),
82 };
83 }
84 tree.into()
85 }
86
87 pub fn linear(iter: impl IntoIterator<Item = Op>) -> Option<Self> {
92 let mut values: Vec<_> = iter.into_iter().collect();
93 let value = values.pop()?;
94 let mut tree = Tree::new(Some(value));
95 let mut node_id = tree.root().id();
96 for value in values.into_iter().rev() {
97 node_id = tree.get_mut(node_id).unwrap().append(Some(value)).id();
98 }
99 tree.get_mut(node_id).unwrap().append(None);
100 Some(tree.into())
101 }
102
103 pub fn arity(&self) -> usize {
107 match self {
108 OpenTree::Comp(tree) => tree.root().boundary().count(),
109 OpenTree::Id(_) => 1,
110 }
111 }
112
113 pub fn size(&self) -> usize {
118 match self {
119 OpenTree::Comp(tree) => tree.nodes().filter(|node| node.value().is_some()).count(),
120 OpenTree::Id(_) => 0,
121 }
122 }
123
124 pub fn is_empty(&self) -> bool {
126 matches!(self, OpenTree::Id(_))
127 }
128
129 pub fn only(self) -> Option<Op> {
133 if let OpenTree::Comp(mut tree) = self
134 && tree.root().children().all(|node| node.value().is_none())
135 {
136 std::mem::take(tree.root_mut().value())
137 } else {
138 None
139 }
140 }
141
142 pub fn is_isomorphic_to(&self, other: &Self) -> bool
149 where
150 Ty: Eq,
151 Op: Eq,
152 {
153 match (self, other) {
154 (OpenTree::Comp(tree1), OpenTree::Comp(tree2)) => tree1.is_isomorphic_to(tree2),
155 (OpenTree::Id(type1), OpenTree::Id(type2)) => *type1 == *type2,
156 _ => false,
157 }
158 }
159
160 pub fn map<CodOp>(self, mut f: impl FnMut(Op) -> CodOp) -> OpenTree<Ty, CodOp> {
162 match self {
163 OpenTree::Comp(tree) => tree.map(|value| value.map(&mut f)).into(),
164 OpenTree::Id(ty) => OpenTree::Id(ty),
165 }
166 }
167}
168
169pub trait OpenNodeRef<T> {
171 fn is_boundary(&self) -> bool;
173
174 fn boundary(&self) -> impl Iterator<Item = Self>;
176
177 fn get_value(&self) -> Option<&T>;
179
180 fn parent_value(&self) -> Option<&T>;
182}
183
184impl<'a, T: 'a> OpenNodeRef<T> for NodeRef<'a, Option<T>> {
185 fn is_boundary(&self) -> bool {
186 let is_null = self.value().is_none();
187 assert!(!(is_null && self.has_children()), "Boundary nodes should be leaves");
188 is_null
189 }
190
191 fn boundary(&self) -> impl Iterator<Item = Self> {
192 self.descendants().filter(|node| node.is_boundary())
193 }
194
195 fn get_value(&self) -> Option<&T> {
196 self.value().as_ref()
197 }
198
199 fn parent_value(&self) -> Option<&T> {
200 self.parent()
201 .map(|p| p.value().as_ref().expect("Inner nodes should not be null"))
202 }
203}
204
205impl<Ty, Op> OpenTree<Ty, OpenTree<Ty, Op>> {
206 pub fn flatten(self) -> OpenTree<Ty, Op> {
208 let mut outer_tree = match self {
210 OpenTree::Id(x) => return OpenTree::Id(x),
211 OpenTree::Comp(tree) => tree,
212 };
213
214 let value = std::mem::take(outer_tree.root_mut().value())
216 .expect("Root node of outer tree should contain a tree");
217 let (mut tree, root_type) = match value {
218 OpenTree::Id(x) => (Tree::new(None), Some(x)),
219 OpenTree::Comp(tree) => (tree, None),
220 };
221
222 let mut queue = VecDeque::new();
223 for (child, leaf) in zip_eq(outer_tree.root().children(), tree.root().boundary()) {
224 queue.push_back((child.id(), leaf.id()));
225 }
226
227 while let Some((outer_id, leaf_id)) = queue.pop_front() {
228 let Some(value) = std::mem::take(outer_tree.get_mut(outer_id).unwrap().value()) else {
229 continue;
230 };
231 match value {
232 OpenTree::Id(_) => {
233 let Ok(outer_parent) =
234 outer_tree.get(outer_id).unwrap().children().exactly_one()
235 else {
236 panic!("Identity tree should have exactly one parent")
237 };
238 queue.push_back((outer_parent.id(), leaf_id));
239 }
240 OpenTree::Comp(inner_tree) => {
241 let subtree_id = tree.extend_tree(inner_tree).id();
242 let value = std::mem::take(tree.get_mut(subtree_id).unwrap().value());
243
244 let mut inner_node = tree.get_mut(leaf_id).unwrap();
245 *inner_node.value() = value;
246 inner_node.reparent_from_id_append(subtree_id);
247
248 let outer_node = outer_tree.get(outer_id).unwrap();
249 let inner_node: NodeRef<_> = inner_node.into();
250 for (child, leaf) in zip_eq(outer_node.children(), inner_node.boundary()) {
251 queue.push_back((child.id(), leaf.id()));
252 }
253 }
254 }
255 }
256
257 if tree.root().value().is_none() {
258 OpenTree::Id(root_type.unwrap())
259 } else {
260 tree.into()
261 }
262 }
263}
264
265#[cfg(test)]
266mod tests {
267 use super::*;
268 use ego_tree::tree;
269
270 type OT = OpenTree<char, char>;
271
272 #[test]
273 fn construct_tree() {
274 assert_eq!(OT::empty('X').arity(), 1);
275
276 let tree = OT::single('f', 2);
277 assert_eq!(tree.arity(), 2);
278 assert_eq!(tree, tree!(Some('f') => { None, None }).into());
279 assert_eq!(tree.only(), Some('f'));
280
281 let tree = tree!(Some('h') => { Some('g') => { Some('f') => { None } } });
282 assert_eq!(OT::linear(vec!['f', 'g', 'h']), Some(tree.into()));
283 }
284
285 #[test]
286 fn flatten_tree() {
287 let tree = OT::from(tree!(
289 Some('f') => {
290 Some('h') => {
291 Some('k') => { None, None},
292 None,
293 },
294 Some('g') => {
295 None,
296 Some('l') => { None, None }
297 },
298 }
299 ));
300 assert!(!tree.is_empty());
301 assert_eq!(tree.size(), 5);
302 assert_eq!(tree.arity(), 6);
303
304 let subtree1 = OT::from(tree!(
305 Some('f') => {
306 None,
307 Some('g') => { None, None },
308 }
309 ));
310 let subtree2 = OT::from(tree!(
311 Some('h') => {
312 Some('k') => { None, None },
313 None
314 }
315 ));
316 let subtree3 = OT::from(tree!(
317 Some('l') => { None, None }
318 ));
319
320 let outer_tree: OpenTree<_, _> = tree!(
321 Some(subtree1.clone()) => {
322 Some(subtree2.clone()) => { None, None, None },
323 None,
324 Some(subtree3.clone()) => { None, None },
325 }
326 )
327 .into();
328 assert!(outer_tree.flatten().is_isomorphic_to(&tree));
329
330 let outer_tree: OpenTree<_, _> = tree!(
331 Some(subtree1) => {
332 Some(OpenTree::Id('X')) => {
333 Some(subtree2) => { None, None, None },
334 },
335 Some(OpenTree::Id('X')) => { None },
336 Some(OpenTree::Id('X')) => {
337 Some(subtree3) => { None, None },
338 },
339 }
340 )
341 .into();
342 assert!(outer_tree.flatten().is_isomorphic_to(&tree));
343
344 let outer_tree: OpenTree<_, _> = OpenTree::Id('X');
346 assert_eq!(outer_tree.flatten(), OT::Id('X'));
347
348 let outer_tree: OpenTree<_, _> = tree!(
350 Some(OT::Id('X')) => { Some(OT::Id('x')) => { None } }
351 )
352 .into();
353 assert_eq!(outer_tree.flatten(), OT::Id('X'));
354 }
355}