use antecedent_core::TemporalNodeKey;
use antecedent_core::{Lag, VariableId};
use crate::algo::bfs_reaches;
use crate::error::GraphError;
use crate::types::{DenseNodeId, MarkedEdge, NodeRef};
use crate::workspace::GraphWorkspace;
#[derive(Clone, Debug)]
pub struct TemporalDag {
nodes: Vec<NodeRef>,
children: Vec<Vec<DenseNodeId>>,
parents: Vec<Vec<DenseNodeId>>,
insert_ws: GraphWorkspace,
}
impl TemporalDag {
#[must_use]
pub fn empty() -> Self {
Self {
nodes: Vec::new(),
children: Vec::new(),
parents: Vec::new(),
insert_ws: GraphWorkspace::default(),
}
}
#[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 { .. } => {}
_ => {
return Err(GraphError::InvalidEndpoints {
message: "TemporalDag accepts only Lagged nodes",
});
}
}
let id = u32::try_from(self.nodes.len()).map_err(|_| GraphError::TooManyNodes)?;
self.nodes.push(node);
self.children.push(Vec::new());
self.parents.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 insert_directed(
&mut self,
from: DenseNodeId,
to: DenseNodeId,
) -> Result<(), GraphError> {
self.validate_node(from)?;
self.validate_node(to)?;
if let (
NodeRef::Lagged { variable: v1, lag: l1 },
NodeRef::Lagged { variable: v2, lag: l2 },
) = (self.nodes[from.as_usize()], self.nodes[to.as_usize()])
{
if v1 == v2 && l1 == l2 && l1.is_contemporaneous() {
return Err(GraphError::ContemporaneousSelfEdge { variable: v1 });
}
if l1 < l2 {
return Err(GraphError::FutureToPast {
from: from.raw(),
to: to.raw(),
from_lag: l1,
to_lag: l2,
});
}
}
if self.children[from.as_usize()].contains(&to) {
return Err(GraphError::DuplicateEdge { from: from.raw(), to: to.raw() });
}
let mut ws = core::mem::take(&mut self.insert_ws);
let cycle = bfs_reaches(&self.children, to, from, &mut ws);
self.insert_ws = ws;
if cycle {
return Err(GraphError::Cycle { from: from.raw(), to: to.raw() });
}
self.children[from.as_usize()].push(to);
self.parents[to.as_usize()].push(from);
Ok(())
}
#[must_use]
pub fn children(&self, id: DenseNodeId) -> &[DenseNodeId] {
&self.children[id.as_usize()]
}
pub fn edges(&self) -> impl Iterator<Item = MarkedEdge> + '_ {
self.children.iter().enumerate().flat_map(|(i, kids)| {
let from = DenseNodeId::from_raw(u32::try_from(i).expect("node fit"));
kids.iter().map(move |&to| MarkedEdge::directed(from, to))
})
}
#[must_use]
pub fn reaches(&self, from: DenseNodeId, to: DenseNodeId) -> bool {
let mut ws = GraphWorkspace::default();
bfs_reaches(&self.children, from, to, &mut ws)
}
pub fn reaches_with(
&self,
from: DenseNodeId,
to: DenseNodeId,
ws: &mut GraphWorkspace,
) -> bool {
bfs_reaches(&self.children, from, to, ws)
}
#[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,
}
}
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::*;
#[test]
fn rejects_contemporaneous_self_edge() {
let mut g = TemporalDag::empty();
let n = g.add_lagged(VariableId::from_raw(0), Lag::CONTEMPORANEOUS).unwrap();
assert!(matches!(g.insert_directed(n, n), Err(GraphError::ContemporaneousSelfEdge { .. })));
}
#[test]
fn allows_lagged_self_edge() {
let mut g = TemporalDag::empty();
let past = g.add_lagged(VariableId::from_raw(0), Lag::from_raw(1)).unwrap();
let now = g.add_lagged(VariableId::from_raw(0), Lag::CONTEMPORANEOUS).unwrap();
g.insert_directed(past, now).unwrap();
assert!(g.reaches(past, now));
}
#[test]
fn rejects_future_to_past_edge() {
let mut g = TemporalDag::empty();
let past = g.add_lagged(VariableId::from_raw(0), Lag::from_raw(1)).unwrap();
let now = g.add_lagged(VariableId::from_raw(1), Lag::CONTEMPORANEOUS).unwrap();
assert!(matches!(g.insert_directed(now, past), Err(GraphError::FutureToPast { .. })));
}
}