use std::{
fmt,
future::Future,
marker::PhantomData,
pin::Pin,
sync::{Arc, OnceLock, Weak},
task::{Context, Poll, Waker},
};
use smallvec::SmallVec;
const INLINE_WAITERS: usize = 32;
pub struct Waiter {
waker: Waker,
shared: OnceLock<Arc<Waker>>,
}
impl Waiter {
pub fn new(waker: Waker) -> Self {
Self {
waker,
shared: OnceLock::new(),
}
}
pub fn noop() -> Self {
Self::new(Waker::noop().clone())
}
pub fn register(&self, list: &mut WaiterList) {
list.register(self);
}
pub fn waker(&self) -> &Waker {
&self.waker
}
fn shared(&self) -> &Arc<Waker> {
self.shared.get_or_init(|| Arc::new(self.waker.clone()))
}
pub fn poll_future<F: Future + ?Sized>(&self, future: Pin<&mut F>) -> Poll<F::Output> {
future.poll(&mut Context::from_waker(self.waker()))
}
}
pub struct WaiterList {
entries: SmallVec<[Weak<Waker>; INLINE_WAITERS]>,
cursor: usize,
}
impl WaiterList {
pub fn new() -> Self {
Self {
entries: SmallVec::new(),
cursor: 0,
}
}
pub fn register(&mut self, waiter: &Waiter) {
let new_weak = Arc::downgrade(waiter.shared());
for _ in 0..self.entries.len().min(2) {
if self.entries[self.cursor].strong_count() == 0 {
self.entries[self.cursor] = new_weak;
return;
}
self.cursor = (self.cursor + 1) % self.entries.len();
}
self.entries.push(new_weak);
}
pub fn take(&mut self) -> Self {
self.cursor = 0;
Self {
entries: std::mem::take(&mut self.entries),
cursor: 0,
}
}
pub fn wake(&mut self) {
self.cursor = 0;
for waker in self.entries.drain(..).filter_map(|w| w.upgrade()) {
waker.wake_by_ref();
}
}
}
impl Default for WaiterList {
fn default() -> Self {
Self::new()
}
}
impl fmt::Debug for WaiterList {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("WaiterList").field("len", &self.entries.len()).finish()
}
}
struct WaiterFn<F, R> {
poll: F,
waiter: Option<Waiter>, _marker: PhantomData<fn() -> R>,
}
pub fn wait<F, R>(poll: F) -> impl Future<Output = R>
where
F: FnMut(&Waiter) -> Poll<R> + Unpin,
{
WaiterFn {
poll,
waiter: None,
_marker: PhantomData,
}
}
impl<F, R> Future for WaiterFn<F, R>
where
F: FnMut(&Waiter) -> Poll<R> + Unpin,
{
type Output = R;
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<R> {
let this = &mut *self;
this.waiter = Some(Waiter::new(cx.waker().clone()));
(this.poll)(this.waiter.as_ref().unwrap())
}
}
#[cfg(all(test, not(loom)))]
mod tests {
use super::*;
#[test]
fn poll_future_bridges_a_std_future() {
let waiter = Waiter::noop();
let fut = std::pin::pin!(std::future::ready(7u8));
assert_eq!(waiter.poll_future(fut), Poll::Ready(7));
let fut = std::pin::pin!(std::future::pending::<u8>());
assert_eq!(waiter.poll_future(fut), Poll::Pending);
let mut boxed: Pin<Box<dyn Future<Output = u8>>> = Box::pin(std::future::ready(9u8));
assert_eq!(waiter.poll_future(boxed.as_mut()), Poll::Ready(9));
}
const fn assert_sync<T: Sync>() {}
const _: () = {
assert_sync::<Waiter>();
assert_sync::<crate::Pending<crate::Consumer<u32>>>();
assert_sync::<crate::Shared<u32>>();
};
#[test]
fn wait_output_need_not_be_unpin() {
struct NotUnpin(#[allow(dead_code)] std::marker::PhantomPinned);
let mut fut = std::pin::pin!(crate::wait(|_| Poll::Ready(NotUnpin(std::marker::PhantomPinned))));
let mut cx = Context::from_waker(Waker::noop());
assert!(fut.as_mut().poll(&mut cx).is_ready());
}
}