use crate::error::{GraphError, GraphResult};
use crate::graph::ComputeGraph;
use crate::node::NodeId;
#[derive(Debug, Clone)]
pub struct DomTree {
idom: Vec<Option<NodeId>>,
children: Vec<Vec<NodeId>>,
n_nodes: usize,
}
impl DomTree {
#[must_use]
pub fn idom(&self, node: NodeId) -> Option<NodeId> {
self.idom.get(node.0 as usize).copied().flatten()
}
pub fn children(&self, node: NodeId) -> &[NodeId] {
self.children
.get(node.0 as usize)
.map(|v| v.as_slice())
.unwrap_or(&[])
}
pub fn roots(&self) -> Vec<NodeId> {
(0..self.n_nodes)
.filter(|&i| self.idom[i].is_none())
.map(|i| NodeId(i as u32))
.collect()
}
#[must_use]
pub fn dominates(&self, dominator: NodeId, node: NodeId) -> bool {
if dominator == node {
return true;
}
let mut cur = node;
loop {
match self.idom(cur) {
None => return false,
Some(p) if p == dominator => return true,
Some(p) => cur = p,
}
}
}
pub fn dominated_by(&self, root: NodeId) -> Vec<NodeId> {
let mut result = Vec::new();
let mut stack = vec![root];
while let Some(n) = stack.pop() {
result.push(n);
for &child in self.children(n) {
stack.push(child);
}
}
result
}
#[must_use]
pub fn depth(&self, node: NodeId) -> usize {
let mut d = 0usize;
let mut cur = node;
loop {
match self.idom(cur) {
None => return d,
Some(p) => {
d += 1;
cur = p;
}
}
}
}
#[must_use]
pub fn lca(&self, a: NodeId, b: NodeId) -> NodeId {
let da = self.depth(a);
let db = self.depth(b);
let mut x = a;
let mut y = b;
let (shallow, deep, diff) = if da <= db {
(x, y, db - da)
} else {
(y, x, da - db)
};
let mut deep = deep;
for _ in 0..diff {
deep = self.idom(deep).unwrap_or(deep);
}
x = shallow;
y = deep;
let mut guard = self.n_nodes + 1;
while x != y {
x = self.idom(x).unwrap_or(x);
y = self.idom(y).unwrap_or(y);
if guard == 0 {
break; }
guard -= 1;
}
x
}
}
fn intersect(mut b1: usize, mut b2: usize, idom_raw: &[usize], rpo: &[usize]) -> usize {
while b1 != b2 {
while rpo[b1] > rpo[b2] {
b1 = idom_raw[b1];
}
while rpo[b2] > rpo[b1] {
b2 = idom_raw[b2];
}
}
b1
}
pub fn analyse(graph: &ComputeGraph) -> GraphResult<DomTree> {
if graph.is_empty() {
return Err(GraphError::EmptyGraph);
}
let n_real = graph.node_count();
let vroot = n_real; let total = n_real + 1;
let mut succ: Vec<Vec<usize>> = vec![Vec::new(); total];
let mut pred: Vec<Vec<usize>> = vec![Vec::new(); total];
for (from, to) in graph.edges() {
let f = from.0 as usize;
let t = to.0 as usize;
succ[f].push(t);
pred[t].push(f);
}
for src in graph.sources() {
succ[vroot].push(src.0 as usize);
pred[src.0 as usize].push(vroot);
}
let mut rpo_order: Vec<usize> = Vec::with_capacity(total);
let mut visited = vec![false; total];
let mut dfs_stack: Vec<(usize, usize)> = vec![(vroot, 0)];
visited[vroot] = true;
while let Some((node, idx)) = dfs_stack.last_mut() {
if *idx < succ[*node].len() {
let child = succ[*node][*idx];
*idx += 1;
if !visited[child] {
visited[child] = true;
dfs_stack.push((child, 0));
}
} else {
let n = *node;
dfs_stack.pop();
rpo_order.push(n);
}
}
rpo_order.reverse();
let mut rpo = vec![usize::MAX; total];
for (i, &node) in rpo_order.iter().enumerate() {
rpo[node] = i;
}
const UNDEF: usize = usize::MAX;
let mut idom_raw = vec![UNDEF; total];
idom_raw[vroot] = vroot;
let mut changed = true;
while changed {
changed = false;
for &b in &rpo_order[1..] {
let mut new_idom = UNDEF;
for &p in &pred[b] {
if idom_raw[p] == UNDEF {
continue; }
if new_idom == UNDEF {
new_idom = p;
} else {
new_idom = intersect(p, new_idom, &idom_raw, &rpo);
}
}
if new_idom != UNDEF && idom_raw[b] != new_idom {
idom_raw[b] = new_idom;
changed = true;
}
}
}
let mut idom_out: Vec<Option<NodeId>> = vec![None; n_real];
for i in 0..n_real {
let raw = idom_raw[i];
if raw == UNDEF || raw == vroot {
idom_out[i] = None; } else {
idom_out[i] = Some(NodeId(raw as u32));
}
}
let mut children: Vec<Vec<NodeId>> = vec![Vec::new(); n_real];
for (i, &idom_opt) in idom_out.iter().enumerate() {
if let Some(parent) = idom_opt {
children[parent.0 as usize].push(NodeId(i as u32));
}
}
Ok(DomTree {
idom: idom_out,
children,
n_nodes: n_real,
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::builder::GraphBuilder;
fn make_chain(n: usize) -> (ComputeGraph, Vec<NodeId>) {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let ids: Vec<NodeId> = (0..n).map(|_| b.add_barrier("x")).collect();
for w in ids.windows(2) {
b.dep(w[0], w[1]);
}
let g = b.build().unwrap();
(g, ids)
}
#[test]
fn dominance_empty_graph_error() {
let g = ComputeGraph::new();
assert!(matches!(analyse(&g), Err(GraphError::EmptyGraph)));
}
#[test]
fn dominance_single_node_is_root() {
let (g, ids) = make_chain(1);
let dt = analyse(&g).unwrap();
assert!(dt.idom(ids[0]).is_none());
assert_eq!(dt.roots(), vec![ids[0]]);
}
#[test]
fn dominance_linear_chain() {
let (g, ids) = make_chain(4);
let dt = analyse(&g).unwrap();
assert!(dt.idom(ids[0]).is_none()); assert_eq!(dt.idom(ids[1]), Some(ids[0]));
assert_eq!(dt.idom(ids[2]), Some(ids[1]));
assert_eq!(dt.idom(ids[3]), Some(ids[2]));
}
#[test]
fn dominance_diamond() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let a = b.add_barrier("a");
let bnode = b.add_barrier("b");
let c = b.add_barrier("c");
let d = b.add_barrier("d");
b.dep(a, bnode).dep(a, c).dep(bnode, d).dep(c, d);
let g = b.build().unwrap();
let dt = analyse(&g).unwrap();
assert!(dt.idom(a).is_none());
assert_eq!(dt.idom(bnode), Some(a));
assert_eq!(dt.idom(c), Some(a));
assert_eq!(dt.idom(d), Some(a));
}
#[test]
fn dominance_dominates_reflexive() {
let (g, ids) = make_chain(3);
let dt = analyse(&g).unwrap();
assert!(dt.dominates(ids[0], ids[0]));
assert!(dt.dominates(ids[1], ids[1]));
}
#[test]
fn dominance_dominates_transitive() {
let (g, ids) = make_chain(4);
let dt = analyse(&g).unwrap();
assert!(dt.dominates(ids[0], ids[1]));
assert!(dt.dominates(ids[0], ids[2]));
assert!(dt.dominates(ids[0], ids[3]));
assert!(!dt.dominates(ids[3], ids[0]));
}
#[test]
fn dominance_dominated_by_subtree() {
let (g, ids) = make_chain(4);
let dt = analyse(&g).unwrap();
let sub = dt.dominated_by(ids[1]);
assert!(sub.contains(&ids[1]));
assert!(sub.contains(&ids[2]));
assert!(sub.contains(&ids[3]));
assert!(!sub.contains(&ids[0]));
}
#[test]
fn dominance_depth_linear() {
let (g, ids) = make_chain(4);
let dt = analyse(&g).unwrap();
assert_eq!(dt.depth(ids[0]), 0);
assert_eq!(dt.depth(ids[1]), 1);
assert_eq!(dt.depth(ids[2]), 2);
assert_eq!(dt.depth(ids[3]), 3);
}
#[test]
fn dominance_lca_same_node() {
let (g, ids) = make_chain(3);
let dt = analyse(&g).unwrap();
assert_eq!(dt.lca(ids[1], ids[1]), ids[1]);
}
#[test]
fn dominance_lca_diamond() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let a = b.add_barrier("a");
let bnode = b.add_barrier("b");
let c = b.add_barrier("c");
let d = b.add_barrier("d");
b.dep(a, bnode).dep(a, c).dep(bnode, d).dep(c, d);
let g = b.build().unwrap();
let dt = analyse(&g).unwrap();
let result = dt.lca(bnode, c);
assert_eq!(result, a);
assert_eq!(dt.lca(bnode, d), a); }
#[test]
fn dominance_children() {
let (g, ids) = make_chain(3);
let dt = analyse(&g).unwrap();
assert_eq!(dt.children(ids[0]), &[ids[1]]);
assert_eq!(dt.children(ids[1]), &[ids[2]]);
assert!(dt.children(ids[2]).is_empty());
}
#[test]
fn dominance_fork_join_children() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let a = b.add_barrier("a");
let b1 = b.add_barrier("b");
let c = b.add_barrier("c");
let d = b.add_barrier("d");
let e = b.add_barrier("e");
b.fan_out(a, &[b1, c, d]);
b.fan_in(&[b1, c, d], e);
let g = b.build().unwrap();
let dt = analyse(&g).unwrap();
assert!(dt.dominates(a, e));
assert_eq!(dt.idom(b1), Some(a));
assert_eq!(dt.idom(c), Some(a));
assert_eq!(dt.idom(d), Some(a));
}
#[test]
fn dominance_two_independent_nodes() {
let mut b = GraphBuilder::new().with_auto_infer_edges(false);
let a = b.add_barrier("a");
let bnode = b.add_barrier("b");
let g = b.build().unwrap();
let dt = analyse(&g).unwrap();
assert!(dt.idom(a).is_none());
assert!(dt.idom(bnode).is_none());
let roots = dt.roots();
assert_eq!(roots.len(), 2);
}
#[test]
fn dominance_longer_chain_all_dominated() {
let (g, ids) = make_chain(8);
let dt = analyse(&g).unwrap();
for &id in &ids[1..] {
assert!(dt.dominates(ids[0], id));
}
for &id in &ids[..7] {
assert!(!dt.dominates(ids[7], id));
}
}
#[test]
fn dominance_dominated_by_includes_self() {
let (g, ids) = make_chain(3);
let dt = analyse(&g).unwrap();
let sub = dt.dominated_by(ids[0]);
assert!(sub.contains(&ids[0]));
}
}