use std::{
fmt,
future::Future,
marker::PhantomData,
pin::Pin,
sync::{Arc, OnceLock, Weak},
task::{Context, Poll, Wake, Waker},
};
use smallvec::SmallVec;
use crate::{
lock::{Lock, WeakLock},
sync::Mutex,
};
const INLINE_WAITERS: usize = 32;
pub struct Waiter {
waker: Waker,
shared: OnceLock<Arc<Waker>>,
}
impl Waiter {
pub fn new(waker: Waker) -> Self {
Self {
waker,
shared: OnceLock::new(),
}
}
pub fn noop() -> Self {
Self::new(Waker::noop().clone())
}
pub fn register(&self, list: &mut WaiterList) {
list.register(self);
}
pub fn waker(&self) -> &Waker {
&self.waker
}
fn shared(&self) -> &Arc<Waker> {
self.shared.get_or_init(|| Arc::new(self.waker.clone()))
}
pub fn poll_future<F: Future + ?Sized>(&self, future: Pin<&mut F>) -> Poll<F::Output> {
future.poll(&mut Context::from_waker(self.waker()))
}
}
impl Clone for Waiter {
fn clone(&self) -> Self {
let shared = self.shared().clone();
Self {
waker: self.waker.clone(),
shared: OnceLock::from(shared),
}
}
}
pub struct WaiterList {
entries: SmallVec<[Weak<Waker>; INLINE_WAITERS]>,
cursor: usize,
}
impl WaiterList {
pub fn new() -> Self {
Self {
entries: SmallVec::new(),
cursor: 0,
}
}
pub fn register(&mut self, waiter: &Waiter) {
let new_weak = Arc::downgrade(waiter.shared());
for _ in 0..self.entries.len().min(2) {
if self.entries[self.cursor].strong_count() == 0 {
self.entries[self.cursor] = new_weak;
return;
}
self.cursor = (self.cursor + 1) % self.entries.len();
}
self.entries.push(new_weak);
}
pub fn take(&mut self) -> Self {
self.cursor = 0;
Self {
entries: std::mem::take(&mut self.entries),
cursor: 0,
}
}
pub fn wake(&mut self) {
self.cursor = 0;
for waker in self.entries.drain(..).filter_map(|w| w.upgrade()) {
waker.wake_by_ref();
}
}
}
impl Default for WaiterList {
fn default() -> Self {
Self::new()
}
}
impl fmt::Debug for WaiterList {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("WaiterList").field("len", &self.entries.len()).finish()
}
}
#[derive(Default)]
pub struct Park(Option<Waiter>);
impl Park {
pub fn new(waiter: Waiter) -> Self {
Self(Some(waiter))
}
pub fn hold(&mut self, cx: &Context<'_>) -> &Waiter {
let reuse = self.0.as_ref().is_some_and(|waiter| {
cx.waker().will_wake(&waiter.waker) && waiter.shared.get().is_none_or(|shared| Arc::weak_count(shared) == 0)
});
if !reuse {
self.0 = Some(Waiter::new(cx.waker().clone()));
}
self.0.as_ref().unwrap()
}
}
impl Clone for Park {
fn clone(&self) -> Self {
Self(None)
}
}
impl fmt::Debug for Park {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Park").field("parked", &self.0.is_some()).finish()
}
}
pub struct Fan<T = WaiterList> {
inner: Arc<FanInner<T>>,
}
enum Target<T> {
Owned(Lock<T>),
Projected(WeakLock<T>),
}
impl<T> Target<T> {
fn upgrade(&self) -> Option<Lock<T>> {
match self {
Self::Owned(lock) => Some(lock.clone()),
Self::Projected(weak) => weak.upgrade(),
}
}
}
struct FanInner<T> {
target: Target<T>,
project: fn(&mut T) -> &mut WaiterList,
defer: Mutex<Defer>,
}
#[derive(Default)]
struct Defer {
held: usize,
owed: bool,
}
impl<T> FanInner<T> {
fn defer_if_held(&self) -> bool {
let mut defer = self.defer.lock().expect("mutex poisoned");
if defer.held == 0 {
return false;
}
defer.owed = true;
true
}
fn notify(&self) {
if self.defer_if_held() {
return;
}
let Some(lock) = self.target.upgrade() else {
return;
};
let mut waiters = {
let mut state = lock.lock();
let mut defer = self.defer.lock().expect("mutex poisoned");
if defer.held > 0 {
defer.owed = true;
return;
}
(self.project)(&mut state).take()
};
waiters.wake();
}
}
impl<T: Send + 'static> Wake for FanInner<T> {
fn wake(self: Arc<Self>) {
self.notify();
}
fn wake_by_ref(self: &Arc<Self>) {
self.notify();
}
}
impl Fan<WaiterList> {
pub fn new() -> Self {
Self::build(Target::Owned(Lock::new(WaiterList::new())), |list| list)
}
}
impl<T: Send + 'static> Fan<T> {
pub fn project(state: &Lock<T>, project: fn(&mut T) -> &mut WaiterList) -> Self {
Self::build(Target::Projected(state.downgrade()), project)
}
fn build(target: Target<T>, project: fn(&mut T) -> &mut WaiterList) -> Self {
Self {
inner: Arc::new(FanInner {
target,
project,
defer: Mutex::new(Defer::default()),
}),
}
}
pub fn register(&self, waiter: &Waiter) {
if let Some(lock) = self.inner.target.upgrade() {
let mut state = lock.lock();
waiter.register((self.inner.project)(&mut state));
}
}
pub fn wake(&self) {
self.inner.notify();
}
pub fn waker(&self) -> Waker {
Waker::from(self.inner.clone())
}
#[must_use = "wakes are held back only while the guard is alive"]
pub fn hold(&self) -> Hold<T> {
self.inner.defer.lock().expect("mutex poisoned").held += 1;
Hold { fan: self.clone() }
}
}
impl Default for Fan<WaiterList> {
fn default() -> Self {
Self::new()
}
}
impl<T> Clone for Fan<T> {
fn clone(&self) -> Self {
Self {
inner: self.inner.clone(),
}
}
}
impl<T> fmt::Debug for Fan<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let defer = self.inner.defer.lock().expect("mutex poisoned");
f.debug_struct("Fan")
.field("held", &defer.held)
.field("owed", &defer.owed)
.finish_non_exhaustive()
}
}
pub struct Hold<T = WaiterList> {
fan: Fan<T>,
}
impl<T> Drop for Hold<T> {
fn drop(&mut self) {
let owed = {
let mut defer = self.fan.inner.defer.lock().expect("mutex poisoned");
defer.held -= 1;
match defer.held {
0 => std::mem::take(&mut defer.owed),
_ => false,
}
};
if owed {
self.fan.inner.notify();
}
}
}
impl<T> fmt::Debug for Hold<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Hold").finish_non_exhaustive()
}
}
struct WaiterFn<F, R> {
poll: F,
park: Park, _marker: PhantomData<fn() -> R>,
}
pub fn wait<F, R>(poll: F) -> impl Future<Output = R>
where
F: FnMut(&Waiter) -> Poll<R> + Unpin,
{
WaiterFn {
poll,
park: Park::default(),
_marker: PhantomData,
}
}
impl<F, R> Future for WaiterFn<F, R>
where
F: FnMut(&Waiter) -> Poll<R> + Unpin,
{
type Output = R;
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<R> {
let this = &mut *self;
let waiter = this.park.hold(cx);
(this.poll)(waiter)
}
}
#[cfg(all(test, not(loom)))]
mod tests {
use super::*;
#[derive(Default)]
struct Flag(std::sync::atomic::AtomicBool);
impl Flag {
fn woken(&self) -> bool {
self.0.load(std::sync::atomic::Ordering::SeqCst)
}
}
impl Wake for Flag {
fn wake(self: Arc<Self>) {
self.wake_by_ref();
}
fn wake_by_ref(self: &Arc<Self>) {
self.0.store(true, std::sync::atomic::Ordering::SeqCst);
}
}
fn flagged(fan: &Fan) -> (Arc<Flag>, Waiter) {
let flag = Arc::new(Flag::default());
let waiter = Waiter::new(Waker::from(flag.clone()));
fan.register(&waiter);
(flag, waiter)
}
#[test]
fn the_waker_fans_out_to_everyone_parked() {
let fan = Fan::new();
let (first, _first_waiter) = flagged(&fan);
let (second, _second_waiter) = flagged(&fan);
fan.waker().wake();
assert!(first.woken() && second.woken(), "one wake must reach every waiter");
}
#[test]
fn a_departed_waiter_releases_its_slot() {
let fan = Fan::new();
let (gone, waiter) = flagged(&fan);
drop(waiter);
let (live, _live_waiter) = flagged(&fan);
fan.wake();
assert!(live.woken());
assert!(!gone.woken(), "a dropped waiter should have no registration left");
}
#[test]
fn waking_does_not_hold_the_lock() {
struct Reentrant(Fan);
impl Wake for Reentrant {
fn wake(self: Arc<Self>) {
self.wake_by_ref();
}
fn wake_by_ref(self: &Arc<Self>) {
self.0.register(&Waiter::noop());
}
}
let fan = Fan::new();
let waiter = Waiter::new(Waker::from(Arc::new(Reentrant(fan.clone()))));
fan.register(&waiter);
let (tx, rx) = std::sync::mpsc::channel();
std::thread::spawn({
let fan = fan.clone();
move || {
fan.wake();
let _ = tx.send(());
}
});
rx.recv_timeout(std::time::Duration::from_secs(5))
.expect("wake reached a waker while holding the lock, and the wake re-entered it");
}
#[test]
fn a_held_wake_lands_when_the_hold_drops() {
let fan = Fan::new();
let (flag, _waiter) = flagged(&fan);
let hold = fan.hold();
fan.wake();
assert!(!flag.woken(), "the wake was delivered while the fan was held");
drop(hold);
assert!(flag.woken(), "the held wake never arrived");
}
#[test]
fn only_the_last_hold_out_delivers() {
let fan = Fan::new();
let (flag, _waiter) = flagged(&fan);
let outer = fan.hold();
let inner = fan.hold();
fan.wake();
drop(inner);
assert!(!flag.woken(), "a hold is still outstanding");
drop(outer);
assert!(flag.woken());
}
#[test]
fn a_quiet_hold_wakes_nobody() {
let fan = Fan::new();
let (flag, _waiter) = flagged(&fan);
drop(fan.hold());
assert!(!flag.woken(), "nothing woke, so nothing was owed");
}
#[test]
fn poll_future_bridges_a_std_future() {
let waiter = Waiter::noop();
let fut = std::pin::pin!(std::future::ready(7u8));
assert_eq!(waiter.poll_future(fut), Poll::Ready(7));
let fut = std::pin::pin!(std::future::pending::<u8>());
assert_eq!(waiter.poll_future(fut), Poll::Pending);
let mut boxed: Pin<Box<dyn Future<Output = u8>>> = Box::pin(std::future::ready(9u8));
assert_eq!(waiter.poll_future(boxed.as_mut()), Poll::Ready(9));
}
const fn assert_sync<T: Sync>() {}
const _: () = {
assert_sync::<Waiter>();
assert_sync::<crate::Pending<crate::Consumer<u32>>>();
assert_sync::<crate::Shared<u32>>();
};
#[test]
fn park_survives_a_poll_that_returns_early() {
let waker = Waker::noop().clone();
let cx = Context::from_waker(&waker);
let mut park = Park::default();
let mut list = WaiterList::new();
fn poll_step(park: &mut Park, cx: &Context<'_>, list: &mut WaiterList, nested: Poll<u8>) -> Poll<u8> {
let waiter = park.hold(cx);
waiter.register(list);
let value = std::task::ready!(nested);
Poll::Ready(value + 1)
}
assert!(poll_step(&mut park, &cx, &mut list, Poll::Pending).is_pending());
assert_eq!(
list.entries[0].strong_count(),
1,
"an early return must leave the registration live"
);
list.wake();
}
#[test]
fn park_reuses_a_drained_waiter_and_retires_a_registered_one() {
let waker = Waker::noop().clone();
let cx = Context::from_waker(&waker);
let mut park = Park::default();
let mut list = WaiterList::new();
let waiter = park.hold(&cx);
waiter.register(&mut list);
let first = waiter.shared().clone();
let waiter = park.hold(&cx);
assert!(!Arc::ptr_eq(&first, waiter.shared()), "a registered waiter was reused");
waiter.register(&mut list);
let second = waiter.shared().clone();
list.wake();
let waiter = park.hold(&cx);
assert!(Arc::ptr_eq(&second, waiter.shared()), "a drained waiter was not reused");
}
#[test]
fn park_retires_a_waiter_for_another_task() {
struct Nop;
impl std::task::Wake for Nop {
fn wake(self: Arc<Self>) {}
}
let waker_a = Waker::from(Arc::new(Nop));
let waker_b = Waker::from(Arc::new(Nop));
let mut park = Park::default();
let first = park.hold(&Context::from_waker(&waker_a)).shared().clone();
let waiter = park.hold(&Context::from_waker(&waker_b));
assert!(
!Arc::ptr_eq(&first, waiter.shared()),
"a waiter for another task was reused"
);
}
#[test]
fn park_new_starts_holding() {
let waker = Waker::noop().clone();
let cx = Context::from_waker(&waker);
let mut list = WaiterList::new();
let waiter = Waiter::new(cx.waker().clone());
waiter.register(&mut list);
let mut park = Park::new(waiter);
assert_eq!(list.entries[0].strong_count(), 1, "a constructed park must hold");
let shared = park.0.as_ref().unwrap().shared().clone();
list.wake();
assert!(
Arc::ptr_eq(&shared, park.hold(&cx).shared()),
"the constructed waiter was not reused"
);
}
#[test]
fn park_clone_is_idle() {
let waker = Waker::noop().clone();
let cx = Context::from_waker(&waker);
let mut park = Park::default();
park.hold(&cx);
assert!(park.0.is_some());
assert!(park.clone().0.is_none(), "a cloned park must start idle");
}
#[test]
fn waiter_clone_shares_identity() {
let waker = Waker::noop().clone();
let cx = Context::from_waker(&waker);
let mut list = WaiterList::new();
let waiter = Waiter::new(cx.waker().clone());
let clone = waiter.clone();
assert!(Arc::ptr_eq(waiter.shared(), clone.shared()));
clone.register(&mut list);
drop(clone);
assert_eq!(list.entries[0].strong_count(), 1, "the original must keep it live");
drop(waiter);
assert_eq!(list.entries[0].strong_count(), 0, "the last handle must release it");
}
#[test]
fn wait_output_need_not_be_unpin() {
struct NotUnpin(#[allow(dead_code)] std::marker::PhantomPinned);
let mut fut = std::pin::pin!(crate::wait(|_| Poll::Ready(NotUnpin(std::marker::PhantomPinned))));
let mut cx = Context::from_waker(Waker::noop());
assert!(fut.as_mut().poll(&mut cx).is_ready());
}
#[test]
fn a_projected_fan_wakes_the_list_inside_the_state() {
#[derive(Default)]
struct State {
waiters: WaiterList,
other: WaiterList,
}
let state = Lock::new(State::default());
let fan = Fan::project(&state, |s| &mut s.waiters);
let flag = Arc::new(Flag::default());
let waiter = Waiter::new(Waker::from(flag.clone()));
waiter.register(&mut state.lock().waiters);
let bystander = Arc::new(Flag::default());
let bystander_waiter = Waiter::new(Waker::from(bystander.clone()));
bystander_waiter.register(&mut state.lock().other);
fan.waker().wake();
assert!(flag.woken(), "the projected list was not woken");
assert!(!bystander.woken(), "only the projected list should be woken");
}
#[test]
fn a_projected_fan_outlives_its_state_harmlessly() {
let state = Lock::new(WaiterList::new());
let fan = Fan::project(&state, |list| list);
let flag = Arc::new(Flag::default());
let waiter = Waiter::new(Waker::from(flag.clone()));
waiter.register(&mut state.lock());
drop(state);
fan.wake();
fan.waker().wake();
assert!(!flag.woken());
}
#[test]
fn a_projected_fan_defers_a_held_wake() {
let state = Lock::new(WaiterList::new());
let fan = Fan::project(&state, |list| list);
let flag = Arc::new(Flag::default());
let waiter = Waiter::new(Waker::from(flag.clone()));
waiter.register(&mut state.lock());
let hold = fan.hold();
fan.wake();
assert!(!flag.woken(), "the wake was delivered while the fan was held");
drop(hold);
assert!(flag.woken(), "the held wake never arrived");
}
}