use std::cell::{Cell, RefCell};
use slotmap::{Key, SlotMap};
use smallvec::SmallVec;
use crate::*;
pub(crate) struct Root {
pub tracker: RefCell<Option<DependencyTracker>>,
pub rev_sorted_buf: RefCell<Vec<NodeId>>,
pub current_node: Cell<NodeId>,
pub root_node: Cell<NodeId>,
pub nodes: RefCell<SlotMap<NodeId, ReactiveNode>>,
pub node_update_queue: RefCell<Vec<NodeId>>,
pub batch_depth: Cell<usize>,
}
thread_local! {
static GLOBAL_ROOT: Cell<Option<&'static Root>> = const { Cell::new(None) };
}
impl Root {
#[cfg_attr(debug_assertions, track_caller)]
pub fn global() -> &'static Root {
GLOBAL_ROOT.with(|root| root.get()).expect("no root found")
}
pub fn set_global(root: Option<&'static Root>) -> Option<&'static Root> {
GLOBAL_ROOT.with(|r| r.replace(root))
}
pub fn new_static() -> &'static Self {
let this = Self {
tracker: RefCell::new(None),
rev_sorted_buf: RefCell::new(Vec::new()),
current_node: Cell::new(NodeId::null()),
root_node: Cell::new(NodeId::null()),
nodes: RefCell::new(SlotMap::default()),
node_update_queue: RefCell::new(Vec::new()),
batch_depth: Cell::new(0),
};
let _ref = Box::leak(Box::new(this));
_ref.reinit();
_ref
}
pub fn reinit(&'static self) {
NodeHandle(self.root_node.get(), self).dispose();
let _ = self.tracker.take();
let _ = self.rev_sorted_buf.take();
let _ = self.node_update_queue.take();
let _ = self.current_node.take();
let _ = self.root_node.take();
let _ = self.nodes.take();
self.batch_depth.set(0);
Root::set_global(Some(self));
let root_node = create_child_scope(|| {});
Root::set_global(None);
self.root_node.set(root_node.0);
self.current_node.set(root_node.0);
}
pub fn create_child_scope(&'static self, f: impl FnOnce()) -> NodeHandle {
let node = create_signal(()).id;
let prev = self.current_node.replace(node);
f();
self.current_node.set(prev);
NodeHandle(node, self)
}
pub fn tracked_scope<T>(&self, f: impl FnOnce() -> T) -> (T, DependencyTracker) {
let prev = self.tracker.replace(Some(DependencyTracker::default()));
let ret = f();
(ret, self.tracker.replace(prev).unwrap())
}
pub fn ensure_node_is_clean(&'static self, node: NodeId) {
let is_clean = self
.nodes
.borrow()
.get(node)
.is_none_or(|node| node.state == NodeState::Clean);
if !is_clean {
self.run_node_update(node);
}
}
fn run_node_update(&'static self, current: NodeId) {
debug_assert_eq!(
self.nodes.borrow()[current].state,
NodeState::Dirty,
"should only update when dirty"
);
let dependencies = std::mem::take(&mut self.nodes.borrow_mut()[current].dependencies);
for dependency in dependencies {
if let Some(node) = self.nodes.borrow_mut().get_mut(dependency) {
node.dependents.retain(|&id| id != current);
}
}
let mut nodes_mut = self.nodes.borrow_mut();
let mut callback = nodes_mut[current].callback.take().unwrap();
let mut value = nodes_mut[current].value.take().unwrap();
drop(nodes_mut);
NodeHandle(current, self).dispose_children();
let prev = self.current_node.replace(current);
let (changed, tracker) = self.tracked_scope(|| callback(&mut value));
self.current_node.set(prev);
tracker.create_dependency_link(self, current);
let mut nodes_mut = self.nodes.borrow_mut();
nodes_mut[current].callback = Some(callback); nodes_mut[current].value = Some(value);
nodes_mut[current].state = NodeState::Clean;
drop(nodes_mut);
if changed {
self.mark_dependents_dirty(current);
}
}
fn mark_dependents_dirty(&self, current: NodeId) {
let mut nodes_mut = self.nodes.borrow_mut();
let dependents = std::mem::take(&mut nodes_mut[current].dependents);
for &dependent in &dependents {
if let Some(dependent) = nodes_mut.get_mut(dependent) {
dependent.state = NodeState::Dirty;
}
}
nodes_mut[current].dependents = dependents;
}
fn propagate_node_updates(&'static self, start_nodes: &[NodeId]) {
let mut rev_sorted = Vec::new();
let mut rev_sorted_buf = self.rev_sorted_buf.try_borrow_mut();
let rev_sorted = if let Ok(rev_sorted_buf) = rev_sorted_buf.as_mut() {
rev_sorted_buf.clear();
rev_sorted_buf
} else {
&mut rev_sorted
};
for &node in start_nodes {
if self.nodes.borrow().get(node).is_none() {
continue;
}
Self::dfs(node, &mut self.nodes.borrow_mut(), rev_sorted);
self.mark_dependents_dirty(node);
}
for &node in rev_sorted.iter().rev() {
let mut nodes_mut = self.nodes.borrow_mut();
if nodes_mut.get(node).is_none() {
continue;
}
let node_state = &mut nodes_mut[node];
node_state.mark = Mark::None;
if nodes_mut[node].state == NodeState::Dirty {
drop(nodes_mut); self.run_node_update(node)
};
}
}
pub fn propagate_updates(&'static self, start_node: NodeId) {
if self.batch_depth.get() > 0 {
self.node_update_queue.borrow_mut().push(start_node);
} else {
let prev = Root::set_global(Some(self));
self.propagate_node_updates(&[start_node]);
Root::set_global(prev);
}
}
fn dfs(current_id: NodeId, nodes: &mut SlotMap<NodeId, ReactiveNode>, buf: &mut Vec<NodeId>) {
let Some(current) = nodes.get_mut(current_id) else {
return;
};
match current.mark {
Mark::Temp => panic!("cyclic reactive dependency"),
Mark::Permanent => return,
Mark::None => {}
}
current.mark = Mark::Temp;
let children = std::mem::take(&mut current.dependents);
for child in &children {
Self::dfs(*child, nodes, buf);
}
nodes[current_id].dependents = children;
nodes[current_id].mark = Mark::Permanent;
buf.push(current_id);
}
fn start_batch(&self) {
self.batch_depth.set(self.batch_depth.get() + 1);
}
fn end_batch(&'static self) {
let depth = self.batch_depth.get();
debug_assert!(depth > 0, "end_batch called without matching start_batch");
self.batch_depth.set(depth - 1);
if depth == 1 {
let nodes = self.node_update_queue.take();
self.propagate_node_updates(&nodes);
}
}
}
#[derive(Clone, Copy)]
pub struct RootHandle {
_ref: &'static Root,
}
impl RootHandle {
pub fn dispose(&self) {
self._ref.reinit();
}
pub fn run_in<T>(&self, f: impl FnOnce() -> T) -> T {
let prev = Root::set_global(Some(self._ref));
let ret = f();
Root::set_global(prev);
ret
}
}
#[derive(Default)]
pub(crate) struct DependencyTracker {
pub dependencies: SmallVec<[NodeId; 1]>,
}
impl DependencyTracker {
pub fn create_dependency_link(self, root: &Root, dependent: NodeId) {
for node in &self.dependencies {
if let Some(node) = root.nodes.borrow_mut().get_mut(*node) {
node.dependents.push(dependent)
}
}
root.nodes.borrow_mut()[dependent].dependencies = self.dependencies;
}
}
#[must_use = "root should be disposed"]
pub fn create_root(f: impl FnOnce()) -> RootHandle {
let _ref = Root::new_static();
#[cfg(not(target_arch = "wasm32"))]
{
#[allow(dead_code)]
struct UnsafeSendPtr<T>(*const T);
unsafe impl<T> Send for UnsafeSendPtr<T> {}
static KEEP_ALIVE: std::sync::Mutex<Vec<UnsafeSendPtr<Root>>> =
std::sync::Mutex::new(Vec::new());
KEEP_ALIVE
.lock()
.unwrap()
.push(UnsafeSendPtr(_ref as *const Root));
}
Root::set_global(Some(_ref));
NodeHandle(_ref.root_node.get(), _ref).run_in(f);
Root::set_global(None);
RootHandle { _ref }
}
#[cfg_attr(debug_assertions, track_caller)]
pub fn create_child_scope(f: impl FnOnce()) -> NodeHandle {
Root::global().create_child_scope(f)
}
#[cfg_attr(debug_assertions, track_caller)]
pub fn on_cleanup(f: impl FnOnce() + 'static) {
let root = Root::global();
if !root.current_node.get().is_null() {
root.nodes.borrow_mut()[root.current_node.get()]
.cleanups
.push(Box::new(f));
}
}
pub fn batch<T>(f: impl FnOnce() -> T) -> T {
let root = Root::global();
root.start_batch();
let ret = f();
root.end_batch();
ret
}
pub fn untrack<T>(f: impl FnOnce() -> T) -> T {
untrack_in_scope(f, Root::global())
}
pub(crate) fn untrack_in_scope<T>(f: impl FnOnce() -> T, root: &'static Root) -> T {
let prev = root.tracker.replace(None);
let ret = f();
root.tracker.replace(prev);
ret
}
pub fn use_current_scope() -> NodeHandle {
let root = Root::global();
NodeHandle(root.current_node.get(), root)
}
pub fn use_global_scope() -> NodeHandle {
let root = Root::global();
NodeHandle(root.root_node.get(), root)
}
#[cfg(test)]
mod tests {
use crate::*;
#[test]
fn test_lazy_vs_eager_updates() {
let _ = create_root(|| {
let counter = create_signal(0);
let trigger = create_signal(0);
create_effect(move || {
let _ = trigger.get();
counter.set(counter.get_untracked() + 1);
});
assert_eq!(counter.get(), 1); trigger.set(1);
assert_eq!(counter.get(), 2);
});
}
#[test]
fn cleanup() {
let _ = create_root(|| {
let cleanup_called = create_signal(false);
let scope = create_child_scope(|| {
on_cleanup(move || {
cleanup_called.set(true);
});
});
assert!(!cleanup_called.get());
scope.dispose();
assert!(cleanup_called.get());
});
}
#[test]
fn cleanup_in_effect() {
let _ = create_root(|| {
let trigger = create_signal(());
let counter = create_signal(0);
create_effect(move || {
trigger.track();
on_cleanup(move || {
counter.set(counter.get() + 1);
});
});
assert_eq!(counter.get(), 0);
trigger.set(());
assert_eq!(counter.get(), 1);
trigger.set(());
assert_eq!(counter.get(), 2);
});
}
#[test]
fn cleanup_is_untracked() {
let _ = create_root(|| {
let trigger = create_signal(());
let counter = create_signal(0);
create_effect(move || {
counter.set(counter.get_untracked() + 1);
on_cleanup(move || {
trigger.track(); });
});
assert_eq!(counter.get(), 1);
trigger.set(());
assert_eq!(counter.get(), 1);
});
}
#[test]
fn batch_memo() {
let _ = create_root(|| {
let state = create_signal(1);
let double = create_memo(move || state.get() * 2);
batch(move || {
state.set(2);
assert_eq!(double.get(), 2);
});
assert_eq!(double.get(), 4);
});
}
#[test]
fn batch_updates_effects_at_end() {
let _ = create_root(|| {
let state1 = create_signal(1);
let state2 = create_signal(2);
let counter = create_signal(0);
create_effect(move || {
counter.set(counter.get_untracked() + 1);
let _ = state1.get() + state2.get();
});
assert_eq!(counter.get(), 1);
state1.set(2);
state2.set(3);
assert_eq!(counter.get(), 3);
batch(move || {
state1.set(3);
assert_eq!(counter.get(), 3);
state2.set(4);
assert_eq!(counter.get(), 3);
});
assert_eq!(counter.get(), 4);
});
}
#[test]
fn nested_batches_compose() {
let _ = create_root(|| {
let state = create_signal("Initial");
let counter = create_signal(0);
create_effect(move || {
counter.set(counter.get_untracked() + 1);
let _ = state.get();
});
assert_eq!(counter.get(), 1);
batch(|| {
state.set("First in outer batch");
assert_eq!(counter.get(), 1);
batch(|| {
state.set("First in inner batch");
assert_eq!(counter.get(), 1); state.set("Last in inner batch");
assert_eq!(counter.get(), 1); });
assert_eq!(counter.get(), 1); state.set("Last in outer batch");
assert_eq!(counter.get(), 1); });
assert_eq!(counter.get(), 2);
});
}
#[test]
fn issue_741_disposing_signal_inside_effect_should_not_panic() {
let _ = create_root(|| {
let a = create_signal(0);
let b = create_signal(0);
create_effect(move || {
a.track();
b.track();
a.dispose();
});
b.set(0);
});
}
#[test]
fn batched_updates_do_not_panic_after_disposal() {
let _ = create_root(|| {
let a = create_signal(0);
batch(|| {
a.set(1);
a.dispose();
})
});
}
}