#![allow(clippy::many_single_char_names)]
use std::sync::Arc;
use antecedent_core::{Lag, TemporalNodeKey, VariableId};
use crate::dag::Dag;
use crate::error::GraphError;
use crate::marked_storage::{self, AdjEntry};
use crate::temporal::TemporalDag;
use crate::types::{DenseNodeId, MarkedEdge, NodeRef};
use crate::workspace::GraphWorkspace;
#[derive(Clone, Debug)]
pub struct Cpdag {
nodes: Vec<NodeRef>,
adj: Vec<Vec<AdjEntry>>,
}
impl Cpdag {
#[must_use]
pub fn empty() -> Self {
Self { nodes: Vec::new(), adj: Vec::new() }
}
#[must_use]
pub fn with_variables(n: u32) -> Self {
let mut g = Self::empty();
for i in 0..n {
let _ = g.add_node(NodeRef::Static(VariableId::from_raw(i)));
}
g
}
pub fn from_named_edges(
schema: &antecedent_core::CausalSchema,
edges: &[(&str, &str)],
) -> Result<Self, GraphError> {
let n = crate::named::schema_node_count(schema)?;
let mut g = Self::with_variables(n);
for &(from_name, to_name) in edges {
let (from, to) = crate::named::resolve_named_edge(schema, from_name, to_name)?;
g.insert_directed(from, to)?;
}
Ok(g)
}
#[must_use]
pub fn node_count(&self) -> usize {
self.nodes.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.nodes.is_empty()
}
#[must_use]
pub fn nodes(&self) -> &[NodeRef] {
&self.nodes
}
pub fn add_node(&mut self, node: NodeRef) -> Result<DenseNodeId, GraphError> {
if !matches!(node, NodeRef::Static(_)) {
return Err(GraphError::InvalidEndpoints {
message: "Cpdag accepts only Static nodes",
});
}
let id = u32::try_from(self.nodes.len()).map_err(|_| GraphError::TooManyNodes)?;
self.nodes.push(node);
self.adj.push(Vec::new());
Ok(DenseNodeId::from_raw(id))
}
pub fn insert_marked(&mut self, edge: MarkedEdge) -> Result<(), GraphError> {
if !edge.is_cpdag_legal() {
return Err(GraphError::InvalidEndpoints {
message: "CPDAG accepts only Tail–Arrow, Tail–Tail, or Conflict–Conflict marks",
});
}
self.validate_node(edge.a)?;
self.validate_node(edge.b)?;
if edge.a == edge.b {
return Err(GraphError::Cycle { from: edge.a.raw(), to: edge.b.raw() });
}
marked_storage::insert_marked_finish(&mut self.adj, edge)
}
pub fn insert_directed(
&mut self,
from: DenseNodeId,
to: DenseNodeId,
) -> Result<(), GraphError> {
self.insert_marked(MarkedEdge::directed(from, to))
}
pub fn insert_undirected(&mut self, a: DenseNodeId, b: DenseNodeId) -> Result<(), GraphError> {
self.insert_marked(MarkedEdge::undirected(a, b))
}
pub fn remove_edge(&mut self, a: DenseNodeId, b: DenseNodeId) -> Result<(), GraphError> {
self.validate_node(a)?;
self.validate_node(b)?;
marked_storage::remove_edge(&mut self.adj, a, b);
Ok(())
}
pub fn orient_undirected(
&mut self,
from: DenseNodeId,
to: DenseNodeId,
) -> Result<(), GraphError> {
self.validate_node(from)?;
self.validate_node(to)?;
marked_storage::orient_undirected_finish(&mut self.adj, from, to)
}
pub fn mark_conflict(&mut self, a: DenseNodeId, b: DenseNodeId) -> Result<(), GraphError> {
self.validate_node(a)?;
self.validate_node(b)?;
marked_storage::mark_conflict_finish(&mut self.adj, a, b)
}
#[must_use]
pub fn has_edge(&self, a: DenseNodeId, b: DenseNodeId) -> bool {
self.edge_between(a, b).is_some()
}
#[must_use]
pub fn edge_between(&self, a: DenseNodeId, b: DenseNodeId) -> Option<MarkedEdge> {
marked_storage::edge_between(&self.adj, a, b)
}
#[must_use]
pub fn edges(&self) -> Vec<MarkedEdge> {
marked_storage::all_marked_edges(&self.adj)
}
#[must_use]
pub fn children(&self, id: DenseNodeId) -> Vec<DenseNodeId> {
marked_storage::directed_children(&self.adj, id).collect()
}
#[must_use]
pub fn parents(&self, id: DenseNodeId) -> Vec<DenseNodeId> {
marked_storage::directed_parents(&self.adj, id).collect()
}
#[must_use]
pub fn undirected_neighbors(&self, id: DenseNodeId) -> Vec<DenseNodeId> {
marked_storage::undirected_neighbors(&self.adj, id).collect()
}
#[must_use]
pub fn adjacent(&self, id: DenseNodeId) -> Vec<DenseNodeId> {
if id.as_usize() >= self.node_count() {
return Vec::new();
}
self.adj[id.as_usize()].iter().map(|e| e.neighbor).collect()
}
pub fn children_iter(&self, id: DenseNodeId) -> impl Iterator<Item = DenseNodeId> + '_ {
marked_storage::directed_children(&self.adj, id)
}
#[must_use]
pub fn from_dag(dag: &Dag) -> Self {
let mut g = Self::empty();
for node in dag.nodes() {
let _ = g.add_node(*node);
}
for e in dag.edges() {
if let Some((from, to)) = e.parent_child() {
let _ = g.insert_directed(from, to);
}
}
g
}
pub fn to_directed_skeleton(&self) -> Result<Dag, GraphError> {
let mut dag = Dag::empty();
for node in &self.nodes {
dag.add_node(*node)?;
}
for e in self.edges() {
if let Some((from, to)) = e.parent_child() {
dag.insert_directed(from, to)?;
}
}
Ok(dag)
}
pub fn try_into_dag(&self) -> Result<Dag, GraphError> {
for e in self.edges() {
if e.is_undirected() {
return Err(GraphError::InvalidEndpoints {
message: "cannot complete Cpdag to Dag while undirected edges remain",
});
}
if e.is_conflict() {
return Err(GraphError::InvalidEndpoints {
message: "cannot complete Cpdag to Dag while conflict edges remain",
});
}
}
self.to_directed_skeleton()
}
#[must_use]
pub fn conflict_edge_count(&self) -> usize {
self.edges().iter().filter(|e| e.is_conflict()).count()
}
#[must_use]
pub fn undirected_edge_count(&self) -> usize {
self.edges().iter().filter(|e| e.is_undirected()).count()
}
#[must_use]
pub fn directed_edge_count(&self) -> usize {
self.edges().iter().filter(|e| e.parent_child().is_some()).count()
}
#[must_use]
pub fn variable_id(&self, id: DenseNodeId) -> Option<VariableId> {
match self.nodes.get(id.as_usize())? {
NodeRef::Static(v) => Some(*v),
_ => None,
}
}
#[must_use]
pub fn reaches_directed_with(
&self,
ws: &mut GraphWorkspace,
from: DenseNodeId,
to: DenseNodeId,
) -> bool {
marked_storage::reaches_directed(&self.adj, ws, from, to)
}
fn validate_node(&self, id: DenseNodeId) -> Result<(), GraphError> {
if id.as_usize() >= self.node_count() {
Err(GraphError::UnknownNode { id: id.raw() })
} else {
Ok(())
}
}
}
#[derive(Clone, Debug)]
pub struct CpdagReview {
pub graph: Cpdag,
pub pending_edges: Arc<[(VariableId, VariableId)]>,
pub pending_undirected: Arc<[(VariableId, VariableId)]>,
pub algorithm: Arc<str>,
}
impl CpdagReview {
#[must_use]
pub fn from_cpdag(graph: Cpdag, algorithm: impl Into<Arc<str>>) -> Self {
let mut pending = Vec::new();
let mut undirected = Vec::new();
for e in graph.edges() {
if let Some((from, to)) = e.parent_child() {
if let (Some(fv), Some(tv)) = (graph.variable_id(from), graph.variable_id(to)) {
pending.push((fv, tv));
}
} else if e.is_undirected() {
if let (Some(av), Some(bv)) = (graph.variable_id(e.a), graph.variable_id(e.b)) {
if av.raw() <= bv.raw() {
undirected.push((av, bv));
} else {
undirected.push((bv, av));
}
}
}
}
Self {
graph,
pending_edges: Arc::from(pending),
pending_undirected: Arc::from(undirected),
algorithm: algorithm.into(),
}
}
#[must_use]
pub fn accept_edge(mut self, from: VariableId, to: VariableId) -> Self {
let pending: Vec<_> =
self.pending_edges.iter().copied().filter(|e| *e != (from, to)).collect();
self.pending_edges = Arc::from(pending);
self
}
pub fn orient_edge(mut self, from: VariableId, to: VariableId) -> Result<Self, GraphError> {
let from_id = self.resolve_var(from)?;
let to_id = self.resolve_var(to)?;
self.graph.orient_undirected(from_id, to_id)?;
let undirected: Vec<_> = self
.pending_undirected
.iter()
.copied()
.filter(|&(a, b)| (a, b) != (from, to) && (a, b) != (to, from))
.collect();
self.pending_undirected = Arc::from(undirected);
if !self.pending_edges.iter().any(|e| *e == (from, to)) {
let mut pending = self.pending_edges.to_vec();
pending.push((from, to));
self.pending_edges = Arc::from(pending);
}
Ok(self)
}
#[must_use]
pub fn is_complete(&self) -> bool {
self.pending_edges.is_empty() && self.pending_undirected.is_empty()
}
pub fn try_into_dag(self) -> Result<Dag, GraphError> {
if !self.is_complete() {
return Err(GraphError::InvalidEndpoints {
message: "CpdagReview is incomplete; accept directed and orient undirected edges first",
});
}
self.graph.try_into_dag()
}
fn resolve_var(&self, var: VariableId) -> Result<DenseNodeId, GraphError> {
for i in 0..self.graph.node_count() {
let id = DenseNodeId::from_raw(u32::try_from(i).expect("fit"));
if self.graph.variable_id(id) == Some(var) {
return Ok(id);
}
}
Err(GraphError::UnknownNode { id: var.raw() })
}
}
#[derive(Clone, Debug)]
pub struct TemporalCpdag {
nodes: Vec<NodeRef>,
adj: Vec<Vec<AdjEntry>>,
}
impl TemporalCpdag {
#[must_use]
pub fn empty() -> Self {
Self { nodes: Vec::new(), adj: Vec::new() }
}
#[must_use]
pub fn node_count(&self) -> usize {
self.nodes.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.nodes.is_empty()
}
#[must_use]
pub fn nodes(&self) -> &[NodeRef] {
&self.nodes
}
pub fn add_node(&mut self, node: NodeRef) -> Result<DenseNodeId, GraphError> {
match node {
NodeRef::Lagged { .. } | NodeRef::Context { .. } => {}
NodeRef::Static(_) => {
return Err(GraphError::InvalidEndpoints {
message: "TemporalCpdag accepts Lagged or Context nodes (not Static)",
});
}
}
let id = u32::try_from(self.nodes.len()).map_err(|_| GraphError::TooManyNodes)?;
self.nodes.push(node);
self.adj.push(Vec::new());
Ok(DenseNodeId::from_raw(id))
}
pub fn add_lagged(
&mut self,
variable: VariableId,
lag: Lag,
) -> Result<DenseNodeId, GraphError> {
self.add_node(NodeRef::Lagged { variable, lag })
}
pub fn add_context(
&mut self,
variable: VariableId,
environment: Option<antecedent_core::EnvironmentId>,
) -> Result<DenseNodeId, GraphError> {
self.add_node(NodeRef::Context { variable, environment })
}
pub fn insert_marked(&mut self, edge: MarkedEdge) -> Result<(), GraphError> {
if !edge.is_cpdag_legal() {
return Err(GraphError::InvalidEndpoints {
message: "CPDAG accepts only Tail–Arrow, Tail–Tail, or Conflict–Conflict marks",
});
}
self.validate_node(edge.a)?;
self.validate_node(edge.b)?;
if let (
NodeRef::Lagged { variable: v1, lag: l1 },
NodeRef::Lagged { variable: v2, lag: l2 },
) = (self.nodes[edge.a.as_usize()], self.nodes[edge.b.as_usize()])
{
if v1 == v2 && l1 == l2 && l1.is_contemporaneous() {
return Err(GraphError::ContemporaneousSelfEdge { variable: v1 });
}
}
if edge.a == edge.b {
return Err(GraphError::Cycle { from: edge.a.raw(), to: edge.b.raw() });
}
marked_storage::insert_marked_finish(&mut self.adj, edge)
}
pub fn insert_directed(
&mut self,
from: DenseNodeId,
to: DenseNodeId,
) -> Result<(), GraphError> {
self.insert_marked(MarkedEdge::directed(from, to))
}
pub fn insert_undirected(&mut self, a: DenseNodeId, b: DenseNodeId) -> Result<(), GraphError> {
self.insert_marked(MarkedEdge::undirected(a, b))
}
pub fn orient_undirected(
&mut self,
from: DenseNodeId,
to: DenseNodeId,
) -> Result<(), GraphError> {
self.validate_node(from)?;
self.validate_node(to)?;
marked_storage::orient_undirected_finish(&mut self.adj, from, to)
}
pub fn mark_conflict(&mut self, a: DenseNodeId, b: DenseNodeId) -> Result<(), GraphError> {
self.validate_node(a)?;
self.validate_node(b)?;
marked_storage::mark_conflict_finish(&mut self.adj, a, b)
}
#[must_use]
pub fn has_edge(&self, a: DenseNodeId, b: DenseNodeId) -> bool {
self.edge_between(a, b).is_some()
}
#[must_use]
pub fn edge_between(&self, a: DenseNodeId, b: DenseNodeId) -> Option<MarkedEdge> {
marked_storage::edge_between(&self.adj, a, b)
}
#[must_use]
pub fn edges(&self) -> Vec<MarkedEdge> {
marked_storage::all_marked_edges(&self.adj)
}
#[must_use]
pub fn children(&self, id: DenseNodeId) -> Vec<DenseNodeId> {
marked_storage::directed_children(&self.adj, id).collect()
}
#[must_use]
pub fn parents(&self, id: DenseNodeId) -> Vec<DenseNodeId> {
marked_storage::directed_parents(&self.adj, id).collect()
}
#[must_use]
pub fn undirected_neighbors(&self, id: DenseNodeId) -> Vec<DenseNodeId> {
marked_storage::undirected_neighbors(&self.adj, id).collect()
}
pub fn children_iter(&self, id: DenseNodeId) -> impl Iterator<Item = DenseNodeId> + '_ {
marked_storage::directed_children(&self.adj, id)
}
#[must_use]
pub fn from_temporal_dag(dag: &TemporalDag) -> Self {
let mut g = Self::empty();
for node in dag.nodes() {
let _ = g.add_node(*node);
}
for e in dag.edges() {
if let Some((from, to)) = e.parent_child() {
let _ = g.insert_directed(from, to);
}
}
g
}
pub fn to_directed_skeleton(&self) -> Result<TemporalDag, GraphError> {
let mut dag = TemporalDag::empty();
for node in &self.nodes {
dag.add_node(*node)?;
}
for e in self.edges() {
if let Some((from, to)) = e.parent_child() {
dag.insert_directed(from, to)?;
}
}
Ok(dag)
}
pub fn try_into_temporal_dag(&self) -> Result<TemporalDag, GraphError> {
for e in self.edges() {
if e.is_undirected() {
return Err(GraphError::InvalidEndpoints {
message: "cannot complete TemporalCpdag to TemporalDag while undirected edges remain",
});
}
if e.is_conflict() {
return Err(GraphError::InvalidEndpoints {
message: "cannot complete TemporalCpdag to TemporalDag while conflict edges remain",
});
}
}
self.to_directed_skeleton()
}
#[must_use]
pub fn conflict_edge_count(&self) -> usize {
self.edges().iter().filter(|e| e.is_conflict()).count()
}
#[must_use]
pub fn temporal_key(&self, id: DenseNodeId) -> Option<TemporalNodeKey> {
match self.nodes.get(id.as_usize())? {
NodeRef::Lagged { variable, lag } => {
let offset = -i32::try_from(lag.raw()).ok()?;
Some(TemporalNodeKey { variable: *variable, offset })
}
_ => None,
}
}
#[must_use]
pub fn undirected_edge_count(&self) -> usize {
self.edges().iter().filter(|e| e.is_undirected()).count()
}
#[must_use]
pub fn directed_edge_count(&self) -> usize {
self.edges().iter().filter(|e| e.parent_child().is_some()).count()
}
#[must_use]
pub fn reaches_directed_with(
&self,
ws: &mut GraphWorkspace,
from: DenseNodeId,
to: DenseNodeId,
) -> bool {
marked_storage::reaches_directed(&self.adj, ws, from, to)
}
fn validate_node(&self, id: DenseNodeId) -> Result<(), GraphError> {
if id.as_usize() >= self.node_count() {
Err(GraphError::UnknownNode { id: id.raw() })
} else {
Ok(())
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::types::Endpoint;
#[test]
fn static_undirected_then_orient() {
let mut g = Cpdag::with_variables(2);
let x = DenseNodeId::from_raw(0);
let y = DenseNodeId::from_raw(1);
g.insert_undirected(x, y).unwrap();
assert!(g.edge_between(x, y).unwrap().is_undirected());
g.orient_undirected(x, y).unwrap();
let e = g.edge_between(x, y).unwrap();
assert_eq!(e.parent_child(), Some((x, y)));
assert_eq!(g.children(x), vec![y]);
assert!(g.undirected_neighbors(x).is_empty());
assert!(g.try_into_dag().is_ok());
}
#[test]
fn static_review_orient_and_accept() {
let mut g = Cpdag::with_variables(2);
let x = DenseNodeId::from_raw(0);
let y = DenseNodeId::from_raw(1);
g.insert_undirected(x, y).unwrap();
let review = CpdagReview::from_cpdag(g, "pc");
assert!(!review.is_complete());
let v0 = VariableId::from_raw(0);
let v1 = VariableId::from_raw(1);
let review = review.orient_edge(v0, v1).unwrap().accept_edge(v0, v1);
assert!(review.is_complete());
let dag = review.try_into_dag().unwrap();
assert!(dag.reaches(x, y));
}
#[test]
fn static_rejects_non_static_nodes() {
let mut g = Cpdag::empty();
assert!(
g.add_node(NodeRef::Lagged {
variable: VariableId::from_raw(0),
lag: Lag::CONTEMPORANEOUS,
})
.is_err()
);
}
#[test]
fn undirected_then_orient() {
let mut g = TemporalCpdag::empty();
let x = g.add_lagged(VariableId::from_raw(0), Lag::CONTEMPORANEOUS).unwrap();
let y = g.add_lagged(VariableId::from_raw(1), Lag::CONTEMPORANEOUS).unwrap();
g.insert_undirected(x, y).unwrap();
assert!(g.edge_between(x, y).unwrap().is_undirected());
g.orient_undirected(x, y).unwrap();
let e = g.edge_between(x, y).unwrap();
assert_eq!(e.parent_child(), Some((x, y)));
assert_eq!(g.children(x), vec![y]);
assert!(g.undirected_neighbors(x).is_empty());
}
#[test]
fn mark_conflict_sets_x_x() {
let mut g = TemporalCpdag::empty();
let x = g.add_lagged(VariableId::from_raw(0), Lag::CONTEMPORANEOUS).unwrap();
let y = g.add_lagged(VariableId::from_raw(1), Lag::CONTEMPORANEOUS).unwrap();
g.insert_undirected(x, y).unwrap();
g.mark_conflict(x, y).unwrap();
let e = g.edge_between(x, y).unwrap();
assert!(e.is_conflict());
assert_eq!(g.conflict_edge_count(), 1);
assert!(g.undirected_neighbors(x).is_empty());
assert!(g.try_into_temporal_dag().is_err());
}
#[test]
fn rejects_circle_marks() {
let mut g = TemporalCpdag::empty();
let x = g.add_lagged(VariableId::from_raw(0), Lag::CONTEMPORANEOUS).unwrap();
let y = g.add_lagged(VariableId::from_raw(1), Lag::CONTEMPORANEOUS).unwrap();
let edge = MarkedEdge {
a: x,
b: y,
at_a: Endpoint::Circle,
at_b: Endpoint::Arrow,
middle: crate::types::MiddleMark::Empty,
};
assert!(matches!(g.insert_marked(edge), Err(GraphError::InvalidEndpoints { .. })));
}
#[test]
fn from_temporal_dag_preserves_directed() {
let mut dag = TemporalDag::empty();
let a = dag.add_lagged(VariableId::from_raw(0), Lag::from_raw(1)).unwrap();
let b = dag.add_lagged(VariableId::from_raw(1), Lag::CONTEMPORANEOUS).unwrap();
dag.insert_directed(a, b).unwrap();
let cpdag = TemporalCpdag::from_temporal_dag(&dag);
assert_eq!(cpdag.edge_between(a, b).unwrap().parent_child(), Some((a, b)));
let back = cpdag.to_directed_skeleton().unwrap();
assert!(back.reaches(a, b));
}
#[test]
fn rejects_contemporaneous_self_edges() {
let mut g = TemporalCpdag::empty();
let x = g.add_lagged(VariableId::from_raw(0), Lag::CONTEMPORANEOUS).unwrap();
assert!(matches!(
g.insert_undirected(x, x),
Err(GraphError::ContemporaneousSelfEdge { .. })
));
assert!(matches!(g.insert_directed(x, x), Err(GraphError::ContemporaneousSelfEdge { .. })));
let x2 = g.add_lagged(VariableId::from_raw(0), Lag::CONTEMPORANEOUS).unwrap();
assert!(matches!(
g.insert_directed(x, x2),
Err(GraphError::ContemporaneousSelfEdge { .. })
));
assert!(g.edges().is_empty());
}
#[test]
fn rejects_lagged_self_loops() {
let mut g = TemporalCpdag::empty();
let x1 = g.add_lagged(VariableId::from_raw(0), Lag::from_raw(1)).unwrap();
assert!(matches!(g.insert_undirected(x1, x1), Err(GraphError::Cycle { .. })));
assert!(matches!(g.insert_directed(x1, x1), Err(GraphError::Cycle { .. })));
assert_eq!(g.undirected_edge_count(), 0);
assert!(g.edges().is_empty());
}
#[test]
fn accepts_context_nodes_without_coercing() {
use antecedent_core::EnvironmentId;
let mut g = TemporalCpdag::empty();
let c = g.add_context(VariableId::from_raw(0), Some(EnvironmentId::from_raw(1))).unwrap();
let y = g.add_lagged(VariableId::from_raw(1), Lag::CONTEMPORANEOUS).unwrap();
match g.nodes()[c.as_usize()] {
NodeRef::Context { variable, environment } => {
assert_eq!(variable, VariableId::from_raw(0));
assert_eq!(environment, Some(EnvironmentId::from_raw(1)));
}
_ => panic!("expected Context node"),
}
g.insert_directed(c, y).unwrap();
assert_eq!(g.edge_between(c, y).unwrap().parent_child(), Some((c, y)));
assert!(g.add_node(NodeRef::Static(VariableId::from_raw(2))).is_err());
}
#[test]
fn marked_edge_undirected_canonical() {
let a = DenseNodeId::from_raw(2);
let b = DenseNodeId::from_raw(1);
let e = MarkedEdge::undirected(a, b);
assert!(e.is_undirected());
assert!(e.is_cpdag_legal());
assert_eq!(e.a.raw(), 1);
assert_eq!(e.b.raw(), 2);
}
}