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}