use std::future::Future;
use std::pin::Pin;
use std::sync::{Arc, Mutex, MutexGuard};
use std::task::{Context, Poll, Waker};
use super::PresentError;
pub(crate) type Value<T> = Result<T, PresentError>;
pub(crate) const DROPPED_WITHOUT_OUTCOME: &str = "presentation ended without an outcome";
struct State<T> {
value: Option<Value<T>>,
waker: Option<Waker>,
closed: bool,
}
struct Inner<T> {
state: Mutex<State<T>>,
}
fn lock<T>(inner: &Inner<T>) -> MutexGuard<'_, State<T>> {
inner
.state
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
}
fn complete<T>(inner: &Inner<T>, generation: u64, value: Value<T>) -> bool {
super::release_if_live(generation);
let waker = {
let mut state = lock(inner);
if state.closed {
return false;
}
state.value = Some(value);
state.waker.take()
};
if let Some(waker) = waker {
waker.wake();
}
true
}
pub(crate) struct Sender<T> {
inner: Arc<Inner<T>>,
generation: u64,
sent: bool,
}
impl<T> Sender<T> {
#[cfg_attr(not(test), allow(dead_code))]
pub(crate) fn generation(&self) -> u64 {
self.generation
}
#[cfg_attr(not(test), allow(dead_code))]
pub(crate) fn send(mut self, value: Value<T>) -> bool {
self.sent = true;
complete(&self.inner, self.generation, value)
}
}
impl<T> Drop for Sender<T> {
fn drop(&mut self) {
if self.sent {
return;
}
complete(
&self.inner,
self.generation,
Err(PresentError::Platform(DROPPED_WITHOUT_OUTCOME.to_string())),
);
}
}
pub(crate) struct Receiver<T> {
inner: Arc<Inner<T>>,
generation: u64,
}
impl<T> Drop for Receiver<T> {
fn drop(&mut self) {
lock(&self.inner).closed = true;
super::release_if_live(self.generation);
}
}
impl<T> Future for Receiver<T> {
type Output = Value<T>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let mut state = lock(&self.inner);
if let Some(value) = state.value.take() {
return Poll::Ready(value);
}
state.waker = Some(cx.waker().clone());
Poll::Pending
}
}
pub(crate) fn channel<T>(generation: u64) -> (Sender<T>, Receiver<T>) {
let inner = Arc::new(Inner {
state: Mutex::new(State {
value: None,
waker: None,
closed: false,
}),
});
(
Sender {
inner: Arc::clone(&inner),
generation,
sent: false,
},
Receiver { inner, generation },
)
}
#[cfg(test)]
pub(crate) mod tests {
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::task::{RawWaker, RawWakerVTable, Waker};
use super::*;
const UNCLAIMED_GENERATION: u64 = u64::MAX;
pub(crate) fn counting_waker() -> (Waker, Arc<AtomicUsize>) {
fn clone(data: *const ()) -> RawWaker {
unsafe { Arc::increment_strong_count(data as *const AtomicUsize) };
RawWaker::new(data, &VTABLE)
}
fn wake(data: *const ()) {
let counter = unsafe { Arc::from_raw(data as *const AtomicUsize) };
counter.fetch_add(1, Ordering::SeqCst);
}
fn wake_by_ref(data: *const ()) {
let counter = unsafe { &*(data as *const AtomicUsize) };
counter.fetch_add(1, Ordering::SeqCst);
}
fn drop_fn(data: *const ()) {
unsafe { drop(Arc::from_raw(data as *const AtomicUsize)) };
}
static VTABLE: RawWakerVTable = RawWakerVTable::new(clone, wake, wake_by_ref, drop_fn);
let counter = Arc::new(AtomicUsize::new(0));
let raw = RawWaker::new(Arc::into_raw(Arc::clone(&counter)) as *const (), &VTABLE);
(unsafe { Waker::from_raw(raw) }, counter)
}
fn poll_once<T>(receiver: &mut Receiver<T>, waker: &Waker) -> Poll<Value<T>> {
let mut cx = Context::from_waker(waker);
Pin::new(receiver).poll(&mut cx)
}
#[test]
fn send_before_poll_is_observed_on_first_poll() {
let (sender, mut receiver) = channel::<&str>(UNCLAIMED_GENERATION);
assert!(sender.send(Ok("done")));
let (waker, counter) = counting_waker();
assert_eq!(poll_once(&mut receiver, &waker), Poll::Ready(Ok("done")));
assert_eq!(counter.load(Ordering::SeqCst), 0);
}
#[test]
fn poll_then_send_wakes_the_registered_waker_once() {
let (sender, mut receiver) = channel::<u32>(UNCLAIMED_GENERATION);
let (waker, counter) = counting_waker();
assert_eq!(poll_once(&mut receiver, &waker), Poll::Pending);
assert_eq!(counter.load(Ordering::SeqCst), 0);
assert!(sender.send(Ok(7)));
assert_eq!(counter.load(Ordering::SeqCst), 1);
assert_eq!(poll_once(&mut receiver, &waker), Poll::Ready(Ok(7)));
}
#[test]
fn the_value_is_delivered_once_then_the_receiver_stays_pending() {
let (sender, mut receiver) = channel::<u32>(UNCLAIMED_GENERATION);
assert!(sender.send(Ok(1)));
let (waker, _counter) = counting_waker();
assert_eq!(poll_once(&mut receiver, &waker), Poll::Ready(Ok(1)));
assert_eq!(poll_once(&mut receiver, &waker), Poll::Pending);
}
#[test]
fn drop_without_send_resolves_platform_error() {
let (sender, mut receiver) = channel::<u32>(UNCLAIMED_GENERATION);
drop(sender);
let (waker, _counter) = counting_waker();
assert_eq!(
poll_once(&mut receiver, &waker),
Poll::Ready(Err(PresentError::Platform(
DROPPED_WITHOUT_OUTCOME.to_string()
)))
);
}
#[test]
fn a_late_send_after_the_receiver_dropped_is_discarded() {
let (sender, receiver) = channel::<u32>(UNCLAIMED_GENERATION);
drop(receiver);
assert!(
!sender.send(Ok(3)),
"a send into a closed channel must report the value discarded"
);
}
#[test]
fn both_halves_carry_the_generation() {
let (sender, receiver) = channel::<u32>(UNCLAIMED_GENERATION);
assert_eq!(sender.generation(), UNCLAIMED_GENERATION);
assert_eq!(receiver.generation, UNCLAIMED_GENERATION);
}
}