use crate::ir::ShapeLabelIdx;
use crate::ir::dg::DependencyGraph;
use petgraph::Outgoing;
use petgraph::algo::tarjan_scc;
use petgraph::prelude::EdgeRef;
use std::collections::{HashMap, HashSet};
use std::fmt::{Display, Formatter};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum ShapeRecursionKind {
NonRecursive,
Positive,
Stratified,
NonStratified,
}
impl Display for ShapeRecursionKind {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
let s = match self {
ShapeRecursionKind::NonRecursive => "not recursive",
ShapeRecursionKind::Positive => "positive recursive",
ShapeRecursionKind::Stratified => "stratified recursive",
ShapeRecursionKind::NonStratified => "non-stratified recursive",
};
write!(f, "{s}")
}
}
impl DependencyGraph {
pub fn shape_recursion_kinds(&self) -> HashMap<ShapeLabelIdx, ShapeRecursionKind> {
let sccs = tarjan_scc(&self.graph);
let mut scc_of: HashMap<ShapeLabelIdx, usize> = HashMap::new();
for (i, component) in sccs.iter().enumerate() {
for &node in component {
scc_of.insert(node, i);
}
}
let mut recursive_scc = vec![false; sccs.len()];
for (i, component) in sccs.iter().enumerate() {
recursive_scc[i] =
component.len() > 1 || (component.len() == 1 && self.graph.contains_edge(component[0], component[0]));
}
let mut scc_successors: Vec<HashSet<usize>> = vec![HashSet::new(); sccs.len()];
let mut internal_negative = vec![false; sccs.len()];
let mut has_negative_edge = vec![false; sccs.len()];
let mut external_negative_edges: Vec<(usize, usize)> = Vec::new();
for node in self.graph.nodes() {
let from_scc = scc_of[&node];
for edge in self.graph.edges_directed(node, Outgoing) {
let to_scc = scc_of[&edge.target()];
if to_scc != from_scc {
scc_successors[from_scc].insert(to_scc);
}
if !edge.weight().value() {
has_negative_edge[from_scc] = true;
if to_scc == from_scc {
internal_negative[from_scc] = true;
} else {
external_negative_edges.push((from_scc, to_scc));
}
}
}
}
fn depends_on_recursive(
i: usize,
recursive_scc: &[bool],
successors: &[HashSet<usize>],
memo: &mut [Option<bool>],
) -> bool {
if let Some(v) = memo[i] {
return v;
}
let result = recursive_scc[i]
|| successors[i]
.iter()
.any(|&succ| depends_on_recursive(succ, recursive_scc, successors, memo));
memo[i] = Some(result);
result
}
let mut memo: Vec<Option<bool>> = vec![None; sccs.len()];
let depends_on_recursive: Vec<bool> = (0..sccs.len())
.map(|i| depends_on_recursive(i, &recursive_scc, &scc_successors, &mut memo))
.collect();
let mut has_unsafe_negative = internal_negative.clone();
for (from, to) in external_negative_edges {
if depends_on_recursive[to] {
has_unsafe_negative[from] = true;
}
}
let mut kind_of_scc = Vec::with_capacity(sccs.len());
for i in 0..sccs.len() {
kind_of_scc.push(if !recursive_scc[i] {
ShapeRecursionKind::NonRecursive
} else if has_unsafe_negative[i] {
ShapeRecursionKind::NonStratified
} else if has_negative_edge[i] {
ShapeRecursionKind::Stratified
} else {
ShapeRecursionKind::Positive
});
}
let mut result = HashMap::new();
for (i, component) in sccs.iter().enumerate() {
for &node in component {
result.insert(node, kind_of_scc[i]);
}
}
result
}
pub fn is_stratified(&self) -> bool {
!self
.shape_recursion_kinds()
.values()
.any(|kind| *kind == ShapeRecursionKind::NonStratified)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::ir::dg::PosNeg;
fn idx(n: usize) -> ShapeLabelIdx {
ShapeLabelIdx::new(n)
}
#[test]
fn non_recursive_shape_is_classified_as_such() {
let mut dg = DependencyGraph::new();
dg.add_edge(idx(0), idx(1), PosNeg::Pos);
let kinds = dg.shape_recursion_kinds();
assert_eq!(kinds[&idx(0)], ShapeRecursionKind::NonRecursive);
assert_eq!(kinds[&idx(1)], ShapeRecursionKind::NonRecursive);
assert!(dg.is_stratified());
}
#[test]
fn self_loop_is_recursive() {
let mut dg = DependencyGraph::new();
dg.add_edge(idx(0), idx(0), PosNeg::Pos);
assert_eq!(dg.shape_recursion_kinds()[&idx(0)], ShapeRecursionKind::Positive);
}
#[test]
fn purely_positive_cycle_is_positive() {
let mut dg = DependencyGraph::new();
dg.add_edge(idx(0), idx(1), PosNeg::Pos);
dg.add_edge(idx(1), idx(0), PosNeg::Pos);
let kinds = dg.shape_recursion_kinds();
assert_eq!(kinds[&idx(0)], ShapeRecursionKind::Positive);
assert_eq!(kinds[&idx(1)], ShapeRecursionKind::Positive);
assert!(dg.is_stratified());
}
#[test]
fn negation_of_a_non_recursive_shape_is_stratified() {
let mut dg = DependencyGraph::new();
dg.add_edge(idx(0), idx(1), PosNeg::Pos);
dg.add_edge(idx(1), idx(0), PosNeg::Pos);
dg.add_edge(idx(0), idx(2), PosNeg::Neg);
let kinds = dg.shape_recursion_kinds();
assert_eq!(kinds[&idx(0)], ShapeRecursionKind::Stratified);
assert_eq!(kinds[&idx(1)], ShapeRecursionKind::Stratified);
assert_eq!(kinds[&idx(2)], ShapeRecursionKind::NonRecursive);
assert!(dg.is_stratified());
}
#[test]
fn negation_embedded_in_the_cycle_itself_is_non_stratified() {
let mut dg = DependencyGraph::new();
dg.add_edge(idx(0), idx(1), PosNeg::Pos);
dg.add_edge(idx(1), idx(0), PosNeg::Neg);
let kinds = dg.shape_recursion_kinds();
assert_eq!(kinds[&idx(0)], ShapeRecursionKind::NonStratified);
assert_eq!(kinds[&idx(1)], ShapeRecursionKind::NonStratified);
assert!(!dg.is_stratified());
}
#[test]
fn negation_of_a_different_recursive_shape_is_non_stratified() {
let mut dg = DependencyGraph::new();
dg.add_edge(idx(0), idx(1), PosNeg::Pos);
dg.add_edge(idx(1), idx(0), PosNeg::Pos);
dg.add_edge(idx(2), idx(3), PosNeg::Pos);
dg.add_edge(idx(3), idx(2), PosNeg::Pos);
dg.add_edge(idx(0), idx(2), PosNeg::Neg);
let kinds = dg.shape_recursion_kinds();
assert_eq!(kinds[&idx(0)], ShapeRecursionKind::NonStratified);
assert_eq!(kinds[&idx(1)], ShapeRecursionKind::NonStratified);
assert_eq!(kinds[&idx(2)], ShapeRecursionKind::Positive);
assert_eq!(kinds[&idx(3)], ShapeRecursionKind::Positive);
assert!(!dg.is_stratified());
}
#[test]
fn negation_of_a_shape_that_transitively_depends_on_recursion_is_non_stratified() {
let mut dg = DependencyGraph::new();
dg.add_edge(idx(0), idx(1), PosNeg::Pos);
dg.add_edge(idx(1), idx(0), PosNeg::Pos);
dg.add_edge(idx(0), idx(2), PosNeg::Neg);
dg.add_edge(idx(2), idx(3), PosNeg::Pos);
dg.add_edge(idx(3), idx(4), PosNeg::Pos);
dg.add_edge(idx(4), idx(3), PosNeg::Pos);
let kinds = dg.shape_recursion_kinds();
assert_eq!(kinds[&idx(0)], ShapeRecursionKind::NonStratified);
assert_eq!(kinds[&idx(1)], ShapeRecursionKind::NonStratified);
assert_eq!(kinds[&idx(2)], ShapeRecursionKind::NonRecursive);
assert_eq!(kinds[&idx(3)], ShapeRecursionKind::Positive);
assert!(!dg.is_stratified());
}
}