use std::any::Any;
use std::cell::RefCell;
use std::rc::Rc;
use super::arena::{GenArena, Key};
#[derive(Copy, Clone, Debug, PartialEq, Eq, PartialOrd, Ord)]
pub(crate) enum NodeState {
Clean = 0,
Check = 1,
Dirty = 2,
}
pub(crate) type EqFn = fn(&dyn Any, &dyn Any) -> bool;
pub(crate) fn eq_any<T: PartialEq + 'static>(a: &dyn Any, b: &dyn Any) -> bool {
match (a.downcast_ref::<T>(), b.downcast_ref::<T>()) {
(Some(a), Some(b)) => a == b,
_ => false,
}
}
pub(crate) enum NodeKind {
Scope,
Signal {
value: Rc<RefCell<Box<dyn Any>>>,
},
Memo {
value: Rc<RefCell<Option<Box<dyn Any>>>>,
compute: Rc<dyn Fn() -> Box<dyn Any>>,
eq: EqFn,
},
Effect {
run: Rc<RefCell<dyn FnMut()>>,
},
}
impl NodeKind {
pub(crate) fn is_effect(&self) -> bool {
matches!(self, NodeKind::Effect { .. })
}
pub(crate) fn is_computation(&self) -> bool {
matches!(self, NodeKind::Memo { .. } | NodeKind::Effect { .. })
}
}
pub(crate) struct Node {
pub kind: NodeKind,
pub state: NodeState,
pub sources: Vec<Key>,
pub source_slots: Vec<u32>,
pub observers: Vec<Key>,
pub observer_slots: Vec<u32>,
pub parent: Option<Key>,
pub owned: Vec<Key>,
pub cleanups: Vec<Box<dyn FnOnce()>>,
pub order: u64,
pub queued: bool,
pub running: bool,
pub run_epoch: u64,
pub seen_epoch: u64,
pub label: Option<&'static str>,
pub flush_runs: u32,
pub flush_stamp: u64,
}
impl Node {
pub(crate) fn new(kind: NodeKind, order: u64) -> Self {
Node {
kind,
state: NodeState::Clean,
sources: Vec::new(),
source_slots: Vec::new(),
observers: Vec::new(),
observer_slots: Vec::new(),
parent: None,
owned: Vec::new(),
cleanups: Vec::new(),
order,
queued: false,
running: false,
run_epoch: 0,
seen_epoch: 0,
label: None,
flush_runs: 0,
flush_stamp: 0,
}
}
pub(crate) fn describe(&self, key: super::arena::Key) -> String {
let kind = match self.kind {
NodeKind::Scope => "scope",
NodeKind::Signal { .. } => "signal",
NodeKind::Memo { .. } => "memo",
NodeKind::Effect { .. } => "effect",
};
match self.label {
Some(l) => format!("{kind} '{l}' (node #{})", key.index),
None => format!("{kind} node #{}", key.index),
}
}
}
pub(crate) type Graph = GenArena<Node>;
pub(crate) fn add_edge(graph: &mut Graph, source: Key, observer: Key) {
let k = {
let obs = graph
.get_mut(observer)
.expect("add_edge: observer vanished");
obs.sources.push(source);
obs.source_slots.push(u32::MAX); obs.sources.len() - 1
};
let j = {
let src = graph.get_mut(source).expect("add_edge: source vanished");
src.observers.push(observer);
src.observer_slots.push(k as u32);
src.observers.len() - 1
};
graph
.get_mut(observer)
.expect("add_edge: observer vanished")
.source_slots[k] = j as u32;
}
pub(crate) fn remove_source_edges(graph: &mut Graph, observer: Key) {
let pairs: Vec<(Key, u32)> = {
let Some(obs) = graph.get_mut(observer) else {
return;
};
obs.sources
.drain(..)
.zip(obs.source_slots.drain(..))
.collect()
};
for (source, slot) in pairs {
let j = slot as usize;
let moved = {
let Some(src) = graph.get_mut(source) else {
continue;
}; if j >= src.observers.len() || src.observers[j] != observer {
continue; }
src.observers.swap_remove(j);
src.observer_slots.swap_remove(j);
if j < src.observers.len() {
Some((src.observers[j], src.observer_slots[j] as usize))
} else {
None
}
};
if let Some((moved_obs, k2)) = moved {
if let Some(o2) = graph.get_mut(moved_obs) {
if k2 < o2.source_slots.len() {
o2.source_slots[k2] = j as u32;
}
}
}
}
}
pub(crate) fn remove_observer_edges(graph: &mut Graph, source: Key) {
let pairs: Vec<(Key, u32)> = {
let Some(src) = graph.get_mut(source) else {
return;
};
src.observers
.drain(..)
.zip(src.observer_slots.drain(..))
.collect()
};
for (observer, slot) in pairs {
let k = slot as usize;
let moved = {
let Some(obs) = graph.get_mut(observer) else {
continue;
};
if k >= obs.sources.len() || obs.sources[k] != source {
continue;
}
obs.sources.swap_remove(k);
obs.source_slots.swap_remove(k);
if k < obs.sources.len() {
Some((obs.sources[k], obs.source_slots[k] as usize))
} else {
None
}
};
if let Some((moved_src, j2)) = moved {
if let Some(s2) = graph.get_mut(moved_src) {
if j2 < s2.observer_slots.len() {
s2.observer_slots[j2] = k as u32;
}
}
}
}
}
#[cfg(test)]
pub(crate) fn check_edge_invariants(graph: &Graph, keys: &[Key]) {
for &key in keys {
let Some(node) = graph.get(key) else { continue };
assert_eq!(node.sources.len(), node.source_slots.len());
assert_eq!(node.observers.len(), node.observer_slots.len());
for (k, (&src, &slot)) in node.sources.iter().zip(&node.source_slots).enumerate() {
let s = graph.get(src).expect("dangling source");
let j = slot as usize;
assert_eq!(s.observers[j], key, "observer back-pointer broken");
assert_eq!(s.observer_slots[j] as usize, k, "slot pairing broken");
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn scope_node(graph: &mut Graph, order: u64) -> Key {
graph.insert(Node::new(NodeKind::Scope, order))
}
#[test]
fn edge_add_remove_repairs_slots() {
let mut g = Graph::new();
let s1 = scope_node(&mut g, 0);
let s2 = scope_node(&mut g, 1);
let o1 = scope_node(&mut g, 2);
let o2 = scope_node(&mut g, 3);
add_edge(&mut g, s1, o1);
add_edge(&mut g, s2, o1);
add_edge(&mut g, s1, o2);
check_edge_invariants(&g, &[s1, s2, o1, o2]);
remove_source_edges(&mut g, o1);
check_edge_invariants(&g, &[s1, s2, o1, o2]);
assert!(g.get(o1).unwrap().sources.is_empty());
assert_eq!(g.get(s1).unwrap().observers, vec![o2]);
assert!(g.get(s2).unwrap().observers.is_empty());
remove_observer_edges(&mut g, s1);
check_edge_invariants(&g, &[s1, s2, o1, o2]);
assert!(g.get(o2).unwrap().sources.is_empty());
}
}