use std::{
future::Future,
panic::Location,
pin::Pin,
sync::Arc,
task::{Context, Poll, Waker},
};
use parking_lot::Mutex;
use crate::flash::{
diag::PrimKind,
flash_ambient,
ids::{Backend, trace_native_from_ambient},
system,
};
pub mod error {
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct RecvError;
impl std::fmt::Display for RecvError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("oneshot channel closed without a value")
}
}
impl std::error::Error for RecvError {}
}
pub use error::RecvError;
struct Inner<T> {
real_waker: Option<Waker>,
value: Option<T>,
receiver_alive: bool,
sender_alive: bool,
}
struct Shared<T> {
backend: Backend,
inner: Mutex<Inner<T>>,
}
enum Parked {
Engine(system::AsyncHandle),
Real,
}
#[must_use]
#[track_caller]
pub fn channel<T>() -> (Sender<T>, Receiver<T>) {
let shared = Arc::new(Shared {
inner: Mutex::new(Inner {
value: None,
sender_alive: true,
receiver_alive: true,
real_waker: None,
}),
backend: if flash_ambient() {
let cvid = system::next_condvar_id();
system::describe_cvid(cvid, PrimKind::Oneshot, Location::caller());
Backend::Engine(cvid)
} else {
Backend::Native
},
});
(
Sender {
shared: Some(Arc::clone(&shared)),
},
Receiver {
shared,
pending: None,
},
)
}
pub struct Sender<T> {
shared: Option<Arc<Shared<T>>>,
}
impl<T> Sender<T> {
pub fn send(mut self, value: T) -> Result<(), T> {
let Some(shared) = self.shared.take() else {
return Err(value);
};
let mut inner = shared.inner.lock();
if !inner.receiver_alive {
return Err(value);
}
inner.value = Some(value);
let waker = inner.real_waker.take();
drop(inner);
match shared.backend {
Backend::Engine(cvid) => system::signal_channel(cvid, false),
Backend::Native => {
trace_native_from_ambient("oneshot", "send");
if let Some(waker) = waker {
waker.wake();
}
}
}
Ok(())
}
}
impl<T> Drop for Sender<T> {
fn drop(&mut self) {
if let Some(shared) = self.shared.take() {
let mut inner = shared.inner.lock();
inner.sender_alive = false;
let waker = inner.real_waker.take();
drop(inner);
match shared.backend {
Backend::Engine(cvid) => system::signal_channel(cvid, false),
Backend::Native => {
trace_native_from_ambient("oneshot", "sender_drop");
if let Some(waker) = waker {
waker.wake();
}
}
}
}
}
}
pub struct Receiver<T> {
shared: Arc<Shared<T>>,
pending: Option<Parked>,
}
impl<T> Future for Receiver<T> {
type Output = Result<T, RecvError>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let this = self.get_mut();
if let Some(Parked::Engine(handle)) = this.pending.as_ref() {
if handle.granted() {
this.pending = None;
} else {
return Poll::Pending;
}
}
let mut inner = this.shared.inner.lock();
if let Some(value) = inner.value.take() {
return Poll::Ready(Ok(value));
}
if !inner.sender_alive {
return Poll::Ready(Err(RecvError));
}
match this.shared.backend {
Backend::Engine(cvid) => {
let (handle, adv) = system::register_channel_async(cvid, cx.waker().clone());
this.pending = Some(Parked::Engine(handle));
drop(inner);
adv.fire();
}
Backend::Native => {
trace_native_from_ambient("oneshot", "recv_park");
inner.real_waker = Some(cx.waker().clone());
this.pending = Some(Parked::Real);
drop(inner);
}
}
Poll::Pending
}
}
impl<T> Drop for Receiver<T> {
fn drop(&mut self) {
let mut inner = self.shared.inner.lock();
inner.receiver_alive = false;
if matches!(self.pending, Some(Parked::Real)) {
inner.real_waker = None;
}
drop(inner);
if let Some(Parked::Engine(handle)) = self.pending.take() {
system::cancel_async_wait(&handle);
}
}
}
#[cfg(test)]
mod tests {
use futures::future::join_all;
use kithara_test_utils::kithara;
use super::{RecvError, channel};
use crate::{flash, tokio::task::spawn};
const ROUNDS: usize = 256;
#[kithara::test(tokio, multi_thread)]
async fn round_trip_no_lost_wakeup() {
flash::reset();
let futs = (0..ROUNDS).map(|r| async move {
let (tx, rx) = channel::<usize>();
drop(spawn(async move {
let _ = tx.send(r * 2);
}));
rx.await.expect("sender delivered")
});
let got: Vec<usize> = join_all(futs).await;
let sum: usize = got.iter().sum();
assert_eq!(sum, (0..ROUNDS).map(|r| r * 2).sum::<usize>());
}
#[kithara::test(tokio, multi_thread)]
async fn dropped_sender_resolves_recv_error() {
flash::reset();
let (tx, rx) = channel::<usize>();
drop(spawn(async move {
drop(tx);
}));
assert_eq!(rx.await, Err(RecvError));
}
}