use crate::hash::FastHashBuilder;
use indexmap::IndexSet;
use papaya::HashMap as PapayaHashMap;
use parking_lot::{Mutex, RwLock};
use slab::Slab;
use std::cell::RefCell;
use std::collections::HashMap;
use std::sync::LazyLock;
use std::sync::atomic::{AtomicU8, Ordering};
#[repr(u8)]
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord)]
pub enum ReactiveState {
Clean = 0,
Check = 1,
Dirty = 2,
}
impl ReactiveState {
pub fn from_u8(v: u8) -> Self {
match v {
0 => ReactiveState::Clean,
1 => ReactiveState::Check,
_ => ReactiveState::Dirty,
}
}
}
use super::SignalId;
#[derive(Clone, Copy, Default)]
pub(crate) struct SignalRelation {
pub(crate) is_source: bool,
pub(crate) is_output: bool,
}
type SignalFilterMapIter<'a> = std::iter::FilterMap<
std::collections::hash_map::Iter<'a, SignalId, SignalRelation>,
fn((&'a SignalId, &'a SignalRelation)) -> Option<SignalId>,
>;
pub struct SourcesIter<'a> {
inner: SignalFilterMapIter<'a>,
}
impl Iterator for SourcesIter<'_> {
type Item = SignalId;
#[inline]
fn next(&mut self) -> Option<Self::Item> {
self.inner.next()
}
#[inline]
fn size_hint(&self) -> (usize, Option<usize>) {
self.inner.size_hint()
}
}
pub struct OutputsIter<'a> {
inner: SignalFilterMapIter<'a>,
}
impl Iterator for OutputsIter<'_> {
type Item = SignalId;
#[inline]
fn next(&mut self) -> Option<Self::Item> {
self.inner.next()
}
#[inline]
fn size_hint(&self) -> (usize, Option<usize>) {
self.inner.size_hint()
}
}
static EFFECT_ARENA: RwLock<Slab<EffectMetadata>> = RwLock::new(Slab::new());
static PENDING_EFFECTS: LazyLock<RwLock<IndexSet<EffectId, FastHashBuilder>>> =
LazyLock::new(|| RwLock::new(IndexSet::default()));
static EFFECT_PARENT: LazyLock<PapayaHashMap<EffectId, EffectId>> =
LazyLock::new(PapayaHashMap::new);
static EFFECT_CHILDREN: LazyLock<PapayaHashMap<EffectId, RwLock<Vec<EffectId>>>> =
LazyLock::new(PapayaHashMap::new);
thread_local! {
static CURRENT_EFFECT: RefCell<Option<EffectId>> = const { RefCell::new(None) };
}
pub fn current_effect() -> Option<EffectId> {
CURRENT_EFFECT.with(|c| *c.borrow())
}
pub fn set_current_effect(effect_id: Option<EffectId>) -> Option<EffectId> {
CURRENT_EFFECT.with(|c| c.replace(effect_id))
}
pub struct CurrentEffectGuard {
previous: Option<EffectId>,
}
impl CurrentEffectGuard {
pub fn new(new_value: Option<EffectId>) -> Self {
let previous = set_current_effect(new_value);
Self { previous }
}
}
impl Drop for CurrentEffectGuard {
fn drop(&mut self) {
set_current_effect(self.previous);
}
}
#[repr(transparent)]
#[derive(Copy, Clone, Eq, PartialEq, Hash, Debug)]
pub struct EffectId(u32);
impl EffectId {
pub fn new(index: u32) -> Self {
Self(index)
}
pub fn index(self) -> usize {
self.0 as usize
}
pub fn with<F, R>(self, f: F) -> Option<R>
where
F: FnOnce(&EffectMetadata) -> R,
{
let arena = EFFECT_ARENA.read();
arena.get(self.index()).map(f)
}
pub fn with_sources<F, R>(self, f: F) -> Option<R>
where
F: FnOnce(SourcesIter<'_>) -> R,
{
self.with(|metadata| metadata.with_sources(f))
}
pub fn add_source(self, source: SignalId) -> Option<()> {
self.with(|metadata| {
metadata.add_source(source);
})
}
pub fn clear_sources(self) -> Option<()> {
self.with(|metadata| {
metadata.clear_sources();
})
}
pub fn remove_source(self, source: SignalId) -> Option<()> {
self.with(|metadata| {
metadata.remove_source(source);
})
}
pub fn has_source(self, signal_id: SignalId) -> bool {
self.with(|metadata| metadata.has_source(signal_id))
.unwrap_or(false)
}
pub fn state(self) -> ReactiveState {
self.with(EffectMetadata::get_state)
.unwrap_or(ReactiveState::Clean)
}
pub fn set_state(self, state: ReactiveState) {
self.with(|metadata| metadata.set_state(state));
}
pub fn upgrade_state(self, new_state: ReactiveState) -> bool {
self.with(|metadata| metadata.upgrade_state(new_state))
.unwrap_or(false)
}
pub fn needs_work(self) -> bool {
self.state() != ReactiveState::Clean
}
pub fn run_callback(self) {
struct CallbackGuard {
effect_id: EffectId,
callback: Option<Box<dyn FnMut() + Send>>,
}
impl CallbackGuard {
fn run(&mut self) {
if let Some(ref mut cb) = self.callback {
cb();
}
}
}
impl Drop for CallbackGuard {
fn drop(&mut self) {
if let Some(cb) = self.callback.take() {
let arena = EFFECT_ARENA.read();
if let Some(meta) = arena.get(self.effect_id.index()) {
*meta.callback.lock() = Some(cb);
}
}
}
}
let callback = {
let arena = EFFECT_ARENA.read();
if let Some(meta) = arena.get(self.index()) {
meta.callback.lock().take()
} else {
None
}
};
if let Some(cb) = callback {
let mut guard = CallbackGuard {
effect_id: self,
callback: Some(cb),
};
guard.run();
}
}
pub fn has_callback(self) -> bool {
self.with(EffectMetadata::has_callback).unwrap_or(false)
}
pub fn parent(self) -> Option<EffectId> {
let guard = EFFECT_PARENT.pin();
guard.get(&self).copied()
}
pub fn with_children<F, R>(self, f: F) -> Option<R>
where
F: FnOnce(&Vec<EffectId>) -> R,
{
let guard = EFFECT_CHILDREN.pin();
guard.get(&self).map(|children| {
let read_guard = children.read();
f(&read_guard)
})
}
pub fn add_child(self, child: EffectId) {
let guard = EFFECT_CHILDREN.pin();
guard
.get_or_insert_with(self, || RwLock::new(Vec::new()))
.write()
.push(child);
}
pub fn remove_child(self, child: EffectId) {
let guard = EFFECT_CHILDREN.pin();
if let Some(children) = guard.get(&self) {
children.write().retain(|c| *c != child);
}
}
pub fn clear_children(self) {
let guard = EFFECT_CHILDREN.pin();
if let Some(children) = guard.get(&self) {
children.write().clear();
}
}
pub fn with_outputs<F, R>(self, f: F) -> Option<R>
where
F: FnOnce(OutputsIter<'_>) -> R,
{
self.with(|metadata| metadata.with_outputs(f))
}
pub fn add_output(self, signal_id: SignalId) -> Option<()> {
self.with(|metadata| {
metadata.add_output(signal_id);
})
}
pub fn clear_signals<F1, F2>(self, on_source: F1, on_output: F2)
where
F1: FnMut(SignalId),
F2: FnMut(SignalId),
{
self.with(|metadata| metadata.clear_signals(on_source, on_output));
}
pub fn is_skippable(self) -> bool {
self.with(EffectMetadata::is_skippable).unwrap_or(false)
}
pub fn is_ui_updates_only(self) -> bool {
self.with(EffectMetadata::is_ui_updates_only)
.unwrap_or(false)
}
pub fn mark_check_recursive(self) {
let current = self.state();
if current == ReactiveState::Clean {
cov_mark::hit!(marking_effect_check_recursive);
self.set_state(ReactiveState::Check);
mark_effect_pending_for_check(self);
self.with_outputs(|outputs| {
for output_signal in outputs {
output_signal.mark_subscribers_check();
}
});
}
}
pub fn mark_dirty(self) {
mark_effect_dirty(self);
}
pub fn update_if_necessary(self, skip_skippable: bool) -> bool {
match self.state() {
ReactiveState::Clean => false,
ReactiveState::Check => {
self.with_sources(|sources| {
for source_id in sources {
source_id.pull(skip_skippable);
if self.state() == ReactiveState::Dirty {
cov_mark::hit!(check_upgraded_to_dirty_by_pull);
break;
}
}
});
if self.state() == ReactiveState::Check {
cov_mark::hit!(check_verified_clean);
self.set_state(ReactiveState::Clean);
return false;
}
cov_mark::hit!(check_became_dirty_running);
self.run_and_mark_children_if_changed()
}
ReactiveState::Dirty => {
cov_mark::hit!(dirty_running);
self.run_and_mark_children_if_changed()
}
}
}
fn run_and_mark_children_if_changed(self) -> bool {
run_single_effect(self);
self.with_outputs(|outputs| {
for output in outputs {
output.with_subscribers(|subs| {
for effect_id in subs {
if effect_id.state() == ReactiveState::Check {
cov_mark::hit!(child_upgraded_check_to_dirty);
}
effect_id.upgrade_state(ReactiveState::Dirty);
}
});
}
});
true
}
}
pub struct EffectMetadata {
pub(crate) state: AtomicU8,
pub(crate) flags: u8,
pub(crate) callback: Mutex<Option<Box<dyn FnMut() + Send>>>,
pub(crate) signals: RwLock<HashMap<SignalId, SignalRelation, FastHashBuilder>>,
}
const FLAG_SKIPPABLE: u8 = 1 << 0;
const FLAG_UI_UPDATES_ONLY: u8 = 1 << 1;
impl EffectMetadata {
pub fn new() -> Self {
Self {
state: AtomicU8::new(ReactiveState::Clean as u8),
flags: 0,
callback: Mutex::new(None),
signals: RwLock::new(HashMap::with_hasher(FastHashBuilder)),
}
}
pub fn new_with_callback(callback: Box<dyn FnMut() + Send>) -> Self {
Self {
state: AtomicU8::new(ReactiveState::Clean as u8),
flags: 0,
callback: Mutex::new(Some(callback)),
signals: RwLock::new(HashMap::with_hasher(FastHashBuilder)),
}
}
pub fn new_with_callback_and_parent(
callback: Box<dyn FnMut() + Send>,
_parent: Option<EffectId>,
) -> Self {
Self {
state: AtomicU8::new(ReactiveState::Clean as u8),
flags: 0,
callback: Mutex::new(Some(callback)),
signals: RwLock::new(HashMap::with_hasher(FastHashBuilder)),
}
}
pub fn new_dirty() -> Self {
Self {
state: AtomicU8::new(ReactiveState::Dirty as u8),
flags: 0,
callback: Mutex::new(None),
signals: RwLock::new(HashMap::with_hasher(FastHashBuilder)),
}
}
pub fn new_with_callback_parent_and_flags(
callback: Box<dyn FnMut() + Send>,
_parent: Option<EffectId>,
skippable: bool,
ui_updates_only: bool,
) -> Self {
let mut flags = 0;
if skippable {
flags |= FLAG_SKIPPABLE;
}
if ui_updates_only {
flags |= FLAG_UI_UPDATES_ONLY;
}
Self {
state: AtomicU8::new(ReactiveState::Clean as u8),
flags,
callback: Mutex::new(Some(callback)),
signals: RwLock::new(HashMap::with_hasher(FastHashBuilder)),
}
}
pub fn new_with_callback_parent_and_skippable(
callback: Box<dyn FnMut() + Send>,
parent: Option<EffectId>,
skippable: bool,
) -> Self {
Self::new_with_callback_parent_and_flags(callback, parent, skippable, false)
}
pub fn get_state(&self) -> ReactiveState {
ReactiveState::from_u8(self.state.load(Ordering::Acquire))
}
pub fn is_skippable(&self) -> bool {
(self.flags & FLAG_SKIPPABLE) != 0
}
pub fn is_ui_updates_only(&self) -> bool {
(self.flags & FLAG_UI_UPDATES_ONLY) != 0
}
pub fn set_state(&self, state: ReactiveState) {
self.state.store(state as u8, Ordering::Release);
}
pub fn replace_state(&self, state: ReactiveState) -> ReactiveState {
ReactiveState::from_u8(self.state.swap(state as u8, Ordering::Release))
}
pub fn upgrade_state(&self, new_state: ReactiveState) -> bool {
let new_state_u8 = new_state as u8;
loop {
let current = self.state.load(Ordering::Acquire);
if current >= new_state_u8 {
return false;
}
match self.state.compare_exchange_weak(
current,
new_state_u8,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => return true,
Err(_) => continue, }
}
}
pub fn needs_work(&self) -> bool {
self.get_state() != ReactiveState::Clean
}
pub fn add_source(&self, source_id: SignalId) {
let mut signals = self.signals.write();
signals.entry(source_id).or_default().is_source = true;
}
pub fn remove_source(&self, source_id: SignalId) {
let mut signals = self.signals.write();
if let Some(relation) = signals.get_mut(&source_id) {
relation.is_source = false;
if !relation.is_source && !relation.is_output {
signals.remove(&source_id);
}
}
}
pub fn has_source(&self, signal_id: SignalId) -> bool {
self.signals
.read()
.get(&signal_id)
.is_some_and(|r| r.is_source)
}
pub fn with_sources<F, R>(&self, f: F) -> R
where
F: FnOnce(SourcesIter<'_>) -> R,
{
fn filter_sources((id, relation): (&SignalId, &SignalRelation)) -> Option<SignalId> {
if relation.is_source { Some(*id) } else { None }
}
let signals = self.signals.read();
let iter = SourcesIter {
inner: signals.iter().filter_map(filter_sources as fn(_) -> _),
};
f(iter)
}
pub fn clear_sources(&self) {
let mut signals = self.signals.write();
signals.retain(|_, relation| {
relation.is_source = false;
relation.is_output });
}
pub fn run_callback(&self) {
let mut cb = self.callback.lock();
if let Some(ref mut f) = *cb {
(*f)();
}
}
pub fn has_callback(&self) -> bool {
self.callback.lock().is_some()
}
pub fn add_output(&self, signal_id: SignalId) {
let mut signals = self.signals.write();
signals.entry(signal_id).or_default().is_output = true;
}
pub fn with_outputs<F, R>(&self, f: F) -> R
where
F: FnOnce(OutputsIter<'_>) -> R,
{
fn filter_outputs((id, relation): (&SignalId, &SignalRelation)) -> Option<SignalId> {
if relation.is_output { Some(*id) } else { None }
}
let signals = self.signals.read();
let iter = OutputsIter {
inner: signals.iter().filter_map(filter_outputs as fn(_) -> _),
};
f(iter)
}
pub fn clear_signals<F1, F2>(&self, mut on_source: F1, mut on_output: F2)
where
F1: FnMut(SignalId),
F2: FnMut(SignalId),
{
let mut signals = self.signals.write();
for (id, relation) in signals.drain() {
if relation.is_source {
on_source(id);
}
if relation.is_output {
on_output(id);
}
}
}
}
impl Default for EffectMetadata {
fn default() -> Self {
Self::new()
}
}
pub fn effect_arena_insert(metadata: EffectMetadata) -> EffectId {
let mut arena = EFFECT_ARENA.write();
let entry = arena.vacant_entry();
let key = entry.key();
entry.insert(metadata);
EffectId::new(key as u32)
}
pub fn set_effect_parent(effect_id: EffectId, parent_id: EffectId) {
let guard = EFFECT_PARENT.pin();
guard.insert(effect_id, parent_id);
}
pub fn effect_arena_remove(id: EffectId) -> Option<EffectMetadata> {
{
let guard = EFFECT_PARENT.pin();
guard.remove(&id);
}
{
let guard = EFFECT_CHILDREN.pin();
guard.remove(&id);
}
let mut arena = EFFECT_ARENA.write();
if arena.contains(id.index()) {
Some(arena.remove(id.index()))
} else {
None
}
}
pub fn mark_effect_pending(effect_id: EffectId) -> bool {
let was_not_dirty = effect_id
.with(|metadata| metadata.replace_state(ReactiveState::Dirty) != ReactiveState::Dirty)
.unwrap_or(false);
if was_not_dirty {
destroy_children(effect_id);
PENDING_EFFECTS.write().insert(effect_id);
}
was_not_dirty
}
pub fn destroy_children(effect_id: EffectId) {
let children: Vec<EffectId> = {
let guard = EFFECT_CHILDREN.pin();
if let Some(rw) = guard.get(&effect_id) {
let children = std::mem::take(&mut *rw.write());
guard.remove(&effect_id);
children
} else {
return;
}
};
for child_id in children {
destroy_children(child_id);
remove_from_pending_set(child_id);
child_id.with_sources(|sources| {
for source_id in sources {
source_id.remove_subscriber(child_id);
}
});
effect_arena_remove(child_id);
}
}
pub fn take_pending_effects() -> Vec<EffectId> {
PENDING_EFFECTS.write().drain(..).collect()
}
pub fn take_pending_effects_split(must_run: &mut Vec<EffectId>, skippable: &mut Vec<EffectId>) {
let mut pending = PENDING_EFFECTS.write();
for effect in pending.drain(..) {
if effect.is_skippable() {
skippable.push(effect);
} else {
must_run.push(effect);
}
}
}
pub fn remove_from_pending_set(effect_id: EffectId) {
PENDING_EFFECTS.write().swap_remove(&effect_id);
}
pub fn mark_effect_pending_for_check(effect_id: EffectId) {
cov_mark::hit!(check_effect_added_to_pending);
PENDING_EFFECTS.write().insert(effect_id);
}
pub fn mark_effect_dirty(effect_id: EffectId) -> bool {
let was_needs_work = effect_id.needs_work();
if was_needs_work {
effect_id.upgrade_state(ReactiveState::Dirty);
} else {
effect_id.set_state(ReactiveState::Dirty);
destroy_children(effect_id);
PENDING_EFFECTS.write().insert(effect_id);
}
!was_needs_work
}
use crate::effect::run_single_effect;
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn stale_access_returns_none() {
let metadata = EffectMetadata::new();
let id = effect_arena_insert(metadata);
effect_arena_remove(id);
assert!(id.with_sources(|_| ()).is_none());
assert_eq!(id.state(), ReactiveState::Clean); assert_eq!(id.add_source(SignalId::new(1)), None);
}
#[test]
fn effect_callback_restored_on_panic() {
use std::sync::Arc;
use std::sync::atomic::{AtomicI32, Ordering};
let run_count = Arc::new(AtomicI32::new(0));
let run_count_clone = run_count.clone();
let callback = Box::new(move || {
let count = run_count_clone.fetch_add(1, Ordering::Relaxed);
if count == 0 {
panic!("test panic in callback");
}
});
let metadata = EffectMetadata::new_with_callback(callback);
let id = effect_arena_insert(metadata);
assert!(id.has_callback());
assert_eq!(run_count.load(Ordering::Relaxed), 0);
let result = std::panic::catch_unwind(|| {
id.run_callback();
});
assert!(result.is_err());
assert_eq!(run_count.load(Ordering::Relaxed), 1);
assert!(id.has_callback());
id.run_callback();
assert_eq!(run_count.load(Ordering::Relaxed), 2);
effect_arena_remove(id);
}
#[test]
fn current_effect_guard_restores_on_panic() {
let effect1 = EffectId::new(10);
let effect2 = EffectId::new(20);
set_current_effect(Some(effect1));
assert_eq!(current_effect(), Some(effect1));
let result = std::panic::catch_unwind(|| {
let _guard = CurrentEffectGuard::new(Some(effect2));
assert_eq!(current_effect(), Some(effect2));
panic!("test panic");
});
assert!(result.is_err());
assert_eq!(current_effect(), Some(effect1));
set_current_effect(None);
}
}