use std::collections::{BTreeMap, VecDeque};
use std::task::Waker;
pub(crate) struct WaitQueue<G> {
waiters: BTreeMap<u64, WaitEntry<G>>,
pending: VecDeque<u64>,
next_id: u64,
}
struct WaitEntry<G> {
waker: Waker,
granted: Option<G>,
}
pub(crate) enum WaiterPoll<G> {
Granted(G),
Pending,
NotRegistered,
}
impl<G> WaitQueue<G> {
pub(crate) const fn new() -> Self {
Self {
waiters: BTreeMap::new(),
pending: VecDeque::new(),
next_id: 0,
}
}
pub(crate) fn is_empty(&self) -> bool {
self.waiters.is_empty()
}
pub(crate) fn register(&mut self, waker: Waker) -> u64 {
let id = self.next_id;
self.next_id += 1;
self.waiters.insert(
id,
WaitEntry {
waker,
granted: None,
},
);
self.pending.push_back(id);
id
}
pub(crate) fn grant_oldest(&mut self, payload: G) -> Option<Waker> {
while let Some(&id) = self.pending.front() {
match self.waiters.get_mut(&id) {
Some(entry) if entry.granted.is_none() => {
entry.granted = Some(payload);
self.pending.pop_front();
return Some(entry.waker.clone());
}
_ => {
self.pending.pop_front();
}
}
}
None
}
pub(crate) fn grant_all(&mut self, payload: G) -> Vec<Waker>
where
G: Clone,
{
let mut wakers = Vec::new();
let pending_ids: Vec<u64> = self.pending.drain(..).collect();
for id in pending_ids {
if let Some(entry) = self.waiters.get_mut(&id) {
if entry.granted.is_none() {
entry.granted = Some(payload.clone());
wakers.push(entry.waker.clone());
}
}
}
wakers
}
pub(crate) fn poll_waiter(&mut self, id: u64, waker: &Waker) -> WaiterPoll<G> {
let Some(entry) = self.waiters.get_mut(&id) else {
return WaiterPoll::NotRegistered;
};
if entry.granted.is_none() {
entry.waker = waker.clone();
return WaiterPoll::Pending;
}
let entry = self
.waiters
.remove(&id)
.expect("invariant: entry present, checked above");
WaiterPoll::Granted(
entry
.granted
.expect("invariant: grant present, checked above"),
)
}
pub(crate) fn deregister(&mut self, id: u64) -> Option<G> {
self.waiters.remove(&id).and_then(|entry| entry.granted)
}
}
#[cfg(test)]
mod tests {
use super::{WaitQueue, WaiterPoll};
use std::task::Waker;
fn noop_waker() -> Waker {
Waker::noop().clone()
}
#[test]
fn cancelled_head_does_not_block_fifo_grant() {
let mut queue = WaitQueue::new();
let cancelled = queue.register(noop_waker());
let active = queue.register(noop_waker());
assert_eq!(queue.deregister(cancelled), None);
assert!(queue.grant_oldest(17_u8).is_some());
assert!(matches!(
queue.poll_waiter(active, &noop_waker()),
WaiterPoll::Granted(17)
));
assert!(queue.is_empty());
}
#[test]
fn grant_all_preserves_every_pending_payload() {
let mut queue = WaitQueue::new();
let first = queue.register(noop_waker());
let second = queue.register(noop_waker());
assert_eq!(queue.grant_all(29_u8).len(), 2);
assert!(matches!(
queue.poll_waiter(first, &noop_waker()),
WaiterPoll::Granted(29)
));
assert!(matches!(
queue.poll_waiter(second, &noop_waker()),
WaiterPoll::Granted(29)
));
assert!(queue.is_empty());
}
}