#![forbid(unsafe_code, missing_docs)]
use std::{
borrow::Borrow,
cmp::{Ordering, Reverse},
collections::BinaryHeap,
fmt,
iter::Iterator,
};
use std::cell::Cell;
use std::collections::{HashMap, HashSet};
use std::collections::hash_map::RandomState;
use std::hash::BuildHasher;
use hashlink::LinkedHashSet;
use slotmap::{DefaultKey, SlotMap};
type TopoOrder = u32;
#[cfg_attr(feature = "serde", derive(serde::Deserialize, serde::Serialize))]
#[cfg_attr(feature = "serde", serde(bound(serialize = "N: serde::Serialize, E: serde::Serialize, H: BuildHasher + Default")))]
#[cfg_attr(feature = "serde", serde(bound(deserialize = "N: serde::Deserialize<'de>, E: serde::Deserialize<'de>, H: BuildHasher + Default")))]
pub struct DAG<N, E, H = RandomState> {
node_info: SlotMap<DefaultKey, NodeInfo<N, H>>,
edge_data: HashMap<(Node, Node), E, H>,
last_topo_order: TopoOrder,
#[cfg_attr(feature = "serde", serde(skip))]
stack_visited_scratch_space: Cell<StackVisitedScratchSpace<Node, H>>,
}
#[repr(transparent)]
#[derive(Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash, Debug)]
#[cfg_attr(feature = "serde", derive(serde::Deserialize, serde::Serialize))]
pub struct Node(DefaultKey);
#[derive(Debug)]
#[cfg_attr(feature = "serde", derive(serde::Deserialize, serde::Serialize))]
#[cfg_attr(feature = "serde", serde(bound(serialize = "N: serde::Serialize, H: BuildHasher + Default")))]
#[cfg_attr(feature = "serde", serde(bound(deserialize = "N: serde::Deserialize<'de>, H: BuildHasher + Default")))]
struct NodeInfo<N, H> {
topo_order: TopoOrder,
data: N,
parents: LinkedHashSet<Node, H>,
children: LinkedHashSet<Node, H>,
}
impl<N, H: BuildHasher + Default> NodeInfo<N, H> {
fn new(topo_order: TopoOrder, data: N) -> Self {
NodeInfo {
topo_order,
data,
parents: Default::default(),
children: Default::default(),
}
}
}
#[derive(Debug, PartialEq, Eq)]
pub enum Error {
NodeMissing,
CycleDetected,
}
impl fmt::Display for Error {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Error::NodeMissing => {
write!(f, "The given node was not found in the topological order")
}
Error::CycleDetected => write!(f, "Cycles of nodes may not be formed in the graph"),
}
}
}
impl std::error::Error for Error {}
impl<N, E, H: BuildHasher + Default> Default for DAG<N, E, H> {
#[inline]
fn default() -> Self {
Self {
last_topo_order: 0,
node_info: SlotMap::default(),
edge_data: Default::default(),
stack_visited_scratch_space: Cell::default(),
}
}
}
impl<N, E> DAG<N, E> {
#[inline]
pub fn new() -> Self { Self::default() }
}
impl<N, E, H: BuildHasher + Default> DAG<N, E, H> {
#[inline]
pub fn add_node(&mut self, data: N) -> Node {
let next_topo_order = self.last_topo_order + 1;
self.last_topo_order = next_topo_order;
let node_info = NodeInfo::new(next_topo_order, data);
Node(self.node_info.insert(node_info))
}
#[inline]
pub fn contains_node(&self, node: impl Borrow<Node>) -> bool {
let node = node.borrow();
self.node_info.contains_key(node.0)
}
#[inline]
pub fn get_node_data(&self, node: impl Borrow<Node>) -> Option<&N> {
let node = node.borrow();
self.node_info.get(node.0).map(|d| &d.data)
}
#[inline]
pub fn get_node_data_mut(&mut self, node: impl Borrow<Node>) -> Option<&mut N> {
let node = node.borrow();
self.node_info.get_mut(node.0).map(|d| &mut d.data)
}
pub fn remove_node(&mut self, node: Node) -> bool {
if !self.node_info.contains_key(node.0) {
return false;
}
let node_info = self.node_info.remove(node.0).unwrap();
for child in &node_info.children {
if let Some(child_node) = self.node_info.get_mut(child.0) {
child_node.parents.remove(&node.into());
}
self.edge_data.remove(&(node, *child));
}
for parent in &node_info.parents {
if let Some(parent_node) = self.node_info.get_mut(parent.0) {
parent_node.children.remove(&node.into());
}
self.edge_data.remove(&(*parent, node));
}
for other_node in self.node_info.values_mut() {
if other_node.topo_order > node_info.topo_order {
other_node.topo_order -= 1;
}
}
self.last_topo_order -= 1;
true
}
pub fn add_edge(
&mut self,
src: impl Borrow<Node>,
dst: impl Borrow<Node>,
data: E,
) -> Result<bool, Error> {
let src = src.borrow();
let dst = dst.borrow();
if !self.node_info.contains_key(src.0) || !self.node_info.contains_key(dst.0) {
return Err(Error::NodeMissing);
}
if src == dst { return Err(Error::CycleDetected);
}
let mut no_prev_edge = self.node_info[src.0].children.insert(*dst);
let upper_bound = self.node_info[src.0].topo_order;
no_prev_edge = no_prev_edge && self.node_info[dst.0].parents.insert(*src);
let lower_bound = self.node_info[dst.0].topo_order;
if !no_prev_edge { return Ok(false);
}
self.edge_data.insert((*src, *dst), data);
if lower_bound < upper_bound {
let mut visited = HashSet::<_, H>::default(); let change_forward = match self.dfs_forward(*dst, &mut visited, upper_bound) {
Ok(change_set) => change_set,
Err(err) => { self.node_info[src.0].children.remove(dst);
self.node_info[dst.0].parents.remove(src);
self.edge_data.remove(&(*src, *dst));
return Err(err);
}
};
let change_backward = self.dfs_backward(*src, &mut visited, lower_bound);
self.reorder_nodes(change_forward, change_backward);
}
Ok(true)
}
#[inline]
pub fn contains_edge(&self, src: impl Borrow<Node>, dst: impl Borrow<Node>) -> bool {
let src = src.borrow();
let dst = dst.borrow();
if !self.node_info.contains_key(src.0) || !self.node_info.contains_key(dst.0) {
return false;
}
self.edge_data.contains_key(&(*src, *dst))
}
pub fn contains_transitive_edge(
&self,
src: impl Borrow<Node>,
dst: impl Borrow<Node>,
) -> bool {
let src = src.borrow();
let dst = dst.borrow();
if !self.node_info.contains_key(src.0) || !self.node_info.contains_key(dst.0) {
return false;
}
if src.0 == dst.0 {
return false;
}
let mut scratch = self.stack_visited_scratch_space.take();
scratch.clear();
scratch.stack.push(*src);
while let Some(key) = scratch.stack.pop() {
if scratch.visited.contains(&key) {
continue;
} else {
scratch.visited.insert(key);
}
let children = &self.node_info.get(key.0).unwrap().children;
if children.contains(dst) {
self.stack_visited_scratch_space.set(scratch);
return true;
} else {
scratch.stack.extend(children.iter());
continue;
}
}
self.stack_visited_scratch_space.set(scratch);
false
}
#[inline]
pub fn get_edge_data(&self, src: impl Borrow<Node>, dst: impl Borrow<Node>) -> Option<&E> {
let src = src.borrow();
let dst = dst.borrow();
self.edge_data.get(&(*src, *dst))
}
#[inline]
pub fn get_edge_data_mut(&mut self, src: impl Borrow<Node>, dst: impl Borrow<Node>) -> Option<&mut E> {
let src = src.borrow();
let dst = dst.borrow();
self.edge_data.get_mut(&(*src, *dst))
}
#[inline]
pub fn get_outgoing_edges(&self, src: impl Borrow<Node>) -> impl Iterator<Item=(&Node, &E)> + '_ {
let src = *src.borrow();
self.node_info.get(src.0)
.into_iter()
.flat_map(|node_info| node_info.children.iter())
.map(move |child_node| (child_node, self.edge_data.get(&(src, *child_node)).unwrap()))
}
#[inline]
pub fn get_outgoing_edge_nodes(&self, src: impl Borrow<Node>) -> impl Iterator<Item=&Node> {
let src = src.borrow();
self.node_info.get(src.0)
.into_iter()
.flat_map(|node_info| node_info.children.iter())
}
#[inline]
pub fn get_outgoing_edge_data(&self, src: impl Borrow<Node>) -> impl Iterator<Item=&E> + '_ {
let src = *src.borrow();
self.node_info.get(src.0)
.into_iter()
.flat_map(|node_info| node_info.children.iter())
.map(move |child_node| self.edge_data.get(&(src, *child_node)).unwrap())
}
#[inline]
pub fn get_outgoing_edge_node_data(&self, src: impl Borrow<Node>) -> impl Iterator<Item=&N> + '_ {
let src = src.borrow();
self.node_info.get(src.0)
.into_iter()
.flat_map(|node_info| node_info.children.iter())
.flat_map(|child_node| self.node_info.get(child_node.0).into_iter())
.map(|node_info| &node_info.data)
}
#[inline]
pub fn get_incoming_edges(&self, dst: impl Borrow<Node>) -> impl Iterator<Item=(&Node, &E)> + '_ {
let dst = *dst.borrow();
self.node_info.get(dst.0)
.into_iter()
.flat_map(|node_info| node_info.parents.iter())
.map(move |parent_node| (parent_node, self.edge_data.get(&(*parent_node, dst)).unwrap()))
}
#[inline]
pub fn get_incoming_edge_nodes(&self, dst: impl Borrow<Node>) -> impl Iterator<Item=&Node> + '_ {
let dst = *dst.borrow();
self.node_info.get(dst.0)
.into_iter()
.flat_map(|node_info| node_info.parents.iter())
}
#[inline]
pub fn get_incoming_edge_data(&self, dst: impl Borrow<Node>) -> impl Iterator<Item=&E> + '_ {
let dst = *dst.borrow();
self.node_info.get(dst.0)
.into_iter()
.flat_map(|node_info| node_info.parents.iter())
.map(move |parent_node| self.edge_data.get(&(*parent_node, dst)).unwrap())
}
#[inline]
pub fn get_incoming_edge_node_data(&self, dst: impl Borrow<Node>) -> impl Iterator<Item=&N> + '_ {
let dst = dst.borrow();
self.node_info.get(dst.0)
.into_iter()
.flat_map(|node_info| node_info.parents.iter())
.flat_map(|parent_node| self.node_info.get(parent_node.0).into_iter())
.map(|node_info| &node_info.data)
}
pub fn remove_edge(&mut self, src: impl Borrow<Node>, dst: impl Borrow<Node>) -> Option<E> {
let src = src.borrow();
let dst = dst.borrow();
if !self.node_info.contains_key(src.0) || !self.node_info.contains_key(dst.0) {
return None;
}
let src_children = &mut self.node_info[src.0].children;
if !src_children.contains(&dst) {
return None;
}
src_children.remove(&dst);
self.node_info[dst.0].parents.remove(&src);
self.edge_data.remove(&(*src, *dst))
}
pub fn remove_outgoing_edges_of_node(&mut self, src: impl Borrow<Node>) -> Option<Vec<(Node, E)>> { let pred_id = src.borrow();
if !self.node_info.contains_key(pred_id.0) {
return None;
}
let children: Vec<_> = self.node_info[pred_id.0].children.drain().collect(); if children.is_empty() {
return None;
}
let mut edge_data = Vec::new();
for succ_id in children {
if let Some(succ) = self.node_info.get_mut(succ_id.0) {
succ.parents.remove(&pred_id);
}
if let Some(data) = self.edge_data.remove(&(*pred_id, succ_id)) {
edge_data.push((succ_id, data));
}
}
Some(edge_data)
}
#[inline]
pub fn len(&self) -> usize {
self.node_info.len()
}
#[inline]
pub fn is_empty(&self) -> bool {
self.len() == 0
}
#[inline]
pub fn iter_unsorted(&self) -> impl Iterator<Item=(TopoOrder, Node)> + '_ {
self.node_info
.iter()
.map(|(key, node)| (node.topo_order, Node(key)))
}
pub fn descendants_unsorted(
&self,
node: impl Borrow<Node>,
) -> Result<DescendantsUnsorted<N, E, H>, Error> {
let node = node.borrow();
if !self.node_info.contains_key(node.0) {
return Err(Error::NodeMissing);
}
let mut stack = Vec::new(); stack.extend(self.node_info[node.0].children.iter());
let visited = HashSet::<_, H>::default();
Ok(DescendantsUnsorted {
dag: self,
stack,
visited,
})
}
pub fn descendants(&self, node: impl Borrow<Node>) -> Result<Descendants<N, E, H>, Error> {
let node = node.borrow();
if !self.node_info.contains_key(node.0) {
return Err(Error::NodeMissing);
}
let mut queue = BinaryHeap::new(); queue.extend(
self.node_info[node.0]
.children
.iter()
.cloned()
.map(|child_node| {
let child_order = self.get_node(child_node).topo_order;
(Reverse(child_order), child_node)
}),
);
let visited = HashSet::<_, H>::default();
Ok(Descendants {
dag: self,
queue,
visited,
})
}
#[inline]
pub fn topo_cmp(&self, node_a: impl Borrow<Node>, node_b: impl Borrow<Node>) -> Ordering {
let node_a = node_a.borrow();
let node_b = node_b.borrow();
self.node_info[node_a.0]
.topo_order
.cmp(&self.node_info[node_b.0].topo_order)
}
fn dfs_forward(
&self,
start_key: Node,
visited: &mut HashSet<Node, H>,
upper_bound: TopoOrder,
) -> Result<HashSet<Node, H>, Error> {
let mut stack = Vec::new(); let mut result = HashSet::<_, H>::default();
stack.push(start_key);
while let Some(next_key) = stack.pop() {
visited.insert(next_key);
result.insert(next_key);
for child_key in self.get_node(next_key).children.iter() {
let child_topo_order = self.get_node(*child_key).topo_order;
if child_topo_order == upper_bound {
return Err(Error::CycleDetected);
}
if !visited.contains(child_key) && child_topo_order < upper_bound {
stack.push(*child_key);
}
}
}
Ok(result)
}
fn dfs_backward(
&self,
start_key: Node,
visited: &mut HashSet<Node, H>,
lower_bound: TopoOrder,
) -> HashSet<Node, H> {
let mut stack = Vec::new(); let mut result = HashSet::<_, H>::default();
stack.push(start_key);
while let Some(next_key) = stack.pop() {
visited.insert(next_key);
result.insert(next_key);
for parent_key in self.get_node(next_key).parents.iter() {
let parent_topo_order = self.get_node(*parent_key).topo_order;
if !visited.contains(parent_key) && lower_bound < parent_topo_order {
stack.push(*parent_key);
}
}
}
result
}
fn reorder_nodes(
&mut self,
change_forward: HashSet<Node, H>,
change_backward: HashSet<Node, H>,
) {
let mut change_forward: Vec<_> = change_forward
.into_iter()
.map(|key| (key, self.get_node(key).topo_order))
.collect(); change_forward.sort_unstable_by_key(|pair| pair.1);
let mut change_backward: Vec<_> = change_backward
.into_iter()
.map(|key| (key, self.get_node(key).topo_order))
.collect(); change_backward.sort_unstable_by_key(|pair| pair.1);
let mut all_keys = Vec::new(); let mut all_topo_orders = Vec::new();
for (key, topo_order) in change_backward {
all_keys.push(key);
all_topo_orders.push(topo_order);
}
for (key, topo_order) in change_forward {
all_keys.push(key);
all_topo_orders.push(topo_order);
}
all_topo_orders.sort_unstable();
for (key, topo_order) in all_keys.into_iter().zip(all_topo_orders.into_iter()) {
self.node_info
.get_mut(key.0)
.unwrap()
.topo_order = topo_order;
}
}
fn get_node(&self, idx: Node) -> &NodeInfo<N, H> {
self.node_info.get(idx.0).unwrap()
}
}
pub struct DescendantsUnsorted<'a, N, E, H> {
dag: &'a DAG<N, E, H>,
stack: Vec<Node>,
visited: HashSet<Node, H>,
}
impl<'a, N, E, H: BuildHasher> Iterator for DescendantsUnsorted<'a, N, E, H> {
type Item = (TopoOrder, Node);
#[inline]
fn next(&mut self) -> Option<Self::Item> {
while let Some(node) = self.stack.pop() {
if self.visited.contains(&node) {
continue;
} else {
self.visited.insert(node);
}
let node_repr = self.dag.node_info.get(node.0).unwrap();
let order = node_repr.topo_order;
self.stack.extend(node_repr.children.iter());
return Some((order, node));
}
None
}
}
pub struct Descendants<'a, N, E, H> {
dag: &'a DAG<N, E, H>,
queue: BinaryHeap<(Reverse<TopoOrder>, Node)>,
visited: HashSet<Node, H>,
}
impl<'a, N, E, H: BuildHasher + Default> Iterator for Descendants<'a, N, E, H> {
type Item = Node;
fn next(&mut self) -> Option<Self::Item> {
loop {
return if let Some((_, node)) = self.queue.pop() {
if self.visited.contains(&node) {
continue;
} else {
self.visited.insert(node);
}
let node_repr = self.dag.node_info.get(node.0).unwrap();
for child in node_repr.children.iter() {
let order = self.dag.get_node(*child).topo_order;
self.queue.push((Reverse(order), *child))
}
Some(node)
} else {
None
};
}
}
}
#[derive(Debug)]
struct StackVisitedScratchSpace<T, H> {
stack: Vec<T>,
visited: HashSet<T, H>,
}
impl<T, H: Default> Default for StackVisitedScratchSpace<T, H> {
fn default() -> Self {
Self { stack: Vec::new(), visited: HashSet::<_, H>::default() }
}
}
impl<T, H> StackVisitedScratchSpace<T, H> {
fn clear(&mut self) {
self.stack.clear();
self.visited.clear();
}
}
#[cfg(test)]
mod tests {
extern crate pretty_env_logger;
use super::*;
fn get_basic_dag() -> Result<([Node; 7], DAG<(), ()>), Error> {
let mut dag = DAG::new();
let dog = dag.add_node(());
let cat = dag.add_node(());
let mouse = dag.add_node(());
let lion = dag.add_node(());
let human = dag.add_node(());
let gazelle = dag.add_node(());
let grass = dag.add_node(());
assert_eq!(dag.len(), 7);
dag.add_edge(lion, human, ())?;
dag.add_edge(lion, gazelle, ())?;
dag.add_edge(human, dog, ())?;
dag.add_edge(human, cat, ())?;
dag.add_edge(dog, cat, ())?;
dag.add_edge(cat, mouse, ())?;
dag.add_edge(gazelle, grass, ())?;
dag.add_edge(mouse, grass, ())?;
Ok(([dog, cat, mouse, lion, human, gazelle, grass], dag))
}
#[test]
fn add_nodes_basic() {
let mut dag = DAG::<_, ()>::new();
let dog = dag.add_node(());
let cat = dag.add_node(());
let mouse = dag.add_node(());
let lion = dag.add_node(());
let human = dag.add_node(());
assert_eq!(dag.len(), 5);
assert!(dag.contains_node(&dog));
assert!(dag.contains_node(&cat));
assert!(dag.contains_node(&mouse));
assert!(dag.contains_node(&lion));
assert!(dag.contains_node(&human));
}
#[test]
fn delete_nodes() {
let mut dag = DAG::<_, ()>::new();
let dog = dag.add_node(());
let cat = dag.add_node(());
let human = dag.add_node(());
assert_eq!(dag.len(), 3);
assert!(dag.contains_node(&dog));
assert!(dag.contains_node(&cat));
assert!(dag.contains_node(&human));
assert!(dag.remove_node(human));
assert_eq!(dag.len(), 2);
assert!(!dag.contains_node(&human));
}
#[test]
fn reject_cycle() {
let mut dag = DAG::new();
let n1 = dag.add_node(());
let n2 = dag.add_node(());
let n3 = dag.add_node(());
assert_eq!(dag.len(), 3);
assert!(dag.add_edge(&n1, &n2, ()).is_ok());
assert!(dag.add_edge(&n2, &n3, ()).is_ok());
assert!(dag.add_edge(&n3, &n1, ()).is_err());
assert!(dag.add_edge(&n1, &n1, ()).is_err());
}
#[test]
fn get_children_unordered() {
let ([dog, cat, mouse, _, human, _, grass], dag) = get_basic_dag().unwrap();
let children: HashSet<_> = dag
.descendants_unsorted(&human)
.unwrap()
.map(|(_, v)| v)
.collect();
let mut expected_children = HashSet::default();
expected_children.extend(vec![dog, cat, mouse, grass]);
assert_eq!(children, expected_children);
let ordered_children: Vec<_> = dag.descendants(human).unwrap().collect();
assert_eq!(ordered_children, vec![dog, cat, mouse, grass])
}
#[test]
fn topo_order_values_no_gaps() {
let ([.., lion, _, _, _], dag) = get_basic_dag().unwrap();
let topo_orders: HashSet<_> = dag
.descendants_unsorted(lion)
.unwrap()
.map(|p| p.0)
.collect();
assert_eq!(topo_orders, (2..=7).collect::<HashSet<_>>())
}
#[test]
fn readme_example() {
let mut dag = DAG::new();
let cat = dag.add_node(());
let dog = dag.add_node(());
let human = dag.add_node(());
assert_eq!(dag.len(), 3);
dag.add_edge(&human, &dog, ()).unwrap();
dag.add_edge(&human, &cat, ()).unwrap();
dag.add_edge(&dog, &cat, ()).unwrap();
let animal_order: Vec<_> = dag.descendants(&human).unwrap().collect();
assert_eq!(animal_order, vec![dog, cat]);
}
#[test]
fn unordered_iter() {
let mut dag = DAG::new();
let cat = dag.add_node(());
let mouse = dag.add_node(());
let dog = dag.add_node(());
let human = dag.add_node(());
assert!(dag.add_edge(&human, &cat, ()).unwrap());
assert!(dag.add_edge(&human, &dog, ()).unwrap());
assert!(dag.add_edge(&dog, &cat, ()).unwrap());
assert!(dag.add_edge(&cat, &mouse, ()).unwrap());
let pairs = dag
.descendants_unsorted(&human)
.unwrap()
.collect::<HashSet<_>>();
let mut expected_pairs = HashSet::default();
expected_pairs.extend(vec![(2, dog), (3, cat), (4, mouse)]);
assert_eq!(pairs, expected_pairs);
}
#[test]
fn topo_cmp() {
use std::cmp::Ordering::*;
let mut dag = DAG::new();
let cat = dag.add_node(());
let mouse = dag.add_node(());
let dog = dag.add_node(());
let human = dag.add_node(());
let horse = dag.add_node(());
assert!(dag.add_edge(&human, &cat, ()).unwrap());
assert!(dag.add_edge(&human, &dog, ()).unwrap());
assert!(dag.add_edge(&dog, &cat, ()).unwrap());
assert!(dag.add_edge(&cat, &mouse, ()).unwrap());
assert_eq!(dag.topo_cmp(&human, &mouse), Less);
assert_eq!(dag.topo_cmp(&cat, &dog), Greater);
assert_eq!(dag.topo_cmp(&cat, &horse), Less);
}
}