use disposition_ir_model::{
edge::EdgeGroups,
entity::{EntityType, EntityTypes},
node::{NodeId, NodeNestingInfos, NodeRank, NodeRanks, NodeRanksNested},
};
use disposition_model_common::{Id, Map};
#[derive(Clone, Copy, Debug)]
pub struct NodeRanksCalculator;
struct SccDag {
adjacency: Vec<Vec<usize>>,
in_degree: Vec<usize>,
}
struct TarjanState {
index_counter: usize,
stack: Vec<usize>,
on_stack: Vec<bool>,
index: Vec<Option<usize>>,
lowlink: Vec<usize>,
scc_ids: Vec<usize>,
scc_counter: usize,
}
impl NodeRanksCalculator {
pub fn calculate<'id>(
edge_groups: &EdgeGroups<'id>,
entity_types: &EntityTypes<'id>,
node_nesting_infos: &NodeNestingInfos<'id>,
) -> NodeRanksNested<'id> {
if node_nesting_infos.is_empty() {
return NodeRanksNested::new();
}
let container_to_children = Self::container_to_children_build(node_nesting_infos);
let dependency_edges = Self::dependency_edges_collect(edge_groups, entity_types);
let lca_level_edges = Self::lca_level_edges_build(&dependency_edges, node_nesting_infos);
let empty_edges: Vec<(NodeId<'id>, NodeId<'id>)> = Vec::new();
let (root, containers) = container_to_children.iter().fold(
(NodeRanks::new(), Map::new()),
|(mut root, mut containers), (container, children)| {
let edges = lca_level_edges.get(container).unwrap_or(&empty_edges);
let ranks = Self::ranks_compute(children, edges);
match container {
None => root = ranks,
Some(container_id) => {
containers.insert(container_id.clone(), ranks);
}
}
(root, containers)
},
);
NodeRanksNested { root, containers }
}
fn container_to_children_build<'id>(
node_nesting_infos: &NodeNestingInfos<'id>,
) -> Map<Option<NodeId<'id>>, Vec<NodeId<'id>>> {
node_nesting_infos.iter().fold(
Map::new(),
|mut container_to_children, (node_id, nesting_info)| {
let chain = &nesting_info.ancestor_chain;
let parent = chain
.len()
.checked_sub(2)
.map(|parent_idx| chain[parent_idx].clone());
container_to_children
.entry(parent)
.or_default()
.push(node_id.clone());
container_to_children
},
)
}
fn lca_level_edges_build<'id>(
dependency_edges: &[(NodeId<'id>, NodeId<'id>)],
node_nesting_infos: &NodeNestingInfos<'id>,
) -> Map<Option<NodeId<'id>>, Vec<(NodeId<'id>, NodeId<'id>)>> {
dependency_edges
.iter()
.filter_map(|(from_id, to_id)| {
Self::lca_level_edge_compute(from_id, to_id, node_nesting_infos)
})
.fold(
Map::new(),
|mut lca_level_edges, (lca_container, divergent_from, divergent_to)| {
lca_level_edges
.entry(lca_container)
.or_default()
.push((divergent_from, divergent_to));
lca_level_edges
},
)
}
fn lca_level_edge_compute<'id>(
from_id: &NodeId<'id>,
to_id: &NodeId<'id>,
node_nesting_infos: &NodeNestingInfos<'id>,
) -> Option<(Option<NodeId<'id>>, NodeId<'id>, NodeId<'id>)> {
let info_from = node_nesting_infos.get(from_id)?;
let info_to = node_nesting_infos.get(to_id)?;
let chain_from = &info_from.ancestor_chain;
let chain_to = &info_to.ancestor_chain;
let lca_depth = chain_from
.iter()
.zip(chain_to.iter())
.take_while(|(a, b)| a == b)
.count();
if lca_depth >= chain_from.len() || lca_depth >= chain_to.len() {
return None;
}
let divergent_from = chain_from[lca_depth].clone();
let divergent_to = chain_to[lca_depth].clone();
if divergent_from == divergent_to {
return None;
}
let lca_container = lca_depth
.checked_sub(1)
.map(|lca_idx| chain_from[lca_idx].clone());
Some((lca_container, divergent_from, divergent_to))
}
fn dependency_edges_collect<'id>(
edge_groups: &EdgeGroups<'id>,
entity_types: &EntityTypes<'id>,
) -> Vec<(NodeId<'id>, NodeId<'id>)> {
edge_groups
.iter()
.filter(|(edge_group_id, _edge_group)| {
Self::edge_group_is_dependency(edge_group_id.as_ref(), entity_types)
})
.flat_map(|(_edge_group_id, edge_group)| edge_group.iter())
.filter(|edge| edge.from != edge.to)
.map(|edge| (edge.from.clone(), edge.to.clone()))
.collect()
}
fn edge_group_is_dependency(edge_group_id: &Id, entity_types: &EntityTypes<'_>) -> bool {
entity_types
.get(edge_group_id)
.map(|types| types.iter().any(Self::entity_type_is_dependency_edge_group))
.unwrap_or(false)
}
fn entity_type_is_dependency_edge_group(entity_type: &EntityType) -> bool {
matches!(
entity_type,
EntityType::DependencyEdgeCyclicDefault
| EntityType::DependencyEdgeSequenceDefault
| EntityType::DependencyEdgeSymmetricDefault
)
}
fn ranks_compute<'id>(
all_node_ids: &[NodeId<'id>],
dependency_edges: &[(NodeId<'id>, NodeId<'id>)],
) -> NodeRanks<'id> {
if all_node_ids.is_empty() {
return NodeRanks::new();
}
if dependency_edges.is_empty() {
return Self::node_ranks_uniform(all_node_ids, 0);
}
let adjacency = Self::ranks_compute_adjacency_build(all_node_ids, dependency_edges);
let node_ranks = Self::ranks_compute_from_adjacency(&adjacency, all_node_ids.len());
all_node_ids
.iter()
.zip(node_ranks)
.map(|(node_id, rank)| (node_id.clone(), NodeRank::new(rank)))
.collect()
}
fn node_ranks_uniform<'id>(all_node_ids: &[NodeId<'id>], rank: u32) -> NodeRanks<'id> {
all_node_ids
.iter()
.map(|node_id| (node_id.clone(), NodeRank::new(rank)))
.collect()
}
fn ranks_compute_adjacency_build<'id>(
all_node_ids: &[NodeId<'id>],
dependency_edges: &[(NodeId<'id>, NodeId<'id>)],
) -> Vec<Vec<usize>> {
let node_to_index: Map<NodeId<'id>, usize> = all_node_ids
.iter()
.enumerate()
.map(|(node_idx, node_id)| (node_id.clone(), node_idx))
.collect();
dependency_edges.iter().fold(
vec![Vec::new(); all_node_ids.len()],
|mut adjacency, (from_id, to_id)| {
if let (Some(&from_idx), Some(&to_idx)) =
(node_to_index.get(from_id), node_to_index.get(to_id))
{
adjacency[from_idx].push(to_idx);
}
adjacency
},
)
}
fn ranks_compute_from_adjacency(adjacency: &[Vec<usize>], node_count: usize) -> Vec<u32> {
let scc_ids = Self::tarjan_scc(adjacency, node_count);
let scc_count = scc_ids.iter().copied().max().map(|m| m + 1).unwrap_or(0);
if scc_count == 0 {
return vec![0; node_count];
}
let scc_dag = Self::scc_dag_build(adjacency, &scc_ids, scc_count);
let scc_ranks =
Self::scc_dag_ranks_compute(&scc_dag.adjacency, &scc_dag.in_degree, scc_count);
scc_ids.iter().map(|&scc_id| scc_ranks[scc_id]).collect()
}
fn scc_dag_build(adjacency: &[Vec<usize>], scc_ids: &[usize], scc_count: usize) -> SccDag {
let mut scc_adjacency: Vec<Vec<usize>> = vec![Vec::new(); scc_count];
for (from_idx, to_indices) in adjacency.iter().enumerate() {
let from_scc = scc_ids[from_idx];
for &to_idx in to_indices {
let to_scc = scc_ids[to_idx];
if from_scc != to_scc {
scc_adjacency[from_scc].push(to_scc);
}
}
}
scc_adjacency.iter_mut().for_each(|neighbours| {
neighbours.sort_unstable();
neighbours.dedup();
});
let mut in_degree: Vec<usize> = vec![0; scc_count];
scc_adjacency.iter().for_each(|neighbours| {
neighbours.iter().for_each(|&to_scc| in_degree[to_scc] += 1);
});
SccDag {
adjacency: scc_adjacency,
in_degree,
}
}
fn tarjan_scc(adjacency: &[Vec<usize>], node_count: usize) -> Vec<usize> {
let mut state = TarjanState {
index_counter: 0,
stack: Vec::new(),
on_stack: vec![false; node_count],
index: vec![None; node_count],
lowlink: vec![0; node_count],
scc_ids: vec![0; node_count],
scc_counter: 0,
};
for node in 0..node_count {
if state.index[node].is_none() {
Self::tarjan_strongconnect_iterative(adjacency, node, &mut state);
}
}
state.scc_ids
}
fn tarjan_strongconnect_iterative(
adjacency: &[Vec<usize>],
start: usize,
state: &mut TarjanState,
) {
let mut call_stack: Vec<(usize, usize)> = Vec::new();
state.index[start] = Some(state.index_counter);
state.lowlink[start] = state.index_counter;
state.index_counter += 1;
state.stack.push(start);
state.on_stack[start] = true;
call_stack.push((start, 0));
while let Some(&mut (v, ref mut ni)) = call_stack.last_mut() {
if *ni < adjacency[v].len() {
let w = adjacency[v][*ni];
*ni += 1;
if state.index[w].is_none() {
state.index[w] = Some(state.index_counter);
state.lowlink[w] = state.index_counter;
state.index_counter += 1;
state.stack.push(w);
state.on_stack[w] = true;
call_stack.push((w, 0));
} else if state.on_stack[w] {
let w_index = state.index[w].unwrap();
if w_index < state.lowlink[v] {
state.lowlink[v] = w_index;
}
}
} else {
if state.lowlink[v] == state.index[v].unwrap() {
let scc_id = state.scc_counter;
state.scc_counter += 1;
while let Some(w) = state.stack.pop() {
state.on_stack[w] = false;
state.scc_ids[w] = scc_id;
if w == v {
break;
}
}
}
call_stack.pop();
if let Some(&mut (caller, _)) = call_stack.last_mut()
&& state.lowlink[v] < state.lowlink[caller]
{
state.lowlink[caller] = state.lowlink[v];
}
}
}
}
fn scc_dag_ranks_compute(
scc_adjacency: &[Vec<usize>],
scc_in_degree: &[usize],
scc_count: usize,
) -> Vec<u32> {
let mut ranks: Vec<u32> = vec![0; scc_count];
let mut in_degree = scc_in_degree.to_vec();
let mut queue: std::collections::VecDeque<usize> = std::collections::VecDeque::new();
in_degree
.iter()
.copied()
.enumerate()
.take(scc_count)
.filter(|(_scc_idx, in_degree_item)| *in_degree_item == 0)
.for_each(|(scc_idx, _in_degree_item)| queue.push_back(scc_idx));
while let Some(scc_idx) = queue.pop_front() {
for &to_scc in &scc_adjacency[scc_idx] {
let candidate_rank = ranks[scc_idx] + 1;
if candidate_rank > ranks[to_scc] {
ranks[to_scc] = candidate_rank;
}
in_degree[to_scc] -= 1;
if in_degree[to_scc] == 0 {
queue.push_back(to_scc);
}
}
}
ranks
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn ranks_from_adjacency_linear_chain_increments() {
let adjacency = vec![vec![1], vec![2], vec![]];
let node_ranks = NodeRanksCalculator::ranks_compute_from_adjacency(&adjacency, 3);
assert_eq!(vec![0, 1, 2], node_ranks);
}
#[test]
fn ranks_from_adjacency_cycle_shares_rank() {
let adjacency = vec![vec![1], vec![2], vec![0]];
let node_ranks = NodeRanksCalculator::ranks_compute_from_adjacency(&adjacency, 3);
assert_eq!(vec![0, 0, 0], node_ranks);
}
#[test]
fn ranks_from_adjacency_diamond_uses_longest_path() {
let adjacency = vec![vec![1, 2], vec![3], vec![3], vec![]];
let node_ranks = NodeRanksCalculator::ranks_compute_from_adjacency(&adjacency, 4);
assert_eq!(vec![0, 1, 1, 2], node_ranks);
}
#[test]
fn ranks_from_adjacency_contracts_cycle_then_continues() {
let adjacency = vec![vec![1], vec![0, 2], vec![]];
let node_ranks = NodeRanksCalculator::ranks_compute_from_adjacency(&adjacency, 3);
assert_eq!(vec![0, 0, 1], node_ranks);
}
#[test]
fn ranks_from_adjacency_no_edges_all_zero() {
let adjacency = vec![vec![], vec![], vec![]];
let node_ranks = NodeRanksCalculator::ranks_compute_from_adjacency(&adjacency, 3);
assert_eq!(vec![0, 0, 0], node_ranks);
}
}