use crate::monotonic::clock_domain::next_identifier_state;
use qubit_collections::map::OrderedIndexMap;
use std::collections::HashMap;
use std::task::{
Context,
Poll,
Waker,
};
use std::time::Duration;
#[must_use = "the allocated identifier must be retained by its registration"]
#[inline]
pub(crate) fn allocate_identifier(
next_identifier: &mut u64,
exhausted_message: &str,
) -> u64 {
let identifier = *next_identifier;
*next_identifier =
next_identifier_state(identifier).expect(exhausted_message);
identifier
}
pub(crate) struct ManualWaiterRegistry {
next_timer_waiter_id: u64,
timer_waiters: OrderedIndexMap<u64, Duration, Option<Waker>>,
next_observer_id: u64,
count_observers: HashMap<u64, (usize, Option<Waker>)>,
deadline_observers: HashMap<u64, Option<Waker>>,
}
impl ManualWaiterRegistry {
#[must_use]
#[inline]
pub(crate) fn new() -> Self {
Self {
next_timer_waiter_id: 1,
timer_waiters: OrderedIndexMap::new(),
next_observer_id: 1,
count_observers: HashMap::new(),
deadline_observers: HashMap::new(),
}
}
#[must_use = "the waiter identifier is required to poll or cancel the wait"]
#[inline]
pub(crate) fn register_timer(&mut self, deadline: Duration) -> u64 {
let waiter_id = allocate_identifier(
&mut self.next_timer_waiter_id,
"manual timer waiter identifiers exhausted",
);
let inserted = self.timer_waiters.try_insert(waiter_id, deadline, None);
assert!(
inserted.is_ok(),
"manual timer waiter identifier must be unique",
);
waiter_id
}
#[inline(always)]
pub(crate) fn unregister_timer(
&mut self,
waiter_id: u64,
) -> Option<Option<Waker>> {
self.timer_waiters
.remove(&waiter_id)
.map(|entry| entry.into_value())
}
pub(crate) fn next_future_deadline(
&self,
elapsed: Duration,
) -> Option<Duration> {
let deadline = self.timer_waiters.first().map(|entry| *entry.order());
debug_assert!(deadline.is_none_or(|deadline| deadline > elapsed));
deadline
}
#[must_use = "due wakers should be invoked after unlocking"]
pub(crate) fn take_due_timer_wakers(
&mut self,
elapsed: Duration,
) -> Vec<Waker> {
let mut wakers = Vec::new();
let mut due_waiters = self.timer_waiters.detach_range(..=elapsed);
while let Some(waiter) = due_waiters.next() {
if let Some(waker) = waiter.into_value_mut().take() {
wakers.push(waker);
}
}
wakers
}
#[inline]
pub(crate) fn register_observer(
&mut self,
expected_count: usize,
count: usize,
) -> Option<u64> {
if count >= expected_count {
return None;
}
let observer_id = allocate_identifier(
&mut self.next_observer_id,
"manual waiter observer identifiers exhausted",
);
self.count_observers
.insert(observer_id, (expected_count, None));
Some(observer_id)
}
#[must_use = "the observer identifier is required to poll or cancel the wait"]
pub(crate) fn register_deadline_observer(&mut self) -> u64 {
let observer_id = allocate_identifier(
&mut self.next_observer_id,
"manual waiter observer identifiers exhausted",
);
self.deadline_observers.insert(observer_id, None);
observer_id
}
#[must_use = "the poll state and detached waker must both be handled"]
pub(crate) fn poll_observer(
&mut self,
observer_id: u64,
context: &Context<'_>,
) -> (Poll<()>, Option<Waker>) {
let Some((_, registered_waker)) =
self.count_observers.get_mut(&observer_id)
else {
return (Poll::Ready(()), None);
};
let replaced_waker = if registered_waker
.as_ref()
.is_none_or(|waker| !waker.will_wake(context.waker()))
{
registered_waker.replace(context.waker().clone())
} else {
None
};
(Poll::Pending, replaced_waker)
}
#[must_use = "the poll state and detached waker must both be handled"]
pub(crate) fn poll_deadline_observer(
&mut self,
observer_id: u64,
elapsed: Duration,
context: &Context<'_>,
) -> (Poll<Duration>, Option<Waker>) {
if !self.deadline_observers.contains_key(&observer_id) {
panic!("manual deadline observer {observer_id} is not registered");
}
if let Some(deadline) = self.next_future_deadline(elapsed) {
let removed_waker =
self.deadline_observers.remove(&observer_id).flatten();
return (Poll::Ready(deadline), removed_waker);
}
let Some(registered_waker) =
self.deadline_observers.get_mut(&observer_id)
else {
unreachable!("deadline observer existence was checked above");
};
let replaced_waker = if registered_waker
.as_ref()
.is_none_or(|waker| !waker.will_wake(context.waker()))
{
registered_waker.replace(context.waker().clone())
} else {
None
};
(Poll::Pending, replaced_waker)
}
#[inline]
pub(crate) fn unregister_observer(
&mut self,
observer_id: u64,
) -> Option<Waker> {
if let Some((_, waker)) = self.count_observers.remove(&observer_id) {
return waker;
}
self.deadline_observers.remove(&observer_id).flatten()
}
#[must_use]
#[inline(always)]
pub(crate) fn contains_observer(&self, observer_id: u64) -> bool {
self.count_observers.contains_key(&observer_id)
}
#[must_use = "the poll state and detached waker must both be handled"]
pub(crate) fn poll_timer(
&mut self,
waiter_id: u64,
elapsed: Duration,
context: &Context<'_>,
) -> (Poll<()>, Option<Waker>) {
let Some(deadline) = self
.timer_waiters
.get_entry(&waiter_id)
.map(|entry| *entry.order())
else {
panic!("manual timer waiter {waiter_id} is not registered");
};
if elapsed < deadline {
let registered_waker = self
.timer_waiters
.get_mut(&waiter_id)
.expect("manual timer waiter must remain registered");
let replaced_waker = if registered_waker
.as_ref()
.is_none_or(|waker| !waker.will_wake(context.waker()))
{
registered_waker.replace(context.waker().clone())
} else {
None
};
return (Poll::Pending, replaced_waker);
}
let removed_waker = self
.timer_waiters
.remove(&waiter_id)
.and_then(|entry| entry.into_value());
(Poll::Ready(()), removed_waker)
}
#[must_use = "reached observer wakers should be invoked after unlocking"]
pub(crate) fn reached_observer_wakers(
&mut self,
elapsed: Duration,
) -> Vec<Waker> {
let count = self.count();
let next_deadline = self.next_future_deadline(elapsed);
let mut wakers = Vec::new();
self.count_observers.retain(|_, (expected_count, waker)| {
if *expected_count <= count {
if let Some(waker) = waker.take() {
wakers.push(waker);
}
false
} else {
true
}
});
if next_deadline.is_some() {
self.deadline_observers.values_mut().for_each(|waker| {
if let Some(waker) = waker.take() {
wakers.push(waker);
}
});
}
wakers
}
#[must_use]
#[inline(always)]
pub(crate) fn count(&self) -> usize {
self.timer_waiters.len()
}
}