use std::{
cell::{Cell, RefCell},
collections::VecDeque,
fmt::Debug,
future::poll_fn,
rc::Rc,
task::{Poll, Waker},
};
use crate::model::{Object, ObjectID, Transition};
enum Op<T> {
Send { value: T, done: Rc<Cell<bool>> },
Recv { slot: Rc<Cell<Option<T>>> },
}
enum Kind {
Send,
Recv { consumed: Option<usize> },
}
struct Request<T> {
transition: Transition,
waker: Waker,
op: Op<T>,
}
enum Record {
Send { value: String },
Recv { consumed: usize, value: String },
}
struct Channel<T> {
id: ObjectID,
seq: usize,
consumer: Option<usize>,
queue: VecDeque<(usize, T)>,
requests: Vec<Request<T>>,
history: Vec<(Transition, Record)>,
}
impl<T: Debug> Channel<T> {
fn new(id: ObjectID) -> Self {
Self {
id,
seq: 0,
consumer: None,
queue: VecDeque::new(),
requests: Vec::new(),
history: Vec::new(),
}
}
fn register(&mut self, op: Op<T>, waker: Waker) {
let transition = Transition::new(self.id, self.seq);
if matches!(op, Op::Recv { .. }) {
match self.consumer {
None => self.consumer = Some(transition.pid),
Some(c) => assert_eq!(
c, transition.pid,
"an MPSC channel has a single consumer; a Receiver must not be shared across processes"
),
}
}
self.seq += 1;
self.requests.push(Request {
transition,
waker,
op,
});
}
fn apply(&mut self, t: Transition) {
let Some(i) = self.requests.iter().position(|r| r.transition == t) else {
panic!("transition must be enabled");
};
let req = self.requests.remove(i);
match req.op {
Op::Send { value, done } => {
self.history.push((
t,
Record::Send {
value: format!("{value:?}"),
},
));
self.queue.push_back((t.seq, value));
done.set(true);
}
Op::Recv { slot } => {
let (send_seq, value) = self.queue.pop_front().expect("recv must be enabled");
self.history.push((
t,
Record::Recv {
consumed: send_seq,
value: format!("{value:?}"),
},
));
slot.set(Some(value));
}
}
req.waker.wake();
}
fn enabled(&self) -> Vec<Transition> {
let mut out = Vec::new();
self.enabled_into(&mut out);
out
}
fn enabled_into(&self, out: &mut Vec<Transition>) {
let queue_nonempty = !self.queue.is_empty();
out.extend(
self.requests
.iter()
.filter(|r| match r.op {
Op::Send { .. } => true,
Op::Recv { .. } => queue_nonempty,
})
.map(|r| r.transition),
);
}
fn kind_of(&self, t: Transition) -> Kind {
if let Some(req) = self.requests.iter().find(|r| r.transition == t) {
return match req.op {
Op::Send { .. } => Kind::Send,
Op::Recv { .. } => Kind::Recv { consumed: None },
};
}
let (_, rec) = self
.history
.iter()
.find(|(tt, _)| *tt == t)
.expect("transition not registered on this channel");
match rec {
Record::Send { .. } => Kind::Send,
Record::Recv { consumed, .. } => Kind::Recv {
consumed: Some(*consumed),
},
}
}
fn depends(&self, t1: Transition, t2: Transition) -> bool {
match (self.kind_of(t1), self.kind_of(t2)) {
(Kind::Send, Kind::Send) => true,
(Kind::Recv { .. }, Kind::Recv { .. }) => false,
(Kind::Send, Kind::Recv { consumed }) => consumed == Some(t1.seq),
(Kind::Recv { consumed }, Kind::Send) => consumed == Some(t2.seq),
}
}
fn label(&self, t: Transition) -> String {
let (_, rec) = self
.history
.iter()
.find(|(tt, _)| *tt == t)
.expect("label called on an unapplied transition");
match rec {
Record::Send { value } => format!("send {value}"),
Record::Recv { consumed, value } => format!("recv -> {value} (#{consumed})"),
}
}
}
pub(crate) struct ChannelHandle<T> {
chan: Rc<RefCell<Channel<T>>>,
}
impl<T> Clone for ChannelHandle<T> {
fn clone(&self) -> Self {
Self {
chan: Rc::clone(&self.chan),
}
}
}
impl<T: Debug> ChannelHandle<T> {
pub(crate) fn new(id: ObjectID) -> Self {
Self {
chan: Rc::new(RefCell::new(Channel::new(id))),
}
}
pub(crate) fn split(&self) -> (Sender<T>, Receiver<T>) {
let chan = Rc::clone(&self.chan);
(
Sender {
chan: Rc::clone(&chan),
},
Receiver { chan },
)
}
}
impl<T: Debug + 'static> Object for ChannelHandle<T> {
fn apply(&mut self, t: Transition) {
self.chan.borrow_mut().apply(t);
}
fn enabled(&self) -> Vec<Transition> {
self.chan.borrow().enabled()
}
fn enabled_into(&self, out: &mut Vec<Transition>) {
self.chan.borrow().enabled_into(out);
}
fn label(&self, t: Transition) -> String {
self.chan.borrow().label(t)
}
fn depends(&self, t1: Transition, t2: Transition) -> bool {
self.chan.borrow().depends(t1, t2)
}
}
pub struct Sender<T> {
chan: Rc<RefCell<Channel<T>>>,
}
impl<T> Clone for Sender<T> {
fn clone(&self) -> Self {
Self {
chan: Rc::clone(&self.chan),
}
}
}
impl<T> std::fmt::Debug for Sender<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self.chan.try_borrow() {
Ok(c) => f
.debug_struct("Sender")
.field("channel", &c.id)
.finish_non_exhaustive(),
Err(_) => f.debug_struct("Sender").finish_non_exhaustive(),
}
}
}
impl<T: Debug> Sender<T> {
pub async fn send(&self, value: T) {
let done = Rc::new(Cell::new(false));
let mut pending = Some(value);
poll_fn(move |cx| {
if done.get() {
return Poll::Ready(());
}
if let Some(value) = pending.take() {
self.chan.borrow_mut().register(
Op::Send {
value,
done: Rc::clone(&done),
},
cx.waker().clone(),
);
}
Poll::Pending
})
.await
}
}
pub struct Receiver<T> {
chan: Rc<RefCell<Channel<T>>>,
}
impl<T> std::fmt::Debug for Receiver<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self.chan.try_borrow() {
Ok(c) => f
.debug_struct("Receiver")
.field("channel", &c.id)
.finish_non_exhaustive(),
Err(_) => f.debug_struct("Receiver").finish_non_exhaustive(),
}
}
}
impl<T: Debug> Receiver<T> {
pub async fn recv(&self) -> T {
let slot = Rc::new(Cell::new(None));
let mut registered = false;
poll_fn(move |cx| {
if let Some(value) = slot.take() {
return Poll::Ready(value);
}
if !registered {
registered = true;
self.chan.borrow_mut().register(
Op::Recv {
slot: Rc::clone(&slot),
},
cx.waker().clone(),
);
}
Poll::Pending
})
.await
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::model::{Executor, ProcessResult};
use std::future::Future;
fn make() -> (ChannelHandle<i32>, Sender<i32>, Receiver<i32>) {
let driver = ChannelHandle::new(0);
let (tx, rx) = driver.split();
(driver, tx, rx)
}
fn drive(exec: &mut Executor, obj: &mut impl Object) {
exec.execute().unwrap();
while let Some(&t) = obj.enabled().first() {
obj.apply(t);
exec.execute().unwrap();
}
}
fn run_single(
body: impl FnOnce(Sender<i32>, Receiver<i32>) -> Box<dyn Future<Output = ProcessResult>>,
) {
let (mut driver, tx, rx) = make();
let mut exec = Executor::default();
exec.schedule(Box::into_pin(body(tx, rx)));
drive(&mut exec, &mut driver);
}
fn enabled_of(handle: &ChannelHandle<i32>, pid: usize) -> Transition {
*handle.enabled().iter().find(|t| t.pid == pid).unwrap()
}
#[test]
fn send_then_recv_returns_value() {
let seen = Rc::new(Cell::new(0));
let dst = seen.clone();
run_single(move |tx, rx| {
Box::new(async move {
tx.send(7).await;
dst.set(rx.recv().await);
Ok(())
})
});
assert_eq!(seen.get(), 7);
}
#[test]
fn fifo_order_for_one_producer() {
let first = Rc::new(Cell::new(0));
let second = Rc::new(Cell::new(0));
let (a, b) = (first.clone(), second.clone());
run_single(move |tx, rx| {
Box::new(async move {
tx.send(1).await;
tx.send(2).await;
a.set(rx.recv().await);
b.set(rx.recv().await);
Ok(())
})
});
assert_eq!(first.get(), 1);
assert_eq!(second.get(), 2);
}
#[test]
#[should_panic(expected = "single consumer")]
fn second_consumer_panics() {
let (_driver, _tx, rx) = make();
let rx = Rc::new(rx);
let mut exec = Executor::default();
let r1 = rx.clone();
exec.schedule(async move {
r1.recv().await;
Ok(())
});
let r2 = rx.clone();
exec.schedule(async move {
r2.recv().await;
Ok(())
});
let _ = exec.execute(); }
#[test]
fn recv_on_empty_blocks() {
let (driver, _tx, rx) = make();
let mut exec = Executor::default();
exec.schedule(async move {
rx.recv().await;
Ok(())
});
exec.execute().unwrap();
assert!(
driver.enabled().is_empty(),
"recv must block on an empty queue"
);
}
#[test]
fn dependency_truth_table() {
let (mut driver, tx, rx) = make();
let tx2 = tx.clone();
let mut exec = Executor::default();
exec.schedule(async move {
tx.send(10).await;
Ok(())
});
exec.schedule(async move {
tx2.send(20).await;
Ok(())
});
exec.schedule(async move {
rx.recv().await;
Ok(())
});
exec.execute().unwrap();
let send0 = enabled_of(&driver, 0);
let send1 = enabled_of(&driver, 1);
assert_eq!(driver.enabled().len(), 2);
assert!(driver.depends(send0, send1), "send/send is dependent");
driver.apply(send0);
exec.execute().unwrap();
let recv = enabled_of(&driver, 2);
driver.apply(recv);
exec.execute().unwrap();
assert!(
driver.depends(send0, recv),
"recv depends on the send it consumed"
);
assert!(
!driver.depends(send1, recv),
"recv is independent of the unconsumed send"
);
assert!(
!driver.depends(recv, send1),
"symmetric: unconsumed send is independent"
);
}
}