use std::{
sync::{
Arc, Mutex, MutexGuard, PoisonError, Weak,
atomic::{AtomicBool, Ordering},
},
task::Waker,
};
pub(super) type SyncWaker = Arc<dyn Fn() + Send + Sync>;
pub(super) enum Slot {
Task(Waker),
Sync(SyncWaker),
}
impl Slot {
fn fire(self) {
match self {
Self::Task(w) => w.wake(),
Self::Sync(f) => f(),
}
}
}
#[derive(Default)]
struct WakerSlots {
slots: Vec<(u64, Slot)>,
fired: bool,
next_id: u64,
}
#[derive(Default)]
pub(super) struct Node {
flag: AtomicBool,
children: Mutex<Vec<Weak<Self>>>,
wakers: Mutex<WakerSlots>,
_parent: Option<Arc<Self>>,
}
impl Node {
pub(super) fn cancel(self: &Arc<Self>) {
if self.flag.swap(true, Ordering::AcqRel) {
return;
}
let drained: Vec<Slot> = {
let mut w = lock(&self.wakers);
w.fired = true;
w.slots.drain(..).map(|(_, slot)| slot).collect()
};
for slot in drained {
slot.fire();
}
let children: Vec<Arc<Self>> = {
let mut kids = lock(&self.children);
let live: Vec<Arc<Self>> = kids.iter().filter_map(Weak::upgrade).collect();
kids.clear();
live
};
for child in children {
child.cancel();
}
}
pub(super) fn child(parent: &Arc<Self>) -> Arc<Self> {
let mut kids = lock(&parent.children);
let born_cancelled = parent.flag.load(Ordering::Acquire);
let node = Arc::new(Self {
flag: AtomicBool::new(born_cancelled),
_parent: Some(Arc::clone(parent)),
children: Mutex::new(Vec::new()),
wakers: Mutex::new(WakerSlots {
fired: born_cancelled,
..WakerSlots::default()
}),
});
kids.retain(|w| w.strong_count() > 0);
kids.push(Arc::downgrade(&node));
node
}
pub(super) fn is_cancelled(&self) -> bool {
self.flag.load(Ordering::Acquire)
}
pub(super) fn refresh_task(&self, id: u64, waker: &Waker) {
let mut w = lock(&self.wakers);
if let Some((_, slot)) = w.slots.iter_mut().find(|(sid, _)| *sid == id) {
*slot = Slot::Task(waker.clone());
} else {
drop(w);
waker.wake_by_ref();
}
}
pub(super) fn register(&self, slot: Slot) -> Option<u64> {
let mut w = lock(&self.wakers);
if !w.fired {
let id = w.next_id;
w.next_id += 1;
w.slots.push((id, slot));
drop(w);
return Some(id);
}
drop(w);
slot.fire();
None
}
pub(super) fn root() -> Arc<Self> {
Arc::new(Self::default())
}
pub(super) fn unregister(&self, id: u64) {
lock(&self.wakers).slots.retain(|(sid, _)| *sid != id);
}
}
fn lock<T>(m: &Mutex<T>) -> MutexGuard<'_, T> {
m.lock().unwrap_or_else(PoisonError::into_inner)
}
#[cfg(test)]
mod tests {
use std::{sync::Arc, time::Duration};
use kithara_test_utils::kithara;
use super::{Node, Slot, lock};
fn noop_sync() -> Slot {
Slot::Sync(Arc::new(|| {}))
}
#[kithara::test(timeout(Duration::from_secs(5)))]
fn register_unregister_churn_leaves_no_slots() {
let node = Node::root();
for _ in 0..1000 {
let id = node
.register(noop_sync())
.expect("a fresh node parks the slot");
node.unregister(id);
}
assert_eq!(
lock(&node.wakers).slots.len(),
0,
"register/unregister churn must not leak slots"
);
}
#[kithara::test(timeout(Duration::from_secs(5)))]
fn child_churn_sweeps_dead_weaks() {
let root = Node::root();
for _ in 0..1000 {
let _ = Node::child(&root);
}
let _survivor = Node::child(&root);
assert_eq!(
lock(&root.children).len(),
1,
"dead-Weak sweep must keep the children vec bounded under churn"
);
}
}