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::with_capacity(self.pending.len());
for id in self.pending.drain(..) {
if let Some(entry) = self.waiters.get_mut(&id)
&& 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() {
if !entry.waker.will_wake(waker) {
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::{
sync::{Arc, Mutex},
task::{Wake, Waker},
};
struct OrderedWake {
id: usize,
order: Arc<Mutex<Vec<usize>>>,
}
impl Wake for OrderedWake {
fn wake(self: Arc<Self>) {
self.order
.lock()
.expect("wake-order lock must remain available")
.push(self.id);
}
fn wake_by_ref(self: &Arc<Self>) {
self.order
.lock()
.expect("wake-order lock must remain available")
.push(self.id);
}
}
fn noop_waker() -> Waker {
Waker::noop().clone()
}
fn ordered_waker(id: usize, order: &Arc<Mutex<Vec<usize>>>) -> Waker {
Waker::from(Arc::new(OrderedWake {
id,
order: Arc::clone(order),
}))
}
#[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());
}
#[test]
fn grant_all_skips_cancelled_ids_and_preserves_wake_order() {
let order = Arc::new(Mutex::new(Vec::new()));
let mut queue = WaitQueue::new();
let first = queue.register(ordered_waker(1, &order));
let cancelled = queue.register(ordered_waker(2, &order));
let third = queue.register(ordered_waker(3, &order));
assert_eq!(queue.deregister(cancelled), None);
for waker in queue.grant_all(31_u8) {
waker.wake();
}
assert_eq!(
*order.lock().expect("wake-order lock must remain available"),
vec![1, 3]
);
assert!(matches!(
queue.poll_waiter(first, &noop_waker()),
WaiterPoll::Granted(31)
));
assert!(matches!(
queue.poll_waiter(cancelled, &noop_waker()),
WaiterPoll::NotRegistered
));
assert!(matches!(
queue.poll_waiter(third, &noop_waker()),
WaiterPoll::Granted(31)
));
assert!(queue.is_empty());
}
}