// Topological sort that doesn't give up on cycles, a port of `sort.ml`.
//
// A --> B
// C --> D gives: [A] [B C] [D]
// B --> C
// C --> B
//
// Sets of nodes are represented as arrays of node indices sorted by node
// identifier, reproducing the iteration order of OCaml's `Set`.
///|
priv struct SortGraph {
/// node index -> neighbors (indices)
forward : Array[Array[Int]]
backward : Array[Array[Int]]
/// node index -> rank of its id in the sorted order of ids
rank : Array[Int]
visited : Array[Bool]
}
///|
/// A set of node indices, kept sorted by rank.
priv struct NodeSet {
elts : Array[Int]
}
///|
fn SortGraph::empty(_self : SortGraph) -> NodeSet {
{ elts: [], }
}
///|
fn SortGraph::singleton(_self : SortGraph, v : Int) -> NodeSet {
{ elts: [v], }
}
///|
fn NodeSet::mem(self : NodeSet, v : Int) -> Bool {
self.elts.contains(v)
}
///|
fn SortGraph::add(self : SortGraph, s : NodeSet, v : Int) -> NodeSet {
if s.mem(v) {
return s
}
let elts = s.elts.copy()
let r = self.rank[v]
let mut i = 0
while i < elts.length() && self.rank[elts[i]] < r {
i += 1
}
elts.insert(i, v)
{ elts, }
}
///|
fn NodeSet::remove(self : NodeSet, v : Int) -> NodeSet {
{ elts: self.elts.filter(x => x != v), }
}
///|
fn NodeSet::diff(self : NodeSet, other : NodeSet) -> NodeSet {
{ elts: self.elts.filter(x => !other.mem(x)), }
}
///|
fn NodeSet::inter(self : NodeSet, other : NodeSet) -> NodeSet {
{ elts: self.elts.filter(x => other.mem(x)), }
}
///|
fn NodeSet::is_empty(self : NodeSet) -> Bool {
self.elts.is_empty()
}
///|
fn NodeSet::pick_one(self : NodeSet) -> (Int, NodeSet)? {
match self.elts {
[] => None
[v, .. rest] => Some((v, { elts: rest.to_owned(), }))
}
}
///|
fn filtered_neighbors(
v : Int,
edges : Array[Array[Int]],
graph_nodes : NodeSet,
) -> Array[Int] {
edges[v].filter(n => graph_nodes.mem(n))
}
///|
fn SortGraph::eliminate_roots_recursively(
self : SortGraph,
edges : Array[Array[Int]],
back_edges : Array[Array[Int]],
nodes : NodeSet,
) -> (Array[(Bool, NodeSet)], NodeSet) {
let sorted = []
let mut graph_nodes = nodes
let mut input_nodes = nodes
while input_nodes.pick_one() is Some((v, rest)) {
input_nodes = rest
if filtered_neighbors(v, back_edges, graph_nodes).is_empty() {
sorted.push(v)
let children = filtered_neighbors(v, edges, graph_nodes)
graph_nodes = graph_nodes.remove(v)
for c in children {
input_nodes = self.add(input_nodes, c)
}
}
}
(sorted.map(v => (false, self.singleton(v))), graph_nodes)
}
///|
fn SortGraph::eliminate_roots(
self : SortGraph,
nodes : NodeSet,
) -> (Array[(Bool, NodeSet)], NodeSet) {
self.eliminate_roots_recursively(self.forward, self.backward, nodes)
}
///|
fn SortGraph::eliminate_leaves(
self : SortGraph,
nodes : NodeSet,
) -> (NodeSet, Array[(Bool, NodeSet)]) {
let (sorted_leaves, remaining) = self.eliminate_roots_recursively(
self.backward,
self.forward,
nodes,
)
(remaining, sorted_leaves.rev())
}
///|
/// Collect all nodes reachable from the start node, excluding the start
/// node unless it can be reached by some cycle.
fn SortGraph::visit(
self : SortGraph,
edges : Array[Array[Int]],
start_node : Int,
nodes : NodeSet,
) -> NodeSet {
let visited = []
let mut acc = self.empty()
fn color(v : Int) -> Unit {
if !self.visited[v] {
self.visited[v] = true
visited.push(v)
for neighbor in edges[v] {
if nodes.mem(neighbor) {
acc = self.add(acc, neighbor)
color(neighbor)
}
}
}
}
color(start_node)
for v in visited {
self.visited[v] = false
}
acc
}
///|
fn SortGraph::sort_subgraph(
self : SortGraph,
nodes : NodeSet,
) -> Array[(Bool, NodeSet)] {
let (sorted_left, nodes) = self.eliminate_roots(nodes)
let (nodes, sorted_right) = self.eliminate_leaves(nodes)
let sorted_middle = match nodes.pick_one() {
None => []
Some((pivot, _)) => self.partition(pivot, nodes)
}
let res = sorted_left
res.append(sorted_middle)
res.append(sorted_right)
res
}
///|
fn SortGraph::partition(
self : SortGraph,
pivot : Int,
nodes : NodeSet,
) -> Array[(Bool, NodeSet)] {
let ancestors = self.visit(self.backward, pivot, nodes)
let descendants = self.visit(self.forward, pivot, nodes)
let strict_ancestors = ancestors.diff(descendants)
let strict_descendants = descendants.diff(ancestors)
let cycle = descendants.inter(ancestors)
let (is_cyclic, pivot_group) = if cycle.is_empty() {
(false, self.singleton(pivot))
} else {
(true, cycle)
}
let other = nodes
.diff(pivot_group)
.diff(strict_ancestors)
.diff(strict_descendants)
let res = self.sort_subgraph(strict_ancestors)
res.push((is_cyclic, pivot_group))
res.append(self.sort_subgraph(strict_descendants))
res.append(self.sort_subgraph(other))
res
}
///|
/// Sort the nodes topologically. Each input element comes with the list of
/// identifiers of the nodes it points to. The result is a list of groups;
/// a group flagged `true` is a cycle. Groups are sorted such that edges only
/// go from a group to itself or to later groups.
pub fn[T, Id : Compare + Hash + Eq] topological_sort(
l : ArrayView[(T, Array[Id])],
id : (T) -> Id,
) -> Array[(Bool, Array[T])] {
let node_tbl : Map[Id, Int] = Map([])
let values : Array[T] = []
let ids : Array[Id] = []
for x in l {
let i = id(x.0)
if !node_tbl.contains(i) {
node_tbl[i] = values.length()
values.push(x.0)
ids.push(i)
}
}
let n = values.length()
let order = Array::makei(n, i => i)
order.sort_by((a, b) => Compare::compare(ids[a], ids[b]))
let rank = Array::make(n, 0)
for r, v in order {
rank[v] = r
}
let forward = Array::makei(n, _ => [])
let backward = Array::makei(n, _ => [])
for x in l {
let v1 = node_tbl[id(x.0)]
for id2 in x.1 {
match node_tbl.get(id2) {
None => ()
Some(v2) => {
forward[v1].push(v2)
backward[v2].push(v1)
}
}
}
}
let graph : SortGraph = {
forward,
backward,
rank,
visited: Array::make(n, false),
}
let groups = graph.sort_subgraph({ elts: order, })
groups.map(g => (g.0, g.1.elts.map(v => values[v])))
}