use crate::sync::ring_deque::{self, RingDeque};
use core::{fmt, marker::PhantomPinned, pin::Pin, task::Poll};
use event_listener_strategy::{
easy_wrapper,
event_listener::{Event, EventListener},
EventListenerFuture, Strategy,
};
use pin_project_lite::pin_project;
use s2n_quic_core::ready;
use std::sync::{
atomic::{AtomicUsize, Ordering},
Arc, Weak,
};
pub use ring_deque::{Closed, Priority};
pub fn new<T>(cap: usize) -> (Sender<T>, Receiver<T>) {
assert!(cap >= 1, "capacity must be at least 2");
let channel = Arc::new(Channel {
queue: RingDeque::new(cap),
recv_ops: Event::new(),
sender_count: AtomicUsize::new(1),
receiver_count: AtomicUsize::new(1),
});
let s = Sender {
channel: channel.clone(),
};
let r = Receiver {
listener: None,
channel,
_pin: PhantomPinned,
};
(s, r)
}
struct Channel<T> {
queue: RingDeque<T>,
recv_ops: Event,
sender_count: AtomicUsize,
receiver_count: AtomicUsize,
}
impl<T> Channel<T> {
fn close(&self) -> Result<(), Closed> {
self.queue.close()?;
self.recv_ops.notify(usize::MAX);
Ok(())
}
}
pub struct Sender<T> {
channel: Arc<Channel<T>>,
}
impl<T> Sender<T> {
#[inline]
pub fn send_back(&self, msg: T) -> Result<Option<T>, Closed> {
let res = self.channel.queue.push_back(msg)?;
self.channel.recv_ops.notify_additional(1);
Ok(res)
}
#[inline]
pub fn send_front(&self, msg: T) -> Result<Option<T>, Closed> {
let res = self.channel.queue.push_front(msg)?;
self.channel.recv_ops.notify_additional(1);
Ok(res)
}
}
impl<T> Drop for Sender<T> {
fn drop(&mut self) {
if self.channel.sender_count.fetch_sub(1, Ordering::AcqRel) == 1 {
let _ = self.channel.close();
}
}
}
impl<T> fmt::Debug for Sender<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "Sender {{ .. }}")
}
}
impl<T> Clone for Sender<T> {
fn clone(&self) -> Sender<T> {
let count = self.channel.sender_count.fetch_add(1, Ordering::Relaxed);
assert!(count < usize::MAX / 2, "too many senders");
Sender {
channel: self.channel.clone(),
}
}
}
pin_project! {
pub struct Receiver<T> {
channel: Arc<Channel<T>>,
listener: Option<EventListener>,
#[pin]
_pin: PhantomPinned
}
impl<T> PinnedDrop for Receiver<T> {
fn drop(this: Pin<&mut Self>) {
let this = this.project();
if this.channel.receiver_count.fetch_sub(1, Ordering::AcqRel) == 1 {
let _ = this.channel.close();
}
}
}
}
impl<T> Receiver<T> {
#[inline]
pub fn try_recv_front(&self) -> Result<Option<T>, Closed> {
self.channel.queue.pop_front()
}
#[inline]
pub fn try_recv_back(&self) -> Result<Option<T>, Closed> {
self.channel.queue.pop_back()
}
#[inline]
pub fn recv_front(&self) -> Recv<'_, T> {
Recv::_new(RecvInner {
receiver: self,
pop_end: PopEnd::Front,
listener: None,
_pin: PhantomPinned,
})
}
#[inline]
pub fn recv_back(&self) -> Recv<'_, T> {
Recv::_new(RecvInner {
receiver: self,
pop_end: PopEnd::Back,
listener: None,
_pin: PhantomPinned,
})
}
#[inline]
pub fn downgrade(&self) -> WeakReceiver<T> {
WeakReceiver {
channel: Arc::downgrade(&self.channel),
}
}
#[inline]
pub fn close(&self) -> Result<(), Closed> {
self.channel.close()
}
}
impl<T> fmt::Debug for Receiver<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "Receiver {{ .. }}")
}
}
impl<T> Clone for Receiver<T> {
fn clone(&self) -> Receiver<T> {
let count = self.channel.receiver_count.fetch_add(1, Ordering::Relaxed);
assert!(count < usize::MAX / 2);
Receiver {
channel: self.channel.clone(),
listener: None,
_pin: PhantomPinned,
}
}
}
#[derive(Clone)]
pub struct WeakReceiver<T> {
channel: Weak<Channel<T>>,
}
impl<T> WeakReceiver<T> {
#[inline]
pub fn pop_front_if<F>(&self, priority: Priority, f: F) -> Result<Option<T>, Closed>
where
F: FnOnce(&T) -> bool,
{
let channel = self.channel.upgrade().ok_or(Closed)?;
channel.queue.pop_front_if(priority, f)
}
#[inline]
pub fn pop_back_if<F>(&self, priority: Priority, f: F) -> Result<Option<T>, Closed>
where
F: FnOnce(&T) -> bool,
{
let channel = self.channel.upgrade().ok_or(Closed)?;
channel.queue.pop_back_if(priority, f)
}
}
easy_wrapper! {
#[derive(Debug)]
#[must_use = "futures do nothing unless you `.await` or poll them"]
pub struct Recv<'a, T>(RecvInner<'a, T> => Result<T, Closed>);
pub(crate) wait();
}
#[derive(Debug)]
enum PopEnd {
Front,
Back,
}
pin_project! {
#[derive(Debug)]
#[project(!Unpin)]
struct RecvInner<'a, T> {
receiver: &'a Receiver<T>,
pop_end: PopEnd,
listener: Option<EventListener>,
#[pin]
_pin: PhantomPinned
}
}
impl<T> EventListenerFuture for RecvInner<'_, T> {
type Output = Result<T, Closed>;
fn poll_with_strategy<'x, S: Strategy<'x>>(
self: Pin<&mut Self>,
strategy: &mut S,
cx: &mut S::Context,
) -> Poll<Result<T, Closed>> {
let this = self.project();
loop {
let message = match this.pop_end {
PopEnd::Front => this.receiver.try_recv_front(),
PopEnd::Back => this.receiver.try_recv_back(),
}?;
if let Some(msg) = message {
return Poll::Ready(Ok(msg));
}
if this.listener.is_some() {
ready!(S::poll(strategy, &mut *this.listener, cx));
} else {
*this.listener = Some(this.receiver.channel.recv_ops.listen());
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::testing::{ext::*, sim, task};
use std::time::Duration;
#[test]
fn test_unlimited() {
sim(|| {
let (tx, rx) = new(2);
async move {
for v in 0u64.. {
if tx.send_back(v).is_err() {
return;
};
task::yield_now().await;
}
}
.primary()
.spawn();
async move {
for expected in 0u64..10 {
let actual = rx.recv_front().await.unwrap();
assert_eq!(actual, expected);
}
}
.primary()
.spawn();
});
}
#[test]
fn test_send_limited() {
sim(|| {
let (tx, rx) = new(2);
async move {
for v in 0u64.. {
if tx.send_back(v).is_err() {
return;
};
Duration::from_millis(1).sleep().await;
}
}
.primary()
.spawn();
async move {
for expected in 0u64..10 {
let actual = rx.recv_front().await.unwrap();
assert_eq!(actual, expected);
}
}
.primary()
.spawn();
});
}
#[test]
fn test_recv_limited() {
sim(|| {
let (tx, rx) = new(2);
async move {
for v in 0u64.. {
match tx.send_back(v) {
Ok(Some(_old)) => {
Duration::from_millis(1).sleep().await;
}
Ok(None) => {
continue;
}
Err(_) => {
return;
}
}
}
}
.primary()
.spawn();
async move {
let mut min = 0;
for _ in 0u64..10 {
let actual = rx.recv_front().await.unwrap();
assert!(actual > min || actual == 0);
min = actual;
Duration::from_millis(1).sleep().await;
}
}
.primary()
.spawn();
});
}
#[test]
fn test_multi_recv() {
sim(|| {
let (tx, rx) = new(2);
async move {
for v in 0u64.. {
if tx.send_back(v).is_err() {
return;
};
task::yield_now().await;
}
}
.primary()
.spawn();
for _ in 0..2 {
let rx = rx.clone();
async move {
let mut min = 0;
for _ in 0u64..10 {
let actual = rx.recv_front().await.unwrap();
assert!(actual > min || actual == 0, "{actual} > {min}");
min = actual;
}
}
.primary()
.spawn();
}
});
}
}