use std::collections::HashSet;
use crate::graph::capability::StableNode;
use crate::graph::Graph;
pub struct Dfs<'r, G: ?Sized, N> {
graph: &'r G,
stack: Vec<N>,
visited: HashSet<N>,
}
impl<'r, G> Dfs<'r, G, G::NodeIx>
where
G: Graph + ?Sized,
{
pub unsafe fn new_unchecked(graph: &'r G, start: G::NodeIx) -> Self
where
G: StableNode,
{
let mut visited = HashSet::new();
let mut stack = Vec::new();
visited.insert(start);
stack.push(start);
Dfs {
graph,
stack,
visited,
}
}
pub fn new(graph: &'r G, start: G::NodeIx) -> Self
where
G: StableNode,
{
unsafe { Self::new_unchecked(graph, start) }
}
pub unsafe fn add_start_unchecked(&mut self, start: G::NodeIx) {
if self.visited.insert(start) {
self.stack.push(start);
}
}
pub fn add_start(&mut self, start: G::NodeIx) {
assert!(Graph::contains_node_index(self.graph, start));
unsafe { self.add_start_unchecked(start) }
}
}
impl<'r, G> Iterator for Dfs<'r, G, G::NodeIx>
where
G: Graph + StableNode + ?Sized,
{
type Item = G::NodeIx;
fn next(&mut self) -> Option<G::NodeIx> {
let node = self.stack.pop()?;
for eix in unsafe { <G as crate::graph::GraphOperation<'_>>::edge_indices_from_unchecked(self.graph, node) } {
for endpoint in unsafe { <G as crate::graph::GraphOperation<'_>>::endpoints_unchecked(self.graph, eix) } {
if endpoint != node && self.visited.insert(endpoint) {
self.stack.push(endpoint);
}
}
}
Some(node)
}
}
pub struct DfsPostOrder<'r, G: ?Sized, N> {
graph: &'r G,
stack: Vec<(N, bool)>,
visited: HashSet<N>,
}
impl<'r, G> DfsPostOrder<'r, G, G::NodeIx>
where
G: Graph + ?Sized,
{
pub fn new(graph: &'r G, start: G::NodeIx) -> Self
where
G: StableNode,
{
unsafe { Self::new_unchecked(graph, start) }
}
pub unsafe fn new_unchecked(graph: &'r G, start: G::NodeIx) -> Self
where
G: StableNode,
{
let mut visited = HashSet::new();
let mut stack = Vec::new();
if Graph::contains_node_index(graph, start) {
visited.insert(start);
stack.push((start, false));
}
DfsPostOrder {
graph,
stack,
visited,
}
}
pub fn add_start(&mut self, start: G::NodeIx) {
if self.visited.insert(start) {
self.stack.push((start, false));
}
}
}
impl<'r, G> Iterator for DfsPostOrder<'r, G, G::NodeIx>
where
G: Graph + StableNode + ?Sized,
{
type Item = G::NodeIx;
fn next(&mut self) -> Option<G::NodeIx> {
loop {
let (node, expanded) = self.stack.last_mut()?;
if *expanded {
let node = self.stack.pop()?.0;
return Some(node);
}
*expanded = true;
let node = *node;
let succs: Vec<G::NodeIx> = unsafe { <G as crate::graph::GraphOperation<'_>>::edge_indices_from_unchecked(self.graph, node) }
.flat_map(|eix| unsafe { <G as crate::graph::GraphOperation<'_>>::endpoints_unchecked(self.graph, eix) }.into_iter())
.filter(|&ep| ep != node)
.collect();
for succ in succs.into_iter().rev() {
if self.visited.insert(succ) {
self.stack.push((succ, false));
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::BTreeGraph;
fn diamond_btree() -> BTreeGraph<u32, &'static str> {
let mut g = BTreeGraph::<_, _>::default();
g.insert_node(0).unwrap();
g.insert_node(1).unwrap();
g.insert_node(2).unwrap();
g.insert_node(3).unwrap();
g.insert_edge("0->1", [0, 1]).unwrap();
g.insert_edge("0->2", [0, 2]).unwrap();
g.insert_edge("1->3", [1, 3]).unwrap();
g.insert_edge("2->3", [2, 3]).unwrap();
g
}
fn linear_btree() -> BTreeGraph<u32, &'static str> {
let mut g = BTreeGraph::<_, _>::default();
g.insert_node(0).unwrap();
g.insert_node(1).unwrap();
g.insert_node(2).unwrap();
g.insert_node(3).unwrap();
g.insert_edge("0->1", [0, 1]).unwrap();
g.insert_edge("1->2", [1, 2]).unwrap();
g.insert_edge("2->3", [2, 3]).unwrap();
g
}
#[test]
fn dfs_diamond() {
let g = diamond_btree();
let order: Vec<u32> = Dfs::new(&g, 0).collect();
assert_eq!(order.len(), 4);
assert_eq!(order[0], 0);
assert!(order.contains(&1));
assert!(order.contains(&2));
assert!(order.contains(&3));
}
#[test]
fn dfs_linear() {
let g = linear_btree();
let order: Vec<u32> = Dfs::new(&g, 0).collect();
assert_eq!(order, vec![0, 1, 2, 3]);
}
#[test]
fn dfs_post_order_diamond() {
let g = diamond_btree();
let order: Vec<u32> = DfsPostOrder::new(&g, 0).collect();
assert_eq!(order.len(), 4);
assert_eq!(order[3], 0);
let pos3 = order.iter().position(|&x| x == 3).unwrap();
let pos1 = order.iter().position(|&x| x == 1).unwrap();
let pos2 = order.iter().position(|&x| x == 2).unwrap();
assert!(pos3 < pos1 || pos3 < pos2);
}
#[test]
fn dfs_post_order_linear() {
let g = linear_btree();
let order: Vec<u32> = DfsPostOrder::new(&g, 0).collect();
assert_eq!(order, vec![3, 2, 1, 0]);
}
#[test]
fn dfs_with_cycle() {
let mut g = BTreeGraph::<_, _>::default();
g.insert_node(0).unwrap();
g.insert_node(1).unwrap();
g.insert_node(2).unwrap();
g.insert_edge("0->1", [0, 1]).unwrap();
g.insert_edge("1->2", [1, 2]).unwrap();
g.insert_edge("2->0", [2, 0]).unwrap();
let order: Vec<u32> = Dfs::new(&g, 0).collect();
assert_eq!(order.len(), 3);
}
}