#[cfg(feature = "trace-error")]
use crate::trace_error;
use crate::{
Context, DebugStore, Effects, MichiuError, MichiuTrace, OptionTraceExt, MichiuTaskSender,
with_context,
};
use slotmap::new_key_type;
use smallvec::SmallVec;
use std::cell::Cell;
use std::marker::PhantomData;
new_key_type! {
pub struct SignalId;
pub struct EffectId;
}
thread_local! {
pub(crate) static ACTIVE_EFFECT: Cell<Option<EffectId>> = const { Cell::new(None) };
pub(crate) static ACTIVE_ELEMENT: Cell<Option<crate::EntityId>> = const { Cell::new(None) };
}
pub struct ActiveElementGuard {
prev: Option<crate::EntityId>,
}
impl ActiveElementGuard {
#[inline]
#[must_use]
pub fn new(id: crate::EntityId) -> Self {
let prev = ACTIVE_ELEMENT.with(|cell| {
let prev = cell.get();
cell.set(Some(id));
prev
});
Self { prev }
}
}
impl Drop for ActiveElementGuard {
#[inline]
fn drop(&mut self) {
ACTIVE_ELEMENT.with(|cell| cell.set(self.prev));
}
}
#[derive(Debug, PartialEq, Eq)]
pub struct ReadSignal<T> {
pub(crate) id: SignalId,
pub(crate) _marker: PhantomData<T>,
}
impl<T> Clone for ReadSignal<T> {
fn clone(&self) -> Self {
*self
}
}
impl<T> Copy for ReadSignal<T> {}
impl<T> ReadSignal<T> {
#[must_use]
pub fn new(id: SignalId) -> Self {
Self {
id,
_marker: PhantomData,
}
}
}
impl<T: 'static> ReadSignal<T> {
#[inline]
fn track(self) {
ACTIVE_EFFECT.with(|cell| {
if let Some(active_effect_id) = cell.get() {
with_context(|cx| {
if let Some(subs) = cx.reactive.react_subscribers.get_mut(self.id) {
if !subs.contains(&active_effect_id) {
subs.push(active_effect_id);
}
} else {
let mut subs: SmallVec<[EffectId; 8]> = SmallVec::new();
subs.push(active_effect_id);
cx.reactive.react_subscribers.insert(self.id, subs);
}
});
}
});
}
#[track_caller]
#[inline]
pub(crate) fn with<U>(self, f: impl FnOnce(&T) -> U) -> U {
self.track();
with_context(|cx| {
let any_val =
cx.reactive
.react_signals
.get(self.id)
.unwrap_or_trace(None, &mut cx.debug, || MichiuError::SignalNotFound {
id: self.id,
});
let val = any_val
.downcast_ref::<T>()
.unwrap_or_trace(None, &mut cx.debug, || MichiuError::DowncastFailed {
expected: std::any::type_name::<T>(),
});
f(val)
})
}
}
impl<T: Clone + 'static> ReadSignal<T> {
#[inline]
#[must_use]
pub fn id(&self) -> SignalId {
self.id
}
#[inline]
#[must_use]
pub fn get(&self) -> T {
self.with(std::clone::Clone::clone)
}
#[inline]
pub fn get_else_by<U: Clone + 'static, F>(
self,
cond_fn: F,
true_val: U,
false_val: U,
) -> impl Fn() -> U + 'static
where
F: Fn(&T) -> bool + 'static,
{
let sig = self;
move || {
let val = sig.get();
if cond_fn(&val) {
true_val.clone()
} else {
false_val.clone()
}
}
}
#[inline]
pub fn get_else_with_by<U: 'static, F, FT, FF>(
self,
cond_fn: F,
true_fn: FT,
false_fn: FF,
) -> impl Fn() -> U + 'static
where
F: Fn(&T) -> bool + 'static,
FT: Fn() -> U + 'static,
FF: Fn() -> U + 'static,
{
let sig = self;
move || {
let val = sig.get();
if cond_fn(&val) { true_fn() } else { false_fn() }
}
}
#[track_caller]
#[inline]
#[must_use]
pub fn get_untracked(&self) -> T {
with_context(|cx| {
let any_val =
cx.reactive
.react_signals
.get(self.id)
.unwrap_or_trace(None, &mut cx.debug, || MichiuError::SignalNotFound {
id: self.id,
});
any_val
.downcast_ref::<T>()
.cloned()
.unwrap_or_trace(None, &mut cx.debug, || MichiuError::DowncastFailed {
expected: std::any::type_name::<T>(),
})
})
}
pub fn map<U, F>(&self, map_fn: F) -> ReadSignal<U>
where
T: Send + Clone + 'static,
U: Send + Clone + PartialEq + 'static,
F: Fn(&T) -> U + Send + Sync + 'static,
{
let source_read = *self;
with_context(|cx| {
let initial_val = map_fn(&source_read.get());
let (read_u, write_u) = cx.create_signal(initial_val);
let write_u_clone = write_u;
create_effect(move |_| {
let s_val = source_read.get();
let u_val = map_fn(&s_val);
if read_u.get_untracked() != u_val {
write_u_clone.set(u_val);
}
});
read_u
})
}
pub fn bi_map<U, F, G>(
&self,
writer: WriteSignal<T>,
map_read: F, map_write: G, ) -> (ReadSignal<U>, WriteSignal<U>)
where
T: Send + Clone + PartialEq + 'static,
U: Send + Clone + PartialEq + 'static,
F: Fn(&T) -> U + Send + Sync + 'static,
G: Fn(U) -> T + Send + Sync + 'static,
{
let source_read = *self;
with_context(|cx| {
let initial_val = map_read(&source_read.get());
let (read_u, write_u) = cx.create_signal(initial_val);
let write_u_clone = write_u;
create_effect(move |_| {
let s_val = source_read.get();
let u_val = map_read(&s_val);
if read_u.get_untracked() != u_val {
write_u_clone.set(u_val);
}
});
let read_u_clone = read_u;
create_effect(move |_| {
let u_val = read_u_clone.get();
let s_val = map_write(u_val);
if source_read.get_untracked() != s_val {
writer.set(s_val);
}
});
(read_u, write_u)
})
}
}
impl ReadSignal<bool> {
#[inline]
pub fn get_else<U: Clone + 'static>(
self,
true_val: U,
false_val: U,
) -> impl Fn() -> U + 'static {
let sig = self;
move || {
if sig.get() {
true_val.clone()
} else {
false_val.clone()
}
}
}
#[inline]
pub fn get_else_with<U: 'static, FT, FF>(
self,
true_fn: FT,
false_fn: FF,
) -> impl Fn() -> U + 'static
where
FT: Fn() -> U + 'static,
FF: Fn() -> U + 'static,
{
let sig = self;
move || {
if sig.get() { true_fn() } else { false_fn() }
}
}
}
#[derive(Debug, PartialEq, Eq)]
pub struct WriteSignal<T> {
pub(crate) id: SignalId,
pub(crate) _marker: PhantomData<T>,
}
impl<T> Clone for WriteSignal<T> {
fn clone(&self) -> Self {
*self
}
}
impl<T> Copy for WriteSignal<T> {}
impl<T> WriteSignal<T> {
#[inline]
#[must_use]
pub fn new(id: SignalId) -> Self {
Self {
id,
_marker: PhantomData,
}
}
}
impl<T: Send + 'static> WriteSignal<T> {
#[inline]
#[must_use]
pub fn id(&self) -> SignalId {
self.id
}
pub fn set(&self, new_value: T) {
let mut effects_to_run = SmallVec::new();
with_context(|cx| {
*cx.reactive.react_signals.get_mut(self.id).unwrap_or_trace(
None,
&mut cx.debug,
|| MichiuError::SignalNotFound { id: self.id },
) = Box::new(new_value);
if let Some(subs) = cx.reactive.react_subscribers.get(self.id) {
effects_to_run.clone_from(subs);
}
});
for effect_id in effects_to_run {
let exists = with_context(|cx| cx.reactive.react_effects.contains_key(effect_id));
if exists {
execute_effect(effect_id);
}
}
}
#[inline]
#[must_use]
pub fn sender(&self) -> SignalSender<T> {
let sender = with_context(|cx| cx.task_sender());
SignalSender {
id: self.id,
sys_task_sender: sender,
_marker: PhantomData,
}
}
#[inline]
#[must_use]
pub fn sender_with(&self, cx: &Context) -> SignalSender<T> {
SignalSender {
id: self.id,
sys_task_sender: cx.task_sender(),
_marker: PhantomData,
}
}
}
pub struct SignalSender<T> {
pub(crate) id: SignalId,
pub(crate) sys_task_sender: MichiuTaskSender,
pub(crate) _marker: PhantomData<T>,
}
impl<T> Clone for SignalSender<T> {
fn clone(&self) -> Self {
Self {
id: self.id,
sys_task_sender: self.sys_task_sender.clone(),
_marker: PhantomData,
}
}
}
impl<T: Send + 'static> SignalSender<T> {
#[inline]
pub fn send(&self, value: T) {
let signal_id = self.id;
let _ = self.sys_task_sender.send(move |_cx| {
let write_signal = WriteSignal::<T> {
id: signal_id,
_marker: PhantomData,
};
write_signal.set(value);
});
}
}
#[track_caller]
#[inline]
pub(crate) fn execute_effect(effect_id: EffectId) {
with_context(|cx| {
let dummy = Effects(Box::new(move |cx| {
#[cfg(feature = "trace-error")]
trace_error!(None, &mut cx.debug, || MichiuTrace::Error {
detail: MichiuError::RecursiveEffectDetected { effect_id },
add: None,
});
}));
let slot = cx
.reactive
.react_effects
.get_mut(effect_id)
.unwrap_or_trace(None, &mut cx.debug, || MichiuError::EffectNotFound {
id: effect_id,
});
let mut effect_closure = std::mem::replace(slot, dummy);
let prev_effect = ACTIVE_EFFECT.with(|cell| {
let prev = cell.get();
cell.set(Some(effect_id));
prev
});
effect_closure.0(cx);
ACTIVE_EFFECT.with(|cell| cell.set(prev_effect));
if let Some(slot) = cx.reactive.react_effects.get_mut(effect_id) {
*slot = effect_closure;
}
});
}
#[inline]
pub(crate) fn create_effect<F>(f: F) -> EffectId
where
F: FnMut(&mut Context) + 'static,
{
let id = with_context(|cx| cx.reactive.react_effects.insert(Effects(Box::new(f))));
execute_effect(id);
id
}