use std::task::Poll;
use generational_box::{Owner, SyncStorage};
use super::{Hook, Hooks};
use crate::{ReactiveHandle, ReactiveMutNoUpdate, ReactiveMutRef, ReactiveRef, SingleWaker};
mod private {
pub trait Sealed {}
impl Sealed for crate::hooks::Hooks<'_, '_> {}
}
pub type State<T> = ReactiveHandle<T, SingleWaker>;
pub type StateRef<'a, T> = ReactiveRef<'a, T, SingleWaker>;
pub type StateMutRef<'a, T> = ReactiveMutRef<'a, T, SingleWaker>;
pub type StateMutNoUpdate<'a, T> = ReactiveMutNoUpdate<'a, T, SingleWaker>;
pub trait UseState: private::Sealed {
fn use_state<T, F>(&mut self, init: F) -> State<T>
where
F: FnOnce() -> T,
T: Unpin + Send + Sync + 'static;
}
struct UseStateImpl<T>
where
T: Unpin + Send + Sync + 'static,
{
state: State<T>,
_storage: Owner<SyncStorage>,
}
impl<T> UseStateImpl<T>
where
T: Unpin + Send + Sync + 'static,
{
pub fn new(initial_value: T) -> Self {
let storage = Owner::default();
UseStateImpl {
state: State::new_in(&storage, initial_value),
_storage: storage,
}
}
}
impl<T> Hook for UseStateImpl<T>
where
T: Unpin + Send + Sync + 'static,
{
fn poll_change(&mut self, cx: &mut std::task::Context) -> std::task::Poll<()> {
self.state.poll_change(None, cx)
}
}
impl UseState for Hooks<'_, '_> {
fn use_state<T, F>(&mut self, init: F) -> State<T>
where
F: FnOnce() -> T,
T: Unpin + Send + Sync + 'static,
{
self.use_hook(move || UseStateImpl::new(init())).state
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn add_and_sub_assign_mutate_value() {
let holder = UseStateImpl::new(0i32);
let mut state = holder.state;
state += 5;
assert_eq!(state.get(), 5);
state -= 2;
assert_eq!(state.get(), 3);
}
#[test]
fn mul_assign_mutates_value() {
let holder = UseStateImpl::new(3i32);
let mut state = holder.state;
state *= 4;
assert_eq!(state.get(), 12);
}
#[test]
fn set_overwrites_and_get_reads() {
let holder = UseStateImpl::new(10i32);
let mut state = holder.state;
state.set(99);
assert_eq!(state.get(), 99);
}
#[test]
fn copy_handles_share_storage() {
let holder = UseStateImpl::new(1i32);
let mut state = holder.state;
let state2 = state;
state += 41;
assert_eq!(state.get(), 42);
assert_eq!(state2.get(), 42);
}
}