use crate::graph::{NodeId, PortIdx};
use std::collections::VecDeque;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) struct Edge {
pub from: (NodeId, PortIdx),
pub to: (NodeId, PortIdx),
}
pub(crate) fn topological_sort(
num_nodes: usize,
edges: &[Edge],
) -> Result<Vec<NodeId>, Vec<NodeId>> {
let mut in_degree = vec![0_usize; num_nodes];
for edge in edges {
if edge.to.0 < num_nodes {
in_degree[edge.to.0] = in_degree[edge.to.0].saturating_add(1);
}
}
let mut queue: VecDeque<NodeId> = (0..num_nodes).filter(|&n| in_degree[n] == 0).collect();
let mut order = Vec::with_capacity(num_nodes);
while let Some(n) = queue.pop_front() {
order.push(n);
for edge in edges {
if edge.from.0 == n && edge.to.0 < num_nodes {
let idx = edge.to.0;
if in_degree[idx] > 0 {
in_degree[idx] -= 1;
if in_degree[idx] == 0 {
queue.push_back(idx);
}
}
}
}
}
if order.len() == num_nodes {
Ok(order)
} else {
let remaining: Vec<NodeId> = (0..num_nodes).filter(|&n| in_degree[n] > 0).collect();
Err(remaining)
}
}
#[cfg(test)]
mod tests {
use super::*;
use proptest::prelude::*;
fn is_valid_topo_order(order: &[NodeId], num_nodes: usize, edges: &[Edge]) -> bool {
if order.len() != num_nodes {
return false;
}
let mut seen = vec![false; num_nodes];
for &node in order {
if node >= num_nodes {
return false;
}
if std::mem::replace(&mut seen[node], true) {
return false;
}
}
for edge in edges {
if edge.from.0 < num_nodes && edge.to.0 < num_nodes {
let Some(p_from) = order.iter().position(|&x| x == edge.from.0) else {
return false;
};
let Some(p_to) = order.iter().position(|&x| x == edge.to.0) else {
return false;
};
if p_from >= p_to {
return false;
}
}
}
true
}
#[test]
fn empty_graph_is_ok() {
let order = topological_sort(0, &[]).unwrap();
assert!(order.is_empty());
}
#[test]
fn no_edges_preserves_all_nodes() {
let order = topological_sort(3, &[]).unwrap();
assert_eq!(order.len(), 3);
assert_eq!(order, vec![0, 1, 2]);
}
#[test]
fn linear_chain_sorts_in_dependency_order() {
let edges = [
Edge {
from: (0, 0),
to: (1, 0),
},
Edge {
from: (1, 0),
to: (2, 0),
},
];
let order = topological_sort(3, &edges).unwrap();
assert_eq!(order, vec![0, 1, 2]);
}
#[test]
fn diamond_dag_respects_dependencies() {
let edges = [
Edge {
from: (0, 0),
to: (1, 0),
},
Edge {
from: (0, 0),
to: (2, 0),
},
Edge {
from: (1, 0),
to: (3, 0),
},
Edge {
from: (2, 0),
to: (3, 0),
},
];
let order = topological_sort(4, &edges).unwrap();
assert_eq!(order[0], 0);
assert_eq!(order[3], 3);
assert!(order[1] == 1 || order[1] == 2);
assert!(order[2] == 1 || order[2] == 2);
}
#[test]
fn disconnected_nodes_are_all_included() {
let edges = [Edge {
from: (0, 0),
to: (1, 0),
}];
let order = topological_sort(4, &edges).unwrap();
assert_eq!(order.len(), 4);
let pos0 = order.iter().position(|&n| n == 0).unwrap();
let pos1 = order.iter().position(|&n| n == 1).unwrap();
assert!(pos0 < pos1);
assert!(order.contains(&2));
assert!(order.contains(&3));
}
#[test]
fn self_loop_is_a_cycle() {
let edges = [Edge {
from: (0, 0),
to: (0, 0),
}];
let result = topological_sort(1, &edges);
let remaining = result.expect_err("self-loop must be a cycle");
assert_eq!(remaining, vec![0]);
}
#[test]
fn two_node_cycle_is_detected() {
let edges = [
Edge {
from: (0, 0),
to: (1, 0),
},
Edge {
from: (1, 0),
to: (0, 0),
},
];
let result = topological_sort(2, &edges);
let remaining = result.expect_err("two-node cycle must be detected");
assert_eq!(remaining.len(), 2);
}
#[test]
fn edges_beyond_num_nodes_are_skipped() {
let edges = [Edge {
from: (0, 0),
to: (5, 0),
}];
let order = topological_sort(2, &edges).unwrap();
assert_eq!(order.len(), 2);
}
proptest! {
#[test]
fn prop_low_to_high_dag_yields_valid_order(
n in 2usize..=16,
seed in prop::collection::vec((proptest::num::u8::ANY, proptest::num::u8::ANY), 0..=30),
) {
let mut edges: Vec<Edge> = Vec::new();
for &(a, b) in &seed {
let from = usize::from(a) % (n - 1);
let span = n - from - 1;
let to = from + 1 + usize::from(b) % span;
edges.push(Edge { from: (from, 0), to: (to, 0) });
}
let order = topological_sort(n, &edges).expect("acyclic graph must sort Ok");
prop_assert!(
is_valid_topo_order(&order, n, &edges),
"order {:?} invalid for n={} edges={:?}",
order,
n,
edges
);
}
}
}