use std::mem;
use std::task::Waker;
use crate::internal::arena::Arena;
use crate::internal::arena::SlotId;
use crate::internal::waker_batch::WakerBatch;
#[derive(Debug)]
pub struct WakerToken(SlotId);
#[derive(Debug)]
pub struct WakerSet {
wakers: Arena<Waker>,
}
impl WakerSet {
pub const fn new() -> Self {
Self {
wakers: Arena::new(),
}
}
pub fn with_capacity(capacity: usize) -> Self {
Self {
wakers: Arena::with_capacity(capacity),
}
}
#[inline]
pub fn drain(&mut self) -> impl Iterator<Item = Waker> + 'static {
let mut wakers = WakerBatch::with_capacity(self.wakers.len());
if self.wakers.is_empty() {
return wakers.into_iter();
}
wakers.extend(self.wakers.drain());
wakers.into_iter()
}
#[inline]
pub fn take_all(&mut self) -> impl Iterator<Item = Waker> + 'static {
self.wakers.take_all()
}
#[inline]
#[must_use = "drop the returned waker after releasing the waker set's state lock"]
pub fn register(&mut self, token: &mut Option<WakerToken>, waker: &Waker) -> Option<Waker> {
if let Some(current) = token.as_ref().map(|token| {
self.wakers
.get_mut(token.0)
.expect("waker token must refer to an occupied slot")
}) {
if current.will_wake(waker) {
return None;
}
return Some(mem::replace(current, waker.clone()));
}
*token = Some(WakerToken(self.wakers.insert(waker.clone())));
None
}
#[inline]
#[must_use = "drop the returned waker after releasing the waker set's state lock"]
pub fn unregister(&mut self, token: &mut Option<WakerToken>) -> Option<Waker> {
token.take().map(|token| self.wakers.remove(token.0))
}
#[cfg(test)]
fn registered_len(&self) -> usize {
self.wakers.len()
}
}
#[cfg(test)]
mod tests {
use std::mem::size_of;
use std::sync::Arc;
use std::sync::atomic::AtomicBool;
use std::sync::atomic::AtomicUsize;
use std::sync::atomic::Ordering;
use std::task::Wake;
use super::*;
struct DropWake {
dropped: Arc<AtomicBool>,
wake_count: AtomicUsize,
}
impl Wake for DropWake {
fn wake(self: Arc<Self>) {
self.wake_count.fetch_add(1, Ordering::Relaxed);
}
}
impl Drop for DropWake {
fn drop(&mut self) {
self.dropped.store(true, Ordering::Relaxed);
}
}
#[test]
fn waker_token_preserves_the_option_niche() {
assert_eq!(size_of::<WakerToken>(), size_of::<usize>());
assert_eq!(size_of::<WakerToken>(), size_of::<Option<WakerToken>>());
}
#[test]
fn unregister_returns_the_waker_for_deferred_drop() {
let dropped = Arc::new(AtomicBool::new(false));
let waker = Waker::from(Arc::new(DropWake {
dropped: dropped.clone(),
wake_count: AtomicUsize::new(0),
}));
let mut wakers = WakerSet::new();
let mut token = None;
drop(wakers.register(&mut token, &waker));
drop(waker);
let removed = wakers.unregister(&mut token);
assert_eq!(wakers.registered_len(), 0);
assert!(!dropped.load(Ordering::Relaxed));
drop(removed);
assert!(dropped.load(Ordering::Relaxed));
}
}