use std::cell::Cell;
use std::future::Future;
use std::rc::Rc;
use std::sync::{Arc, Mutex};
use reactive_graph::owner::{Owner, on_cleanup};
use reactive_graph::signal::RwSignal;
use reactive_graph::spawn_local_scoped_with_cancellation;
use reactive_graph::traits::{Set, Update};
use tokio::task::AbortHandle;
use crate::runtime::ReactiveRuntime;
pub type TaskError = Arc<dyn std::error::Error + Send + Sync>;
#[derive(Default)]
pub enum AsyncValue<T> {
#[default]
Idle,
Loading(Option<T>),
Ready(T),
Error(TaskError),
}
impl<T> AsyncValue<T> {
pub fn is_idle(&self) -> bool {
matches!(self, AsyncValue::Idle)
}
pub fn is_loading(&self) -> bool {
matches!(self, AsyncValue::Loading(_))
}
pub fn is_ready(&self) -> bool {
matches!(self, AsyncValue::Ready(_))
}
pub fn is_error(&self) -> bool {
matches!(self, AsyncValue::Error(_))
}
pub fn ready(&self) -> Option<&T> {
match self {
AsyncValue::Ready(t) => Some(t),
_ => None,
}
}
pub fn value(&self) -> Option<&T> {
match self {
AsyncValue::Ready(t) => Some(t),
AsyncValue::Loading(prev) => prev.as_ref(),
_ => None,
}
}
pub fn error(&self) -> Option<&TaskError> {
match self {
AsyncValue::Error(e) => Some(e),
_ => None,
}
}
}
impl<T: Clone> Clone for AsyncValue<T> {
fn clone(&self) -> Self {
match self {
AsyncValue::Idle => AsyncValue::Idle,
AsyncValue::Loading(prev) => AsyncValue::Loading(prev.clone()),
AsyncValue::Ready(t) => AsyncValue::Ready(t.clone()),
AsyncValue::Error(e) => AsyncValue::Error(e.clone()),
}
}
}
impl<T: std::fmt::Debug> std::fmt::Debug for AsyncValue<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
AsyncValue::Idle => f.write_str("Idle"),
AsyncValue::Loading(prev) => f.debug_tuple("Loading").field(prev).finish(),
AsyncValue::Ready(t) => f.debug_tuple("Ready").field(t).finish(),
AsyncValue::Error(e) => f.debug_tuple("Error").field(e).finish(),
}
}
}
pub struct UseTask<T: Send + Sync + 'static> {
signal: RwSignal<AsyncValue<T>>,
run: Rc<dyn Fn()>,
}
impl<T: Send + Sync + 'static> UseTask<T> {
pub fn signal(&self) -> RwSignal<AsyncValue<T>> {
self.signal
}
pub fn restart(&self) {
(self.run)();
}
}
impl<T: Send + Sync + 'static> Clone for UseTask<T> {
fn clone(&self) -> Self {
UseTask {
signal: self.signal,
run: self.run.clone(),
}
}
}
struct Coordinator<T: Send + Sync + 'static> {
signal: RwSignal<AsyncValue<T>>,
generation: Cell<u64>,
bg_abort: Arc<Mutex<Option<AbortHandle>>>,
}
pub fn use_task<T, E, Fut, F>(fetch: F) -> UseTask<T>
where
T: Send + Sync + 'static,
E: std::error::Error + Send + Sync + 'static,
Fut: Future<Output = Result<T, E>> + Send + 'static,
F: Fn() -> Fut + 'static,
{
let coord = Rc::new(Coordinator {
signal: RwSignal::new(AsyncValue::Idle),
generation: Cell::new(0),
bg_abort: Arc::new(Mutex::new(None)),
});
{
let bg_abort = coord.bg_abort.clone();
on_cleanup(move || {
if let Some(handle) = bg_abort.lock().expect("bg_abort poisoned").take() {
handle.abort();
}
});
}
let signal = coord.signal;
let fetch = Rc::new(fetch);
let owner = Owner::current();
let run: Rc<dyn Fn()> = {
let coord = coord.clone();
let fetch = fetch.clone();
Rc::new(move || {
let go = || run_once(&coord, &fetch);
match &owner {
Some(owner) => owner.with(go),
None => go(),
}
})
};
run();
UseTask { signal, run }
}
fn run_once<T, E, Fut, F>(coord: &Rc<Coordinator<T>>, fetch: &Rc<F>)
where
T: Send + Sync + 'static,
E: std::error::Error + Send + Sync + 'static,
Fut: Future<Output = Result<T, E>> + Send + 'static,
F: Fn() -> Fut + 'static,
{
let rt = ReactiveRuntime::get().expect(
"frust-reactive: use_task called before ReactiveRuntime::init — \
this is a wiring bug: initialize the reactive runtime (the shell does \
this on startup) before mounting components that use use_task",
);
if let Some(handle) = coord.bg_abort.lock().expect("bg_abort poisoned").take() {
handle.abort();
}
let generation = coord.generation.get().wrapping_add(1);
coord.generation.set(generation);
coord.signal.update(|state| {
let prev = match std::mem::take(state) {
AsyncValue::Ready(t) => Some(t),
AsyncValue::Loading(prev) => prev,
_ => None,
};
*state = AsyncValue::Loading(prev);
});
let join = rt.handle().spawn((fetch)());
*coord.bg_abort.lock().expect("bg_abort poisoned") = Some(join.abort_handle());
let coord = coord.clone();
spawn_local_scoped_with_cancellation(async move {
let outcome = join.await;
if coord.generation.get() != generation {
return;
}
match outcome {
Ok(Ok(value)) => coord.signal.set(AsyncValue::Ready(value)),
Ok(Err(err)) => {
let err: TaskError = Arc::new(err);
coord.signal.set(AsyncValue::Error(err));
}
Err(_join_err) => {}
}
});
}
#[cfg(test)]
mod tests {
use super::*;
use reactive_graph::owner::Owner;
use reactive_graph::traits::GetUntracked;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::mpsc;
use std::time::{Duration, Instant};
use crate::ReactiveRuntime;
fn noop_waker() -> crate::FrameWaker {
Arc::new(|| {})
}
fn init_rt() -> &'static ReactiveRuntime {
ReactiveRuntime::init(noop_waker())
}
fn pump_until(rt: &ReactiveRuntime, timeout: Duration, mut cond: impl FnMut() -> bool) -> bool {
let start = Instant::now();
loop {
rt.pump_local();
if cond() {
return true;
}
if start.elapsed() >= timeout {
return false;
}
std::thread::sleep(Duration::from_millis(1));
}
}
fn pump_a_few(rt: &ReactiveRuntime) {
for _ in 0..4 {
rt.pump_local();
}
}
#[derive(Debug)]
struct TestError(&'static str);
impl std::fmt::Display for TestError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.0)
}
}
impl std::error::Error for TestError {}
#[test]
fn use_task_resolves_to_ready() {
let _guard = crate::WAKER_TEST_LOCK
.lock()
.unwrap_or_else(|e| e.into_inner());
let rt = init_rt();
let owner = Owner::new();
let task = owner.with(|| use_task(|| async { crate::spawn_blocking(|| 6 * 7).await }));
assert!(
task.signal().get_untracked().is_loading(),
"state must be Loading immediately after use_task"
);
assert!(
pump_until(rt, Duration::from_secs(5), || task
.signal()
.get_untracked()
.is_ready()),
"task should reach Ready after pumping"
);
assert_eq!(task.signal().get_untracked().ready().copied(), Some(42));
owner.cleanup();
}
#[test]
fn use_task_reports_error() {
let _guard = crate::WAKER_TEST_LOCK
.lock()
.unwrap_or_else(|e| e.into_inner());
let rt = init_rt();
let owner = Owner::new();
let task = owner.with(|| {
use_task(|| async {
crate::spawn_blocking(|| ()).await.expect("join");
Result::<i32, TestError>::Err(TestError("boom"))
})
});
assert!(
pump_until(rt, Duration::from_secs(5), || task
.signal()
.get_untracked()
.is_error()),
"task should reach Error after pumping"
);
assert_eq!(
task.signal().get_untracked().error().map(|e| e.to_string()),
Some("boom".to_string())
);
owner.cleanup();
}
#[test]
fn restart_carries_previous_value_in_loading() {
let _guard = crate::WAKER_TEST_LOCK
.lock()
.unwrap_or_else(|e| e.into_inner());
let rt = init_rt();
let seq = Arc::new(AtomicUsize::new(0));
let owner = Owner::new();
let task = {
let seq = seq.clone();
owner.with(|| {
use_task(move || {
let seq = seq.clone();
async move {
crate::spawn_blocking(move || seq.fetch_add(1, Ordering::SeqCst)).await
}
})
})
};
assert!(pump_until(rt, Duration::from_secs(5), || task
.signal()
.get_untracked()
.is_ready()));
assert_eq!(task.signal().get_untracked().ready().copied(), Some(0));
task.restart();
assert_eq!(
task.signal().get_untracked().value().copied(),
Some(0),
"Loading must carry the previous Ready value for flicker-free refresh"
);
assert!(task.signal().get_untracked().is_loading());
owner.cleanup();
}
#[test]
fn unmount_while_loading_never_writes_disposed_signal() {
let _guard = crate::WAKER_TEST_LOCK
.lock()
.unwrap_or_else(|e| e.into_inner());
let rt = init_rt();
for i in 0..1_000u32 {
let ran = Arc::new(AtomicUsize::new(0));
let owner = Owner::new();
let task = {
let ran = ran.clone();
owner.with(|| {
use_task(move || {
let ran = ran.clone();
async move {
crate::spawn_blocking(move || {
ran.fetch_add(1, Ordering::SeqCst);
i
})
.await
}
})
})
};
owner.cleanup();
drop(task);
pump_a_few(rt);
}
}
#[test]
fn task_completing_after_teardown_is_safe() {
let _guard = crate::WAKER_TEST_LOCK
.lock()
.unwrap_or_else(|e| e.into_inner());
let rt = init_rt();
for _ in 0..1_000u32 {
let (tx, rx) = mpsc::channel::<()>();
let rx = Arc::new(Mutex::new(rx));
let owner = Owner::new();
let task = {
let rx = rx.clone();
owner.with(|| {
use_task(move || {
let rx = rx.clone();
async move {
crate::spawn_blocking(move || {
let _ = rx.lock().expect("rx").recv();
1u32
})
.await
}
})
})
};
owner.cleanup();
drop(task);
let _ = tx.send(());
pump_a_few(rt);
}
}
#[test]
fn restart_storm_settles_on_latest() {
let _guard = crate::WAKER_TEST_LOCK
.lock()
.unwrap_or_else(|e| e.into_inner());
let rt = init_rt();
let counter = Arc::new(AtomicUsize::new(0));
let owner = Owner::new();
let task = {
let counter = counter.clone();
owner.with(|| {
use_task(move || {
let seq = counter.fetch_add(1, Ordering::SeqCst);
async move { crate::spawn_blocking(move || seq).await }
})
})
};
for _ in 0..1_000u32 {
task.restart();
}
let last_seq = counter.load(Ordering::SeqCst) - 1;
assert!(
pump_until(rt, Duration::from_secs(10), || {
matches!(
task.signal().get_untracked().ready().copied(),
Some(seq) if seq == last_seq
)
}),
"restart storm must settle on the newest fetch's value (last-write-wins)"
);
owner.cleanup();
}
#[test]
fn out_of_order_completion_respects_generation() {
let _guard = crate::WAKER_TEST_LOCK
.lock()
.unwrap_or_else(|e| e.into_inner());
let rt = init_rt();
for _ in 0..1_000u32 {
let seq = Arc::new(AtomicUsize::new(0));
let (release_tx, release_rx) = mpsc::channel::<()>();
let release_rx = Arc::new(Mutex::new(release_rx));
let owner = Owner::new();
let task = {
let seq = seq.clone();
let release_rx = release_rx.clone();
owner.with(|| {
use_task(move || {
let n = seq.fetch_add(1, Ordering::SeqCst);
let release_rx = release_rx.clone();
async move {
crate::spawn_blocking(move || {
if n == 0 {
let _ = release_rx.lock().expect("rx").recv();
}
n
})
.await
}
})
})
};
task.restart();
assert!(
pump_until(rt, Duration::from_secs(5), || matches!(
task.signal().get_untracked().ready().copied(),
Some(1)
)),
"the newer fetch must resolve to 1"
);
let _ = release_tx.send(());
pump_a_few(rt);
assert_eq!(
task.signal().get_untracked().ready().copied(),
Some(1),
"a stale, later-completing fetch must not clobber the newer result"
);
owner.cleanup();
}
}
}