use std::cell::RefCell;
use crate::{ReadSignal, Root, create_empty_signal, create_signal};
#[cfg_attr(debug_assertions, track_caller)]
pub fn create_selector_with<T>(
mut f: impl FnMut() -> T + 'static,
mut eq: impl FnMut(&T, &T) -> bool + 'static,
) -> ReadSignal<T> {
let root = Root::global();
let signal = create_empty_signal();
let prev = root.current_node.replace(signal.id);
let (initial, tracker) = root.tracked_scope(&mut f);
root.current_node.set(prev);
tracker.create_dependency_link(root, signal.id);
let mut signal_mut = signal.get_mut();
signal_mut.value = Some(Box::new(initial));
signal_mut.callback = Some(Box::new(move |value| {
let value = value.downcast_mut().expect("wrong memo type");
let new = f();
if eq(&new, value) {
false
} else {
*value = new;
true
}
}));
*signal
}
#[cfg_attr(debug_assertions, track_caller)]
pub fn create_memo<T>(f: impl FnMut() -> T + 'static) -> ReadSignal<T> {
create_selector_with(f, |_, _| false)
}
#[cfg_attr(debug_assertions, track_caller)]
pub fn create_selector<T>(f: impl FnMut() -> T + 'static) -> ReadSignal<T>
where
T: PartialEq,
{
create_selector_with(f, PartialEq::eq)
}
#[cfg_attr(debug_assertions, track_caller)]
pub fn create_reducer<T, Msg>(
initial: T,
reduce: impl FnMut(&T, Msg) -> T,
) -> (ReadSignal<T>, impl Fn(Msg)) {
let reduce = RefCell::new(reduce);
let signal = create_signal(initial);
let dispatch = move |msg| signal.update(|value| *value = reduce.borrow_mut()(value, msg));
(*signal, dispatch)
}
#[cfg(test)]
mod tests {
use crate::*;
#[test]
fn memo() {
let _ = create_root(|| {
let state = create_signal(0);
let double = create_memo(move || state.get() * 2);
assert_eq!(double.get(), 0);
state.set(1);
assert_eq!(double.get(), 2);
state.set(2);
assert_eq!(double.get(), 4);
});
}
#[test]
fn memo_only_run_once() {
let _ = create_root(|| {
let state = create_signal(0);
let counter = create_signal(0);
let double = create_memo(move || {
counter.set_silent(counter.get_untracked() + 1);
state.get() * 2
});
assert_eq!(counter.get(), 1); state.set(2);
assert_eq!(counter.get(), 2);
assert_eq!(double.get(), 4);
assert_eq!(counter.get(), 2); });
}
#[test]
fn dependency_on_memo() {
let _ = create_root(|| {
let state = create_signal(0);
let double = create_memo(move || state.get() * 2);
let quadruple = create_memo(move || double.get() * 2);
assert_eq!(quadruple.get(), 0);
state.set(1);
assert_eq!(quadruple.get(), 4);
});
}
#[test]
fn untracked_memo() {
let _ = create_root(|| {
let state = create_signal(1);
let double = create_memo(move || state.get_untracked() * 2);
assert_eq!(double.get(), 2);
state.set(2);
assert_eq!(double.get(), 2); });
}
#[test]
fn memos_should_recreate_dependencies_each_time() {
let _ = create_root(|| {
let condition = create_signal(true);
let state1 = create_signal(0);
let state2 = create_signal(1);
let counter = create_signal(0);
create_memo(move || {
counter.set_silent(counter.get_untracked() + 1);
if condition.get() {
state1.track();
} else {
state2.track();
}
});
assert_eq!(counter.get(), 1);
state1.set(1);
assert_eq!(counter.get(), 2);
state2.set(1);
assert_eq!(counter.get(), 2);
condition.set(false);
assert_eq!(counter.get(), 3);
state1.set(2);
assert_eq!(counter.get(), 3);
state2.set(2);
assert_eq!(counter.get(), 4); });
}
#[test]
fn destroy_memos_on_scope_dispose() {
let _ = create_root(|| {
let counter = create_signal(0);
let trigger = create_signal(());
let child_scope = create_child_scope(move || {
let _ = create_memo(move || {
trigger.track();
counter.set_silent(counter.get_untracked() + 1);
});
});
assert_eq!(counter.get(), 1);
trigger.set(());
assert_eq!(counter.get(), 2);
child_scope.dispose();
trigger.set(());
assert_eq!(counter.get(), 2); });
}
#[test]
fn selector() {
let _ = create_root(|| {
let state = create_signal(0);
let double = create_selector(move || state.get() * 2);
let counter = create_signal(0);
create_effect(move || {
counter.set(counter.get_untracked() + 1);
double.track();
});
assert_eq!(double.get(), 0);
assert_eq!(counter.get(), 1);
state.set(0);
state.set(0);
state.set(0);
assert_eq!(double.get(), 0);
assert_eq!(counter.get(), 1);
state.set(2);
assert_eq!(double.get(), 4);
assert_eq!(counter.get(), 2);
});
}
#[test]
fn reducer() {
let _ = create_root(|| {
enum Msg {
Increment,
Decrement,
}
let (state, dispatch) = create_reducer(0, |state, msg: Msg| match msg {
Msg::Increment => *state + 1,
Msg::Decrement => *state - 1,
});
assert_eq!(state.get(), 0);
dispatch(Msg::Increment);
assert_eq!(state.get(), 1);
dispatch(Msg::Decrement);
assert_eq!(state.get(), 0);
dispatch(Msg::Increment);
dispatch(Msg::Increment);
assert_eq!(state.get(), 2);
});
}
#[test]
fn memo_reducer() {
let _ = create_root(|| {
enum Msg {
Increment,
Decrement,
}
let (state, dispatch) = create_reducer(0, |state, msg: Msg| match msg {
Msg::Increment => *state + 1,
Msg::Decrement => *state - 1,
});
let doubled = create_memo(move || state.get() * 2);
assert_eq!(doubled.get(), 0);
dispatch(Msg::Increment);
assert_eq!(doubled.get(), 2);
dispatch(Msg::Decrement);
assert_eq!(doubled.get(), 0);
});
}
}