use std::collections::HashSet;
use petgraph::data::DataMap;
use petgraph::graph::NodeIndex;
use petgraph::visit::IntoNeighborsDirected;
use petgraph::visit::Reversed;
use petgraph::visit::{GraphBase, IntoNeighbors, IntoNodeIdentifiers, Visitable};
use petgraph::Incoming;
#[derive(Clone)]
pub struct StableTopo<G> {
graph: G,
ordered: HashSet<NodeIndex>,
tovisit: Vec<NodeIndex>,
}
impl<G> StableTopo<G>
where
G: IntoNeighborsDirected + IntoNodeIdentifiers + Visitable,
G: GraphBase<NodeId = NodeIndex>,
{
pub fn new(graph: G) -> Self {
let mut topo = StableTopo {
graph,
ordered: HashSet::new(),
tovisit: Vec::new(),
};
topo.extend_with_initials();
topo
}
pub fn extend_with_initials(&mut self) {
self.tovisit.extend(
self.graph
.node_identifiers()
.filter(|&a| self.graph.neighbors_directed(a, Incoming).next().is_none()),
);
}
}
impl<G> Iterator for StableTopo<G>
where
G: IntoNeighborsDirected + IntoNodeIdentifiers + Visitable + DataMap,
G: GraphBase<NodeId = NodeIndex>,
G::NodeWeight: Ord,
{
type Item = NodeIndex;
fn next(&mut self) -> Option<Self::Item> {
self.tovisit.sort_unstable_by(|a, b| {
self.graph
.node_weight(*a)
.unwrap()
.cmp(self.graph.node_weight(*b).unwrap())
});
while let Some(nix) = self.tovisit.pop() {
if self.ordered.contains(&nix) {
continue;
}
self.ordered.insert(nix);
let mut neighbors = Vec::new();
for neigh in self.graph.neighbors(nix) {
if Reversed(&self.graph)
.neighbors(neigh)
.all(|b| self.ordered.contains(&b))
{
neighbors.push(neigh);
}
}
neighbors.sort_unstable_by(|a, b| {
self.graph
.node_weight(*a)
.unwrap()
.cmp(self.graph.node_weight(*b).unwrap())
});
self.tovisit.extend(neighbors);
return Some(nix);
}
None
}
}
#[cfg(test)]
mod tests {
use petgraph::prelude::*;
use super::*;
#[test]
fn test_stable_topo() {
let mut graph: Graph<&str, (), Directed> = Graph::new();
let node1 = graph.add_node("Node 1");
let node2 = graph.add_node("Node 2");
let node3 = graph.add_node("Node 3");
let node4 = graph.add_node("Node 4");
graph.add_edge(node1, node2, ());
graph.add_edge(node2, node3, ());
graph.add_edge(node2, node4, ());
graph.add_edge(node3, node4, ());
let stable_topo = StableTopo::new(&graph);
let topo_order: Vec<NodeIndex> = stable_topo.collect();
assert_eq!(topo_order, vec![node1, node2, node3, node4]);
}
#[test]
fn test_stable_topo_with_weights() {
let mut graph: Graph<&str, (), Directed> = Graph::new();
let node4 = graph.add_node("Node 4");
let node3 = graph.add_node("Node 3");
let node2 = graph.add_node("Node 2");
let node1 = graph.add_node("Node 1");
graph.add_edge(node1, node2, ());
graph.add_edge(node2, node3, ());
graph.add_edge(node2, node4, ());
graph.add_edge(node3, node4, ());
let stable_topo = StableTopo::new(&graph);
let topo_order: Vec<NodeIndex> = stable_topo.collect();
assert_eq!(topo_order, vec![node1, node2, node3, node4]);
}
}