use core::{
cell::Cell,
convert::Infallible,
fmt,
ops::{Deref, DerefMut},
};
use either::Either::{self, *};
use fairly_unsafe_cell::*;
use frugal_async::{Mutex, TakeCell};
use crate::prelude::*;
use crate::queues::Queue;
pub struct State<Q, F> {
queue: Mutex<Q>,
buffered_final_value: FairlyUnsafeCell<Option<F>>,
len: Cell<usize>,
notify_the_sender: TakeCell<()>,
notify_the_receiver: TakeCell<()>,
did_any_endpoint_drop_yet: Cell<bool>,
}
impl<Q, F> fmt::Debug for State<Q, F>
where
Q: fmt::Debug,
F: fmt::Debug,
{
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("State")
.field("queue_item_count", &self.len)
.field("queue", &self.queue)
.field("buffered_final_value", &self.buffered_final_value)
.finish()
}
}
impl<Q: Queue, F> State<Q, F> {
pub fn new(queue: Q) -> Self {
State {
len: Cell::new(queue.len()),
queue: Mutex::new(queue),
buffered_final_value: FairlyUnsafeCell::new(None),
notify_the_sender: TakeCell::new_with(()),
notify_the_receiver: TakeCell::new(),
did_any_endpoint_drop_yet: Cell::new(false),
}
}
fn len(&self) -> usize {
self.len.get()
}
fn is_empty(&self) -> bool {
self.len.get() == 0
}
fn close(&self, fin: F) {
let mut last = unsafe { self.buffered_final_value.borrow_mut() };
*last = Some(fin);
self.notify_the_receiver.set(());
}
}
pub fn new_sssr<R, Q, F>(state_ref: R) -> (Sender<R, Q, F>, Receiver<R, Q, F>)
where
R: Deref<Target = State<Q, F>> + Clone,
{
(
Sender {
state: state_ref.clone(),
},
Receiver { state: state_ref },
)
}
#[derive(Debug)]
pub struct Sender<R, Q, F>
where
R: Deref<Target = State<Q, F>>,
{
state: R,
}
impl<R, Q, F> Drop for Sender<R, Q, F>
where
R: Deref<Target = State<Q, F>>,
{
fn drop(&mut self) {
self.state.did_any_endpoint_drop_yet.set(true)
}
}
impl<R, Q, F> Sender<R, Q, F>
where
R: Deref<Target = State<Q, F>>,
Q: Queue,
{
pub fn len(&self) -> usize {
self.state.len()
}
pub fn is_empty(&self) -> bool {
self.state.is_empty()
}
pub fn is_receiver_dropped(&self) -> bool {
self.state.did_any_endpoint_drop_yet.get()
}
}
impl<R: Deref<Target = State<Q, F>>, Q: Queue, F> Consumer for Sender<R, Q, F> {
type Item = Q::Item;
type Final = F;
type Error = Infallible;
async fn consume(&mut self, val: Either<Self::Item, Self::Final>) -> Result<(), Self::Error> {
match val {
Left(mut item) => {
loop {
let did_it_work = {
self.state.queue.write().await.deref_mut().enqueue(item)
};
match did_it_work {
Some(item_) => {
let () = self.state.notify_the_sender.take().await;
item = item_;
}
None => {
self.state.len.set(self.state.len.get() + 1);
self.state.notify_the_receiver.set(());
return Ok(());
}
}
}
}
Right(fin) => {
self.state.close(fin);
Ok(())
}
}
}
async fn flush(&mut self) -> Result<(), Self::Error> {
Ok(()) }
}
impl<R: Deref<Target = State<Q, F>>, Q: Queue, F> BulkConsumer for Sender<R, Q, F> {
async fn expose_slots_gracefully<Fun, Ret>(&mut self, f: Fun) -> Result<Ret, (Fun, Self::Error)>
where
Fun: AsyncFnOnce(&mut [Self::Item]) -> (usize, Ret),
{
let mut f = Some(f);
loop {
let ret = self
.state
.queue
.write()
.await
.deref_mut()
.expose_slots(async |queue_slots| {
if queue_slots.is_empty() {
(0, None)
} else {
let (amount, ret) = (f.take().expect(
"Running this branch only once, we return after having called f",
))(queue_slots)
.await;
self.state.len.set(self.state.len.get() + amount);
self.state.notify_the_receiver.set(());
(amount, Some(ret))
}
})
.await;
match ret {
None => {
let () = self.state.notify_the_sender.take().await;
}
Some(ret) => return Ok(ret),
}
}
}
}
#[derive(Debug)]
pub struct Receiver<R, Q, F>
where
R: Deref<Target = State<Q, F>>,
{
state: R,
}
impl<R, Q, F> Drop for Receiver<R, Q, F>
where
R: Deref<Target = State<Q, F>>,
{
fn drop(&mut self) {
self.state.did_any_endpoint_drop_yet.set(true)
}
}
impl<R: Deref<Target = State<Q, F>>, Q: Queue, F> Receiver<R, Q, F> {
pub fn len(&self) -> usize {
self.state.len()
}
pub fn is_empty(&self) -> bool {
self.state.is_empty()
}
pub fn is_sender_dropped(&self) -> bool {
self.state.did_any_endpoint_drop_yet.get()
}
}
impl<R: Deref<Target = State<Q, F>>, Q: Queue, F> Producer for Receiver<R, Q, F> {
type Item = Q::Item;
type Final = F;
type Error = Infallible;
async fn produce(&mut self) -> Result<Either<Self::Item, Self::Final>, Self::Error> {
loop {
match self.state.queue.write().await.deref_mut().dequeue() {
Some(item) => {
self.state.len.set(self.state.len.get() - 1);
self.state.notify_the_sender.set(());
return Ok(Left(item));
}
None => {
match unsafe { self.state.buffered_final_value.borrow_mut().take() } {
Some(fin) => {
return Ok(Right(fin));
}
None => {
}
}
}
}
let () = self.state.notify_the_receiver.take().await;
}
}
async fn slurp(&mut self) -> Result<(), Self::Error> {
Ok(()) }
}
impl<R: Deref<Target = State<Q, F>>, Q: Queue, F> BulkProducer for Receiver<R, Q, F> {
async fn expose_items_gracefully<Fun, Ret>(
&mut self,
f: Fun,
) -> Result<Either<Ret, (Fun, Self::Final)>, (Fun, Self::Error)>
where
Fun: AsyncFnOnce(&[Self::Item]) -> (usize, Ret),
{
let mut f = Some(f);
loop {
let ret = self
.state
.queue
.write()
.await
.expose_items(async |queue_items| {
if queue_items.is_empty() {
match unsafe { self.state.buffered_final_value.borrow_mut().take() } {
Some(fin) => (0, Some(Right(fin))),
None => (0, None),
}
} else {
let (amount, ret) = (f.take().expect(
"Running this branch only once, we return after having called f",
))(queue_items)
.await;
self.state.len.set(self.state.len.get() - amount);
self.state.notify_the_sender.set(());
(amount, Some(Left(ret)))
}
})
.await;
match ret {
None => {
let () = self.state.notify_the_receiver.take().await;
}
Some(Left(ret)) => return Ok(Left(ret)),
Some(Right(fin)) => {
return Ok(Right((
f.take()
.expect("Branch guarded by early return of Ok(Left(_))"),
fin,
)))
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use futures::join;
use crate::queues::new_static;
#[test]
fn test_spsc_sufficient_capacity() {
let state = State::new(new_static::<i16, 99>());
let (mut sender, mut receiver) = new_sssr(&state);
pollster::block_on(async {
assert!(sender.consume_item(300).await.is_ok());
assert!(sender.consume_item(400).await.is_ok());
assert!(sender.consume_item(500).await.is_ok());
assert!(sender.consume_final(-17).await.is_ok());
assert_eq!(300, receiver.produce().await.unwrap().unwrap_left());
assert_eq!(400, receiver.produce().await.unwrap().unwrap_left());
assert_eq!(500, receiver.produce().await.unwrap().unwrap_left());
assert_eq!(-17, receiver.produce().await.unwrap().unwrap_right());
});
}
#[test]
fn test_spsc_low_capacity() {
pollster::block_on(async {
let state = State::new(new_static::<i16, 2>());
let (mut sender, mut receiver) = new_sssr(&state);
let send_things = async {
assert!(sender.consume_item(300).await.is_ok());
assert!(sender.consume_item(400).await.is_ok());
assert!(sender.consume_item(500).await.is_ok());
assert!(sender.consume_final(-17).await.is_ok());
};
let receive_things = async {
assert_eq!(300, receiver.produce().await.unwrap().unwrap_left());
assert_eq!(400, receiver.produce().await.unwrap().unwrap_left());
assert_eq!(500, receiver.produce().await.unwrap().unwrap_left());
assert_eq!(-17, receiver.produce().await.unwrap().unwrap_right());
};
join!(send_things, receive_things);
});
}
#[test]
fn test_spsc_immediate_final() {
pollster::block_on(async {
let state = State::new(new_static::<i16, 3>());
let (mut sender, mut receiver) = new_sssr(&state);
let send_things = async {
assert!(sender.consume_final(-17).await.is_ok());
};
let receive_things = async {
assert_eq!(-17, receiver.produce().await.unwrap().unwrap_right());
};
join!(send_things, receive_things);
});
}
#[test]
fn test_spsc_receive_then_send_concurrently() {
pollster::block_on(async {
let state = State::new(new_static::<i16, 2>());
let (mut sender, mut receiver) = new_sssr(&state);
let send_things = async {
assert!(sender.consume_item(300).await.is_ok());
assert!(sender.consume_item(400).await.is_ok());
assert!(sender.consume_item(500).await.is_ok());
assert!(sender.consume_final(-17).await.is_ok());
};
let receive_things = async {
assert_eq!(300, receiver.produce().await.unwrap().unwrap_left());
assert_eq!(400, receiver.produce().await.unwrap().unwrap_left());
assert_eq!(500, receiver.produce().await.unwrap().unwrap_left());
assert_eq!(-17, receiver.produce().await.unwrap().unwrap_right());
};
join!(receive_things, send_things);
});
}
#[test]
fn test_spsc_capacity_1() {
pollster::block_on(async {
let state = State::new(new_static::<i16, 1>());
let (mut sender, mut receiver) = new_sssr(&state);
let send_things = async {
assert!(sender.consume_item(300).await.is_ok());
assert!(sender.consume_item(400).await.is_ok());
assert!(sender.consume_item(500).await.is_ok());
assert!(sender.consume_final(-17).await.is_ok());
};
let receive_things = async {
assert_eq!(300, receiver.produce().await.unwrap().unwrap_left());
assert_eq!(400, receiver.produce().await.unwrap().unwrap_left());
assert_eq!(500, receiver.produce().await.unwrap().unwrap_left());
assert_eq!(-17, receiver.produce().await.unwrap().unwrap_right());
};
join!(receive_things, send_things);
});
}
}