use std::collections::{HashMap, VecDeque};
use std::fmt::Debug;
#[derive(Clone)]
pub struct TopoSort<T: Eq + std::hash::Hash + Copy, V> {
dependencies: HashMap<T, HashMap<T, V>>,
}
impl<T: Debug, V: Debug> Debug for TopoSort<T, V>
where
T: Eq + std::hash::Hash + Copy,
{
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let Self { dependencies } = self;
f.debug_struct("TopoSort")
.field("dependencies", dependencies)
.finish()
}
}
impl<T: Eq + std::hash::Hash + Copy, V> Default for TopoSort<T, V> {
fn default() -> Self {
Self {
dependencies: HashMap::new(),
}
}
}
impl<T: Eq + std::hash::Hash + Copy, V> TopoSort<T, V> {
pub fn new() -> Self {
Self::default()
}
pub fn add_dependency(&mut self, element: T, depends_on: T, value: V) {
self.dependencies
.entry(element)
.or_default()
.insert(depends_on, value);
}
pub fn retain(&mut self, filter: impl Fn(&T, &T, &V) -> bool) {
self.dependencies.retain(|element, deps| {
deps.retain(|depends_on, value| filter(element, depends_on, value));
!deps.is_empty()
});
}
pub fn sort_elements(&self, elements: &[T]) -> TopoSortIter<T> {
TopoSortIter::new(&self.dependencies, elements)
}
}
pub struct TopoSortIter<T: Eq + std::hash::Hash + Copy> {
in_degree: HashMap<T, usize>,
dependents: HashMap<T, Vec<T>>,
ready: VecDeque<T>,
}
impl<T: Eq + std::hash::Hash + Copy> TopoSortIter<T> {
pub fn into_unordered_vec(self) -> Vec<T> {
let mut result: Vec<T> = self.ready.into_iter().collect();
for (node, count) in self.in_degree {
if count != 0 {
result.push(node);
}
}
result
}
}
impl<T: Eq + std::hash::Hash + Copy> TopoSortIter<T> {
fn new<V>(dependencies: &HashMap<T, HashMap<T, V>>, all_nodes: &[T]) -> Self {
let mut in_degree: HashMap<T, usize> = HashMap::with_capacity(all_nodes.len());
let mut dependents: HashMap<T, Vec<T>> = HashMap::with_capacity(all_nodes.len());
for node in all_nodes {
in_degree.insert(*node, 0);
}
for (node, deps) in dependencies {
if !in_degree.contains_key(node) {
continue;
}
let mut number_of_deps = 0;
for dep in deps.keys() {
if !in_degree.contains_key(dep) {
continue;
}
dependents.entry(*dep).or_default().push(*node);
number_of_deps += 1;
}
in_degree.insert(*node, number_of_deps);
}
let ready: VecDeque<T> = all_nodes
.iter()
.copied()
.filter(|node| in_degree.get(node).copied().unwrap_or(0) == 0)
.collect();
Self {
in_degree,
dependents,
ready,
}
}
}
impl<T: Eq + std::hash::Hash + Copy> Iterator for TopoSortIter<T> {
type Item = T;
fn next(&mut self) -> Option<Self::Item> {
let node = self.ready.pop_front()?;
if let Some(deps) = self.dependents.remove(&node) {
for dependent in deps {
if let Some(count) = self.in_degree.get_mut(&dependent) {
*count -= 1;
if *count == 0 {
self.ready.push_back(dependent);
}
}
}
}
Some(node)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_empty() {
let topo: TopoSort<i32, ()> = TopoSort::new();
let elements = [1, 2, 3, 4, 5];
let result: Vec<_> = topo.sort_elements(&elements).collect();
assert_eq!(result, elements);
}
#[test]
fn test_single_dependency() {
let mut topo = TopoSort::new();
let elements = [2, 1];
topo.add_dependency(2, 1, ()); let result: Vec<_> = topo.sort_elements(&elements).collect();
assert_eq!(result, vec![1, 2]);
}
#[test]
fn test_chain() {
let mut topo = TopoSort::new();
let elements = [3, 2, 1];
topo.add_dependency(3, 2, ()); topo.add_dependency(2, 1, ()); let result: Vec<_> = topo.sort_elements(&elements).collect();
assert_eq!(result, vec![1, 2, 3]);
}
#[test]
fn test_diamond() {
let mut topo = TopoSort::new();
let elements = [4, 2, 3, 1];
topo.add_dependency(4, 2, ()); topo.add_dependency(4, 3, ()); topo.add_dependency(2, 1, ()); topo.add_dependency(3, 1, ()); let result: Vec<_> = topo.sort_elements(&elements).collect();
assert_eq!(result[0], 1);
assert_eq!(result[3], 4);
assert!(result.contains(&2));
assert!(result.contains(&3));
}
#[test]
fn circular_dependency() {
let elements = ['A', 'B', 'C', 'D'];
let mut topo = TopoSort::new();
topo.add_dependency('B', 'A', ());
topo.add_dependency('C', 'A', ());
topo.add_dependency('D', 'A', ());
topo.add_dependency('C', 'B', ());
topo.add_dependency('D', 'C', ());
topo.add_dependency('B', 'D', ());
let mut iter = topo.sort_elements(&elements);
let first = iter.next().unwrap();
assert_eq!(first, 'A');
let second = iter.next();
assert!(second.is_none());
let unordered = iter.into_unordered_vec();
assert_eq!(unordered.len(), 3);
assert!(unordered.contains(&'B'));
assert!(unordered.contains(&'C'));
assert!(unordered.contains(&'D'));
}
#[test]
fn test_missing_dependencies() {
let mut topo = TopoSort::new();
let elements = [3, 2, 1];
topo.add_dependency(3, 4, ()); topo.add_dependency(2, 1, ());
topo.add_dependency(3, 2, ());
let result: Vec<_> = topo.sort_elements(&elements).collect();
assert_eq!(result, vec![1, 2, 3]);
}
}