use std::{
future::Future,
panic::Location,
pin::Pin,
sync::{
Arc,
atomic::{AtomicUsize, Ordering},
},
task::{Context, Poll, Waker},
};
use parking_lot::Mutex;
use crate::{
flash::{
diag::PrimKind,
flash_ambient,
ids::{Backend, trace_native_from_ambient},
system,
},
native::tokio::sync::broadcast as inner,
};
pub mod error {
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RecvError {
Closed,
Lagged(u64),
}
impl std::fmt::Display for RecvError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Closed => f.write_str("channel closed"),
Self::Lagged(n) => write!(f, "channel lagged by {n}"),
}
}
}
impl std::error::Error for RecvError {}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum TryRecvError {
Empty,
Closed,
Lagged(u64),
}
impl std::fmt::Display for TryRecvError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Empty => f.write_str("channel empty"),
Self::Closed => f.write_str("channel closed"),
Self::Lagged(n) => write!(f, "channel lagged by {n}"),
}
}
}
impl std::error::Error for TryRecvError {}
pub struct SendError<T>(pub T);
impl<T> std::fmt::Debug for SendError<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("SendError(..)")
}
}
impl<T> std::fmt::Display for SendError<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("sending on a channel with no receivers")
}
}
impl<T> std::error::Error for SendError<T> {}
}
pub use error::{RecvError, SendError, TryRecvError};
struct Shared {
senders: AtomicUsize,
backend: Backend,
gate: Mutex<Gate>,
}
#[derive(Default)]
struct Gate {
wakers: Vec<Waker>,
closed: bool,
}
impl Shared {
fn close(&self) {
let mut gate = self.gate.lock();
gate.closed = true;
let drained = std::mem::take(&mut gate.wakers);
drop(gate);
match self.backend {
Backend::Engine(cvid) => system::signal_channel(cvid, true),
Backend::Native => {
trace_native_from_ambient("broadcast", "close");
for waker in drained {
waker.wake();
}
}
}
}
fn signal(&self) {
let mut gate = self.gate.lock();
let drained = std::mem::take(&mut gate.wakers);
drop(gate);
match self.backend {
Backend::Engine(cvid) => system::signal_channel(cvid, true),
Backend::Native => {
trace_native_from_ambient("broadcast", "send");
for waker in drained {
waker.wake();
}
}
}
}
}
#[must_use]
#[track_caller]
pub fn channel<T: Clone>(capacity: usize) -> (Sender<T>, Receiver<T>) {
let (tx, rx) = inner::channel(capacity);
let shared = Arc::new(Shared {
gate: Mutex::new(Gate::default()),
senders: AtomicUsize::new(1),
backend: if flash_ambient() {
let cvid = system::next_condvar_id();
system::describe_cvid(cvid, PrimKind::Broadcast, Location::caller());
Backend::Engine(cvid)
} else {
Backend::Native
},
});
(
Sender {
inner: tx,
shared: Arc::clone(&shared),
},
Receiver {
shared,
inner: rx,
pending: None,
},
)
}
pub struct Sender<T> {
shared: Arc<Shared>,
inner: inner::Sender<T>,
}
impl<T> Clone for Sender<T> {
fn clone(&self) -> Self {
self.shared.senders.fetch_add(1, Ordering::AcqRel);
Self {
inner: self.inner.clone(),
shared: Arc::clone(&self.shared),
}
}
}
impl<T> Drop for Sender<T> {
fn drop(&mut self) {
if self.shared.senders.fetch_sub(1, Ordering::AcqRel) == 1 {
self.shared.close();
}
}
}
impl<T> Sender<T> {
#[must_use]
pub fn receiver_count(&self) -> usize {
self.inner.receiver_count()
}
#[must_use]
pub fn subscribe(&self) -> Receiver<T> {
Receiver {
inner: self.inner.subscribe(),
shared: Arc::clone(&self.shared),
pending: None,
}
}
}
impl<T: Clone> Sender<T> {
pub fn send(&self, value: T) -> Result<usize, SendError<T>> {
let result = self.inner.send(value);
self.shared.signal();
result.map_err(|e| SendError(e.0))
}
}
pub struct Receiver<T> {
shared: Arc<Shared>,
pending: Option<Parked>,
inner: inner::Receiver<T>,
}
enum Parked {
Engine(system::AsyncHandle),
Real(Waker),
}
impl<T: Clone> Receiver<T> {
pub fn recv(&mut self) -> Recv<'_, T> {
Recv { rx: self }
}
pub fn try_recv(&mut self) -> Result<T, TryRecvError> {
match self.inner.try_recv() {
Ok(value) => Ok(value),
Err(inner::error::TryRecvError::Empty) => Err(TryRecvError::Empty),
Err(inner::error::TryRecvError::Closed) => Err(TryRecvError::Closed),
Err(inner::error::TryRecvError::Lagged(n)) => Err(TryRecvError::Lagged(n)),
}
}
}
pub struct Recv<'a, T> {
rx: &'a mut Receiver<T>,
}
impl<T: Clone> Future for Recv<'_, T> {
type Output = Result<T, RecvError>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let rx = &mut *self.get_mut().rx;
match rx.pending.as_ref() {
Some(Parked::Engine(handle)) => {
if handle.granted() {
rx.pending = None;
} else {
return Poll::Pending;
}
}
Some(Parked::Real(_)) => rx.pending = None,
None => {}
}
let mut gate = rx.shared.gate.lock();
match rx.inner.try_recv() {
Ok(value) => {
drop(gate);
Poll::Ready(Ok(value))
}
Err(inner::error::TryRecvError::Lagged(n)) => {
drop(gate);
Poll::Ready(Err(RecvError::Lagged(n)))
}
Err(inner::error::TryRecvError::Closed) => {
drop(gate);
Poll::Ready(Err(RecvError::Closed))
}
Err(inner::error::TryRecvError::Empty) => {
if gate.closed {
drop(gate);
return Poll::Ready(Err(RecvError::Closed));
}
match rx.shared.backend {
Backend::Engine(cvid) => {
let (handle, adv) =
system::register_channel_async(cvid, cx.waker().clone());
rx.pending = Some(Parked::Engine(handle));
drop(gate);
adv.fire();
}
Backend::Native => {
trace_native_from_ambient("broadcast", "recv_park");
let waker = cx.waker().clone();
gate.wakers.push(waker.clone());
rx.pending = Some(Parked::Real(waker));
drop(gate);
}
}
Poll::Pending
}
}
}
}
impl<T> Drop for Recv<'_, T> {
fn drop(&mut self) {
match self.rx.pending.take() {
Some(Parked::Real(waker)) => {
self.rx
.shared
.gate
.lock()
.wakers
.retain(|w| !w.will_wake(&waker));
}
Some(Parked::Engine(handle)) => system::cancel_async_wait(&handle),
None => {}
}
}
}
#[cfg(test)]
mod tests {
use kithara_test_utils::kithara;
use super::{RecvError, TryRecvError, channel};
use crate::{
flash,
tokio::task::{spawn, yield_now},
};
struct Consts;
impl Consts {
const MSGS: usize = 100;
const SUBS: usize = 4;
}
#[kithara::test(tokio, multi_thread)]
async fn fan_out_no_lost_wakeup() {
flash::reset();
let (tx, _rx0) = channel::<usize>(Consts::MSGS + 1);
let handles: Vec<_> = (0..Consts::SUBS)
.map(|_| {
let mut rx = tx.subscribe();
spawn(async move {
let mut got = Vec::new();
while let Ok(value) = rx.recv().await {
got.push(value);
if got.len() == Consts::MSGS {
break;
}
}
got
})
})
.collect();
for i in 0..Consts::MSGS {
tx.send(i).expect("subscribers present");
}
for handle in handles {
assert_eq!(
handle.await.expect("task joined"),
(0..Consts::MSGS).collect::<Vec<_>>()
);
}
}
#[kithara::test(tokio, multi_thread)]
async fn drop_senders_closes_receiver() {
flash::reset();
let (tx, mut rx) = channel::<usize>(4);
let waiter = spawn(async move { rx.recv().await });
yield_now().await;
drop(tx);
assert_eq!(waiter.await.expect("task joined"), Err(RecvError::Closed));
}
#[kithara::test(tokio, multi_thread)]
async fn overflow_reports_lagged() {
flash::reset();
let (tx, mut rx) = channel::<usize>(2);
for i in 0..5 {
tx.send(i).expect("receiver present");
}
assert_eq!(rx.try_recv(), Err(TryRecvError::Lagged(3)));
assert_eq!(rx.recv().await, Ok(3));
assert_eq!(rx.recv().await, Ok(4));
}
}