use std::collections::VecDeque;
use itertools::enumerate;
use super::graph::{Label, Tree, TreeIndex};
#[derive(Clone, Copy, Debug)]
pub struct DfsNodeData {
pub depth: usize,
pub index: TreeIndex,
pub n_remaining: usize,
}
impl DfsNodeData {
pub fn extract(self) -> (usize, TreeIndex, usize) {
(self.depth, self.index, self.n_remaining)
}
}
pub trait TraversalMut {
type Item;
fn new<N, const K: usize>(tree: &Tree<N, K>, root: TreeIndex) -> Self;
fn iter<N, const K: usize>(tree: &Tree<N, K>, root: TreeIndex) -> TraversalIter<'_, Self, N, K>
where
Self: Sized,
{
TraversalIter::from(Self::new(tree, root), tree)
}
fn next<N, const K: usize>(&mut self, tree: &Tree<N, K>) -> Option<Self::Item>;
fn skip_subtree(&mut self);
fn size_hint(&self) -> (usize, Option<usize>);
}
#[derive(Debug)]
pub struct TraversalIter<'a, T: TraversalMut, N, const K: usize> {
traversal: T,
tree: &'a Tree<N, K>,
}
impl<'a, T: TraversalMut, N, const K: usize> TraversalIter<'a, T, N, K> {
pub fn new(tree: &'a Tree<N, K>, root: TreeIndex) -> TraversalIter<'a, T, N, K> {
TraversalIter {
traversal: T::new(tree, root),
tree,
}
}
pub fn from(traversal: T, tree: &'a Tree<N, K>) -> TraversalIter<'a, T, N, K> {
TraversalIter { traversal, tree }
}
pub fn skip_subtree(&mut self) {
self.traversal.skip_subtree()
}
}
impl<T: TraversalMut, N, const K: usize> Iterator for TraversalIter<'_, T, N, K> {
type Item = T::Item;
fn next(&mut self) -> Option<Self::Item> {
self.traversal.next(self.tree)
}
fn size_hint(&self) -> (usize, Option<usize>) {
self.traversal.size_hint()
}
}
#[derive(Clone, Debug)]
pub struct DfsPre {
pub(super) stack: Vec<DfsNodeData>,
pub(super) last_push: usize,
pub size_lb: usize,
pub size_ub: usize,
}
impl TraversalMut for DfsPre {
type Item = DfsNodeData;
#[inline]
fn new<N, const K: usize>(tree: &Tree<N, K>, root: TreeIndex) -> DfsPre {
DfsPre {
stack: vec![DfsNodeData {
depth: 0,
index: root,
n_remaining: 0,
}],
last_push: 0,
size_lb: if root == tree.get_root_idx() {
tree.len()
} else {
0
},
size_ub: tree.len(),
}
}
fn skip_subtree(&mut self) {
self.size_lb = self.stack.len();
self.size_ub -= self.last_push;
for _ in 0..self.last_push {
self.stack.pop();
}
}
fn next<N, const K: usize>(&mut self, tree: &Tree<N, K>) -> Option<DfsNodeData> {
let data = self.stack.pop()?;
let node = tree
.tree_node(data.index)
.expect("node indicies should stay valid while traversing the tree");
self.last_push = 0;
for (n_remaining, child) in node.children.iter().rev().flatten().enumerate() {
self.stack.push(DfsNodeData {
depth: data.depth + 1,
index: *child,
n_remaining,
});
self.last_push += 1;
}
self.size_lb = self.size_lb.saturating_sub(1);
self.size_ub = self.size_ub.saturating_sub(1);
Some(data)
}
fn size_hint(&self) -> (usize, Option<usize>) {
(self.size_lb, Some(self.size_ub))
}
}
#[derive(Debug)]
pub struct EdgeData {
pub src: TreeIndex,
pub label: Label,
pub dest: TreeIndex,
}
#[derive(Clone, Debug)]
pub struct DfsEdge {
pub(super) stack: Vec<(usize, TreeIndex, Label, TreeIndex)>,
pub(super) last_push: usize,
pub size_lb: usize,
pub size_ub: usize,
}
impl TraversalMut for DfsEdge {
type Item = EdgeData;
#[inline]
fn new<N, const K: usize>(tree: &Tree<N, K>, root: TreeIndex) -> DfsEdge {
let mut stack = Vec::with_capacity(K);
let mut last_push = 0;
for ed in tree.children(tree.get_root_idx()).rev() {
stack.push((1, root, ed.label, ed.target_idx));
last_push += 1;
}
DfsEdge {
stack,
last_push,
size_lb: if root == tree.get_root_idx() {
tree.len()
} else {
0
},
size_ub: tree.len(),
}
}
fn skip_subtree(&mut self) {
self.size_lb = self.stack.len();
self.size_ub -= self.last_push;
for _ in 0..self.last_push {
self.stack.pop();
}
}
fn next<N, const K: usize>(&mut self, tree: &Tree<N, K>) -> Option<Self::Item> {
let (depth, src_idx, label, dest_idx) = self.stack.pop()?;
let node = tree.tree_node(dest_idx).unwrap();
self.last_push = 0;
for (label, child) in enumerate(node.children.iter()).rev() {
if let Some(val) = child {
self.stack.push((depth + 1, dest_idx, label, *val));
self.last_push += 1;
}
}
self.size_lb = self.size_lb.saturating_sub(1);
self.size_ub = self.size_ub.saturating_sub(1);
Some(EdgeData {
src: src_idx,
label,
dest: dest_idx,
})
}
fn size_hint(&self) -> (usize, Option<usize>) {
(self.size_lb, Some(self.size_ub))
}
}
#[derive(Clone, Debug)]
pub struct Bfs {
pub(super) queue: VecDeque<DfsNodeData>,
pub(super) last_push: usize,
pub size_lb: usize,
pub size_ub: usize,
}
impl TraversalMut for Bfs {
type Item = DfsNodeData;
#[inline]
fn new<N, const K: usize>(tree: &Tree<N, K>, root: TreeIndex) -> Bfs {
Bfs {
queue: VecDeque::from([DfsNodeData {
depth: 0,
index: root,
n_remaining: 0,
}]),
last_push: 0,
size_lb: if root == tree.get_root_idx() {
tree.len()
} else {
0
},
size_ub: tree.len(),
}
}
fn skip_subtree(&mut self) {
self.size_lb = self.queue.len();
self.size_ub -= self.last_push;
for _ in 0..self.last_push {
self.queue.pop_back();
}
}
fn next<N, const K: usize>(&mut self, tree: &Tree<N, K>) -> Option<DfsNodeData> {
let data = self.queue.pop_front()?;
let node = tree.tree_node(data.index).unwrap();
self.last_push = 0;
for (n_remaining, child) in node.children.iter().flatten().enumerate() {
self.queue.push_back(DfsNodeData {
depth: data.depth + 1,
index: *child,
n_remaining,
});
self.last_push += 1;
}
self.size_lb = self.size_lb.saturating_sub(1);
self.size_ub = self.size_ub.saturating_sub(1);
Some(data)
}
fn size_hint(&self) -> (usize, Option<usize>) {
(self.size_lb, Some(self.size_ub))
}
}
#[allow(unused_variables)]
#[cfg(test)]
mod test {
use assertables::*;
use super::*;
use crate::tree::iter::{Bfs, DfsEdge, DfsPre, TraversalMut};
#[test]
pub fn test_dfs_node_order() {
let mut tree = Tree::<usize, 2>::new();
let z = tree.add_root(10); let c0 = tree.add_child_node(z, 0, 11).unwrap(); let c1 = tree.add_child_node(z, 1, 12).unwrap(); let l0 = tree.add_child_node(c0, 0, 13).unwrap(); let l1 = tree.add_child_node(c0, 1, 14).unwrap(); let r0 = tree.add_child_node(c1, 0, 15).unwrap(); let r1 = tree.add_child_node(c1, 1, 16).unwrap(); let l2 = tree.add_child_node(l0, 1, 17).unwrap(); let l2r = tree.add_child_node(l2, 0, 18).unwrap(); let l2l = tree.add_child_node(l2, 1, 19).unwrap();
let iter = DfsPre::iter(&tree, z);
let nodes = Vec::from_iter(iter.map(|data| data.index));
assert_eq!(nodes, vec![z, c0, l0, l2, l2r, l2l, l1, c1, r0, r1]);
}
#[test]
pub fn test_dfs_skip_subtree() {
let mut tree = Tree::<usize, 2>::new();
let z = tree.add_root(10); let c0 = tree.add_child_node(z, 0, 11).unwrap(); let c1 = tree.add_child_node(z, 1, 12).unwrap(); let l0 = tree.add_child_node(c0, 0, 13).unwrap(); let l1 = tree.add_child_node(c0, 1, 14).unwrap(); let r0 = tree.add_child_node(c1, 0, 15).unwrap(); let r1 = tree.add_child_node(c1, 1, 16).unwrap(); let l2 = tree.add_child_node(l0, 1, 17).unwrap(); let l2r = tree.add_child_node(l2, 0, 18).unwrap(); let l2l = tree.add_child_node(l2, 1, 19).unwrap();
let mut iter = DfsPre::iter(&tree, z);
iter.next(); iter.next();
iter.skip_subtree();
let nodes = Vec::from_iter(iter.map(|data| data.index));
assert_eq!(nodes, vec![c1, r0, r1]);
}
#[test]
pub fn test_dfs_remaining() {
let mut tree = Tree::<usize, 2>::new();
let z = tree.add_root(10); let c0 = tree.add_child_node(z, 0, 11).unwrap(); let c1 = tree.add_child_node(z, 1, 12).unwrap(); let l0 = tree.add_child_node(c0, 0, 13).unwrap(); let l1 = tree.add_child_node(c0, 1, 14).unwrap(); let r0 = tree.add_child_node(c1, 0, 15).unwrap(); let r1 = tree.add_child_node(c1, 1, 16).unwrap(); let l2 = tree.add_child_node(l0, 1, 17).unwrap(); let l2r = tree.add_child_node(l2, 0, 18).unwrap(); let l2l = tree.add_child_node(l2, 1, 19).unwrap();
let iter = DfsPre::iter(&tree, z);
let remaining = Vec::from_iter(iter.map(|data| data.n_remaining));
assert_eq!(remaining, vec![0, 1, 1, 0, 1, 0, 0, 0, 1, 0]);
}
#[test]
pub fn test_dfs_size_hint() {
let mut tree = Tree::<usize, 2>::new();
let z = tree.add_root(10); let c0 = tree.add_child_node(z, 0, 11).unwrap(); let c1 = tree.add_child_node(z, 1, 12).unwrap(); let l0 = tree.add_child_node(c0, 0, 13).unwrap(); let l1 = tree.add_child_node(c0, 1, 14).unwrap(); let r0 = tree.add_child_node(c1, 0, 15).unwrap(); let r1 = tree.add_child_node(c1, 1, 16).unwrap(); let l2 = tree.add_child_node(l0, 1, 17).unwrap(); let l2r = tree.add_child_node(l2, 0, 18).unwrap(); let l2l = tree.add_child_node(l2, 1, 19).unwrap();
let mut iter = DfsPre::iter(&tree, z);
iter.next(); iter.next();
assert_eq!(iter.size_hint(), (8, Some(8)));
iter.skip_subtree();
assert_le!(iter.size_hint().0, 3);
assert_ge!(iter.size_hint().1.unwrap(), 3);
}
#[test]
pub fn test_dfs_edge_iter() {
let mut tree = Tree::<(), 2>::new();
let z = tree.add_root(()); let c0 = tree.add_child_node(z, 0, ()).unwrap(); let c1 = tree.add_child_node(z, 1, ()).unwrap(); let l0 = tree.add_child_node(c0, 0, ()).unwrap(); let l1 = tree.add_child_node(c0, 1, ()).unwrap(); let r0 = tree.add_child_node(c1, 0, ()).unwrap(); let r1 = tree.add_child_node(c1, 1, ()).unwrap(); let rr1 = tree.add_child_node(r1, 1, ()).unwrap();
let iter = DfsEdge::iter(&tree, z);
let nodes = Vec::from_iter(iter.map(|edge| (edge.src, edge.label, edge.dest)));
assert_eq!(
nodes,
vec![
(z, 0, c0),
(c0, 0, l0),
(c0, 1, l1),
(z, 1, c1),
(c1, 0, r0),
(c1, 1, r1),
(r1, 1, rr1)
]
);
let mut iter = DfsEdge::iter(&tree, z);
iter.next();
iter.skip_subtree();
assert_eq!(iter.next().unwrap().dest, c1);
}
#[test]
fn test_bfs_node_order() {
let mut tree = Tree::<usize, 2>::new();
let z = tree.add_root(10); let c0 = tree.add_child_node(z, 0, 11).unwrap(); let c1 = tree.add_child_node(z, 1, 12).unwrap(); let l0 = tree.add_child_node(c0, 0, 13).unwrap(); let l1 = tree.add_child_node(c0, 1, 14).unwrap(); let r0 = tree.add_child_node(c1, 0, 15).unwrap(); let r1 = tree.add_child_node(c1, 1, 16).unwrap(); let l2 = tree.add_child_node(l0, 1, 17).unwrap(); let l2r = tree.add_child_node(l2, 0, 18).unwrap(); let l2l = tree.add_child_node(l2, 1, 19).unwrap();
let iter = Bfs::iter(&tree, z);
let nodes = Vec::from_iter(iter.map(|data| data.index));
assert_eq!(nodes, vec![z, c0, c1, l0, l1, r0, r1, l2, l2r, l2l]);
}
#[test]
pub fn test_bfs_skip_subtree() {
let mut tree = Tree::<usize, 2>::new();
let z = tree.add_root(10); let c0 = tree.add_child_node(z, 0, 11).unwrap(); let c1 = tree.add_child_node(z, 1, 12).unwrap(); let l0 = tree.add_child_node(c0, 0, 13).unwrap(); let l1 = tree.add_child_node(c0, 1, 14).unwrap(); let r0 = tree.add_child_node(c1, 0, 15).unwrap(); let r1 = tree.add_child_node(c1, 1, 16).unwrap(); let l2 = tree.add_child_node(l0, 1, 17).unwrap(); let l2r = tree.add_child_node(l2, 0, 18).unwrap(); let l2l = tree.add_child_node(l2, 1, 19).unwrap();
let mut iter = Bfs::iter(&tree, z);
iter.next(); iter.next();
iter.skip_subtree();
let nodes = Vec::from_iter(iter.map(|data| data.index));
assert_eq!(nodes, vec![c1, r0, r1]);
}
}