use std::{collections::HashMap, default::Default, num::NonZeroUsize, task::Waker};
#[repr(transparent)]
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub(crate) struct Token(NonZeroUsize);
#[derive(Debug)]
pub(crate) struct WakerSet {
wakers: HashMap<Token, Waker>,
driving_waker: Option<Token>,
next_token: NonZeroUsize,
}
impl WakerSet {
#[must_use]
#[inline]
pub fn new() -> Self {
Self {
wakers: HashMap::new(),
next_token: NonZeroUsize::new(1).unwrap(),
driving_waker: None,
}
}
#[must_use]
pub fn add_waker(&mut self, waker: Waker) -> Token {
let token = Token(self.next_token);
self.next_token = self
.next_token
.get()
.checked_add(1)
.and_then(NonZeroUsize::new)
.expect("Overflow when creating token");
self.wakers.insert(token, waker);
self.driving_waker = Some(token);
token
}
pub fn replace_waker(&mut self, token: Token, waker: &Waker) {
let current_waker = self
.wakers
.get_mut(&token)
.expect("No matching token in wakerset");
if !current_waker.will_wake(&waker) {
current_waker.clone_from(&waker)
}
self.driving_waker = Some(token);
}
pub fn wake_driver(&self) {
if let Some(ref token) = self.driving_waker {
self.wakers
.get(token)
.expect("Driving waker is not present in waker set; this shouldn't be possible")
.wake_by_ref();
}
}
pub fn wake_all(&self) {
self.wakers.values().for_each(|waker| waker.wake_by_ref());
}
pub fn discard_and_wake(&mut self, token: Token) {
self.wakers.remove(&token);
if self.driving_waker == Some(token) || self.driving_waker.is_none() {
match self.wakers.iter().next() {
None => self.driving_waker = None,
Some((&token, waker)) => {
self.driving_waker = Some(token);
waker.wake_by_ref();
}
}
}
}
pub fn discard_wake_all(&mut self, token: Token) {
self.wakers.remove(&token);
self.wake_all()
}
}
impl Default for WakerSet {
fn default() -> Self {
Self::new()
}
}