use std::pin::Pin;
use futures::channel::oneshot::Canceled;
use futures::channel::oneshot::Receiver;
use futures::task::Context;
use futures::task::Poll;
use futures::Future;
use pin_project::pin_project;
use thiserror::Error;
#[pin_project]
pub struct ConservativeReceiver<T>(#[pin] Receiver<T>);
impl<T> ConservativeReceiver<T> {
pub fn new(recv: Receiver<T>) -> Self {
ConservativeReceiver(recv)
}
}
impl<T> Future for ConservativeReceiver<T> {
type Output = Result<T, ConservativeReceiverError>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let mut this = self.project();
match this.0.as_mut().poll(cx) {
Poll::Ready(Ok(output)) => Poll::Ready(Ok(output)),
Poll::Ready(Err(Canceled)) => Poll::Ready(Err(ConservativeReceiverError::Canceled)),
Poll::Pending => Poll::Ready(Err(ConservativeReceiverError::ReceiveBeforeSend)),
}
}
}
#[derive(Error, Debug)]
pub enum ConservativeReceiverError {
#[error("oneshot canceled")]
Canceled,
#[error("recv called on channel before send")]
ReceiveBeforeSend,
}
#[cfg(test)]
mod test {
use assert_matches::assert_matches;
use futures::channel::oneshot::channel;
use super::*;
#[tokio::test]
async fn recv_after_send() {
let (send, recv) = channel();
let recv = ConservativeReceiver::new(recv);
send.send(42).expect("Failed to send");
assert_matches!(recv.await, Ok(42));
}
#[tokio::test]
async fn recv_before_send() {
let (send, recv) = channel();
let recv = ConservativeReceiver::new(recv);
assert_matches!(
recv.await,
Err(ConservativeReceiverError::ReceiveBeforeSend)
);
send.send(42).expect_err("Should fail to send");
}
#[tokio::test]
async fn recv_canceled_send() {
let (_, recv) = channel::<()>();
let recv = ConservativeReceiver::new(recv);
assert_matches!(recv.await, Err(ConservativeReceiverError::Canceled));
}
}