use fnv::{FnvHashMap, FnvHashSet};
use generational_arena::{Arena, Index};
use std::borrow::Borrow;
use std::collections::{hash_map, VecDeque};
use std::hash::Hash;
use std::ops::RangeBounds;
#[cfg(any(test, feature = "test-utils"))]
pub mod naive;
pub struct Dag<N> {
entries: Arena<Entry>,
nodes: FnvHashMap<N, Index>,
cycles: FnvHashSet<Cycle>,
next_order: u64,
}
#[derive(Default)]
struct Entry {
forward: FnvHashSet<Index>,
backward: FnvHashSet<Index>,
order: u64,
}
#[derive(Debug, Eq, PartialEq)]
pub enum Error {
NotFound,
SelfLoop,
}
struct DagTraverser {
stack: VecDeque<Index>,
visited: FnvHashSet<Index>,
direction: Direction,
}
enum Direction {
Forward,
Backward,
}
#[derive(Copy, Clone, Eq, PartialEq)]
enum ControlFlow {
Stop,
Continue,
}
struct SearchVisitor {
target: Index,
found: bool,
}
#[derive(Default)]
struct CollectVisitor {
collected: Vec<(Index, u64)>,
}
trait Visitor {
fn visit(&mut self, index: &Index, order: u64) -> ControlFlow;
}
#[derive(Hash, Eq, PartialEq, Debug)]
struct Cycle(Index, Index);
enum GraphChange {
DeleteNode(Index),
DeleteEdge(Index, Index),
}
impl<N> Dag<N>
where
N: Hash + Eq,
{
pub fn new() -> Self {
Self {
entries: Arena::new(),
nodes: FnvHashMap::default(),
cycles: FnvHashSet::default(),
next_order: 0,
}
}
pub fn insert(&mut self, node: N) -> bool {
if let hash_map::Entry::Vacant(e) = self.nodes.entry(node) {
let entry = Entry {
order: self.next_order,
..Entry::default()
};
let index = self.entries.insert(entry);
e.insert(index);
self.next_order += 1;
true
} else {
false
}
}
pub fn remove<Q: ?Sized>(&mut self, node: &Q) -> Option<N>
where
N: Borrow<Q>,
Q: Hash + Eq,
{
let (node, index) = if let Some((n, index)) = self.nodes.remove_entry(node) {
(n, index)
} else {
return None;
};
let entry = self.entries.remove(index).unwrap();
for i in entry.forward {
let entry = self.entries.get_mut(i).unwrap();
entry.backward.remove(&index);
}
for i in entry.backward {
let entry = self.entries.get_mut(i).unwrap();
entry.forward.remove(&index);
}
if !self.cycles.is_empty() {
self.update_cycles(GraphChange::DeleteNode(index));
}
Some(node)
}
pub fn connect<Q: ?Sized>(&mut self, v: &Q, u: &Q) -> Result<bool, Error>
where
N: Borrow<Q>,
Q: Hash + Eq,
{
let v_index = *self.nodes.get(v).ok_or(Error::NotFound)?;
let u_index = *self.nodes.get(u).ok_or(Error::NotFound)?;
if v_index == u_index {
return Err(Error::SelfLoop);
}
if self
.entries
.get(v_index)
.unwrap()
.forward
.contains(&u_index)
{
return Ok(false);
}
self.add_edge_helper(v_index, u_index, false);
let tmp = self.entries.get2_mut(v_index, u_index);
let v_entry = tmp.0.unwrap();
let u_entry = tmp.1.unwrap();
v_entry.forward.insert(u_index);
u_entry.backward.insert(v_index);
Ok(true)
}
pub fn disconnect<Q: ?Sized>(&mut self, v: &Q, u: &Q) -> Result<bool, Error>
where
N: Borrow<Q>,
Q: Hash + Eq,
{
let v_index = self.nodes.get(v).ok_or(Error::NotFound)?;
let u_index = self.nodes.get(u).ok_or(Error::NotFound)?;
if v_index == u_index {
return Err(Error::SelfLoop);
}
{
let tmp = self.entries.get2_mut(*v_index, *u_index);
let v_entry = tmp.0.unwrap();
let u_entry = tmp.1.unwrap();
if !v_entry.forward.remove(u_index) {
return Ok(false);
}
u_entry.backward.remove(v_index);
}
if !self.cycles.is_empty() {
self.update_cycles(GraphChange::DeleteEdge(*v_index, *u_index));
}
Ok(true)
}
pub fn contains<Q: ?Sized>(&self, v: &Q) -> bool
where
N: Borrow<Q>,
Q: Hash + Eq,
{
self.nodes.contains_key(v)
}
pub fn is_connected<Q: ?Sized>(&self, v: &Q, u: &Q) -> bool
where
N: Borrow<Q>,
Q: Hash + Eq,
{
let (v_index, u_index) = match (self.nodes.get(v), self.nodes.get(u)) {
(Some(v_index), Some(u_index)) => (v_index, u_index),
_ => return false,
};
self.entries
.get(*u_index)
.unwrap()
.backward
.contains(v_index)
}
pub fn is_reachable<Q: ?Sized>(&self, v: &Q, u: &Q) -> bool
where
N: Borrow<Q>,
Q: Hash + Eq,
{
let (v_index, u_index) = match (self.nodes.get(v), self.nodes.get(u)) {
(Some(v_index), Some(u_index)) => (v_index, u_index),
_ => return false,
};
if v_index == u_index {
return false;
}
let v_entry = self.entries.get(*v_index).unwrap();
let u_entry = self.entries.get(*u_index).unwrap();
let mut visitor = SearchVisitor::new(*u_index);
let mut traverser = DagTraverser::new(Direction::Forward);
traverser.push_index(*v_index);
if v_entry.order < u_entry.order {
traverser.traverse(self, 0..=u_entry.order, &mut visitor);
visitor.found
} else if self.cycles.is_empty() {
false
} else {
traverser.traverse(self, 0..=u64::MAX, &mut visitor);
visitor.found
}
}
#[inline(always)]
fn update_cycles(&mut self, change: GraphChange) {
let cycles = std::mem::take(&mut self.cycles);
for cycle in cycles {
if change.should_remove(&cycle) {
continue;
}
let v = cycle.0;
let u = cycle.1;
let mut visitor = SearchVisitor::new(u);
let mut traverser = DagTraverser::new(Direction::Forward);
traverser.push_index(v);
traverser.traverse(self, 0..=u64::MAX, &mut visitor);
if !visitor.found {
continue;
}
self.add_edge_helper(v, u, true);
}
}
fn add_edge_helper(&mut self, v_index: Index, u_index: Index, visit_all: bool) {
let (v_order, u_order) = {
let v_entry = self.entries.get(v_index).unwrap();
let u_entry = self.entries.get(u_index).unwrap();
(v_entry.order, u_entry.order)
};
let mut traverser = DagTraverser::new(Direction::Forward);
let mut visited_forward = CollectVisitor::default();
let mut visited_backward = CollectVisitor::default();
let range = if self.cycles.is_empty() && !visit_all {
0..=v_order
} else {
0..=u64::MAX
};
traverser.push_index(u_index);
traverser.traverse(self, range, &mut visited_forward);
if traverser.has_visited(&v_index) {
self.cycles.insert(Cycle(v_index, u_index));
} else {
traverser.direction = Direction::Backward;
traverser.push_index(v_index);
traverser.traverse(self, (u_order + 1).., &mut visited_backward);
let visited_forward = visited_forward.collected;
let visited_backward = visited_backward.collected;
self.reorder(visited_forward, visited_backward);
}
}
fn reorder(
&mut self,
mut visited_forward: Vec<(Index, u64)>,
mut visited_backward: Vec<(Index, u64)>,
) {
visited_forward.sort_by_key(|(_, order)| *order);
visited_backward.sort_by_key(|(_, order)| *order);
let len1 = visited_forward.len();
let len2 = visited_backward.len();
let mut i1 = 0usize;
let mut i2 = 0usize;
let mut index_iter = visited_backward.iter().chain(visited_forward.iter());
while i1 < len1 && i2 < len2 {
let (_, o1) = visited_forward[i1];
let (_, o2) = visited_backward[i2];
let index = index_iter.next().unwrap().0;
self.entries.get_mut(index).unwrap().order = if o1 < o2 {
i1 += 1;
o1
} else {
i2 += 1;
o2
};
}
while i1 < len1 {
let index = index_iter.next().unwrap().0;
self.entries.get_mut(index).unwrap().order = visited_forward[i1].1;
i1 += 1;
}
while i2 < len2 {
let index = index_iter.next().unwrap().0;
self.entries.get_mut(index).unwrap().order = visited_backward[i2].1;
i2 += 1;
}
}
}
impl DagTraverser {
pub fn new(direction: Direction) -> Self {
Self {
direction,
stack: VecDeque::new(),
visited: FnvHashSet::default(),
}
}
#[inline(always)]
pub fn has_visited(&self, node: &Index) -> bool {
self.visited.contains(node)
}
#[inline(always)]
pub fn push_index(&mut self, index: Index) {
self.stack.push_front(index);
}
pub fn traverse<N, R: RangeBounds<u64>, V: Visitor>(
&mut self,
dag: &Dag<N>,
_range: R,
visitor: &mut V,
) {
while let Some(index) = self.stack.pop_back() {
let entry = dag.entries.get(index).unwrap();
if !self.visited.insert(index) {
continue;
}
match visitor.visit(&index, entry.order) {
ControlFlow::Continue => {}
ControlFlow::Stop => break,
}
let to_visit = match self.direction {
Direction::Forward => &entry.forward,
Direction::Backward => &entry.backward,
};
let mut new_items = 0;
for v_index in to_visit {
if self.visited.contains(v_index) {
continue;
}
new_items += 1;
}
self.stack.reserve(new_items);
for v_index in to_visit {
if self.visited.contains(v_index) {
continue;
}
self.stack.push_front(*v_index);
}
}
}
}
impl SearchVisitor {
pub fn new(target: Index) -> Self {
SearchVisitor {
target,
found: false,
}
}
}
impl Visitor for SearchVisitor {
#[inline(always)]
fn visit(&mut self, index: &Index, _order: u64) -> ControlFlow {
if self.found || &self.target == index {
self.found = true;
ControlFlow::Stop
} else {
ControlFlow::Continue
}
}
}
impl Visitor for CollectVisitor {
#[inline(always)]
fn visit(&mut self, index: &Index, order: u64) -> ControlFlow {
self.collected.push((*index, order));
ControlFlow::Continue
}
}
impl Visitor for () {
#[inline(always)]
fn visit(&mut self, _index: &Index, _order: u64) -> ControlFlow {
ControlFlow::Continue
}
}
impl<N> Default for Dag<N>
where
N: Hash + Eq,
{
fn default() -> Self {
Self::new()
}
}
impl GraphChange {
#[inline(always)]
pub fn should_remove(&self, cycle: &Cycle) -> bool {
match self {
Self::DeleteEdge(v, u) => &cycle.0 == v && &cycle.1 == u,
Self::DeleteNode(v) => &cycle.0 == v || &cycle.1 == v,
}
}
}