use std::{
collections::VecDeque,
future::Future,
panic::Location,
pin::Pin,
sync::Arc,
task::{Context, Poll, Waker},
};
use error::{SendError, TryRecvError, TrySendError};
use parking_lot::Mutex;
pub use super::unbounded::UnboundedSender;
use crate::flash::{
diag::PrimKind,
flash_ambient,
ids::{CvId, trace_native_from_ambient},
system,
};
pub use crate::native::tokio::sync::mpsc::error;
pub(super) struct Inner<T> {
data_waker: Option<Waker>,
space_wakers: Vec<Waker>,
queue: VecDeque<T>,
receiver_alive: bool,
senders: usize,
}
#[derive(Clone, Copy)]
enum Backend {
Engine {
data: CvId,
space: CvId,
},
Native,
}
pub(super) struct Shared<T> {
backend: Backend,
inner: Mutex<Inner<T>>,
capacity: Option<usize>,
}
impl<T> Shared<T> {
#[track_caller]
fn new(capacity: Option<usize>) -> Arc<Self> {
Arc::new(Self {
capacity,
inner: Mutex::new(Inner {
queue: VecDeque::new(),
senders: 1,
receiver_alive: true,
data_waker: None,
space_wakers: Vec::new(),
}),
backend: if flash_ambient() {
let data = system::next_condvar_id();
let space = system::next_condvar_id();
system::describe_cvid(data, PrimKind::MpscData, Location::caller());
system::describe_cvid(space, PrimKind::MpscSpace, Location::caller());
Backend::Engine { data, space }
} else {
Backend::Native
},
})
}
pub(super) fn add_sender(&self) {
self.inner.lock().senders += 1;
}
fn wake_data(&self, waker: Option<Waker>) {
match self.backend {
Backend::Engine { data, .. } => system::signal_channel(data, false),
Backend::Native => {
trace_native_from_ambient("mpsc", "wake_data");
if let Some(waker) = waker {
waker.wake();
}
}
}
}
fn wake_space(&self, all: bool, wakers: Vec<Waker>) {
match self.backend {
Backend::Engine { space, .. } => system::signal_channel(space, all),
Backend::Native => {
trace_native_from_ambient("mpsc", "wake_space");
for w in wakers {
w.wake();
}
}
}
}
}
enum Parked {
Engine(system::AsyncHandle),
Real(Waker),
}
#[must_use]
#[track_caller]
pub fn channel<T>(capacity: usize) -> (Sender<T>, Receiver<T>) {
let capacity = capacity.max(1);
let shared = Shared::new(Some(capacity));
(
Sender {
shared: Arc::clone(&shared),
capacity,
},
Receiver {
shared,
pending: None,
},
)
}
#[must_use]
#[track_caller]
pub fn unbounded_channel<T>() -> (UnboundedSender<T>, UnboundedReceiver<T>) {
let shared = Shared::new(None);
(
UnboundedSender {
shared: Arc::clone(&shared),
},
UnboundedReceiver {
shared,
pending: None,
},
)
}
pub(super) fn push_unbounded<T>(shared: &Shared<T>, value: T) -> Result<(), SendError<T>> {
let mut inner = shared.inner.lock();
if !inner.receiver_alive {
return Err(SendError(value));
}
inner.queue.push_back(value);
let waker = inner.data_waker.take();
drop(inner);
shared.wake_data(waker);
Ok(())
}
fn take_one_space_waker<T>(backend: Backend, inner: &mut Inner<T>) -> Vec<Waker> {
if matches!(backend, Backend::Engine { .. }) || inner.space_wakers.is_empty() {
Vec::new()
} else {
vec![inner.space_wakers.remove(0)]
}
}
fn poll_recv_inner<T>(
shared: &Shared<T>,
pending: &mut Option<Parked>,
cx: &mut Context<'_>,
) -> Poll<Option<T>> {
if let Some(Parked::Engine(handle)) = pending.as_ref() {
if handle.granted() {
*pending = None;
} else {
return Poll::Pending;
}
}
let mut inner = shared.inner.lock();
if let Some(value) = inner.queue.pop_front() {
let bounded = shared.capacity.is_some();
let wakers = if bounded {
take_one_space_waker(shared.backend, &mut inner)
} else {
Vec::new()
};
drop(inner);
if bounded {
shared.wake_space(false, wakers);
}
return Poll::Ready(Some(value));
}
if inner.senders == 0 {
return Poll::Ready(None);
}
match shared.backend {
Backend::Engine { data, .. } => {
let (handle, adv) = system::register_channel_async(data, cx.waker().clone());
*pending = Some(Parked::Engine(handle));
drop(inner);
adv.fire();
}
Backend::Native => {
trace_native_from_ambient("mpsc", "recv_park");
let waker = cx.waker().clone();
inner.data_waker = Some(waker.clone());
*pending = Some(Parked::Real(waker));
drop(inner);
}
}
Poll::Pending
}
fn try_recv_inner<T>(shared: &Shared<T>) -> Result<T, TryRecvError> {
let mut inner = shared.inner.lock();
let bounded = shared.capacity.is_some();
match inner.queue.pop_front() {
Some(value) => {
let wakers = if bounded {
take_one_space_waker(shared.backend, &mut inner)
} else {
Vec::new()
};
drop(inner);
if bounded {
shared.wake_space(false, wakers);
}
Ok(value)
}
None if inner.senders == 0 => Err(TryRecvError::Disconnected),
None => Err(TryRecvError::Empty),
}
}
fn close_receiver<T>(shared: &Shared<T>, pending: &mut Option<Parked>) {
let mut inner = shared.inner.lock();
inner.receiver_alive = false;
let wakers = std::mem::take(&mut inner.space_wakers);
if matches!(pending, Some(Parked::Real(_))) {
inner.data_waker = None;
}
drop(inner);
shared.wake_space(true, wakers);
if let Some(Parked::Engine(handle)) = pending.take() {
system::cancel_async_wait(&handle);
}
}
pub(super) fn drop_sender<T>(shared: &Shared<T>) {
let mut inner = shared.inner.lock();
inner.senders -= 1;
let last = inner.senders == 0;
let waker = if last { inner.data_waker.take() } else { None };
drop(inner);
if last {
shared.wake_data(waker);
}
}
pub struct Sender<T> {
shared: Arc<Shared<T>>,
capacity: usize,
}
impl<T> Sender<T> {
pub fn send(&self, value: T) -> Send<'_, T> {
Send {
shared: &self.shared,
value: Some(value),
pending: None,
}
}
pub fn try_send(&self, value: T) -> Result<(), TrySendError<T>> {
let mut inner = self.shared.inner.lock();
if !inner.receiver_alive {
return Err(TrySendError::Closed(value));
}
if inner.queue.len() >= self.capacity {
return Err(TrySendError::Full(value));
}
inner.queue.push_back(value);
let waker = inner.data_waker.take();
drop(inner);
self.shared.wake_data(waker);
Ok(())
}
}
impl<T> Clone for Sender<T> {
fn clone(&self) -> Self {
self.shared.add_sender();
Self {
shared: Arc::clone(&self.shared),
capacity: self.capacity,
}
}
}
impl<T> Drop for Sender<T> {
fn drop(&mut self) {
drop_sender(&self.shared);
}
}
pub struct Send<'a, T> {
shared: &'a Shared<T>,
pending: Option<Parked>,
value: Option<T>,
}
impl<T> Unpin for Send<'_, T> {}
impl<T> Future for Send<'_, T> {
type Output = Result<(), SendError<T>>;
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 !inner.receiver_alive {
drop(inner);
return Poll::Ready(
this.value
.take()
.map_or_else(|| Ok(()), |value| Err(SendError(value))),
);
}
let cap = this.shared.capacity.unwrap_or(usize::MAX);
if inner.queue.len() < cap {
if let Some(value) = this.value.take() {
inner.queue.push_back(value);
let waker = inner.data_waker.take();
drop(inner);
this.shared.wake_data(waker);
}
return Poll::Ready(Ok(()));
}
match this.shared.backend {
Backend::Engine { space, .. } => {
let (handle, adv) = system::register_channel_async(space, cx.waker().clone());
this.pending = Some(Parked::Engine(handle));
drop(inner);
adv.fire();
}
Backend::Native => {
trace_native_from_ambient("mpsc", "send_park");
let waker = cx.waker().clone();
inner.space_wakers.push(waker.clone());
this.pending = Some(Parked::Real(waker));
drop(inner);
}
}
Poll::Pending
}
}
impl<T> Drop for Send<'_, T> {
fn drop(&mut self) {
match self.pending.take() {
Some(Parked::Real(waker)) => {
let mut inner = self.shared.inner.lock();
inner.space_wakers.retain(|w| !w.will_wake(&waker));
}
Some(Parked::Engine(handle)) => system::cancel_async_wait(&handle),
None => {}
}
}
}
pub struct Receiver<T> {
shared: Arc<Shared<T>>,
pending: Option<Parked>,
}
impl<T> Receiver<T> {
pub fn poll_recv(&mut self, cx: &mut Context<'_>) -> Poll<Option<T>> {
poll_recv_inner(&self.shared, &mut self.pending, cx)
}
pub fn recv(&mut self) -> Recv<'_, T> {
Recv {
shared: &self.shared,
pending: &mut self.pending,
}
}
pub fn try_recv(&mut self) -> Result<T, TryRecvError> {
try_recv_inner(&self.shared)
}
}
impl<T> Drop for Receiver<T> {
fn drop(&mut self) {
close_receiver(&self.shared, &mut self.pending);
}
}
pub struct UnboundedReceiver<T> {
shared: Arc<Shared<T>>,
pending: Option<Parked>,
}
impl<T> UnboundedReceiver<T> {
pub fn poll_recv(&mut self, cx: &mut Context<'_>) -> Poll<Option<T>> {
poll_recv_inner(&self.shared, &mut self.pending, cx)
}
pub fn recv(&mut self) -> Recv<'_, T> {
Recv {
shared: &self.shared,
pending: &mut self.pending,
}
}
pub fn try_recv(&mut self) -> Result<T, TryRecvError> {
try_recv_inner(&self.shared)
}
}
impl<T> Drop for UnboundedReceiver<T> {
fn drop(&mut self) {
close_receiver(&self.shared, &mut self.pending);
}
}
pub struct Recv<'a, T> {
pending: &'a mut Option<Parked>,
shared: &'a Shared<T>,
}
impl<T> Future for Recv<'_, T> {
type Output = Option<T>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<T>> {
let this = self.get_mut();
poll_recv_inner(this.shared, this.pending, cx)
}
}
#[cfg(test)]
mod tests {
use kithara_test_utils::kithara;
use super::{channel, error::TrySendError, unbounded_channel};
use crate::{flash, tokio::task::spawn};
struct Consts;
impl Consts {
const PER_PRODUCER: usize = 200;
const PRODUCERS: usize = 8;
}
#[kithara::test(tokio, multi_thread)]
async fn bounded_fan_in_no_lost_wakeup() {
flash::reset();
let (tx, mut rx) = channel::<usize>(4);
for p in 0..Consts::PRODUCERS {
let tx = tx.clone();
drop(spawn(async move {
for i in 0..Consts::PER_PRODUCER {
tx.send(p * Consts::PER_PRODUCER + i)
.await
.expect("receiver alive");
}
}));
}
drop(tx);
let mut seen = 0usize;
let mut sum = 0u64;
while let Some(v) = rx.recv().await {
seen += 1;
sum += v as u64;
}
assert_eq!(seen, Consts::PRODUCERS * Consts::PER_PRODUCER);
let n = (Consts::PRODUCERS * Consts::PER_PRODUCER) as u64;
assert_eq!(sum, n * (n - 1) / 2);
}
#[kithara::test(tokio, multi_thread)]
async fn unbounded_fan_in_no_lost_wakeup() {
flash::reset();
let (tx, mut rx) = unbounded_channel::<usize>();
for p in 0..Consts::PRODUCERS {
let tx = tx.clone();
drop(spawn(async move {
for i in 0..Consts::PER_PRODUCER {
tx.send(p * Consts::PER_PRODUCER + i)
.expect("receiver alive");
}
}));
}
drop(tx);
let mut seen = 0usize;
while (rx.recv().await).is_some() {
seen += 1;
}
assert_eq!(seen, Consts::PRODUCERS * Consts::PER_PRODUCER);
}
#[kithara::test(tokio, multi_thread)]
async fn drop_senders_closes_receiver() {
flash::reset();
let (tx, mut rx) = channel::<usize>(2);
let handle = spawn(async move {
tx.send(1).await.expect("alive");
tx.send(2).await.expect("alive");
drop(tx);
});
assert_eq!(rx.recv().await, Some(1));
assert_eq!(rx.recv().await, Some(2));
assert_eq!(rx.recv().await, None);
handle.await.expect("producer joined");
}
#[kithara::test(tokio, multi_thread)]
async fn bounded_try_send_reports_full_and_closed() {
flash::reset();
let (tx, mut rx) = channel::<usize>(1);
tx.try_send(1).expect("receiver alive");
assert!(matches!(tx.try_send(2), Err(TrySendError::Full(2))));
assert_eq!(rx.try_recv(), Ok(1));
drop(rx);
assert!(matches!(tx.try_send(3), Err(TrySendError::Closed(3))));
}
}