use std::any::Any;
use std::collections::HashSet;
use std::error::Error;
use std::fmt;
use std::future::Future;
use std::marker::PhantomData;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use parking_lot::Mutex;
use tokio::sync::watch;
use tokio::task::JoinHandle;
use crate::context::{GraphNode, SlotId};
#[cfg(not(feature = "vec_edges"))]
type EdgeVec = smallvec::SmallVec<[SlotId; 4]>;
#[cfg(feature = "vec_edges")]
type EdgeVec = Vec<SlotId>;
fn edge_insert(edges: &mut EdgeVec, id: SlotId) -> bool {
if edges.contains(&id) {
false
} else {
edges.push(id);
true
}
}
fn edge_remove(edges: &mut EdgeVec, id: SlotId) -> bool {
if let Some(pos) = edges.iter().position(|x| *x == id) {
edges.swap_remove(pos);
true
} else {
false
}
}
type AsyncAny = dyn Any + Send + Sync;
type BoxedAsyncFuture = Pin<Box<dyn Future<Output = Arc<AsyncAny>> + Send>>;
type AsyncComputeFn = dyn Fn(AsyncComputeContext) -> BoxedAsyncFuture + Send + Sync;
type AsyncEqualsFn = dyn Fn(&AsyncAny, &AsyncAny) -> bool + Send + Sync;
type BoxedEffectFuture = Pin<Box<dyn Future<Output = Option<BoxedCleanupFn>> + Send>>;
type BoxedCleanupFn = Box<dyn FnOnce() + Send>;
type AsyncEffectFn = dyn Fn(AsyncComputeContext) -> BoxedEffectFuture + Send + Sync;
static NEXT_ASYNC_CONTEXT_ID: AtomicU64 = AtomicU64::new(0);
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct AsyncContextId(u64);
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum AsyncSlotStateView {
None,
Empty,
Computing { revision: u64 },
Resolved,
Error,
}
#[derive(Debug)]
pub enum AsyncSlotState {
Empty,
Computing {
revision: u64,
handle: JoinHandle<()>,
},
Resolved,
Error,
}
impl fmt::Display for AsyncSlotState {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Empty => write!(f, "Empty"),
Self::Computing { revision, .. } => {
write!(f, "Computing(revision={revision})")
}
Self::Resolved => write!(f, "Resolved"),
Self::Error => write!(f, "Error"),
}
}
}
#[allow(dead_code)]
pub(crate) enum TransitionOutcome {
Accepted,
Stale,
Unchanged,
}
#[allow(dead_code)]
pub(crate) enum InvalidationResult {
HadInFlight(JoinHandle<()>),
WasResolved,
WasError,
AlreadyEmpty,
}
pub(crate) struct AsyncSlotNode {
pub(crate) state: AsyncSlotState,
pub(crate) value: Option<Arc<AsyncAny>>,
pub(crate) error: Option<Arc<dyn Error + Send + Sync>>,
pub(crate) revision: u64,
pub(crate) compute: Arc<AsyncComputeFn>,
pub(crate) equals: Option<Arc<AsyncEqualsFn>>,
pub(crate) dependencies: EdgeVec,
pub(crate) dependents: EdgeVec,
pub(crate) notifier: Option<watch::Sender<AsyncCompletion>>,
}
#[derive(Clone)]
pub(crate) enum AsyncCompletion {
Pending,
Resolved(Arc<dyn Any + Send + Sync>),
Error(Arc<dyn Error + Send + Sync>),
}
impl AsyncSlotNode {
pub(crate) fn transition_to_computing(
&mut self,
handle: JoinHandle<()>,
) -> Option<JoinHandle<()>> {
let old = std::mem::replace(
&mut self.state,
AsyncSlotState::Computing {
revision: self.revision,
handle,
},
);
match old {
AsyncSlotState::Computing { handle, .. } => Some(handle),
_ => None,
}
}
pub(crate) fn transition_to_resolved(
&mut self,
revision: u64,
value: Arc<AsyncAny>,
) -> TransitionOutcome {
match &self.state {
AsyncSlotState::Computing {
revision: current_revision,
..
} if *current_revision == revision => {}
_ => return TransitionOutcome::Stale,
}
let is_new = match (&self.value, &self.equals) {
(Some(old), Some(eq)) => !eq(old.as_ref(), value.as_ref()),
_ => true,
};
self.state = AsyncSlotState::Resolved;
self.value = Some(value);
self.error = None;
if is_new {
TransitionOutcome::Accepted
} else {
TransitionOutcome::Unchanged
}
}
pub(crate) fn transition_to_error(
&mut self,
revision: u64,
error: Arc<dyn Error + Send + Sync>,
) -> TransitionOutcome {
match &self.state {
AsyncSlotState::Computing {
revision: current_revision,
..
} if *current_revision == revision => {}
_ => return TransitionOutcome::Stale,
}
self.state = AsyncSlotState::Error;
self.error = Some(error);
self.value = None;
TransitionOutcome::Accepted
}
pub(crate) fn invalidate(&mut self) -> InvalidationResult {
self.revision += 1;
match std::mem::replace(&mut self.state, AsyncSlotState::Empty) {
AsyncSlotState::Computing { handle, .. } => InvalidationResult::HadInFlight(handle),
AsyncSlotState::Resolved => InvalidationResult::WasResolved,
AsyncSlotState::Error => InvalidationResult::WasError,
AsyncSlotState::Empty => InvalidationResult::AlreadyEmpty,
}
}
pub(crate) fn clear(&mut self) -> Option<JoinHandle<()>> {
self.revision += 1;
self.value = None;
self.error = None;
match std::mem::replace(&mut self.state, AsyncSlotState::Empty) {
AsyncSlotState::Computing { handle, .. } => Some(handle),
_ => None,
}
}
}
pub(crate) struct AsyncCellNode {
pub(crate) value: Arc<AsyncAny>,
pub(crate) dependents: EdgeVec,
}
pub(crate) struct AsyncEffectNode {
pub(crate) effect_fn: Arc<AsyncEffectFn>,
pub(crate) cleanup: Option<BoxedCleanupFn>,
pub(crate) dependencies: EdgeVec,
pub(crate) dependents: EdgeVec,
pub(crate) force_run: bool,
pub(crate) in_flight: Option<JoinHandle<()>>,
}
#[allow(dead_code)]
pub(crate) enum AsyncNode {
Slot(AsyncSlotNode),
Cell(AsyncCellNode),
Effect(AsyncEffectNode),
}
pub(crate) struct AsyncContextInner {
pub(crate) nodes: Vec<Option<AsyncNode>>,
generations: Vec<u64>,
next_id: u64,
free_ids: Vec<u64>,
pub(crate) context_id: AsyncContextId,
batch_depth: usize,
batched_cells: HashSet<SlotId>,
pending_async_effects: Vec<SlotId>,
scheduled_async_effects: HashSet<SlotId>,
deps_pool: Vec<Arc<Mutex<HashSet<SlotId>>>>,
}
const DEPS_POOL_CAP: usize = 32;
impl AsyncContextInner {
fn take_deps(&mut self) -> Arc<Mutex<HashSet<SlotId>>> {
self.deps_pool
.pop()
.unwrap_or_else(|| Arc::new(Mutex::new(HashSet::new())))
}
fn recycle_deps(&mut self, mut deps: Arc<Mutex<HashSet<SlotId>>>, mut set: HashSet<SlotId>) {
if self.deps_pool.len() >= DEPS_POOL_CAP {
return;
}
let Some(tracker) = Arc::get_mut(&mut deps) else {
return;
};
set.clear();
*tracker.get_mut() = set;
self.deps_pool.push(deps);
}
pub(crate) fn alloc_id(&mut self) -> SlotId {
match self.free_ids.pop() {
Some(id) => SlotId(id),
None => {
let id = SlotId(self.next_id);
self.next_id += 1;
id
}
}
}
fn node_index(id: SlotId) -> Option<usize> {
usize::try_from(id.0).ok()
}
pub(crate) fn generation(&self, id: SlotId) -> u64 {
Self::node_index(id)
.and_then(|idx| self.generations.get(idx))
.copied()
.unwrap_or(0)
}
pub(crate) fn get_node(&self, id: SlotId) -> Option<&AsyncNode> {
Self::node_index(id)
.and_then(|idx| self.nodes.get(idx))
.and_then(|opt| opt.as_ref())
}
pub(crate) fn get_node_mut(&mut self, id: SlotId) -> Option<&mut AsyncNode> {
Self::node_index(id)
.and_then(|idx| self.nodes.get_mut(idx))
.and_then(|opt| opt.as_mut())
}
pub(crate) fn insert_node(&mut self, id: SlotId, node: AsyncNode) {
let index = Self::node_index(id).expect("SlotId does not fit usize");
if self.nodes.len() <= index {
self.nodes.resize_with(index + 1, || None);
}
if self.generations.len() <= index {
self.generations.resize(index + 1, 0);
}
self.nodes[index] = Some(node);
}
}
fn register_dependency_locked(
inner: &mut AsyncContextInner,
dependency_id: SlotId,
dependent_id: SlotId,
) {
if dependency_id == dependent_id {
return;
}
if let Some(node) = inner.get_node_mut(dependent_id) {
match node {
AsyncNode::Slot(s) => {
edge_insert(&mut s.dependencies, dependency_id);
}
AsyncNode::Effect(e) => {
edge_insert(&mut e.dependencies, dependency_id);
}
AsyncNode::Cell(_) => {}
}
}
if let Some(node) = inner.get_node_mut(dependency_id) {
match node {
AsyncNode::Slot(s) => {
edge_insert(&mut s.dependents, dependent_id);
}
AsyncNode::Cell(c) => {
edge_insert(&mut c.dependents, dependent_id);
}
AsyncNode::Effect(_) => {}
}
}
}
#[cfg(feature = "instrumentation")]
pub type Window1Hook = Arc<dyn Fn() -> Pin<Box<dyn Future<Output = ()> + Send>> + Send + Sync>;
impl<T> crate::context::sealed::Sealed for AsyncSlotHandle<T> {}
impl<T> GraphNode for AsyncSlotHandle<T> {
fn node_id(&self) -> SlotId {
self.id
}
}
impl<T> crate::context::sealed::Sealed for AsyncCellHandle<T> {}
impl<T> GraphNode for AsyncCellHandle<T> {
fn node_id(&self) -> SlotId {
self.id
}
}
impl crate::context::sealed::Sealed for AsyncEffectHandle {}
impl GraphNode for AsyncEffectHandle {
fn node_id(&self) -> SlotId {
self.id
}
}
pub struct AsyncTeardownScope {
ctx: AsyncContext,
owned: Mutex<Vec<SlotId>>,
}
impl AsyncTeardownScope {
pub fn computed_async<T, F, Fut>(&self, compute: F) -> AsyncSlotHandle<T>
where
T: Clone + Send + Sync + 'static,
F: Fn(AsyncComputeContext) -> Fut + Send + Sync + 'static,
Fut: Future<Output = T> + Send + 'static,
{
let handle = self.ctx.computed_async(compute);
self.owned.lock().push(handle.id);
handle
}
pub fn memo_async<T, F, Fut>(&self, compute: F) -> AsyncSlotHandle<T>
where
T: PartialEq + Clone + Send + Sync + 'static,
F: Fn(AsyncComputeContext) -> Fut + Send + Sync + 'static,
Fut: Future<Output = T> + Send + 'static,
{
let handle = self.ctx.memo_async(compute);
self.owned.lock().push(handle.id);
handle
}
pub fn cell<T>(&self, value: T) -> AsyncCellHandle<T>
where
T: PartialEq + Clone + Send + Sync + 'static,
{
let handle = self.ctx.cell(value);
self.owned.lock().push(handle.id);
handle
}
pub fn effect_async<F, Fut, C>(&self, effect: F) -> AsyncEffectHandle
where
F: Fn(AsyncComputeContext) -> Fut + Send + Sync + 'static,
Fut: Future<Output = Option<C>> + Send + 'static,
C: FnOnce() + Send + 'static,
{
let handle = self.ctx.effect_async(effect);
self.owned.lock().push(handle.id);
handle
}
pub fn context(&self) -> &AsyncContext {
&self.ctx
}
pub fn len(&self) -> usize {
self.owned.lock().len()
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
pub fn disarm(self) {
self.owned.lock().clear();
}
}
impl Drop for AsyncTeardownScope {
fn drop(&mut self) {
let owned = std::mem::take(&mut *self.owned.lock());
for id in owned.into_iter().rev() {
self.ctx.dispose_id(id);
}
}
}
pub struct AsyncContext {
inner: Arc<Mutex<AsyncContextInner>>,
#[cfg(feature = "instrumentation")]
window1_hook: Mutex<Option<Window1Hook>>,
#[cfg(feature = "instrumentation")]
window1_resolved_hits: AtomicU64,
}
pub struct AsyncSlotHandle<T> {
pub(crate) id: SlotId,
pub(crate) _marker: PhantomData<T>,
}
impl<T> Clone for AsyncSlotHandle<T> {
fn clone(&self) -> Self {
*self
}
}
impl<T> Copy for AsyncSlotHandle<T> {}
pub struct AsyncCellHandle<T> {
pub(crate) id: SlotId,
pub(crate) _marker: PhantomData<T>,
}
impl<T> Clone for AsyncCellHandle<T> {
fn clone(&self) -> Self {
*self
}
}
impl<T> Copy for AsyncCellHandle<T> {}
pub struct AsyncEffectHandle {
pub(crate) id: SlotId,
}
impl Clone for AsyncEffectHandle {
fn clone(&self) -> Self {
*self
}
}
impl Copy for AsyncEffectHandle {}
pub struct AsyncSignalHandle<T> {
pub(crate) slot: AsyncSlotHandle<T>,
pub(crate) effect: AsyncEffectHandle,
}
impl<T> AsyncSignalHandle<T> {
pub fn get(&self, ctx: &AsyncContext) -> Option<T>
where
T: Clone + Send + Sync + 'static,
{
ctx.get_signal(self)
}
pub async fn get_async(&self, ctx: &AsyncContext) -> T
where
T: Clone + Send + Sync + 'static,
{
ctx.get_signal_async(self).await
}
pub fn dispose(&self, ctx: &AsyncContext) {
ctx.dispose_signal(self);
}
pub fn is_active(&self, ctx: &AsyncContext) -> bool {
ctx.is_signal_active(self)
}
}
impl<T> Clone for AsyncSignalHandle<T> {
fn clone(&self) -> Self {
*self
}
}
impl<T> Copy for AsyncSignalHandle<T> {}
pub struct AsyncComputeContext {
pub(crate) _context_id: AsyncContextId,
pub(crate) _node_id: SlotId,
pub(crate) _node_gen: u64,
pub(crate) inner: Arc<Mutex<AsyncContextInner>>,
pub(crate) dependencies: Arc<Mutex<HashSet<SlotId>>>,
}
impl AsyncComputeContext {
pub fn get_cell<T>(&self, handle: &AsyncCellHandle<T>) -> T
where
T: Clone + Send + Sync + 'static,
{
self.dependencies.lock().insert(handle.id);
let mut inner = self.inner.lock();
if inner.generation(self._node_id) == self._node_gen {
register_dependency_locked(&mut inner, handle.id, self._node_id);
}
match inner.get_node(handle.id) {
Some(AsyncNode::Cell(cell)) => cell
.value
.as_ref()
.downcast_ref::<T>()
.expect("type mismatch in async compute get_cell")
.clone(),
_ => panic!("AsyncCellHandle does not point to a Cell node"),
}
}
pub fn get_async<T>(
&self,
handle: &AsyncSlotHandle<T>,
) -> impl Future<Output = T> + Send + use<T>
where
T: Clone + Send + Sync + 'static,
{
self.dependencies.lock().insert(handle.id);
{
let mut inner = self.inner.lock();
if inner.generation(self._node_id) == self._node_gen {
register_dependency_locked(&mut inner, handle.id, self._node_id);
}
}
let inner_arc = self.inner.clone();
let handle = *handle;
async move {
let ctx = AsyncContext {
inner: inner_arc,
#[cfg(feature = "instrumentation")]
window1_hook: Mutex::new(None),
#[cfg(feature = "instrumentation")]
window1_resolved_hits: AtomicU64::new(0),
};
ctx.get_async(&handle).await
}
}
pub fn get_signal_async<T>(
&self,
handle: &AsyncSignalHandle<T>,
) -> impl Future<Output = T> + Send + use<T>
where
T: Clone + Send + Sync + 'static,
{
self.get_async(&handle.slot)
}
}
fn spawn_async_compute(ctx: &AsyncContext, slot_id: SlotId) -> watch::Receiver<AsyncCompletion> {
let inner_arc: Arc<Mutex<AsyncContextInner>> = ctx.inner.clone();
let mut inner = inner_arc.lock();
let (compute, context_id, spawn_revision) = match inner.get_node(slot_id) {
Some(AsyncNode::Slot(slot)) => {
if let AsyncSlotState::Computing { .. } = &slot.state {
return slot
.notifier
.as_ref()
.expect("computing without notifier")
.subscribe();
}
(slot.compute.clone(), inner.context_id, slot.revision)
}
_ => panic!("spawn_async_compute: not a slot node"),
};
let (tx, rx) = watch::channel(AsyncCompletion::Pending);
let inner_for_compute = inner_arc.clone();
let tx_clone = tx.clone();
let deps_arc = inner.take_deps();
let deps_for_extract = deps_arc.clone();
let slot_gen = inner.generation(slot_id);
let join_handle = tokio::spawn(async move {
let compute_ctx = AsyncComputeContext {
_context_id: context_id,
_node_id: slot_id,
_node_gen: slot_gen,
inner: inner_for_compute.clone(),
dependencies: deps_arc,
};
let result = compute(compute_ctx).await;
let deps = std::mem::take(&mut *deps_for_extract.lock());
{
let mut inner = inner_for_compute.lock();
let current_revision = match inner.get_node(slot_id) {
Some(AsyncNode::Slot(s)) => s.revision,
_ => {
inner.recycle_deps(deps_for_extract, deps);
drop(inner);
let _ = tx_clone.send(AsyncCompletion::Error(Arc::new(std::io::Error::other(
"slot node removed during compute",
))));
return;
}
};
if current_revision != spawn_revision {
inner.recycle_deps(deps_for_extract, deps);
return;
}
AsyncContext::update_dependencies(&mut inner, slot_id, &deps);
inner.recycle_deps(deps_for_extract, deps);
if let Some(AsyncNode::Slot(slot)) = inner.get_node_mut(slot_id) {
slot.transition_to_resolved(spawn_revision, result.clone());
}
}
let _ = tx_clone.send(AsyncCompletion::Resolved(result));
});
if let Some(AsyncNode::Slot(slot)) = inner.get_node_mut(slot_id) {
slot.transition_to_computing(join_handle);
slot.notifier = Some(tx);
}
rx
}
impl Default for AsyncContext {
fn default() -> Self {
Self::new()
}
}
impl AsyncContext {
pub fn new() -> Self {
let context_id = AsyncContextId(NEXT_ASYNC_CONTEXT_ID.fetch_add(1, Ordering::Relaxed));
Self {
inner: Arc::new(Mutex::new(AsyncContextInner {
nodes: Vec::new(),
generations: Vec::new(),
next_id: 0,
free_ids: Vec::new(),
context_id,
batch_depth: 0,
batched_cells: HashSet::new(),
pending_async_effects: Vec::new(),
scheduled_async_effects: HashSet::new(),
deps_pool: Vec::new(),
})),
#[cfg(feature = "instrumentation")]
window1_hook: Mutex::new(None),
#[cfg(feature = "instrumentation")]
window1_resolved_hits: AtomicU64::new(0),
}
}
#[cfg(feature = "instrumentation")]
pub fn __install_window1_hook(&self, hook: Window1Hook) {
*self.window1_hook.lock() = Some(hook);
}
#[cfg(feature = "instrumentation")]
pub fn __window1_resolved_hits(&self) -> u64 {
self.window1_resolved_hits.load(Ordering::Relaxed)
}
pub fn cell<T>(&self, value: T) -> AsyncCellHandle<T>
where
T: PartialEq + Clone + Send + Sync + 'static,
{
let id;
{
let mut inner = self.inner.lock();
id = inner.alloc_id();
let node = AsyncCellNode {
value: Arc::new(value),
dependents: EdgeVec::new(),
};
inner.insert_node(id, AsyncNode::Cell(node));
}
AsyncCellHandle {
id,
_marker: PhantomData,
}
}
pub fn get_cell<T>(&self, handle: &AsyncCellHandle<T>) -> T
where
T: Clone + Send + Sync + 'static,
{
let inner = self.inner.lock();
match inner.get_node(handle.id) {
Some(AsyncNode::Cell(cell)) => cell
.value
.as_ref()
.downcast_ref::<T>()
.expect("type mismatch in AsyncContext::get_cell")
.clone(),
_ => panic!("AsyncCellHandle does not point to a Cell node"),
}
}
pub fn set_cell<T>(&self, handle: &AsyncCellHandle<T>, value: T)
where
T: PartialEq + Clone + Send + Sync + 'static,
{
let dependents;
let is_batching;
{
let mut inner = self.inner.lock();
is_batching = inner.batch_depth > 0;
match inner.get_node_mut(handle.id) {
Some(AsyncNode::Cell(cell)) => {
let changed = !(*cell
.value
.as_ref()
.downcast_ref::<T>()
.expect("type mismatch in AsyncContext::set_cell")
== value);
if changed {
cell.value = Arc::new(value);
if is_batching {
inner.batched_cells.insert(handle.id);
return;
}
dependents = cell.dependents.clone();
} else {
return;
}
}
_ => panic!("AsyncCellHandle does not point to a Cell node"),
}
}
self.invalidate_frontier_async(&dependents);
}
pub(crate) fn get_slot_state(&self, id: SlotId) -> AsyncSlotStateView {
let inner = self.inner.lock();
match inner.get_node(id) {
Some(AsyncNode::Slot(slot)) => match &slot.state {
AsyncSlotState::Empty => AsyncSlotStateView::Empty,
AsyncSlotState::Computing { revision, .. } => AsyncSlotStateView::Computing {
revision: *revision,
},
AsyncSlotState::Resolved => AsyncSlotStateView::Resolved,
AsyncSlotState::Error => AsyncSlotStateView::Error,
},
_ => AsyncSlotStateView::None,
}
}
pub(crate) fn get_slot_revision(&self, id: SlotId) -> Option<u64> {
let inner = self.inner.lock();
match inner.get_node(id) {
Some(AsyncNode::Slot(slot)) => Some(slot.revision),
_ => None,
}
}
pub(crate) fn register_dependency(&self, dependency_id: SlotId, dependent_id: SlotId) {
let mut inner = self.inner.lock();
register_dependency_locked(&mut inner, dependency_id, dependent_id);
}
fn update_dependencies(
inner: &mut AsyncContextInner,
node_id: SlotId,
new_deps: &HashSet<SlotId>,
) {
let old_deps = match inner.get_node(node_id) {
Some(AsyncNode::Slot(s)) => s.dependencies.iter().copied().collect::<HashSet<_>>(),
_ => return,
};
for old_id in old_deps.difference(new_deps) {
if let Some(AsyncNode::Slot(s)) = inner.get_node_mut(*old_id) {
edge_remove(&mut s.dependents, node_id);
}
if let Some(AsyncNode::Cell(c)) = inner.get_node_mut(*old_id) {
edge_remove(&mut c.dependents, node_id);
}
}
if let Some(AsyncNode::Slot(s)) = inner.get_node_mut(node_id) {
s.dependencies = new_deps.iter().copied().collect();
}
for new_id in new_deps {
if let Some(AsyncNode::Slot(s)) = inner.get_node_mut(*new_id) {
edge_insert(&mut s.dependents, node_id);
}
if let Some(AsyncNode::Cell(c)) = inner.get_node_mut(*new_id) {
edge_insert(&mut c.dependents, node_id);
}
}
}
fn invalidate_frontier_async(&self, roots: &[SlotId]) {
let (effects, in_flight_handles, notifiers) = {
let mut inner = self.inner.lock();
let mut stack: smallvec::SmallVec<[SlotId; 16]> = smallvec::SmallVec::new();
stack.extend_from_slice(roots);
let mut effects: EdgeVec = EdgeVec::new();
let mut in_flight_handles: smallvec::SmallVec<[JoinHandle<()>; 4]> =
smallvec::SmallVec::new();
let mut notifiers: smallvec::SmallVec<[watch::Sender<AsyncCompletion>; 4]> =
smallvec::SmallVec::new();
let mut visited: EdgeVec = EdgeVec::new();
while let Some(id) = stack.pop() {
if !edge_insert(&mut visited, id) {
continue;
}
match inner.get_node_mut(id) {
Some(AsyncNode::Slot(slot)) => {
if let InvalidationResult::HadInFlight(handle) = slot.invalidate() {
in_flight_handles.push(handle);
}
if let Some(notifier) = slot.notifier.take() {
notifiers.push(notifier);
}
let dependents = slot.dependents.clone();
for dep_id in &dependents {
match inner.get_node(*dep_id) {
Some(AsyncNode::Effect(_)) => effects.push(*dep_id),
_ => stack.push(*dep_id),
}
}
}
Some(AsyncNode::Effect(_)) => {
effects.push(id);
}
_ => {}
}
}
(effects, in_flight_handles, notifiers)
};
for handle in in_flight_handles {
handle.abort();
}
drop(notifiers);
let has_effects = !effects.is_empty();
if has_effects {
let mut inner = self.inner.lock();
for effect_id in &effects {
Self::schedule_async_effect(&mut inner, *effect_id);
}
}
if has_effects {
self.flush_async_effects();
}
}
pub fn computed_async<T, F, Fut>(&self, compute: F) -> AsyncSlotHandle<T>
where
T: Clone + Send + Sync + 'static,
F: Fn(AsyncComputeContext) -> Fut + Send + Sync + 'static,
Fut: Future<Output = T> + Send + 'static,
{
self.slot_async_with_equals(compute, None)
}
pub fn memo_async<T, F, Fut>(&self, compute: F) -> AsyncSlotHandle<T>
where
T: PartialEq + Clone + Send + Sync + 'static,
F: Fn(AsyncComputeContext) -> Fut + Send + Sync + 'static,
Fut: Future<Output = T> + Send + 'static,
{
let equals: Arc<AsyncEqualsFn> = Arc::new(|old: &AsyncAny, new: &AsyncAny| -> bool {
let old_val = old.downcast_ref::<T>();
let new_val = new.downcast_ref::<T>();
match (old_val, new_val) {
(Some(o), Some(n)) => o == n,
_ => false,
}
});
self.slot_async_with_equals(compute, Some(equals))
}
fn slot_async_with_equals<T, F, Fut>(
&self,
compute: F,
equals: Option<Arc<AsyncEqualsFn>>,
) -> AsyncSlotHandle<T>
where
T: Clone + Send + Sync + 'static,
F: Fn(AsyncComputeContext) -> Fut + Send + Sync + 'static,
Fut: Future<Output = T> + Send + 'static,
{
let compute_arc: Arc<AsyncComputeFn> = Arc::new(move |ctx| {
let fut = compute(ctx);
Box::pin(async move { Arc::new(fut.await) as Arc<AsyncAny> })
});
let id;
{
let mut inner = self.inner.lock();
id = inner.alloc_id();
let node = AsyncSlotNode {
state: AsyncSlotState::Empty,
value: None,
error: None,
revision: 0,
compute: compute_arc,
equals,
dependencies: EdgeVec::new(),
dependents: EdgeVec::new(),
notifier: None,
};
inner.insert_node(id, AsyncNode::Slot(node));
}
AsyncSlotHandle {
id,
_marker: PhantomData,
}
}
pub fn get<T>(&self, handle: &AsyncSlotHandle<T>) -> Option<T>
where
T: Clone + Send + Sync + 'static,
{
let inner = self.inner.lock();
match inner.get_node(handle.id) {
Some(AsyncNode::Slot(slot)) => match &slot.state {
AsyncSlotState::Resolved => {
let val = slot.value.as_ref().expect("resolved without value");
Some(
val.downcast_ref::<T>()
.expect("type mismatch in get")
.clone(),
)
}
_ => None,
},
_ => None,
}
}
pub async fn get_async<T>(&self, handle: &AsyncSlotHandle<T>) -> T
where
T: Clone + Send + Sync + 'static,
{
loop {
if let Some(val) = self.get(handle) {
return val;
}
#[cfg(feature = "instrumentation")]
{
let hook = self.window1_hook.lock().take();
if let Some(hook) = hook {
hook().await;
}
}
let mut recv = {
let inner = self.inner.lock();
match inner.get_node(handle.id) {
Some(AsyncNode::Slot(slot)) => match &slot.state {
AsyncSlotState::Computing { .. } => slot
.notifier
.as_ref()
.expect("computing without notifier")
.subscribe(),
AsyncSlotState::Error | AsyncSlotState::Empty => {
drop(inner);
spawn_async_compute(self, handle.id)
}
AsyncSlotState::Resolved => {
#[cfg(feature = "instrumentation")]
self.window1_resolved_hits.fetch_add(1, Ordering::Relaxed);
let val = slot.value.as_ref().expect("resolved without value");
return val
.downcast_ref::<T>()
.expect("type mismatch in get_async")
.clone();
}
},
_ => panic!("AsyncSlotHandle does not point to a Slot node"),
}
};
'await_completion: loop {
if recv.changed().await.is_err() {
break 'await_completion;
}
let completion = recv.borrow_and_update().clone();
match completion {
AsyncCompletion::Resolved(val) => {
return val
.downcast_ref::<T>()
.expect("type mismatch in get_async completion")
.clone();
}
AsyncCompletion::Error(_err) => {
recv = spawn_async_compute(self, handle.id);
continue 'await_completion;
}
AsyncCompletion::Pending => continue 'await_completion,
}
}
}
}
pub fn batch<F, R>(&self, run: F) -> R
where
F: FnOnce(&AsyncContext) -> R,
{
{
let mut inner = self.inner.lock();
inner.batch_depth += 1;
}
let result = run(self);
{
let mut inner = self.inner.lock();
inner.batch_depth -= 1;
if inner.batch_depth == 0 {
let batched = inner.batched_cells.drain().collect::<Vec<_>>();
drop(inner);
let mut roots: EdgeVec = EdgeVec::new();
for cell_id in &batched {
let inner = self.inner.lock();
if let Some(AsyncNode::Cell(c)) = inner.get_node(*cell_id) {
roots.extend_from_slice(&c.dependents);
}
}
if !roots.is_empty() {
self.invalidate_frontier_async(&roots);
}
}
}
result
}
pub fn effect_async<F, Fut, C>(&self, effect: F) -> AsyncEffectHandle
where
F: Fn(AsyncComputeContext) -> Fut + Send + Sync + 'static,
Fut: Future<Output = Option<C>> + Send + 'static,
C: FnOnce() + Send + 'static,
{
let id;
{
let mut inner = self.inner.lock();
id = inner.alloc_id();
let effect_fn: Arc<AsyncEffectFn> = Arc::new(move |ctx| {
let fut = effect(ctx);
Box::pin(async move { fut.await.map(|c| Box::new(c) as BoxedCleanupFn) })
});
let node = AsyncEffectNode {
effect_fn,
cleanup: None,
dependencies: EdgeVec::new(),
dependents: EdgeVec::new(),
force_run: true,
in_flight: None,
};
inner.insert_node(id, AsyncNode::Effect(node));
Self::schedule_async_effect(&mut inner, id);
}
let handle = AsyncEffectHandle { id };
self.flush_async_effects();
handle
}
pub fn dependent_count(&self, node: &impl GraphNode) -> usize {
let inner = self.inner.lock();
match inner.get_node(node.node_id()) {
Some(AsyncNode::Slot(slot)) => slot.dependents.len(),
Some(AsyncNode::Cell(cell)) => cell.dependents.len(),
Some(AsyncNode::Effect(_)) | None => 0,
}
}
pub fn dependency_count(&self, node: &impl GraphNode) -> usize {
let inner = self.inner.lock();
match inner.get_node(node.node_id()) {
Some(AsyncNode::Slot(slot)) => slot.dependencies.len(),
Some(AsyncNode::Effect(effect)) => effect.dependencies.len(),
Some(AsyncNode::Cell(_)) | None => 0,
}
}
pub fn is_async_effect_active(&self, handle: &AsyncEffectHandle) -> bool {
let inner = self.inner.lock();
matches!(inner.get_node(handle.id), Some(AsyncNode::Effect(_)))
}
fn detach_edges_locked(inner: &mut AsyncContextInner, id: SlotId, node: &AsyncNode) {
let (dependencies, dependents) = match node {
AsyncNode::Slot(s) => (s.dependencies.clone(), s.dependents.clone()),
AsyncNode::Cell(c) => (EdgeVec::new(), c.dependents.clone()),
AsyncNode::Effect(e) => (e.dependencies.clone(), EdgeVec::new()),
};
for dep_id in &dependencies {
match inner.get_node_mut(*dep_id) {
Some(AsyncNode::Slot(s)) => {
edge_remove(&mut s.dependents, id);
}
Some(AsyncNode::Cell(c)) => {
edge_remove(&mut c.dependents, id);
}
_ => {}
}
}
for dependent_id in &dependents {
match inner.get_node_mut(*dependent_id) {
Some(AsyncNode::Slot(s)) => {
edge_remove(&mut s.dependencies, id);
}
Some(AsyncNode::Effect(e)) => {
edge_remove(&mut e.dependencies, id);
}
_ => {}
}
}
}
fn retire_id_locked(inner: &mut AsyncContextInner, id: SlotId) {
let Some(idx) = AsyncContextInner::node_index(id) else {
return;
};
if idx < inner.nodes.len() {
inner.nodes[idx] = None;
}
if idx < inner.generations.len() {
inner.generations[idx] += 1;
}
inner.free_ids.push(id.0);
}
fn invalidate_disposed_dependents(&self, roots: &[SlotId]) {
if roots.is_empty() {
return;
}
let (in_flight_handles, notifiers) = {
let mut inner = self.inner.lock();
let mut stack: smallvec::SmallVec<[SlotId; 16]> = smallvec::SmallVec::new();
stack.extend_from_slice(roots);
let mut in_flight_handles: smallvec::SmallVec<[JoinHandle<()>; 4]> =
smallvec::SmallVec::new();
let mut notifiers: smallvec::SmallVec<[watch::Sender<AsyncCompletion>; 4]> =
smallvec::SmallVec::new();
let mut visited: EdgeVec = EdgeVec::new();
while let Some(id) = stack.pop() {
if !edge_insert(&mut visited, id) {
continue;
}
if let Some(AsyncNode::Slot(slot)) = inner.get_node_mut(id) {
if let InvalidationResult::HadInFlight(handle) = slot.invalidate() {
in_flight_handles.push(handle);
}
if let Some(notifier) = slot.notifier.take() {
notifiers.push(notifier);
}
let dependents = slot.dependents.clone();
for dep_id in &dependents {
if !matches!(inner.get_node(*dep_id), Some(AsyncNode::Effect(_))) {
stack.push(*dep_id);
}
}
}
}
(in_flight_handles, notifiers)
};
for handle in in_flight_handles {
handle.abort();
}
drop(notifiers);
}
pub fn dispose_slot<T>(&self, handle: &AsyncSlotHandle<T>) {
let (in_flight, notifier, stale) = {
let mut inner = self.inner.lock();
if !matches!(inner.get_node(handle.id), Some(AsyncNode::Slot(_))) {
return;
}
let Some(idx) = AsyncContextInner::node_index(handle.id) else {
return;
};
let Some(node) = inner.nodes[idx].take() else {
return;
};
Self::detach_edges_locked(&mut inner, handle.id, &node);
let AsyncNode::Slot(mut slot) = node else {
return;
};
let stale = slot.dependents.clone();
let in_flight = slot.clear();
let notifier = slot.notifier.take();
Self::retire_id_locked(&mut inner, handle.id);
(in_flight, notifier, stale)
};
if let Some(in_flight) = in_flight {
in_flight.abort();
}
drop(notifier);
self.invalidate_disposed_dependents(&stale);
}
pub fn dispose_cell<T>(&self, handle: &AsyncCellHandle<T>) {
let stale = {
let mut inner = self.inner.lock();
if !matches!(inner.get_node(handle.id), Some(AsyncNode::Cell(_))) {
return;
}
let Some(idx) = AsyncContextInner::node_index(handle.id) else {
return;
};
let Some(node) = inner.nodes[idx].take() else {
return;
};
Self::detach_edges_locked(&mut inner, handle.id, &node);
let AsyncNode::Cell(cell) = node else {
return;
};
let stale = cell.dependents.clone();
Self::retire_id_locked(&mut inner, handle.id);
stale
};
self.invalidate_disposed_dependents(&stale);
}
fn dispose_id(&self, id: SlotId) {
let kind = match self.inner.lock().get_node(id) {
Some(AsyncNode::Slot(_)) => 0u8,
Some(AsyncNode::Cell(_)) => 1,
Some(AsyncNode::Effect(_)) => 2,
None => return,
};
let marker = std::marker::PhantomData;
match kind {
0 => self.dispose_slot(&AsyncSlotHandle::<()> {
id,
_marker: marker,
}),
1 => self.dispose_cell(&AsyncCellHandle::<()> {
id,
_marker: marker,
}),
_ => self.dispose_async_effect(&AsyncEffectHandle { id }),
}
}
pub(crate) fn handle(&self) -> AsyncContext {
AsyncContext {
inner: Arc::clone(&self.inner),
#[cfg(feature = "instrumentation")]
window1_hook: Mutex::new(None),
#[cfg(feature = "instrumentation")]
window1_resolved_hits: AtomicU64::new(0),
}
}
pub fn scope(&self) -> AsyncTeardownScope {
AsyncTeardownScope {
ctx: self.handle(),
owned: Mutex::new(Vec::new()),
}
}
pub fn dispose_async_effect(&self, handle: &AsyncEffectHandle) {
let (cleanup, in_flight) = {
let mut inner = self.inner.lock();
inner.pending_async_effects.retain(|&id| id != handle.id);
inner.scheduled_async_effects.remove(&handle.id);
let (cleanup, in_flight) = match inner.get_node_mut(handle.id) {
Some(AsyncNode::Effect(e)) => {
let deps = e.dependencies.clone();
let prior_cleanup = e.cleanup.take();
let prior_in_flight = e.in_flight.take();
for dep_id in &deps {
match inner.get_node_mut(*dep_id) {
Some(AsyncNode::Slot(s)) => {
edge_remove(&mut s.dependents, handle.id);
}
Some(AsyncNode::Cell(c)) => {
edge_remove(&mut c.dependents, handle.id);
}
_ => {}
}
}
(prior_cleanup, prior_in_flight)
}
_ => return,
};
let index = usize::try_from(handle.id.0).ok();
if let Some(idx) = index
&& idx < inner.nodes.len()
{
inner.nodes[idx] = None;
if idx < inner.generations.len() {
inner.generations[idx] += 1;
}
inner.free_ids.push(handle.id.0);
}
(cleanup, in_flight)
};
if let Some(in_flight) = in_flight {
in_flight.abort();
}
if let Some(cleanup) = cleanup {
cleanup();
}
}
pub fn signal_async<T, F, Fut>(&self, compute: F) -> AsyncSignalHandle<T>
where
T: PartialEq + Clone + Send + Sync + 'static,
F: Fn(AsyncComputeContext) -> Fut + Send + Sync + 'static,
Fut: Future<Output = T> + Send + 'static,
{
let slot = self.memo_async(compute);
let effect = self.effect_async(move |ctx: AsyncComputeContext| {
let fut = ctx.get_async(&slot);
async move {
let _ = fut.await;
None::<fn()>
}
});
AsyncSignalHandle { slot, effect }
}
pub fn get_signal<T: Clone + Send + Sync + 'static>(
&self,
handle: &AsyncSignalHandle<T>,
) -> Option<T> {
self.get(&handle.slot)
}
pub async fn get_signal_async<T: Clone + Send + Sync + 'static>(
&self,
handle: &AsyncSignalHandle<T>,
) -> T {
self.get_async(&handle.slot).await
}
pub fn dispose_signal<T>(&self, handle: &AsyncSignalHandle<T>) {
self.dispose_async_effect(&handle.effect);
}
pub fn is_signal_active<T>(&self, handle: &AsyncSignalHandle<T>) -> bool {
let inner = self.inner.lock();
matches!(inner.get_node(handle.effect.id), Some(AsyncNode::Effect(_)))
}
fn schedule_async_effect(inner: &mut AsyncContextInner, id: SlotId) {
if let Some(AsyncNode::Effect(e)) = inner.get_node_mut(id) {
e.force_run = true;
}
if inner.scheduled_async_effects.insert(id) {
inner.pending_async_effects.push(id);
}
}
fn flush_async_effects(&self) {
let effect_ids: Vec<SlotId>;
{
let mut inner = self.inner.lock();
effect_ids = inner.pending_async_effects.drain(..).collect();
inner.scheduled_async_effects.clear();
}
let ctx_inner = self.inner.clone();
for effect_id in effect_ids {
let should_run = {
let inner = self.inner.lock();
match inner.get_node(effect_id) {
Some(AsyncNode::Effect(e)) => e.force_run,
_ => false,
}
};
if !should_run {
continue;
}
{
let mut inner = self.inner.lock();
if let Some(AsyncNode::Effect(e)) = inner.get_node_mut(effect_id) {
e.force_run = false;
}
}
let effect_fn = {
let inner = self.inner.lock();
match inner.get_node(effect_id) {
Some(AsyncNode::Effect(e)) => Some(e.effect_fn.clone()),
_ => None,
}
};
if let Some(fn_arc) = effect_fn {
{
let mut inner = self.inner.lock();
if let Some(AsyncNode::Effect(e)) = inner.get_node_mut(effect_id)
&& let Some(prior) = e.in_flight.take()
{
prior.abort();
}
}
let inner_for_ctx = ctx_inner.clone();
let effect_gen = {
let inner = self.inner.lock();
inner.generation(effect_id)
};
let join = tokio::spawn(async move {
{
let mut inner = inner_for_ctx.lock();
if inner.generation(effect_id) == effect_gen
&& let Some(AsyncNode::Effect(e)) = inner.get_node_mut(effect_id)
&& let Some(cleanup) = e.cleanup.take()
{
drop(inner);
cleanup();
}
}
let (context_id, deps_arc) = {
let mut inner = inner_for_ctx.lock();
let deps = inner.take_deps();
(inner.context_id, deps)
};
let deps_for_extract = deps_arc.clone();
let compute_ctx = AsyncComputeContext {
_context_id: context_id,
_node_id: effect_id,
_node_gen: effect_gen,
inner: inner_for_ctx.clone(),
dependencies: deps_arc,
};
let cleanup = fn_arc(compute_ctx).await;
let deps = std::mem::take(&mut *deps_for_extract.lock());
{
let mut inner = inner_for_ctx.lock();
if inner.generation(effect_id) == effect_gen {
AsyncContext::update_effect_dependencies(&mut inner, effect_id, &deps);
inner.recycle_deps(deps_for_extract, deps);
if let Some(AsyncNode::Effect(e)) = inner.get_node_mut(effect_id) {
e.cleanup = cleanup;
}
} else {
inner.recycle_deps(deps_for_extract, deps);
drop(inner);
if let Some(cleanup) = cleanup {
cleanup();
}
}
}
});
let mut inner = self.inner.lock();
if inner.generation(effect_id) == effect_gen
&& let Some(AsyncNode::Effect(e)) = inner.get_node_mut(effect_id)
{
e.in_flight = Some(join);
}
}
}
}
fn update_effect_dependencies(
inner: &mut AsyncContextInner,
effect_id: SlotId,
new_deps: &HashSet<SlotId>,
) {
let old_deps = match inner.get_node(effect_id) {
Some(AsyncNode::Effect(e)) => e.dependencies.iter().copied().collect::<HashSet<_>>(),
_ => return,
};
for old_id in old_deps.difference(new_deps) {
match inner.get_node_mut(*old_id) {
Some(AsyncNode::Slot(s)) => {
edge_remove(&mut s.dependents, effect_id);
}
Some(AsyncNode::Cell(c)) => {
edge_remove(&mut c.dependents, effect_id);
}
_ => {}
}
}
if let Some(AsyncNode::Effect(e)) = inner.get_node_mut(effect_id) {
e.dependencies = new_deps.iter().copied().collect();
}
for new_id in new_deps {
match inner.get_node_mut(*new_id) {
Some(AsyncNode::Slot(s)) => {
edge_insert(&mut s.dependents, effect_id);
}
Some(AsyncNode::Cell(c)) => {
edge_insert(&mut c.dependents, effect_id);
}
_ => {}
}
}
}
}
impl crate::reactive_graph::Teardown for AsyncTeardownScope {
fn len(&self) -> usize {
AsyncTeardownScope::len(self)
}
fn disarm(self) {
AsyncTeardownScope::disarm(self);
}
}
impl crate::reactive_graph::ReactiveGraph for AsyncContext {
type SlotHandle<T> = AsyncSlotHandle<T>;
type CellHandle<T> = AsyncCellHandle<T>;
type EffectHandle = AsyncEffectHandle;
type Scope<'a> = AsyncTeardownScope;
fn dispose_slot<T: 'static>(&self, handle: &Self::SlotHandle<T>) {
AsyncContext::dispose_slot(self, handle);
}
fn dispose_cell<T: 'static>(&self, handle: &Self::CellHandle<T>) {
AsyncContext::dispose_cell(self, handle);
}
fn dispose_effect(&self, handle: &Self::EffectHandle) {
AsyncContext::dispose_async_effect(self, handle);
}
fn scope(&self) -> Self::Scope<'_> {
AsyncContext::scope(self)
}
fn batch<R>(&self, run: impl FnOnce(&Self) -> R) -> R {
AsyncContext::batch(self, run)
}
fn dependent_count(&self, node: &impl GraphNode) -> usize {
AsyncContext::dependent_count(self, node)
}
fn dependency_count(&self, node: &impl GraphNode) -> usize {
AsyncContext::dependency_count(self, node)
}
}
impl crate::reactive_graph::AsyncReactiveGraph for AsyncContext {
fn cell<T>(&self, value: T) -> Self::CellHandle<T>
where
T: PartialEq + Clone + Send + Sync + 'static,
{
AsyncContext::cell(self, value)
}
fn get_cell<T>(&self, handle: &Self::CellHandle<T>) -> T
where
T: Clone + Send + Sync + 'static,
{
AsyncContext::get_cell(self, handle)
}
fn set_cell<T>(&self, handle: &Self::CellHandle<T>, value: T)
where
T: PartialEq + Clone + Send + Sync + 'static,
{
AsyncContext::set_cell(self, handle, value);
}
fn get<T>(&self, handle: &Self::SlotHandle<T>) -> impl Future<Output = T> + Send
where
T: Clone + Send + Sync + 'static,
{
AsyncContext::get_async(self, handle)
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Arc;
use std::sync::atomic::AtomicU64;
use tokio::runtime::Runtime;
#[test]
fn dispose_slot_cancels_an_in_flight_recompute_and_discards_its_result() {
let rt = Runtime::new().unwrap();
let _guard = rt.enter();
let ctx = AsyncContext::new();
let ran = Arc::new(AtomicU64::new(0));
let finished = Arc::new(AtomicU64::new(0));
let r = Arc::clone(&ran);
let f = Arc::clone(&finished);
let slow = ctx.computed_async(move |_c| {
let r = Arc::clone(&r);
let f = Arc::clone(&f);
Box::pin(async move {
r.fetch_add(1, Ordering::SeqCst);
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
f.fetch_add(1, Ordering::SeqCst);
1i64
}) as std::pin::Pin<Box<dyn Future<Output = i64> + Send>>
});
rt.block_on(async {
let handle = {
let ctx = ctx.handle();
tokio::spawn(async move { ctx.get_async(&slow).await })
};
tokio::time::sleep(std::time::Duration::from_millis(20)).await;
ctx.dispose_slot(&slow);
tokio::time::sleep(std::time::Duration::from_millis(300)).await;
handle.abort();
});
assert_eq!(ran.load(Ordering::SeqCst), 1, "the compute did start");
assert_eq!(
finished.load(Ordering::SeqCst),
0,
"an in-flight compute must be cancelled by disposal, not allowed to finish"
);
}
#[test]
fn dispose_bumps_the_generation_before_recycling_the_id() {
let rt = Runtime::new().unwrap();
let _guard = rt.enter();
let ctx = AsyncContext::new();
let cell = ctx.cell(1i64);
let id = cell.id;
let before = ctx.inner.lock().generation(id);
ctx.dispose_cell(&cell);
let after = ctx.inner.lock().generation(id);
assert_eq!(after, before + 1, "generation must advance on disposal");
assert!(
ctx.inner.lock().free_ids.contains(&id.0),
"the id must be recycled after the bump, not before"
);
}
#[test]
fn dispose_detaches_both_directions_and_is_kind_checked() {
let rt = Runtime::new().unwrap();
let _guard = rt.enter();
let ctx = AsyncContext::new();
let src = ctx.cell(4i64);
let derived = ctx.computed_async(move |c| {
Box::pin(async move { c.get_cell(&src) })
as std::pin::Pin<Box<dyn Future<Output = i64> + Send>>
});
rt.block_on(ctx.get_async(&derived));
assert_eq!(ctx.dependent_count(&src), 1);
assert_eq!(ctx.dependency_count(&derived), 1);
ctx.dispose_slot(&derived);
assert_eq!(ctx.dependent_count(&src), 0);
assert_eq!(ctx.dependency_count(&derived), 0);
ctx.dispose_slot(&derived);
let stale = AsyncCellHandle::<i64> {
id: derived.id,
_marker: std::marker::PhantomData,
};
ctx.dispose_cell(&stale);
}
#[test]
fn teardown_scope_disposes_on_drop_and_disarm_cancels() {
let rt = Runtime::new().unwrap();
let _guard = rt.enter();
let ctx = AsyncContext::new();
let topic = ctx.cell(1i64);
{
let scope = ctx.scope();
let a = scope.computed_async(move |c| {
Box::pin(async move { c.get_cell(&topic) + 1 })
as std::pin::Pin<Box<dyn Future<Output = i64> + Send>>
});
assert_eq!(scope.len(), 1);
assert_eq!(rt.block_on(ctx.get_async(&a)), 2);
assert_eq!(ctx.dependent_count(&topic), 1);
}
assert_eq!(ctx.dependent_count(&topic), 0);
}
fn stub_compute(_ctx: AsyncComputeContext) -> BoxedAsyncFuture {
Box::pin(async { Arc::new(()) as Arc<AsyncAny> })
}
fn make_slot_node(revision: u64) -> AsyncSlotNode {
AsyncSlotNode {
state: AsyncSlotState::Empty,
value: None,
error: None,
revision,
compute: Arc::new(stub_compute),
equals: None,
dependencies: EdgeVec::new(),
dependents: EdgeVec::new(),
notifier: None,
}
}
fn make_slot_node_with_memo(revision: u64, value: Option<Arc<AsyncAny>>) -> AsyncSlotNode {
AsyncSlotNode {
state: AsyncSlotState::Empty,
value,
error: None,
revision,
compute: Arc::new(stub_compute),
equals: Some(Arc::new(|old: &AsyncAny, new: &AsyncAny| -> bool {
let old_val = old.downcast_ref::<i32>();
let new_val = new.downcast_ref::<i32>();
match (old_val, new_val) {
(Some(o), Some(n)) => o == n,
_ => false,
}
})),
dependencies: EdgeVec::new(),
dependents: EdgeVec::new(),
notifier: None,
}
}
#[test]
fn async_slot_state_starts_empty() {
let ctx = AsyncContext::new();
let id;
{
let mut inner = ctx.inner.lock();
id = inner.alloc_id();
inner.insert_node(id, AsyncNode::Slot(make_slot_node(0)));
}
let state = ctx.get_slot_state(id);
assert!(matches!(state, AsyncSlotStateView::Empty));
}
#[test]
fn empty_to_computing_transition() {
let rt = Runtime::new().unwrap();
let handle = rt.spawn(async {});
let mut node = make_slot_node(0);
let old = node.transition_to_computing(handle);
assert!(old.is_none());
assert!(matches!(
node.state,
AsyncSlotState::Computing { revision: 0, .. }
));
}
#[test]
fn computing_to_resolved_transition() {
let rt = Runtime::new().unwrap();
let handle = rt.spawn(async {});
let mut node = make_slot_node(0);
node.transition_to_computing(handle);
let result = node.transition_to_resolved(0, Arc::new(42i32));
assert!(matches!(result, TransitionOutcome::Accepted));
assert!(matches!(node.state, AsyncSlotState::Resolved));
assert_eq!(
node.value.as_ref().unwrap().downcast_ref::<i32>().unwrap(),
&42
);
}
#[test]
fn computing_to_error_transition() {
let rt = Runtime::new().unwrap();
let handle = rt.spawn(async {});
let mut node = make_slot_node(0);
node.transition_to_computing(handle);
let err: Arc<dyn Error + Send + Sync> = Arc::new(std::io::Error::other("test error"));
let result = node.transition_to_error(0, err);
assert!(matches!(result, TransitionOutcome::Accepted));
assert!(matches!(node.state, AsyncSlotState::Error));
assert!(node.error.is_some());
assert!(node.value.is_none());
}
#[test]
fn stale_completion_is_rejected() {
let rt = Runtime::new().unwrap();
let handle = rt.spawn(async {});
let mut node = make_slot_node(1);
node.transition_to_computing(handle);
let result = node.transition_to_resolved(0, Arc::new(42i32));
assert!(matches!(result, TransitionOutcome::Stale));
}
#[test]
fn computing_to_computing_stale_returns_old_handle() {
let rt = Runtime::new().unwrap();
let handle1 = rt.spawn(async {});
let handle2 = rt.spawn(async {});
let mut node = make_slot_node(0);
node.transition_to_computing(handle1);
node.revision = 1;
let old = node.transition_to_computing(handle2);
assert!(old.is_some());
assert!(matches!(
node.state,
AsyncSlotState::Computing { revision: 1, .. }
));
}
#[test]
fn resolved_to_computing_via_invalidation() {
let rt = Runtime::new().unwrap();
let handle = rt.spawn(async {});
let mut node = make_slot_node(0);
node.transition_to_computing(handle);
node.transition_to_resolved(0, Arc::new(42i32));
assert!(matches!(node.state, AsyncSlotState::Resolved));
let result = node.invalidate();
assert!(matches!(result, InvalidationResult::WasResolved));
assert!(matches!(node.state, AsyncSlotState::Empty));
assert_eq!(node.revision, 1);
}
#[test]
fn error_to_computing_via_invalidation() {
let mut node = AsyncSlotNode {
state: AsyncSlotState::Error,
value: None,
error: Some(Arc::new(std::io::Error::other("test"))),
revision: 0,
compute: Arc::new(stub_compute),
equals: None,
dependencies: EdgeVec::new(),
dependents: EdgeVec::new(),
notifier: None,
};
let result = node.invalidate();
assert!(matches!(result, InvalidationResult::WasError));
assert!(matches!(node.state, AsyncSlotState::Empty));
assert_eq!(node.revision, 1);
}
#[test]
fn clear_aborts_in_flight() {
let rt = Runtime::new().unwrap();
let handle = rt.spawn(async { std::future::pending::<()>().await });
let mut node = make_slot_node(0);
node.transition_to_computing(handle);
let old_handle = node.clear();
assert!(old_handle.is_some());
old_handle.unwrap().abort();
assert!(matches!(node.state, AsyncSlotState::Empty));
assert!(node.value.is_none());
assert_eq!(node.revision, 1);
}
#[test]
fn memo_unchanged_transition() {
let rt = Runtime::new().unwrap();
let handle = rt.spawn(async {});
let mut node = make_slot_node_with_memo(0, Some(Arc::new(42i32)));
node.transition_to_computing(handle);
let result = node.transition_to_resolved(0, Arc::new(42i32));
assert!(matches!(result, TransitionOutcome::Unchanged));
}
#[test]
fn async_context_cell_basic() {
let ctx = AsyncContext::new();
let cell = ctx.cell(10i32);
assert_eq!(ctx.get_cell(&cell), 10);
ctx.set_cell(&cell, 20);
assert_eq!(ctx.get_cell(&cell), 20);
}
#[test]
fn async_context_cell_noop_on_equal() {
let ctx = AsyncContext::new();
let cell = ctx.cell(10i32);
ctx.set_cell(&cell, 10);
assert_eq!(ctx.get_cell(&cell), 10);
}
#[test]
fn async_context_id_unique() {
let ctx1 = AsyncContext::new();
let ctx2 = AsyncContext::new();
let id1 = ctx1.inner.lock().context_id;
let id2 = ctx2.inner.lock().context_id;
assert_ne!(id1, id2);
}
#[tokio::test]
async fn computed_async_basic() {
let ctx = AsyncContext::new();
let slot = ctx.computed_async(|_ctx| async move { 42i32 });
let val = ctx.get_async(&slot).await;
assert_eq!(val, 42);
}
#[tokio::test]
async fn computed_async_reads_cell() {
let ctx = AsyncContext::new();
let cell = ctx.cell(10i32);
let slot = ctx.computed_async(move |ctx| {
let val = ctx.get_cell(&cell);
async move { val + 1 }
});
let val = ctx.get_async(&slot).await;
assert_eq!(val, 11);
}
#[tokio::test]
async fn computed_async_cached() {
let ctx = AsyncContext::new();
let count = Arc::new(AtomicU64::new(0));
let count_clone = count.clone();
let slot = ctx.computed_async(move |_| {
let c = count_clone.clone();
async move {
c.fetch_add(1, Ordering::Relaxed);
42i32
}
});
let v1 = ctx.get_async(&slot).await;
let v2 = ctx.get_async(&slot).await;
assert_eq!(v1, 42);
assert_eq!(v2, 42);
assert_eq!(count.load(Ordering::Relaxed), 1);
}
#[tokio::test]
async fn computed_async_invalidation() {
let ctx = AsyncContext::new();
let cell = ctx.cell(1i32);
let slot = ctx.computed_async(move |ctx| {
let val = ctx.get_cell(&cell);
async move { val * 2 }
});
assert_eq!(ctx.get_async(&slot).await, 2);
ctx.set_cell(&cell, 5);
assert_eq!(ctx.get_async(&slot).await, 10);
}
#[tokio::test]
async fn memo_async_suppresses_equal() {
let ctx = AsyncContext::new();
let cell = ctx.cell(1i32);
let count = Arc::new(AtomicU64::new(0));
let count_clone = count.clone();
let slot = ctx.memo_async(move |ctx| {
let val = ctx.get_cell(&cell);
let c = count_clone.clone();
async move {
c.fetch_add(1, Ordering::Relaxed);
val / val
}
});
assert_eq!(ctx.get_async(&slot).await, 1);
ctx.set_cell(&cell, 2);
assert_eq!(ctx.get_async(&slot).await, 1);
assert_eq!(count.load(Ordering::Relaxed), 2);
}
#[tokio::test]
async fn batch_defers_invalidation() {
let ctx = AsyncContext::new();
let cell = ctx.cell(1i32);
let slot = ctx.computed_async(move |ctx| {
let val = ctx.get_cell(&cell);
async move { val * 10 }
});
assert_eq!(ctx.get_async(&slot).await, 10);
ctx.batch(|ctx| {
ctx.set_cell(&cell, 2);
ctx.set_cell(&cell, 3);
});
assert_eq!(ctx.get_async(&slot).await, 30);
}
#[tokio::test]
async fn concurrent_get_async_deduplicates() {
let ctx = AsyncContext::new();
let count = Arc::new(AtomicU64::new(0));
let count_clone = count.clone();
let slot = ctx.computed_async(move |_| {
let c = count_clone.clone();
async move {
c.fetch_add(1, Ordering::Relaxed);
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
99i32
}
});
let (v1, v2) = tokio::join!(ctx.get_async(&slot), ctx.get_async(&slot));
assert_eq!(v1, 99);
assert_eq!(v2, 99);
assert_eq!(count.load(Ordering::Relaxed), 1);
}
#[tokio::test]
async fn async_slot_reads_async_slot() {
let ctx = AsyncContext::new();
let cell = ctx.cell(5i32);
let base = ctx.computed_async(move |ctx| {
let v = ctx.get_cell(&cell);
async move { v + 10 }
});
let derived = ctx.computed_async(move |ctx| {
let base_handle = base;
async move {
let v = ctx.get_async(&base_handle).await;
v * 2
}
});
assert_eq!(ctx.get_async(&derived).await, 30);
}
#[tokio::test]
async fn async_chain_invalidation() {
let ctx = AsyncContext::new();
let cell = ctx.cell(1i32);
let cell_clone = cell;
let base = ctx.computed_async(move |ctx| {
let v = ctx.get_cell(&cell_clone);
async move { v + 10 }
});
let derived = ctx.computed_async(move |ctx| {
let bh = base;
async move {
let v = ctx.get_async(&bh).await;
v * 2
}
});
assert_eq!(ctx.get_async(&derived).await, 22);
ctx.set_cell(&cell_clone, 3);
assert_eq!(ctx.get_async(&derived).await, 26);
}
#[tokio::test]
async fn async_chain_three_levels() {
let ctx = AsyncContext::new();
let cell = ctx.cell(1i32);
let a = ctx.computed_async(move |ctx| {
let v = ctx.get_cell(&cell);
async move { v + 1 }
});
let b = ctx.computed_async(move |ctx| {
let ah = a;
async move { ctx.get_async(&ah).await + 1 }
});
let c = ctx.computed_async(move |ctx| {
let bh = b;
async move { ctx.get_async(&bh).await + 1 }
});
assert_eq!(ctx.get_async(&c).await, 4);
ctx.set_cell(&cell, 10);
assert_eq!(ctx.get_async(&c).await, 13);
}
#[tokio::test]
async fn async_dependency_tracks_slot_edges() {
let ctx = AsyncContext::new();
let cell = ctx.cell(3i32);
let slot = ctx.computed_async(move |ctx| {
let v = ctx.get_cell(&cell);
async move { v * 2 }
});
let _ = ctx.get_async(&slot).await;
{
let inner = ctx.inner.lock();
if let Some(AsyncNode::Slot(s)) = inner.get_node(slot.id) {
assert!(s.dependencies.contains(&cell.id));
}
}
{
let inner = ctx.inner.lock();
if let Some(AsyncNode::Cell(c)) = inner.get_node(cell.id) {
assert!(c.dependents.contains(&slot.id));
}
}
}
#[tokio::test]
async fn async_dependency_updates_on_rerun() {
let ctx = AsyncContext::new();
let cell_a = ctx.cell(1i32);
let cell_b = ctx.cell(100i32);
let flag = ctx.cell(true);
let slot = ctx.computed_async(move |ctx| {
let f = ctx.get_cell(&flag);
let v = if f {
ctx.get_cell(&cell_a)
} else {
ctx.get_cell(&cell_b)
};
async move { v }
});
assert_eq!(ctx.get_async(&slot).await, 1);
{
let inner = ctx.inner.lock();
if let Some(AsyncNode::Slot(s)) = inner.get_node(slot.id) {
assert!(s.dependencies.contains(&cell_a.id));
assert!(!s.dependencies.contains(&cell_b.id));
}
}
ctx.set_cell(&flag, false);
assert_eq!(ctx.get_async(&slot).await, 100);
{
let inner = ctx.inner.lock();
if let Some(AsyncNode::Slot(s)) = inner.get_node(slot.id) {
assert!(!s.dependencies.contains(&cell_a.id));
assert!(s.dependencies.contains(&cell_b.id));
}
}
}
#[tokio::test]
async fn dependency_trackers_are_pooled_and_reused() {
let ctx = AsyncContext::new();
let cell = ctx.cell(0i32);
let slot = ctx.computed_async(move |cctx| {
let v = cctx.get_cell(&cell);
async move { v }
});
assert_eq!(ctx.get_async(&slot).await, 0);
let pooled = {
let inner = ctx.inner.lock();
assert_eq!(
inner.deps_pool.len(),
1,
"completed run returns its tracker"
);
assert!(
inner.deps_pool[0].lock().capacity() > 0,
"table capacity is recycled with the tracker, not just the Arc"
);
Arc::as_ptr(&inner.deps_pool[0])
};
for i in 1..40i32 {
ctx.set_cell(&cell, i);
assert_eq!(ctx.get_async(&slot).await, i);
assert!(ctx.inner.lock().deps_pool.len() <= DEPS_POOL_CAP);
}
let inner = ctx.inner.lock();
assert_eq!(inner.deps_pool.len(), 1, "steady state reuses one tracker");
assert_eq!(
Arc::as_ptr(&inner.deps_pool[0]),
pooled,
"the same allocation is recycled across recomputes"
);
}
#[tokio::test]
async fn leaked_dependency_tracker_is_not_recycled() {
let ctx = AsyncContext::new();
let cell = ctx.cell(1i32);
type LeakedTrackers = Arc<Mutex<Vec<Arc<Mutex<HashSet<SlotId>>>>>>;
let leaked: LeakedTrackers = Arc::new(Mutex::new(Vec::new()));
let leaked_for_compute = leaked.clone();
let slot = ctx.computed_async(move |cctx| {
leaked_for_compute.lock().push(cctx.dependencies.clone());
let v = cctx.get_cell(&cell);
async move { v }
});
assert_eq!(ctx.get_async(&slot).await, 1);
assert_eq!(leaked.lock().len(), 1);
assert!(
ctx.inner.lock().deps_pool.is_empty(),
"a tracker still referenced elsewhere must not re-enter the pool"
);
}
#[tokio::test]
async fn dispose_bumps_generation_then_id_recycles_fresh() {
let ctx = Arc::new(AsyncContext::new());
let cell = ctx.cell(1i32);
let effect_a = ctx.effect_async(move |c| {
let _v = c.get_cell(&cell);
async move { None::<fn()> }
});
tokio::time::sleep(std::time::Duration::from_millis(20)).await;
let a_id = effect_a.id;
let g0 = ctx.inner.lock().generation(a_id);
let cell_b = ctx.cell(2i32);
ctx.dispose_async_effect(&effect_a);
let g1 = ctx.inner.lock().generation(a_id);
assert_eq!(g1, g0 + 1, "dispose must bump the node generation");
let effect_b = ctx.effect_async(move |c| {
let _v = c.get_cell(&cell_b);
async move { None::<fn()> }
});
assert_eq!(
effect_b.id, a_id,
"free_ids LIFO should recycle A's id for B"
);
assert_eq!(
ctx.inner.lock().generation(effect_b.id),
g1,
"recycled node keeps the bumped generation",
);
}
#[tokio::test]
async fn stale_generation_context_does_not_alias_recycled_effect() {
let ctx = Arc::new(AsyncContext::new());
let cell = ctx.cell(1i32);
let effect_a = ctx.effect_async(move |c| {
let _v = c.get_cell(&cell);
async move { None::<fn()> }
});
tokio::time::sleep(std::time::Duration::from_millis(20)).await;
let a_id = effect_a.id;
let ctx_id = ctx.inner.lock().context_id;
let a_gen = ctx.inner.lock().generation(a_id);
let stale_ctx = AsyncComputeContext {
_context_id: ctx_id,
_node_id: a_id,
_node_gen: a_gen,
inner: ctx.inner.clone(),
dependencies: Arc::new(Mutex::new(HashSet::new())),
};
let cell_b = ctx.cell(99i32);
ctx.dispose_async_effect(&effect_a);
let effect_b = ctx.effect_async(move |c| {
let _v = c.get_cell(&cell_b);
async move { None::<fn()> }
});
tokio::time::sleep(std::time::Duration::from_millis(20)).await;
assert_eq!(
effect_b.id, a_id,
"B must reuse A's recycled id for the test"
);
let _ = stale_ctx.get_cell(&cell);
let inner = ctx.inner.lock();
match inner.get_node(a_id) {
Some(AsyncNode::Effect(e)) => {
assert!(
e.dependencies.contains(&cell_b.id),
"B's real dependency must stay intact",
);
assert!(
!e.dependencies.contains(&cell.id),
"stale-generation write must not alias B with A's old dependency",
);
}
_ => panic!("recycled node should be B's effect"),
}
if let Some(AsyncNode::Cell(c)) = inner.get_node(cell.id) {
assert!(
!c.dependents.contains(&a_id),
"`cell` must not gain a phantom dependent via the stale write",
);
}
}
#[tokio::test]
async fn invalidation_aborts_in_flight() {
let ctx = Arc::new(AsyncContext::new());
let cell = ctx.cell(1i32);
let compute_count = Arc::new(AtomicU64::new(0));
let count_clone = compute_count.clone();
let slot = ctx.computed_async(move |ctx| {
let _v = ctx.get_cell(&cell);
let c = count_clone.clone();
async move {
c.fetch_add(1, Ordering::Relaxed);
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
42i32
}
});
let ctx_clone = ctx.clone();
let handle = tokio::spawn(async move { ctx_clone.get_async(&slot).await });
tokio::time::sleep(std::time::Duration::from_millis(20)).await;
ctx.set_cell(&cell, 2);
tokio::time::sleep(std::time::Duration::from_millis(20)).await;
let val = ctx.get_async(&slot).await;
assert_eq!(val, 42);
let _ = handle.await;
}
#[tokio::test]
async fn stale_revision_prevents_publish() {
let ctx = Arc::new(AsyncContext::new());
let cell = ctx.cell(1i32);
let slot = ctx.computed_async(move |ctx| {
let v = ctx.get_cell(&cell);
async move { v + 1 }
});
let ctx1 = ctx.clone();
let h1 = tokio::spawn(async move { ctx1.get_async(&slot).await });
let v1 = h1.await.unwrap();
assert_eq!(v1, 2);
let state = ctx.get_slot_state(slot.id);
assert!(matches!(state, AsyncSlotStateView::Resolved));
ctx.set_cell(&cell, 10);
let state = ctx.get_slot_state(slot.id);
assert!(matches!(state, AsyncSlotStateView::Empty));
let v2 = ctx.get_async(&slot).await;
assert_eq!(v2, 11);
}
#[tokio::test]
async fn dropping_one_waiter_does_not_cancel_shared_compute() {
let ctx = AsyncContext::new();
let compute_count = Arc::new(AtomicU64::new(0));
let count_clone = compute_count.clone();
let slot = ctx.computed_async(move |_| {
let c = count_clone.clone();
async move {
c.fetch_add(1, Ordering::Relaxed);
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
99i32
}
});
let (v2, v3) = tokio::join!(ctx.get_async(&slot), ctx.get_async(&slot));
assert_eq!(v2, 99);
assert_eq!(v3, 99);
assert_eq!(compute_count.load(Ordering::Relaxed), 1);
}
#[tokio::test]
async fn effect_async_runs_on_creation() {
let ctx = AsyncContext::new();
let cell = ctx.cell(10i32);
let result = Arc::new(Mutex::new(0i32));
let result_clone = result.clone();
ctx.effect_async(move |ctx| {
let v = ctx.get_cell(&cell);
let r = result_clone.clone();
async move {
*r.lock() = v;
None::<fn()>
}
});
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
assert_eq!(*result.lock(), 10);
}
#[tokio::test(flavor = "multi_thread")]
async fn effect_async_reruns_on_cell_change() {
let ctx = AsyncContext::new();
let cell = ctx.cell(1i32);
let count = Arc::new(AtomicU64::new(0));
let count_clone = count.clone();
ctx.effect_async(move |ctx| {
let _v = ctx.get_cell(&cell);
let c = count_clone.clone();
async move {
c.fetch_add(1, Ordering::Relaxed);
None::<fn()>
}
});
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
assert_eq!(count.load(Ordering::Relaxed), 1);
ctx.set_cell(&cell, 2);
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
assert!(count.load(Ordering::Relaxed) >= 2);
}
#[tokio::test]
async fn effect_async_cleanup_runs_on_rerun() {
let ctx = AsyncContext::new();
let cell = ctx.cell(1i32);
let cleanup_count = Arc::new(AtomicU64::new(0));
let cleanup_clone = cleanup_count.clone();
ctx.effect_async(move |ctx| {
let _v = ctx.get_cell(&cell);
let c = cleanup_clone.clone();
async move {
Some(move || {
c.fetch_add(1, Ordering::Relaxed);
})
}
});
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
assert_eq!(cleanup_count.load(Ordering::Relaxed), 0);
ctx.set_cell(&cell, 2);
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
assert!(cleanup_count.load(Ordering::Relaxed) >= 1);
}
#[tokio::test]
async fn effect_rerun_aborts_prior_inflight() {
let ctx = AsyncContext::new();
let cell = ctx.cell(1i32);
let done_count = Arc::new(AtomicU64::new(0));
let cleanup_count = Arc::new(AtomicU64::new(0));
let done_clone = done_count.clone();
let cleanup_clone = cleanup_count.clone();
let handle = ctx.effect_async(move |ctx| {
let _v = ctx.get_cell(&cell);
let d = done_clone.clone();
let c = cleanup_clone.clone();
async move {
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
d.fetch_add(1, Ordering::Relaxed);
let c = c.clone();
Some(move || {
c.fetch_add(1, Ordering::Relaxed);
})
}
});
tokio::time::sleep(std::time::Duration::from_millis(20)).await;
ctx.set_cell(&cell, 2); tokio::time::sleep(std::time::Duration::from_millis(250)).await;
assert_eq!(
done_count.load(Ordering::Relaxed),
1,
"prior in-flight effect must be aborted on re-run, not double-executed"
);
ctx.dispose_async_effect(&handle);
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
assert_eq!(
cleanup_count.load(Ordering::Relaxed),
1,
"the surviving run's cleanup must run exactly once on dispose"
);
}
#[tokio::test]
async fn dispose_async_effect_removes_it() {
let ctx = AsyncContext::new();
let cell = ctx.cell(1i32);
let count = Arc::new(AtomicU64::new(0));
let count_clone = count.clone();
let handle = ctx.effect_async(move |ctx| {
let _v = ctx.get_cell(&cell);
let c = count_clone.clone();
async move {
c.fetch_add(1, Ordering::Relaxed);
None::<fn()>
}
});
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
let after_first = count.load(Ordering::Relaxed);
assert!(after_first >= 1);
ctx.dispose_async_effect(&handle);
ctx.set_cell(&cell, 2);
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
assert_eq!(count.load(Ordering::Relaxed), after_first);
}
#[test]
fn sync_get_returns_none_for_empty_slot() {
let ctx = AsyncContext::new();
let slot = ctx.computed_async(|_| async { 42i32 });
assert!(ctx.get(&slot).is_none());
}
#[tokio::test]
async fn sync_get_returns_some_after_resolve() {
let ctx = AsyncContext::new();
let slot = ctx.computed_async(|_| async { 42i32 });
let val = ctx.get_async(&slot).await;
assert_eq!(val, 42);
assert_eq!(ctx.get(&slot), Some(42));
}
#[tokio::test]
async fn sync_get_returns_none_after_invalidation() {
let ctx = AsyncContext::new();
let cell = ctx.cell(1i32);
let slot = ctx.computed_async(move |ctx| {
let v = ctx.get_cell(&cell);
async move { v * 2 }
});
let _ = ctx.get_async(&slot).await;
assert_eq!(ctx.get(&slot), Some(2));
ctx.set_cell(&cell, 5);
assert!(ctx.get(&slot).is_none());
}
#[tokio::test]
async fn sync_get_avoids_spawn_overhead() {
let ctx = AsyncContext::new();
let count = Arc::new(AtomicU64::new(0));
let count_clone = count.clone();
let slot = ctx.computed_async(move |_| {
let c = count_clone.clone();
async move {
c.fetch_add(1, Ordering::Relaxed);
99i32
}
});
let v1 = ctx.get_async(&slot).await;
assert_eq!(v1, 99);
assert_eq!(count.load(Ordering::Relaxed), 1);
let v2 = ctx.get(&slot);
assert_eq!(v2, Some(99));
assert_eq!(count.load(Ordering::Relaxed), 1);
}
#[tokio::test]
async fn sync_get_with_memo_returns_cached() {
let ctx = AsyncContext::new();
let cell = ctx.cell(3i32);
let slot = ctx.memo_async(move |ctx| {
let v = ctx.get_cell(&cell);
async move { v.abs() }
});
assert_eq!(ctx.get_async(&slot).await, 3);
assert_eq!(ctx.get(&slot), Some(3));
}
#[tokio::test]
async fn get_async_uses_sync_fast_path() {
let ctx = AsyncContext::new();
let cell = ctx.cell(10i32);
let count = Arc::new(AtomicU64::new(0));
let count_clone = count.clone();
let slot = ctx.computed_async(move |ctx| {
let v = ctx.get_cell(&cell);
let c = count_clone.clone();
async move {
c.fetch_add(1, Ordering::Relaxed);
v + 1
}
});
let v1 = ctx.get_async(&slot).await;
assert_eq!(v1, 11);
assert_eq!(count.load(Ordering::Relaxed), 1);
let v2 = ctx.get_async(&slot).await;
assert_eq!(v2, 11);
assert_eq!(count.load(Ordering::Relaxed), 1);
}
#[test]
fn async_schedule_effect_dedupes_pending_queue() {
let ctx = AsyncContext::new();
let rt = Runtime::new().unwrap();
let _guard = rt.enter();
let cell = ctx.cell(0i32);
let effect = ctx.effect_async(move |ctx| {
let _ = ctx.get_cell(&cell);
async { None::<fn()> }
});
rt.block_on(async { tokio::time::sleep(std::time::Duration::from_millis(20)).await });
{
let mut inner = ctx.inner.lock();
AsyncContext::schedule_async_effect(&mut inner, effect.id);
AsyncContext::schedule_async_effect(&mut inner, effect.id);
AsyncContext::schedule_async_effect(&mut inner, effect.id);
let count = inner
.pending_async_effects
.iter()
.filter(|&&id| id == effect.id)
.count();
assert_eq!(
count, 1,
"pending_async_effects must dedupe the same effect id; got {:?}",
inner.pending_async_effects
);
assert!(inner.scheduled_async_effects.contains(&effect.id));
}
ctx.flush_async_effects();
{
let inner = ctx.inner.lock();
assert!(
!inner.scheduled_async_effects.contains(&effect.id),
"flush must clear scheduled_async_effects"
);
}
rt.block_on(async { tokio::time::sleep(std::time::Duration::from_millis(20)).await });
}
}