use std::future::poll_fn;
use std::panic::{AssertUnwindSafe, catch_unwind};
use std::sync::{Arc, Mutex, PoisonError};
use std::task::{Poll, Waker};
#[derive(Default)]
struct Counters {
issued: u64,
completed: u64,
waiters: Vec<(u64, Waker)>,
}
#[derive(Default)]
pub(super) struct Fence {
counters: Mutex<Counters>,
}
pub(super) struct Ticket(Arc<Fence>);
impl Fence {
pub(super) fn issue(self: &Arc<Self>) -> Ticket {
self.counters
.lock()
.unwrap_or_else(PoisonError::into_inner)
.issued += 1;
Ticket(Arc::clone(self))
}
pub(super) async fn settled(&self) {
let target = self
.counters
.lock()
.unwrap_or_else(PoisonError::into_inner)
.issued;
poll_fn(|cx| {
let mut counters = self.counters.lock().unwrap_or_else(PoisonError::into_inner);
if counters.completed >= target {
return Poll::Ready(());
}
if !counters
.waiters
.iter()
.any(|(waiting_for, waker)| *waiting_for == target && waker.will_wake(cx.waker()))
{
counters.waiters.push((target, cx.waker().clone()));
}
Poll::Pending
})
.await;
}
}
impl Drop for Ticket {
fn drop(&mut self) {
let due: Vec<Waker> = {
let mut counters = self
.0
.counters
.lock()
.unwrap_or_else(PoisonError::into_inner);
counters.completed += 1;
let completed = counters.completed;
let (due, waiting): (Vec<_>, Vec<_>) = counters
.waiters
.drain(..)
.partition(|(target, _)| *target <= completed);
counters.waiters = waiting;
due.into_iter().map(|(_, waker)| waker).collect()
};
for waker in due {
let _contained = catch_unwind(AssertUnwindSafe(|| waker.wake()));
}
}
}