use core::{cell::UnsafeCell, task::Waker};
pub(crate) struct EndpointWaiter {
waker: UnsafeCell<Option<Waker>>,
}
impl EndpointWaiter {
#[inline]
pub(crate) const fn empty() -> Self {
Self {
waker: UnsafeCell::new(None),
}
}
#[inline]
pub(crate) fn replace(&self, waker: Waker) -> Option<Waker> {
unsafe { (&mut *self.waker.get()).replace(waker) }
}
#[inline]
pub(crate) fn take(&self) -> Option<Waker> {
unsafe { (&mut *self.waker.get()).take() }
}
#[inline]
pub(crate) fn is_empty(&self) -> bool {
unsafe { (&*self.waker.get()).is_none() }
}
}
impl Drop for EndpointWaiter {
fn drop(&mut self) {
if self.waker.get_mut().is_some() {
crate::invariant();
}
}
}
#[cfg(test)]
mod tests {
use super::EndpointWaiter;
use std::{
cell::Cell,
task::{RawWaker, RawWakerVTable, Waker},
};
unsafe fn clone_count_waker(data: *const ()) -> RawWaker {
RawWaker::new(data, &COUNT_WAKER_VTABLE)
}
unsafe fn wake_count_waker(data: *const ()) {
let count = unsafe { &*data.cast::<Cell<usize>>() };
count.set(count.get() + 1);
}
unsafe fn drop_count_waker(_: *const ()) {}
static COUNT_WAKER_VTABLE: RawWakerVTable = RawWakerVTable::new(
clone_count_waker,
wake_count_waker,
wake_count_waker,
drop_count_waker,
);
fn counting_waker(count: &Cell<usize>) -> Waker {
let data = core::ptr::from_ref(count).cast::<()>();
unsafe { Waker::from_raw(RawWaker::new(data, &COUNT_WAKER_VTABLE)) }
}
#[test]
fn replacement_moves_displaced_owner_out_of_the_lease_record() {
let first = Cell::new(0);
let second = Cell::new(0);
let waiter = EndpointWaiter::empty();
assert!(waiter.replace(counting_waker(&first)).is_none());
let displaced = waiter.replace(counting_waker(&second));
assert!(displaced.is_some());
drop(displaced);
assert_eq!(first.get(), 0);
if let Some(waker) = waiter.take() {
waker.wake();
}
assert_eq!(second.get(), 1);
assert!(waiter.take().is_none());
}
#[test]
fn drop_rejects_a_registered_wake_owner() {
let count = Cell::new(0);
let waiter = EndpointWaiter::empty();
assert!(waiter.replace(counting_waker(&count)).is_none());
let rejected = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| drop(waiter)));
assert!(rejected.is_err());
}
}