use std::cmp::Reverse;
use std::collections::{BinaryHeap, HashMap, HashSet, VecDeque};
use std::fmt;
use std::hash::Hash;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct CycleError;
impl fmt::Display for CycleError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("edge would create a cycle")
}
}
impl std::error::Error for CycleError {}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Dag<K: Hash + Eq + Clone> {
forward: HashMap<K, HashSet<K>>,
reverse: HashMap<K, HashSet<K>>,
edge_count: usize,
}
impl<K: Hash + Eq + Clone> Dag<K> {
#[inline]
pub fn new() -> Self {
Self {
forward: HashMap::new(),
reverse: HashMap::new(),
edge_count: 0,
}
}
#[inline]
pub fn node_count(&self) -> usize {
self.forward.len()
}
#[inline]
pub fn edge_count(&self) -> usize {
self.edge_count
}
#[inline]
pub fn is_empty(&self) -> bool {
self.forward.is_empty()
}
#[inline]
pub fn contains_node(&self, key: &K) -> bool {
self.forward.contains_key(key)
}
#[inline]
pub fn contains_edge(&self, from: &K, to: &K) -> bool {
self.forward.get(from).is_some_and(|s| s.contains(to))
}
#[inline]
pub fn nodes(&self) -> impl Iterator<Item = &K> {
self.forward.keys()
}
pub fn edges(&self) -> impl Iterator<Item = (&K, &K)> {
self.forward
.iter()
.flat_map(|(from, tos)| tos.iter().map(move |to| (from, to)))
}
pub fn sources(&self) -> Vec<&K> {
self.reverse
.iter()
.filter(|(_, deps)| deps.is_empty())
.map(|(k, _)| k)
.collect()
}
pub fn sinks(&self) -> Vec<&K> {
self.forward
.iter()
.filter(|(_, deps)| deps.is_empty())
.map(|(k, _)| k)
.collect()
}
#[inline]
pub fn insert_node(&mut self, key: K) -> bool {
if self.forward.contains_key(&key) {
return false;
}
self.forward.insert(key.clone(), HashSet::new());
self.reverse.insert(key, HashSet::new());
true
}
pub fn add_edge(&mut self, from: K, to: K) -> Result<bool, CycleError> {
if from == to {
return Err(CycleError);
}
if self.contains_edge(&from, &to) {
return Ok(false);
}
if self.is_reachable(&to, &from) {
return Err(CycleError);
}
self.insert_node(from.clone());
self.insert_node(to.clone());
self.forward.get_mut(&from).unwrap().insert(to.clone());
self.reverse.get_mut(&to).unwrap().insert(from);
self.edge_count += 1;
Ok(true)
}
pub fn remove_node(&mut self, key: &K) -> bool {
let Some(dependents) = self.forward.remove(key) else {
return false;
};
self.edge_count -= dependents.len();
for dep in &dependents {
if let Some(rev) = self.reverse.get_mut(dep) {
rev.remove(key);
}
}
if let Some(dependencies) = self.reverse.remove(key) {
self.edge_count -= dependencies.len();
for dep in &dependencies {
if let Some(fwd) = self.forward.get_mut(dep) {
fwd.remove(key);
}
}
}
true
}
#[inline]
pub fn remove_edge(&mut self, from: &K, to: &K) -> bool {
let removed = self
.forward
.get_mut(from)
.is_some_and(|s| s.remove(to));
if removed {
self.reverse.get_mut(to).unwrap().remove(from);
self.edge_count -= 1;
}
removed
}
#[inline]
pub fn dependents(&self, key: &K) -> Option<&HashSet<K>> {
self.forward.get(key)
}
#[inline]
pub fn dependencies(&self, key: &K) -> Option<&HashSet<K>> {
self.reverse.get(key)
}
pub fn transitive_dependents(&self, key: &K) -> HashSet<K> {
self.walk(&self.forward, key)
}
pub fn transitive_dependencies(&self, key: &K) -> HashSet<K> {
self.walk(&self.reverse, key)
}
pub fn topological_order(&self) -> Vec<K>
where
K: Ord,
{
let in_degree: HashMap<&K, usize> = self
.reverse
.iter()
.map(|(k, deps)| (k, deps.len()))
.collect();
self.kahn(in_degree, None)
}
pub fn topological_order_from(&self, starts: &[K]) -> Vec<K>
where
K: Ord,
{
let mut affected: HashSet<&K> = HashSet::new();
for start in starts {
if let Some((key, _)) = self.forward.get_key_value(start) {
affected.insert(key);
for dep in self.transitive_dependents(start) {
let (key, _) = self.forward.get_key_value(&dep).unwrap();
affected.insert(key);
}
}
}
let in_degree: HashMap<&K, usize> = affected
.iter()
.map(|&k| {
let inside = self.reverse[k]
.iter()
.filter(|dep| affected.contains(dep))
.count();
(k, inside)
})
.collect();
self.kahn(in_degree, Some(&affected))
}
fn kahn(&self, mut in_degree: HashMap<&K, usize>, within: Option<&HashSet<&K>>) -> Vec<K>
where
K: Ord,
{
let mut ready: BinaryHeap<Reverse<&K>> = in_degree
.iter()
.filter(|(_, degree)| **degree == 0)
.map(|(&k, _)| Reverse(k))
.collect();
let mut order = Vec::with_capacity(in_degree.len());
while let Some(Reverse(current)) = ready.pop() {
order.push(current.clone());
for dependent in &self.forward[current] {
if within.is_some_and(|set| !set.contains(dependent)) {
continue;
}
let degree = in_degree.get_mut(dependent).unwrap();
*degree -= 1;
if *degree == 0 {
ready.push(Reverse(dependent));
}
}
}
order
}
fn walk(&self, adj: &HashMap<K, HashSet<K>>, start: &K) -> HashSet<K> {
let mut result = HashSet::new();
let mut queue = VecDeque::new();
if let Some(neighbours) = adj.get(start) {
for n in neighbours {
if result.insert(n.clone()) {
queue.push_back(n);
}
}
}
while let Some(current) = queue.pop_front() {
if let Some(neighbours) = adj.get(current) {
for n in neighbours {
if result.insert(n.clone()) {
queue.push_back(n);
}
}
}
}
result
}
pub fn is_reachable(&self, start: &K, target: &K) -> bool {
let Some(neighbours) = self.forward.get(start) else {
return false;
};
let mut visited = HashSet::new();
let mut queue = VecDeque::new();
for n in neighbours {
if n == target {
return true;
}
if visited.insert(n) {
queue.push_back(n);
}
}
while let Some(current) = queue.pop_front() {
if let Some(next) = self.forward.get(current) {
for n in next {
if n == target {
return true;
}
if visited.insert(n) {
queue.push_back(n);
}
}
}
}
false
}
}
impl<K: Hash + Eq + Clone> Default for Dag<K> {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn empty_graph() {
let dag = Dag::<u32>::new();
assert!(dag.is_empty());
assert_eq!(dag.node_count(), 0);
assert_eq!(dag.edge_count(), 0);
assert!(!dag.contains_node(&1));
assert!(!dag.contains_edge(&1, &2));
}
#[test]
fn insert_node() {
let mut dag = Dag::new();
assert!(dag.insert_node(1));
assert!(!dag.insert_node(1));
assert_eq!(dag.node_count(), 1);
assert!(dag.contains_node(&1));
assert!(!dag.contains_node(&2));
}
#[test]
fn add_edge() {
let mut dag = Dag::new();
assert_eq!(dag.add_edge(1, 2), Ok(true));
assert_eq!(dag.add_edge(1, 2), Ok(false));
assert_eq!(dag.node_count(), 2);
assert_eq!(dag.edge_count(), 1);
assert!(dag.contains_edge(&1, &2));
assert!(!dag.contains_edge(&2, &1));
}
#[test]
fn add_edge_auto_inserts_nodes() {
let mut dag = Dag::new();
dag.add_edge(10, 20).unwrap();
assert!(dag.contains_node(&10));
assert!(dag.contains_node(&20));
}
#[test]
fn self_loop_rejected() {
let mut dag = Dag::new();
assert_eq!(dag.add_edge(1, 1), Err(CycleError));
assert_eq!(dag.edge_count(), 0);
}
#[test]
fn direct_cycle_rejected() {
let mut dag = Dag::new();
dag.add_edge(1, 2).unwrap();
assert_eq!(dag.add_edge(2, 1), Err(CycleError));
assert_eq!(dag.edge_count(), 1);
}
#[test]
fn transitive_cycle_rejected() {
let mut dag = Dag::new();
dag.add_edge(1, 2).unwrap();
dag.add_edge(2, 3).unwrap();
dag.add_edge(3, 4).unwrap();
assert_eq!(dag.add_edge(4, 1), Err(CycleError));
assert_eq!(dag.add_edge(4, 2), Err(CycleError));
assert_eq!(dag.edge_count(), 3);
}
#[test]
fn diamond_is_valid() {
let mut dag = Dag::new();
dag.add_edge(1, 2).unwrap();
dag.add_edge(1, 3).unwrap();
dag.add_edge(2, 4).unwrap();
dag.add_edge(3, 4).unwrap();
assert_eq!(dag.node_count(), 4);
assert_eq!(dag.edge_count(), 4);
assert_eq!(dag.add_edge(4, 1), Err(CycleError));
}
#[test]
fn remove_node() {
let mut dag = Dag::new();
dag.add_edge(1, 2).unwrap();
dag.add_edge(2, 3).unwrap();
dag.add_edge(1, 3).unwrap();
assert!(dag.remove_node(&2));
assert!(!dag.contains_node(&2));
assert_eq!(dag.node_count(), 2);
assert_eq!(dag.edge_count(), 1);
assert!(dag.contains_edge(&1, &3));
assert!(!dag.contains_edge(&1, &2));
assert!(!dag.contains_edge(&2, &3));
}
#[test]
fn remove_node_nonexistent() {
let mut dag = Dag::<u32>::new();
assert!(!dag.remove_node(&99));
}
#[test]
fn remove_edge() {
let mut dag = Dag::new();
dag.add_edge(1, 2).unwrap();
dag.add_edge(2, 3).unwrap();
assert!(dag.remove_edge(&1, &2));
assert!(!dag.contains_edge(&1, &2));
assert_eq!(dag.edge_count(), 1);
assert!(dag.contains_node(&1));
assert!(dag.contains_node(&2));
}
#[test]
fn remove_edge_nonexistent() {
let mut dag = Dag::<u32>::new();
assert!(!dag.remove_edge(&1, &2));
}
#[test]
fn remove_edge_enables_previously_cyclic_edge() {
let mut dag = Dag::new();
dag.add_edge(1, 2).unwrap();
dag.add_edge(2, 3).unwrap();
assert_eq!(dag.add_edge(3, 1), Err(CycleError));
dag.remove_edge(&1, &2);
assert_eq!(dag.add_edge(3, 1), Ok(true));
}
#[test]
fn direct_dependents() {
let mut dag = Dag::new();
dag.add_edge(1, 2).unwrap();
dag.add_edge(1, 3).unwrap();
let deps = dag.dependents(&1).unwrap();
assert_eq!(deps.len(), 2);
assert!(deps.contains(&2));
assert!(deps.contains(&3));
}
#[test]
fn direct_dependencies() {
let mut dag = Dag::new();
dag.add_edge(1, 3).unwrap();
dag.add_edge(2, 3).unwrap();
let deps = dag.dependencies(&3).unwrap();
assert_eq!(deps.len(), 2);
assert!(deps.contains(&1));
assert!(deps.contains(&2));
}
#[test]
fn dependents_of_unknown_node() {
let dag = Dag::<u32>::new();
assert!(dag.dependents(&1).is_none());
}
#[test]
fn dependencies_of_unknown_node() {
let dag = Dag::<u32>::new();
assert!(dag.dependencies(&1).is_none());
}
#[test]
fn dependents_of_leaf_node() {
let mut dag = Dag::new();
dag.add_edge(1, 2).unwrap();
let deps = dag.dependents(&2).unwrap();
assert!(deps.is_empty());
}
#[test]
fn transitive_dependents_chain() {
let mut dag = Dag::new();
dag.add_edge(1, 2).unwrap();
dag.add_edge(2, 3).unwrap();
dag.add_edge(3, 4).unwrap();
let td = dag.transitive_dependents(&1);
assert_eq!(td, HashSet::from([2, 3, 4]));
}
#[test]
fn transitive_dependencies_chain() {
let mut dag = Dag::new();
dag.add_edge(1, 2).unwrap();
dag.add_edge(2, 3).unwrap();
dag.add_edge(3, 4).unwrap();
let td = dag.transitive_dependencies(&4);
assert_eq!(td, HashSet::from([1, 2, 3]));
}
#[test]
fn transitive_dependents_diamond() {
let mut dag = Dag::new();
dag.add_edge(1, 2).unwrap();
dag.add_edge(1, 3).unwrap();
dag.add_edge(2, 4).unwrap();
dag.add_edge(3, 4).unwrap();
let td = dag.transitive_dependents(&1);
assert_eq!(td, HashSet::from([2, 3, 4]));
let td2 = dag.transitive_dependents(&2);
assert_eq!(td2, HashSet::from([4]));
}
#[test]
fn transitive_dependencies_diamond() {
let mut dag = Dag::new();
dag.add_edge(1, 2).unwrap();
dag.add_edge(1, 3).unwrap();
dag.add_edge(2, 4).unwrap();
dag.add_edge(3, 4).unwrap();
let td = dag.transitive_dependencies(&4);
assert_eq!(td, HashSet::from([1, 2, 3]));
let td2 = dag.transitive_dependencies(&2);
assert_eq!(td2, HashSet::from([1]));
}
#[test]
fn transitive_queries_on_empty_graph() {
let dag = Dag::<u32>::new();
assert!(dag.transitive_dependents(&1).is_empty());
assert!(dag.transitive_dependencies(&1).is_empty());
}
#[test]
fn transitive_queries_on_isolated_node() {
let mut dag = Dag::new();
dag.insert_node(1);
assert!(dag.transitive_dependents(&1).is_empty());
assert!(dag.transitive_dependencies(&1).is_empty());
}
#[test]
fn string_keys() {
let mut dag = Dag::new();
dag.add_edge("a".to_string(), "b".to_string()).unwrap();
dag.add_edge("b".to_string(), "c".to_string()).unwrap();
assert!(dag.contains_edge(&"a".to_string(), &"b".to_string()));
let td = dag.transitive_dependents(&"a".to_string());
assert_eq!(td.len(), 2);
assert!(td.contains("b"));
assert!(td.contains("c"));
}
#[test]
fn wide_fan_out() {
let mut dag = Dag::new();
for i in 1..=100 {
dag.add_edge(0u32, i).unwrap();
}
assert_eq!(dag.node_count(), 101);
assert_eq!(dag.edge_count(), 100);
let deps = dag.dependents(&0).unwrap();
assert_eq!(deps.len(), 100);
let td = dag.transitive_dependents(&0);
assert_eq!(td.len(), 100);
}
#[test]
fn wide_fan_in() {
let mut dag = Dag::new();
for i in 1..=100 {
dag.add_edge(i, 0u32).unwrap();
}
let deps = dag.dependencies(&0).unwrap();
assert_eq!(deps.len(), 100);
let td = dag.transitive_dependencies(&0);
assert_eq!(td.len(), 100);
}
#[test]
fn default_is_empty() {
let dag: Dag<u32> = Dag::default();
assert!(dag.is_empty());
}
#[test]
fn cycle_error_display() {
assert_eq!(CycleError.to_string(), "edge would create a cycle");
}
#[test]
fn nodes_and_edges() {
let mut dag = Dag::new();
dag.add_edge(1, 2).unwrap();
dag.add_edge(1, 3).unwrap();
dag.insert_node(4);
let mut nodes: Vec<u32> = dag.nodes().copied().collect();
nodes.sort();
assert_eq!(nodes, vec![1, 2, 3, 4]);
let mut edges: Vec<(u32, u32)> = dag.edges().map(|(a, b)| (*a, *b)).collect();
edges.sort();
assert_eq!(edges, vec![(1, 2), (1, 3)]);
}
#[test]
fn sources_and_sinks() {
let mut dag = Dag::new();
dag.add_edge(1, 2).unwrap();
dag.add_edge(1, 3).unwrap();
dag.add_edge(2, 4).unwrap();
dag.add_edge(3, 4).unwrap();
dag.insert_node(5);
let mut sources: Vec<u32> = dag.sources().into_iter().copied().collect();
sources.sort();
assert_eq!(sources, vec![1, 5]);
let mut sinks: Vec<u32> = dag.sinks().into_iter().copied().collect();
sinks.sort();
assert_eq!(sinks, vec![4, 5]);
}
#[test]
fn sources_and_sinks_of_empty_graph() {
let dag = Dag::<u32>::new();
assert!(dag.sources().is_empty());
assert!(dag.sinks().is_empty());
}
#[test]
fn reachability() {
let mut dag = Dag::new();
dag.add_edge(1, 2).unwrap();
dag.add_edge(2, 3).unwrap();
assert!(dag.is_reachable(&1, &3));
assert!(dag.is_reachable(&1, &2));
assert!(!dag.is_reachable(&3, &1));
assert!(!dag.is_reachable(&1, &1));
assert!(!dag.is_reachable(&1, &99));
assert!(!dag.is_reachable(&99, &1));
}
#[test]
fn topological_order_chain() {
let mut dag = Dag::new();
dag.add_edge(3, 2).unwrap();
dag.add_edge(2, 1).unwrap();
assert_eq!(dag.topological_order(), vec![3, 2, 1]);
}
#[test]
fn topological_order_diamond_breaks_ties_by_key() {
let mut dag = Dag::new();
dag.add_edge(1, 3).unwrap();
dag.add_edge(1, 2).unwrap();
dag.add_edge(3, 4).unwrap();
dag.add_edge(2, 4).unwrap();
assert_eq!(dag.topological_order(), vec![1, 2, 3, 4]);
}
#[test]
fn topological_order_is_stable_across_calls() {
let mut dag = Dag::new();
for i in (1..=200u32).rev() {
dag.add_edge(0, i).unwrap();
}
let expected: Vec<u32> = (0..=200).collect();
assert_eq!(dag.topological_order(), expected);
assert_eq!(dag.topological_order(), expected);
}
#[test]
fn topological_order_places_isolated_nodes_by_key() {
let mut dag = Dag::new();
dag.add_edge(2, 3).unwrap();
dag.insert_node(1);
assert_eq!(dag.topological_order(), vec![1, 2, 3]);
}
#[test]
fn topological_order_of_empty_graph() {
assert!(Dag::<u32>::new().topological_order().is_empty());
}
#[test]
fn topological_order_from_includes_start_and_dependents_only() {
let mut dag = Dag::new();
dag.add_edge(1, 2).unwrap();
dag.add_edge(1, 3).unwrap();
dag.add_edge(2, 4).unwrap();
dag.add_edge(3, 4).unwrap();
dag.add_edge(4, 5).unwrap();
assert_eq!(dag.topological_order_from(&[2]), vec![2, 4, 5]);
assert_eq!(dag.topological_order_from(&[2, 3]), vec![2, 3, 4, 5]);
assert_eq!(dag.topological_order_from(&[1]), vec![1, 2, 3, 4, 5]);
assert_eq!(dag.topological_order_from(&[5]), vec![5]);
}
#[test]
fn topological_order_from_treats_outside_dependencies_as_satisfied() {
let mut dag = Dag::new();
dag.add_edge(1, 3).unwrap();
dag.add_edge(2, 3).unwrap();
assert_eq!(dag.topological_order_from(&[2]), vec![2, 3]);
}
#[test]
fn topological_order_from_unknown_start() {
let mut dag = Dag::new();
dag.add_edge(1, 2).unwrap();
assert!(dag.topological_order_from(&[99]).is_empty());
assert_eq!(dag.topological_order_from(&[99, 1]), vec![1, 2]);
}
#[test]
fn equality() {
let mut a = Dag::new();
a.add_edge(1, 2).unwrap();
a.add_edge(2, 3).unwrap();
let mut b = Dag::new();
b.add_edge(2, 3).unwrap();
b.add_edge(1, 2).unwrap();
assert_eq!(a, b);
}
}