use std::{
future::poll_fn,
sync::Mutex,
task::{Context, Poll, Waker},
};
use tokio::sync::mpsc;
#[derive(Debug)]
pub(crate) struct SharedRecv<T> {
inner: Mutex<Inner<T>>,
}
#[derive(Debug)]
struct Inner<T> {
rx: mpsc::Receiver<T>,
wakers: Vec<Waker>,
}
impl<T> SharedRecv<T> {
pub fn new(rx: mpsc::Receiver<T>) -> Self {
Self {
inner: Mutex::new(Inner {
rx,
wakers: Vec::new(),
}),
}
}
pub fn poll_recv(&self, cx: &mut Context<'_>) -> Poll<Option<T>> {
let mut inner = self.inner.lock().unwrap();
match inner.rx.poll_recv(cx) {
Poll::Ready(value) => {
let wakers = std::mem::take(&mut inner.wakers);
drop(inner);
for waker in wakers {
waker.wake();
}
Poll::Ready(value)
}
Poll::Pending => {
if !inner.wakers.iter().any(|w| w.will_wake(cx.waker())) {
inner.wakers.push(cx.waker().clone());
}
Poll::Pending
}
}
}
pub async fn recv(&self) -> Option<T> {
poll_fn(|cx| self.poll_recv(cx)).await
}
}