use super::TreeChromosome;
use crate::collections::{Tree, TreeNode};
use std::{collections::VecDeque, marker::PhantomData};
pub trait TreeIterator<T> {
fn iter_pre_order(&self) -> PreOrderIterator<'_, T>;
fn iter_post_order(&self) -> PostOrderIterator<'_, T>;
fn iter_breadth_first(&self) -> TreeBreadthFirstIterator<'_, T>;
fn apply<F: Fn(&mut TreeNode<T>)>(&mut self, visit_fn: F);
}
impl<T> TreeIterator<T> for TreeNode<T> {
fn iter_pre_order(&self) -> PreOrderIterator<'_, T> {
PreOrderIterator { stack: vec![self] }
}
fn iter_post_order(&self) -> PostOrderIterator<'_, T> {
PostOrderIterator {
stack: vec![(self, false)],
}
}
fn iter_breadth_first(&self) -> TreeBreadthFirstIterator<'_, T> {
TreeBreadthFirstIterator {
queue: vec![self].into_iter().collect(),
}
}
fn apply<F: Fn(&mut TreeNode<T>)>(&mut self, visit_fn: F) {
let visitor = TreeVisitor::new(visit_fn);
visitor.visit(self);
}
}
impl<T> TreeIterator<T> for Tree<T> {
fn iter_pre_order(&self) -> PreOrderIterator<'_, T> {
PreOrderIterator {
stack: self
.root()
.map_or(Vec::new(), |root| vec![root].into_iter().collect()),
}
}
fn iter_post_order(&self) -> PostOrderIterator<'_, T> {
PostOrderIterator {
stack: self
.root()
.map_or(Vec::new(), |root| vec![(root, false)].into_iter().collect()),
}
}
fn iter_breadth_first(&self) -> TreeBreadthFirstIterator<'_, T> {
TreeBreadthFirstIterator {
queue: self
.root()
.map_or(VecDeque::new(), |root| vec![root].into_iter().collect()),
}
}
fn apply<F: Fn(&mut TreeNode<T>)>(&mut self, visit_fn: F) {
let visitor = TreeVisitor::new(visit_fn);
if let Some(root) = self.root_mut() {
visitor.visit(root);
}
}
}
impl<T> TreeIterator<T> for TreeChromosome<T> {
fn iter_pre_order(&self) -> PreOrderIterator<'_, T> {
self.root().iter_pre_order()
}
fn iter_post_order(&self) -> PostOrderIterator<'_, T> {
self.root().iter_post_order()
}
fn iter_breadth_first(&self) -> TreeBreadthFirstIterator<'_, T> {
self.root().iter_breadth_first()
}
fn apply<F: Fn(&mut TreeNode<T>)>(&mut self, visit_fn: F) {
self.root_mut().apply(visit_fn);
}
}
pub struct PreOrderIterator<'a, T> {
stack: Vec<&'a TreeNode<T>>,
}
impl<'a, T> Iterator for PreOrderIterator<'a, T> {
type Item = &'a TreeNode<T>;
fn next(&mut self) -> Option<Self::Item> {
self.stack.pop().inspect(|node| {
if let Some(children) = node.children() {
for child in children.iter().rev() {
self.stack.push(child);
}
}
})
}
}
pub struct PostOrderIterator<'a, T> {
stack: Vec<(&'a TreeNode<T>, bool)>,
}
impl<'a, T> Iterator for PostOrderIterator<'a, T> {
type Item = &'a TreeNode<T>;
fn next(&mut self) -> Option<Self::Item> {
while let Some((node, visited)) = self.stack.pop() {
if visited {
return Some(node);
}
self.stack.push((node, true));
if let Some(children) = node.children() {
for child in children.iter().rev() {
self.stack.push((child, false));
}
}
}
None
}
}
pub struct TreeBreadthFirstIterator<'a, T> {
queue: VecDeque<&'a TreeNode<T>>,
}
impl<'a, T> Iterator for TreeBreadthFirstIterator<'a, T> {
type Item = &'a TreeNode<T>;
fn next(&mut self) -> Option<Self::Item> {
let node = self.queue.pop_front()?;
if let Some(children) = node.children() {
self.queue.extend(children.iter());
}
Some(node)
}
}
pub struct TreeVisitor<T, F>
where
F: Fn(&mut TreeNode<T>),
{
visitor: F,
_marker: PhantomData<T>,
}
impl<T, F> TreeVisitor<T, F>
where
F: Fn(&mut TreeNode<T>),
{
pub fn new(visitor: F) -> Self {
TreeVisitor {
visitor,
_marker: PhantomData,
}
}
pub fn visit(&self, node: &mut TreeNode<T>) {
(self.visitor)(node);
if let Some(children) = node.children_mut() {
for child in children.iter_mut() {
self.visit(child);
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::Op;
use crate::collections::{Tree, TreeNode};
use crate::node::Node;
#[test]
fn test_tree_traversal() {
let leaf = Op::constant(4.0);
let node2 = TreeNode::with_children(Op::constant(2.0), vec![TreeNode::new(leaf)]);
let node3 = TreeNode::new(Op::constant(3.0));
let root = Tree::new(TreeNode::with_children(
Op::constant(1.0),
vec![node2, node3],
));
let pre_order: Vec<f32> = root
.iter_pre_order()
.map(|n| match &n.value() {
Op::Const(_, v) => *v,
_ => panic!("Expected constant but got {:?}", n.value()),
})
.collect();
assert_eq!(pre_order, vec![1.0, 2.0, 4.0, 3.0]);
let post_order: Vec<f32> = root
.iter_post_order()
.map(|n| match &n.value() {
Op::Const(_, v) => *v,
_ => panic!("Expected constant"),
})
.collect();
assert_eq!(post_order, vec![4.0, 2.0, 3.0, 1.0]);
let bfs: Vec<f32> = root
.iter_breadth_first()
.map(|n| match &n.value() {
Op::Const(_, v) => *v,
_ => panic!("Expected constant"),
})
.collect();
assert_eq!(bfs, vec![1.0, 2.0, 3.0, 4.0]);
}
}