Skip to main content

shuttle_std/sync/
mpsc.rs

1//! Multi-producer, single-consumer FIFO queue communication primitives.
2
3use crate::sync::{ResourceSignature, ResourceType};
4use shuttle_engine::runtime::execution::ExecutionState;
5use shuttle_engine::runtime::task::clock::VectorClock;
6use shuttle_engine::runtime::task::{TaskId, DEFAULT_INLINE_TASKS};
7use shuttle_engine::runtime::thread;
8use smallvec::SmallVec;
9use std::cell::RefCell;
10use std::fmt::Debug;
11use std::rc::Rc;
12use std::result::Result;
13pub use std::sync::mpsc::{RecvError, RecvTimeoutError, SendError, TryRecvError, TrySendError};
14use std::sync::Arc;
15use std::time::Duration;
16use tracing::trace;
17
18const MAX_INLINE_MESSAGES: usize = 32;
19
20/// Create an unbounded channel
21#[track_caller]
22pub fn channel<T>() -> (Sender<T>, Receiver<T>) {
23    let channel = Arc::new(Channel::new(None));
24    let sender = Sender {
25        inner: Arc::clone(&channel),
26    };
27    let receiver = Receiver {
28        inner: Arc::clone(&channel),
29    };
30    (sender, receiver)
31}
32
33/// Create a bounded channel
34#[track_caller]
35pub fn sync_channel<T>(bound: usize) -> (SyncSender<T>, Receiver<T>) {
36    let channel = Arc::new(Channel::new(Some(bound)));
37    let sender = SyncSender {
38        inner: Arc::clone(&channel),
39    };
40    let receiver = Receiver {
41        inner: Arc::clone(&channel),
42    };
43    (sender, receiver)
44}
45
46#[derive(Debug)]
47struct Channel<T> {
48    bound: Option<usize>, // None for an unbounded channel, Some(k) for a bounded channel of size k
49    state: Rc<RefCell<ChannelState<T>>>,
50    #[allow(unused)]
51    signature: ResourceSignature,
52}
53
54// For tracking causality on channels, we timestamp each message with the clock of the sender.
55// When the receiver gets the message, it updates its clock with the the associated timestamp.
56// For unbounded channels, that's all the work we need to do.
57//
58// For bounded and rendezvous channels, things get a bit more interesting.
59// Consider a bounded channel of depth K.  As soon as the sender successfully sends its K+1'th
60// message, it knows that the receiver has received at least 1 message.  At this point, the
61// first receive event causally precedes the (K+1)'th send.  By the rule for vector clocks,
62//  (clock of the first receive)  <  (clock of the K+1'th send)
63// In order to ensure this ordering, we add a return queue of depth K to bounded channels.
64// Initially, this queue contains K empty vector clocks.  On each receive, we push the
65// receiver's clock at the time of the receive to the end of this queue.  Whenever the sender
66// successfully sends a message, it pops the clock at the front of the queue, and updates its
67// own clock with this value.  Thus, on the (K+1)'th send, the sender's clock will be updated
68// with the clock at the first receive, as needed.
69//
70// The story is similar for rendezvous channels, except we have to handle things a bit more
71// specially because K=0.
72
73struct TimestampedValue<T> {
74    value: T,
75    clock: VectorClock,
76}
77
78impl<T> TimestampedValue<T> {
79    fn new(value: T, clock: VectorClock) -> Self {
80        Self { value, clock }
81    }
82}
83
84// Note: The channels in std::sync::mpsc only support a single Receiver (which cannot be
85// cloned).  The state below admits a more general use case, where multiple Senders
86// and Receivers can share a single channel.
87struct ChannelState<T> {
88    messages: SmallVec<[TimestampedValue<T>; MAX_INLINE_MESSAGES]>, // messages in the channel
89    receiver_clock: Option<SmallVec<[VectorClock; MAX_INLINE_MESSAGES]>>, // receiver vector clocks for bounded case
90    known_senders: usize,                                           // number of senders referencing this channel
91    known_receivers: usize,                                         // number or receivers referencing this channel
92    waiting_senders: SmallVec<[TaskId; DEFAULT_INLINE_TASKS]>,      // list of currently blocked senders
93    waiting_receivers: SmallVec<[TaskId; DEFAULT_INLINE_TASKS]>,    // list of currently blocked receivers
94}
95
96impl<T> Debug for ChannelState<T> {
97    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
98        write!(f, "Channel {{ ")?;
99        write!(f, "num_messages: {} ", self.messages.len())?;
100        write!(
101            f,
102            "known_senders {} known_receivers {} ",
103            self.known_senders, self.known_receivers
104        )?;
105        write!(f, "waiting_senders: [{:?}] ", self.waiting_senders)?;
106        write!(f, "waiting_receivers: [{:?}] ", self.waiting_receivers)?;
107        write!(f, "}}")
108    }
109}
110
111impl<T> Channel<T> {
112    #[track_caller]
113    fn new(bound: Option<usize>) -> Self {
114        let receiver_clock = if let Some(bound) = bound {
115            let mut s = SmallVec::with_capacity(bound);
116            for _ in 0..bound {
117                s.push(VectorClock::new());
118            }
119            Some(s)
120        } else {
121            None
122        };
123        Self {
124            bound,
125            state: Rc::new(RefCell::new(ChannelState {
126                messages: SmallVec::new(),
127                receiver_clock,
128                known_senders: 1,
129                known_receivers: 1,
130                waiting_senders: SmallVec::new(),
131                waiting_receivers: SmallVec::new(),
132            })),
133            signature: ExecutionState::new_resource_signature(ResourceType::MpscChannel),
134        }
135    }
136
137    fn try_send(&self, message: T) -> Result<(), TrySendError<T>> {
138        self.send_internal(message, false)
139    }
140
141    fn send(&self, message: T) -> Result<(), SendError<T>> {
142        self.send_internal(message, true).map_err(|e| match e {
143            TrySendError::Full(_) => unreachable!(),
144            TrySendError::Disconnected(m) => SendError(m),
145        })
146    }
147
148    fn is_rendezvous(&self) -> bool {
149        self.bound == Some(0)
150    }
151
152    fn sender_must_block(&self, state: &ChannelState<T>) -> bool {
153        let (is_rendezvous, is_full) = if let Some(bound) = self.bound {
154            // For a rendezvous channel (bound = 0), "is_full" holds when there is a message in the channel.
155            // For a non-rendezvous channel (bound > 0), "is_full" holds when the capacity is reached.
156            // We cover both these cases at once using max(bound, 1) below.
157            (bound == 0, state.messages.len() >= std::cmp::max(bound, 1))
158        } else {
159            (false, false)
160        };
161
162        // The sender should block in any of the following situations:
163        //    the channel is full (as defined above)
164        //    there are already waiting senders
165        //    this is a rendezvous channel and there are no waiting receivers
166        is_full || !state.waiting_senders.is_empty() || (is_rendezvous && state.waiting_receivers.is_empty())
167    }
168
169    fn send_internal(&self, message: T, can_block: bool) -> Result<(), TrySendError<T>> {
170        // Because channels are always fair wrt. waiting senders (waiting senders is an *ordered* list),
171        // blocking sends do not commute thus must always provide a switch before blocking
172        thread::switch();
173
174        let me = ExecutionState::me();
175        let mut state = self.state.borrow_mut();
176        let should_block = self.sender_must_block(&state);
177
178        trace!(
179            state = ?state,
180            "sender {:?} starting send on channel {:p}",
181            me,
182            self,
183        );
184        if state.known_receivers == 0 {
185            // No receivers are left, so the channel is disconnected.  Stop and return failure.
186            return Err(TrySendError::Disconnected(message));
187        }
188
189        if should_block {
190            if !can_block {
191                return Err(TrySendError::Full(message));
192            }
193
194            state.waiting_senders.push(me);
195            trace!(
196                state = ?state,
197                "blocking sender {:?} on channel {:p}",
198                me,
199                self,
200            );
201            ExecutionState::with(|s| s.current_mut().block(false));
202            drop(state);
203
204            thread::switch();
205
206            state = self.state.borrow_mut();
207            trace!(
208                state = ?state,
209                "unblocked sender {:?} on channel {:p}",
210                me,
211                self,
212            );
213
214            // Check again that we still have a receiver; if not, return with error.
215            // We repeat this check because the receivers may have disconnected while the sender was blocked.
216            if state.known_receivers == 0 {
217                state.waiting_senders.retain(|t| *t != me);
218                // No receivers are left, so the channel is disconnected.  Stop and return failure.
219                return Err(TrySendError::Disconnected(message));
220            }
221
222            let head = state.waiting_senders.remove(0);
223            assert_eq!(head, me);
224        }
225
226        ExecutionState::with(|s| {
227            let clock = s.increment_clock();
228            state.messages.push(TimestampedValue::new(message, clock.clone()));
229        });
230
231        // The sender has just added a message to the channel, so unblock the first waiting receiver if any
232        if let Some(&tid) = state.waiting_receivers.first() {
233            ExecutionState::with(|s| {
234                s.get_mut(tid).unblock();
235
236                // When a sender successfully sends on a rendezvous channel, it knows that the receiver will perform
237                // the matching receive, so we need to update the sender's clock with the receiver's.
238                if self.is_rendezvous() {
239                    let recv_clock = s.get_clock(tid).clone();
240                    s.update_clock(&recv_clock);
241                }
242            });
243        }
244        // Check and unblock the next the waiting sender, if eligible
245        if let Some(&tid) = state.waiting_senders.first() {
246            let bound = self.bound.expect("can't have waiting senders on an unbounded channel");
247            if state.messages.len() < bound {
248                ExecutionState::with(|s| s.get_mut(tid).unblock());
249            }
250        }
251
252        if !self.is_rendezvous() {
253            if let Some(receiver_clock) = &mut state.receiver_clock {
254                let recv_clock = receiver_clock.remove(0);
255                ExecutionState::with(|s| s.update_clock(&recv_clock));
256            }
257        }
258
259        Ok(())
260    }
261
262    fn recv(&self) -> Result<T, RecvError> {
263        self.recv_internal(true).map_err(|e| match e {
264            TryRecvError::Disconnected => RecvError,
265            TryRecvError::Empty => unreachable!(),
266        })
267    }
268
269    fn try_recv(&self) -> Result<T, TryRecvError> {
270        self.recv_internal(false)
271    }
272
273    fn receiver_must_block(&self, state: &ChannelState<T>) -> bool {
274        // The receiver should block in any of the following situations:
275        //    the channel is empty
276        //    there are waiting receivers
277        state.messages.is_empty() || !state.waiting_receivers.is_empty()
278    }
279
280    fn recv_internal(&self, can_block: bool) -> Result<T, TryRecvError> {
281        // Because channels are always fair wrt. waiting receivers (waiting receivers is an *ordered* list),
282        // blocking receives do not commute thus must always provide a switch before blocking
283        thread::switch();
284
285        let me = ExecutionState::me();
286        let mut state = self.state.borrow_mut();
287        let should_block = self.receiver_must_block(&state);
288
289        trace!(
290            state = ?state,
291            "starting recv on channel {:p}",
292            self,
293        );
294        // Check if there are any senders left; if not, and the channel is empty, fail with error
295        // (If there are no senders, but the channel is nonempty, the receiver can successfully consume that message.)
296        if state.messages.is_empty() && state.known_senders == 0 {
297            return Err(TryRecvError::Disconnected);
298        }
299
300        // If this is a rendezvous channel, and the channel is empty, and there are waiting senders,
301        // notify the first waiting sender
302        if self.is_rendezvous() && state.messages.is_empty() {
303            if let Some(&tid) = state.waiting_senders.first() {
304                // Note: another receiver may have unblocked the sender already
305                ExecutionState::with(|s| s.get_mut(tid).unblock());
306            } else if !can_block {
307                // Nobody to rendezvous with
308                return Err(TryRecvError::Empty);
309            }
310        }
311
312        // Handle the try_recv case, accounting for the number of msgs available and already waiting receivers.
313        if !self.is_rendezvous() && !can_block && state.waiting_receivers.len() >= state.messages.len() {
314            return Err(TryRecvError::Empty);
315        }
316
317        // Pre-increment the receiver's clock before continuing
318        //
319        // Note: The reason for pre-incrementing the receiver's clock is to deal properly with rendezvous channels.
320        // Here's the scenario we have to handle:
321        //   1. the receiver arrives at a rendezvous channel and blocks
322        //   2. the sender arrives, sees the receiver is waiting and does not block
323        //   3. the sender drops the message in the channel and updates its clock with the receiver's clock and continues
324        //   4. later, the receiver unblocks and picks up the message and updates its clock with the sender's
325        // Without the pre-increment, in step 3, the sender would update its clock with the receiver's clock before
326        // it is incremented.  (The increment records the fact that the receiver arrived at the synchronization point.)
327        ExecutionState::with(|s| {
328            let _ = s.increment_clock();
329        });
330
331        if should_block {
332            state.waiting_receivers.push(me);
333            trace!(
334                state = ?state,
335                "blocking receiver {:?} on channel {:p}",
336                me,
337                self,
338            );
339            ExecutionState::with(|s| s.current_mut().block(false));
340            drop(state);
341
342            thread::switch();
343
344            state = self.state.borrow_mut();
345            trace!(
346                state = ?state,
347                "unblocked receiver {:?} on channel {:p}",
348                me,
349                self,
350            );
351
352            // Check again if there are any senders left; if not, and the channel is empty, fail with error
353            // (If there are no senders, but the channel is nonempty, the receiver can successfully consume that message.)
354            // We repeat this check because the senders may have disconnected while the receiver was blocked.
355            if state.messages.is_empty() && state.known_senders == 0 {
356                state.waiting_receivers.retain(|t| *t != me);
357                return Err(TryRecvError::Disconnected);
358            }
359
360            let head = state.waiting_receivers.remove(0);
361            assert_eq!(head, me);
362        }
363
364        let item = state.messages.remove(0);
365        // The receiver has just removed an element from the channel.  Check if any waiting senders
366        // need to be notified.
367        if let Some(&tid) = state.waiting_senders.first() {
368            let bound = self.bound.expect("can't have waiting senders on an unbounded channel");
369            // Unblock the first waiting sender provided one of the following conditions hold:
370            // - this is a non-rendezvous bounded channel (bound > 0)
371            // - this is a rendezvous channel and we have additional waiting receivers
372            if bound > 0 || !state.waiting_receivers.is_empty() {
373                ExecutionState::with(|s| s.get_mut(tid).unblock());
374            }
375        }
376        // Check and unblock the next the waiting receiver, if eligible
377        // Note: this is a no-op for mpsc channels, since there can only be one receiver
378        if let Some(&tid) = state.waiting_receivers.first() {
379            if !state.messages.is_empty() {
380                ExecutionState::with(|s| s.get_mut(tid).unblock());
381            }
382        }
383
384        // Update receiver clock from the clock attached to the message received
385        let TimestampedValue { value, clock } = item;
386        ExecutionState::with(|s| {
387            // Since we already incremented the receiver's clock above, just update it here
388            s.get_clock_mut(me).update(&clock);
389
390            // If this is a (non-rendezvous) bounded channel, propagate causality backwards to sender
391            if let Some(receiver_clock) = &mut state.receiver_clock {
392                let bound = self.bound.expect("unexpected internal error"); // must be defined for bounded channels
393                if bound > 0 {
394                    // non-rendezvous
395                    assert!(receiver_clock.len() < bound);
396                    receiver_clock.push(s.get_clock(me).clone());
397                }
398            }
399        });
400        Ok(value)
401    }
402}
403
404// Safety: A Channel is never actually passed across true threads, only across continuations. The
405// Rc<RefCell<_>> type therefore can't be preempted mid-bookkeeping-operation.
406// TODO We use this workaround in several places in Shuttle.  Maybe there's a cleaner solution.
407unsafe impl<T: Send> Send for Channel<T> {}
408unsafe impl<T: Send> Sync for Channel<T> {}
409
410/// The receiving half of Rust's [`channel`] (or [`sync_channel`]) type.
411/// This half can only be owned by one thread.
412#[derive(Debug)]
413pub struct Receiver<T> {
414    inner: Arc<Channel<T>>,
415}
416
417impl<T> Receiver<T> {
418    /// Attempts to wait for a value on this receiver, returning an error if the
419    /// corresponding channel has hung up.
420    pub fn recv(&self) -> Result<T, RecvError> {
421        self.inner.recv()
422    }
423
424    /// Attempts to wait for a value on this receiver, returning an error if the
425    /// corresponding channel has hung up.
426    pub fn try_recv(&self) -> Result<T, TryRecvError> {
427        self.inner.try_recv()
428    }
429
430    /// Attempts to wait for a value on this receiver, returning an error if the
431    /// corresponding channel has hung up, or if it waits more than timeout.
432    pub fn recv_timeout(&self, _timeout: Duration) -> Result<T, RecvTimeoutError> {
433        // TODO support the timeout case -- this method never times out
434        self.inner.recv().map_err(|_| RecvTimeoutError::Disconnected)
435    }
436
437    /// Returns an iterator that will block waiting for messages, but never
438    /// [`panic!`]. It will return [`None`] when the channel has hung up.
439    pub fn iter(&self) -> Iter<'_, T> {
440        Iter { rx: self }
441    }
442
443    /// Returns an iterator that will attempt to yield all pending values.
444    /// It will return `None` if there are no more pending values or if the
445    /// channel has hung up. The iterator will never [`panic!`] or block the
446    /// user by waiting for values.
447    pub fn try_iter(&self) -> TryIter<'_, T> {
448        TryIter { rx: self }
449    }
450}
451
452impl<T> Drop for Receiver<T> {
453    fn drop(&mut self) {
454        if ExecutionState::should_stop() {
455            return;
456        }
457        let mut state = self.inner.state.borrow_mut();
458        assert!(state.known_receivers > 0);
459        state.known_receivers -= 1;
460        if state.known_receivers == 0 {
461            // Last receiver was dropped; wake up all senders
462            for &tid in state.waiting_senders.iter() {
463                ExecutionState::with(|s| s.get_mut(tid).unblock());
464            }
465        }
466    }
467}
468
469/// An iterator over messages on a [`Receiver`], created by [`iter`].
470///
471/// This iterator will block whenever [`next`] is called,
472/// waiting for a new message, and [`None`] will be returned
473/// when the corresponding channel has hung up.
474///
475/// [`iter`]: Receiver::iter
476/// [`next`]: Iterator::next
477#[derive(Debug)]
478pub struct Iter<'a, T: 'a> {
479    rx: &'a Receiver<T>,
480}
481
482/// An iterator that attempts to yield all pending values for a [`Receiver`],
483/// created by [`try_iter`].
484///
485/// [`None`] will be returned when there are no pending values remaining or
486/// if the corresponding channel has hung up.
487///
488/// This iterator will never block the caller in order to wait for data to
489/// become available. Instead, it will return [`None`].
490///
491/// [`try_iter`]: Receiver::try_iter
492#[derive(Debug)]
493pub struct TryIter<'a, T: 'a> {
494    rx: &'a Receiver<T>,
495}
496
497/// An owning iterator over messages on a [`Receiver`],
498/// created by [`into_iter`].
499///
500/// This iterator will block whenever [`next`]
501/// is called, waiting for a new message, and [`None`] will be
502/// returned if the corresponding channel has hung up.
503///
504/// [`into_iter`]: Receiver::into_iter
505/// [`next`]: Iterator::next
506#[derive(Debug)]
507pub struct IntoIter<T> {
508    rx: Receiver<T>,
509}
510
511impl<T> Iterator for Iter<'_, T> {
512    type Item = T;
513
514    fn next(&mut self) -> Option<T> {
515        self.rx.recv().ok()
516    }
517}
518
519impl<T> Iterator for TryIter<'_, T> {
520    type Item = T;
521
522    fn next(&mut self) -> Option<T> {
523        self.rx.try_recv().ok()
524    }
525}
526
527impl<'a, T> IntoIterator for &'a Receiver<T> {
528    type Item = T;
529    type IntoIter = Iter<'a, T>;
530
531    fn into_iter(self) -> Iter<'a, T> {
532        self.iter()
533    }
534}
535
536impl<T> Iterator for IntoIter<T> {
537    type Item = T;
538    fn next(&mut self) -> Option<T> {
539        self.rx.recv().ok()
540    }
541}
542
543impl<T> IntoIterator for Receiver<T> {
544    type Item = T;
545    type IntoIter = IntoIter<T>;
546
547    fn into_iter(self) -> IntoIter<T> {
548        IntoIter { rx: self }
549    }
550}
551
552/// The sending-half of Rust's asynchronous [`channel`] type. This half can only be
553/// owned by one thread, but it can be cloned to send to other threads.
554#[derive(Debug)]
555pub struct Sender<T> {
556    inner: Arc<Channel<T>>,
557}
558
559impl<T> Sender<T> {
560    /// Attempts to send a value on this channel, returning it back if it could
561    /// not be sent.
562    pub fn send(&self, t: T) -> Result<(), SendError<T>> {
563        self.inner.send(t)
564    }
565}
566
567impl<T> Clone for Sender<T> {
568    fn clone(&self) -> Self {
569        let mut state = self.inner.state.borrow_mut();
570        state.known_senders += 1;
571        drop(state);
572        Self {
573            inner: self.inner.clone(),
574        }
575    }
576}
577
578impl<T> Drop for Sender<T> {
579    fn drop(&mut self) {
580        if ExecutionState::should_stop() {
581            return;
582        }
583        let mut state = self.inner.state.borrow_mut();
584        assert!(state.known_senders > 0);
585        state.known_senders -= 1;
586        if state.known_senders == 0 {
587            // Last sender was dropped; wake up all receivers
588            for &tid in state.waiting_receivers.iter() {
589                ExecutionState::with(|s| s.get_mut(tid).unblock());
590            }
591        }
592    }
593}
594
595/// The sending-half of Rust's synchronous [`sync_channel`] type.
596///
597/// Messages can be sent through this channel with [`SyncSender::send`] or \[`try_send`\] (TODO)
598///
599/// [`SyncSender::send`] will block if there is no space in the internal buffer.
600#[derive(Debug)]
601pub struct SyncSender<T> {
602    inner: Arc<Channel<T>>,
603}
604
605impl<T> SyncSender<T> {
606    /// Sends a value on this synchronous channel.
607    ///
608    /// This function will *block* until space in the internal buffer becomes
609    /// available or a receiver is available to hand off the message to.
610    pub fn send(&self, t: T) -> Result<(), SendError<T>> {
611        self.inner.send(t)
612    }
613
614    /// Attempts to send a value on this channel without blocking.
615    ///
616    /// This method differs from [`send`] by returning immediately if the
617    /// channel's buffer is full or no receiver is waiting to acquire some
618    /// data. Compared with [`send`], this function has two failure cases
619    /// instead of one (one for disconnection, one for a full buffer).
620    ///
621    /// [`send`]: Self::send
622    pub fn try_send(&self, t: T) -> Result<(), TrySendError<T>> {
623        self.inner.try_send(t)
624    }
625}
626
627impl<T> Clone for SyncSender<T> {
628    fn clone(&self) -> Self {
629        let mut state = self.inner.state.borrow_mut();
630        state.known_senders += 1;
631        drop(state);
632        Self {
633            inner: self.inner.clone(),
634        }
635    }
636}
637
638impl<T> Drop for SyncSender<T> {
639    fn drop(&mut self) {
640        if ExecutionState::should_stop() {
641            return;
642        }
643        let mut state = self.inner.state.borrow_mut();
644        assert!(state.known_senders > 0);
645        state.known_senders -= 1;
646        if state.known_senders == 0 {
647            // Last sender was dropped; wake up any receivers
648            for &tid in state.waiting_receivers.iter() {
649                ExecutionState::with(|s| s.get_mut(tid).unblock());
650            }
651        }
652    }
653}
654
655#[cfg(test)]
656mod tests {
657    use super::*;
658
659    #[test]
660    fn unique_resource_signature_mpsc() {
661        shuttle_schedulers::check_random(
662            || {
663                let (sender1, _) = channel::<i32>();
664                let (sender2, _) = channel::<i32>();
665                assert_ne!(sender1.inner.signature, sender2.inner.signature);
666            },
667            1,
668        );
669    }
670}