Skip to main content

tokio/sync/mpsc/
chan.rs

1use crate::loom::cell::UnsafeCell;
2use crate::loom::future::AtomicWaker;
3use crate::loom::sync::atomic::AtomicUsize;
4use crate::loom::sync::Arc;
5use crate::runtime::park::CachedParkThread;
6use crate::sync::mpsc::error::TryRecvError;
7use crate::sync::mpsc::{bounded, list, unbounded};
8use crate::sync::notify::Notify;
9use crate::util::cacheline::CachePadded;
10
11use std::fmt;
12use std::panic;
13use std::process;
14use std::sync::atomic::Ordering::{AcqRel, Acquire, Relaxed, Release};
15use std::task::Poll::{Pending, Ready};
16use std::task::{ready, Context, Poll};
17
18/// Channel sender.
19pub(crate) struct Tx<T, S> {
20    inner: Arc<Chan<T, S>>,
21}
22
23impl<T, S: fmt::Debug> fmt::Debug for Tx<T, S> {
24    fn fmt(&self, fmt: &mut fmt::Formatter<'_>) -> fmt::Result {
25        fmt.debug_struct("Tx").field("inner", &self.inner).finish()
26    }
27}
28
29/// Channel receiver.
30pub(crate) struct Rx<T, S: Semaphore> {
31    inner: Arc<Chan<T, S>>,
32}
33
34impl<T, S: Semaphore + fmt::Debug> fmt::Debug for Rx<T, S> {
35    fn fmt(&self, fmt: &mut fmt::Formatter<'_>) -> fmt::Result {
36        fmt.debug_struct("Rx").field("inner", &self.inner).finish()
37    }
38}
39
40pub(crate) trait Semaphore {
41    fn is_idle(&self) -> bool;
42
43    fn add_permit(&self);
44
45    fn add_permits(&self, n: usize);
46
47    fn close(&self);
48
49    fn is_closed(&self) -> bool;
50}
51
52pub(super) struct Chan<T, S> {
53    /// Handle to the push half of the lock-free list.
54    tx: CachePadded<list::Tx<T>>,
55
56    /// Receiver waker. Notified when a value is pushed into the channel.
57    rx_waker: CachePadded<AtomicWaker>,
58
59    /// Notifies all tasks listening for the receiver being dropped.
60    notify_rx_closed: Notify,
61
62    /// Coordinates access to channel's capacity.
63    semaphore: S,
64
65    /// Tracks the number of outstanding sender handles.
66    ///
67    /// When this drops to zero, the send half of the channel is closed.
68    tx_count: AtomicUsize,
69
70    /// Tracks the number of outstanding weak sender handles.
71    tx_weak_count: AtomicUsize,
72
73    /// Only accessed by `Rx` handle.
74    rx_fields: UnsafeCell<RxFields<T>>,
75}
76
77impl<T, S> fmt::Debug for Chan<T, S>
78where
79    S: fmt::Debug,
80{
81    fn fmt(&self, fmt: &mut fmt::Formatter<'_>) -> fmt::Result {
82        fmt.debug_struct("Chan")
83            .field("tx", &*self.tx)
84            .field("semaphore", &self.semaphore)
85            .field("rx_waker", &*self.rx_waker)
86            .field("tx_count", &self.tx_count)
87            .field("rx_fields", &"...")
88            .finish()
89    }
90}
91
92/// Fields only accessed by `Rx` handle.
93struct RxFields<T> {
94    /// Channel receiver. This field is only accessed by the `Receiver` type.
95    list: list::Rx<T>,
96
97    /// `true` if `Rx::close` is called.
98    rx_closed: bool,
99}
100
101impl<T> fmt::Debug for RxFields<T> {
102    fn fmt(&self, fmt: &mut fmt::Formatter<'_>) -> fmt::Result {
103        fmt.debug_struct("RxFields")
104            .field("list", &self.list)
105            .field("rx_closed", &self.rx_closed)
106            .finish()
107    }
108}
109
110unsafe impl<T: Send, S: Send> Send for Chan<T, S> {}
111unsafe impl<T: Send, S: Sync> Sync for Chan<T, S> {}
112impl<T, S> panic::RefUnwindSafe for Chan<T, S> {}
113impl<T, S> panic::UnwindSafe for Chan<T, S> {}
114
115pub(crate) fn channel<T, S: Semaphore>(semaphore: S) -> (Tx<T, S>, Rx<T, S>) {
116    let (tx, rx) = list::channel();
117    channel_from_list(tx, rx, semaphore)
118}
119
120#[cfg(all(test, not(loom)))]
121pub(crate) fn channel_from_index<T, S: Semaphore>(
122    start_index: usize,
123    semaphore: S,
124) -> (Tx<T, S>, Rx<T, S>) {
125    let (tx, rx) = list::channel_from_index(start_index);
126    channel_from_list(tx, rx, semaphore)
127}
128
129fn channel_from_list<T, S: Semaphore>(
130    tx: list::Tx<T>,
131    rx: list::Rx<T>,
132    semaphore: S,
133) -> (Tx<T, S>, Rx<T, S>) {
134    let chan = Arc::new(Chan {
135        notify_rx_closed: Notify::new(),
136        tx: CachePadded::new(tx),
137        semaphore,
138        rx_waker: CachePadded::new(AtomicWaker::new()),
139        tx_count: AtomicUsize::new(1),
140        tx_weak_count: AtomicUsize::new(0),
141        rx_fields: UnsafeCell::new(RxFields {
142            list: rx,
143            rx_closed: false,
144        }),
145    });
146
147    (Tx::new(chan.clone()), Rx::new(chan))
148}
149
150// ===== impl Tx =====
151
152impl<T, S> Tx<T, S> {
153    fn new(chan: Arc<Chan<T, S>>) -> Tx<T, S> {
154        Tx { inner: chan }
155    }
156
157    pub(super) fn strong_count(&self) -> usize {
158        self.inner.tx_count.load(Acquire)
159    }
160
161    pub(super) fn weak_count(&self) -> usize {
162        self.inner.tx_weak_count.load(Relaxed)
163    }
164
165    pub(super) fn downgrade(&self) -> Arc<Chan<T, S>> {
166        self.inner.increment_weak_count();
167
168        self.inner.clone()
169    }
170
171    // Returns the upgraded channel or None if the upgrade failed.
172    pub(super) fn upgrade(chan: Arc<Chan<T, S>>) -> Option<Self> {
173        let mut tx_count = chan.tx_count.load(Acquire);
174
175        loop {
176            if tx_count == 0 {
177                // channel is closed
178                return None;
179            }
180
181            match chan
182                .tx_count
183                .compare_exchange_weak(tx_count, tx_count + 1, AcqRel, Acquire)
184            {
185                Ok(_) => return Some(Tx { inner: chan }),
186                Err(prev_count) => tx_count = prev_count,
187            }
188        }
189    }
190
191    pub(super) fn semaphore(&self) -> &S {
192        &self.inner.semaphore
193    }
194
195    /// Send a message and notify the receiver.
196    pub(crate) fn send(&self, value: T) {
197        self.inner.send(value);
198    }
199
200    /// Wake the receive half
201    pub(crate) fn wake_rx(&self) {
202        self.inner.rx_waker.wake();
203    }
204
205    /// Returns `true` if senders belong to the same channel.
206    pub(crate) fn same_channel(&self, other: &Self) -> bool {
207        Arc::ptr_eq(&self.inner, &other.inner)
208    }
209}
210
211impl<T, S: Semaphore> Tx<T, S> {
212    pub(crate) fn is_closed(&self) -> bool {
213        self.inner.semaphore.is_closed()
214    }
215
216    pub(crate) async fn closed(&self) {
217        // In order to avoid a race condition, we first request a notification,
218        // **then** check whether the semaphore is closed. If the semaphore is
219        // closed the notification request is dropped.
220        let notified = self.inner.notify_rx_closed.notified();
221
222        if self.inner.semaphore.is_closed() {
223            return;
224        }
225        notified.await;
226    }
227}
228
229impl<T, S> Clone for Tx<T, S> {
230    fn clone(&self) -> Tx<T, S> {
231        // Using a Relaxed ordering here is sufficient as the caller holds a
232        // strong ref to `self`, preventing a concurrent decrement to zero.
233        self.inner.tx_count.fetch_add(1, Relaxed);
234
235        Tx {
236            inner: self.inner.clone(),
237        }
238    }
239}
240
241impl<T, S> Drop for Tx<T, S> {
242    fn drop(&mut self) {
243        if self.inner.tx_count.fetch_sub(1, AcqRel) != 1 {
244            return;
245        }
246
247        // Close the list, which sends a `Close` message
248        self.inner.tx.close();
249
250        // Notify the receiver
251        self.wake_rx();
252    }
253}
254
255// ===== impl Rx =====
256
257impl<T, S: Semaphore> Rx<T, S> {
258    fn new(chan: Arc<Chan<T, S>>) -> Rx<T, S> {
259        Rx { inner: chan }
260    }
261
262    pub(crate) fn close(&mut self) {
263        self.inner.rx_fields.with_mut(|rx_fields_ptr| {
264            let rx_fields = unsafe { &mut *rx_fields_ptr };
265
266            if rx_fields.rx_closed {
267                return;
268            }
269
270            rx_fields.rx_closed = true;
271        });
272
273        self.inner.semaphore.close();
274        self.inner.notify_rx_closed.notify_waiters();
275    }
276
277    pub(crate) fn is_closed(&self) -> bool {
278        // There two internal states that can represent a closed channel
279        //
280        //  1. When `close` is called.
281        //  In this case, the inner semaphore will be closed.
282        //
283        //  2. When all senders are dropped.
284        //  In this case, the semaphore remains unclosed, and the `index` in the list won't
285        //  reach the tail position. It is necessary to check the list if the last block is
286        //  `closed`.
287        self.inner.semaphore.is_closed() || self.inner.tx_count.load(Acquire) == 0
288    }
289
290    pub(crate) fn is_empty(&self) -> bool {
291        self.inner.rx_fields.with(|rx_fields_ptr| {
292            let rx_fields = unsafe { &*rx_fields_ptr };
293            rx_fields.list.is_empty(&self.inner.tx)
294        })
295    }
296
297    pub(crate) fn len(&self) -> usize {
298        self.inner.rx_fields.with(|rx_fields_ptr| {
299            let rx_fields = unsafe { &*rx_fields_ptr };
300            rx_fields.list.len(&self.inner.tx)
301        })
302    }
303
304    /// Receive the next value
305    pub(crate) fn recv(&mut self, cx: &mut Context<'_>) -> Poll<Option<T>> {
306        use super::block::Read;
307
308        ready!(crate::trace::trace_leaf(cx));
309
310        // Keep track of task budget
311        let coop = ready!(crate::task::coop::poll_proceed(cx));
312
313        self.inner.rx_fields.with_mut(|rx_fields_ptr| {
314            let rx_fields = unsafe { &mut *rx_fields_ptr };
315
316            macro_rules! try_recv {
317                () => {
318                    match rx_fields.list.pop(&self.inner.tx) {
319                        Some(Read::Value(value)) => {
320                            self.inner.semaphore.add_permit();
321                            coop.made_progress();
322                            return Ready(Some(value));
323                        }
324                        Some(Read::Closed) => {
325                            // A channel is closed when all tx handles are
326                            // dropped. Dropping a tx handle releases memory,
327                            // which ensures that if dropping the tx handle is
328                            // visible, then all messages sent are also visible.
329                            debug_assert!(self.inner.semaphore.is_idle());
330                            coop.made_progress();
331                            return Ready(None);
332                        }
333                        None => {} // fall through
334                    }
335                };
336            }
337
338            try_recv!();
339
340            self.inner.rx_waker.register_by_ref(cx.waker());
341
342            // It is possible that a value was pushed between attempting to read
343            // and registering the task, so we have to check the channel a
344            // second time here.
345            try_recv!();
346
347            if rx_fields.rx_closed && self.inner.semaphore.is_idle() {
348                coop.made_progress();
349                Ready(None)
350            } else {
351                Pending
352            }
353        })
354    }
355
356    /// Receives up to `limit` values into `buffer`
357    ///
358    /// For `limit > 0`, receives up to limit values into `buffer`.
359    /// For `limit == 0`, immediately returns Ready(0).
360    pub(crate) fn recv_many(
361        &mut self,
362        cx: &mut Context<'_>,
363        buffer: &mut Vec<T>,
364        limit: usize,
365    ) -> Poll<usize> {
366        use super::block::Read;
367
368        ready!(crate::trace::trace_leaf(cx));
369
370        // Keep track of task budget
371        let coop = ready!(crate::task::coop::poll_proceed(cx));
372
373        if limit == 0 {
374            coop.made_progress();
375            return Ready(0usize);
376        }
377
378        let mut remaining = limit;
379        let initial_length = buffer.len();
380
381        self.inner.rx_fields.with_mut(|rx_fields_ptr| {
382            let rx_fields = unsafe { &mut *rx_fields_ptr };
383            macro_rules! try_recv {
384                () => {
385                    while remaining > 0 {
386                        match rx_fields.list.pop(&self.inner.tx) {
387                            Some(Read::Value(value)) => {
388                                remaining -= 1;
389                                buffer.push(value);
390                            }
391
392                            Some(Read::Closed) => {
393                                let number_added = buffer.len() - initial_length;
394                                if number_added > 0 {
395                                    self.inner.semaphore.add_permits(number_added);
396                                }
397                                // A channel is closed when all tx handles are
398                                // dropped. Dropping a tx handle releases memory,
399                                // which ensures that if dropping the tx handle is
400                                // visible, then all messages sent are also visible.
401                                debug_assert!(self.inner.semaphore.is_idle());
402                                coop.made_progress();
403                                return Ready(number_added);
404                            }
405
406                            None => {
407                                break; // fall through
408                            }
409                        }
410                    }
411                    let number_added = buffer.len() - initial_length;
412                    if number_added > 0 {
413                        self.inner.semaphore.add_permits(number_added);
414                        coop.made_progress();
415                        return Ready(number_added);
416                    }
417                };
418            }
419
420            try_recv!();
421
422            self.inner.rx_waker.register_by_ref(cx.waker());
423
424            // It is possible that a value was pushed between attempting to read
425            // and registering the task, so we have to check the channel a
426            // second time here.
427            try_recv!();
428
429            if rx_fields.rx_closed && self.inner.semaphore.is_idle() {
430                debug_assert_eq!(buffer.len(), initial_length);
431                coop.made_progress();
432                Ready(0usize)
433            } else {
434                Pending
435            }
436        })
437    }
438
439    /// Try to receive the next value.
440    pub(crate) fn try_recv(&mut self) -> Result<T, TryRecvError> {
441        use super::list::TryPopResult;
442
443        self.inner.rx_fields.with_mut(|rx_fields_ptr| {
444            let rx_fields = unsafe { &mut *rx_fields_ptr };
445
446            macro_rules! try_recv {
447                () => {
448                    match rx_fields.list.try_pop(&self.inner.tx) {
449                        TryPopResult::Ok(value) => {
450                            self.inner.semaphore.add_permit();
451                            return Ok(value);
452                        }
453                        TryPopResult::Closed => return Err(TryRecvError::Disconnected),
454                        // If close() was called, an empty queue should report Disconnected.
455                        TryPopResult::Empty
456                            if rx_fields.rx_closed && self.inner.semaphore.is_idle() =>
457                        {
458                            return Err(TryRecvError::Disconnected)
459                        }
460                        TryPopResult::Empty => return Err(TryRecvError::Empty),
461                        TryPopResult::Busy => {} // fall through
462                    }
463                };
464            }
465
466            try_recv!();
467
468            // If a previous `poll_recv` call has set a waker, we wake it here.
469            // This allows us to put our own CachedParkThread waker in the
470            // AtomicWaker slot instead.
471            //
472            // This is not a spurious wakeup to `poll_recv` since we just got a
473            // Busy from `try_pop`, which only happens if there are messages in
474            // the queue.
475            self.inner.rx_waker.wake();
476
477            // Park the thread until the problematic send has completed.
478            let mut park = CachedParkThread::new();
479            let waker = park.waker().unwrap();
480            loop {
481                self.inner.rx_waker.register_by_ref(&waker);
482                // It is possible that the problematic send has now completed,
483                // so we have to check for messages again.
484                try_recv!();
485                park.park();
486            }
487        })
488    }
489
490    pub(super) fn semaphore(&self) -> &S {
491        &self.inner.semaphore
492    }
493
494    pub(super) fn sender_strong_count(&self) -> usize {
495        self.inner.tx_count.load(Acquire)
496    }
497
498    pub(super) fn sender_weak_count(&self) -> usize {
499        self.inner.tx_weak_count.load(Relaxed)
500    }
501}
502
503impl<T, S: Semaphore> Drop for Rx<T, S> {
504    fn drop(&mut self) {
505        use super::block::Read::Value;
506
507        self.close();
508
509        self.inner.rx_fields.with_mut(|rx_fields_ptr| {
510            let rx_fields = unsafe { &mut *rx_fields_ptr };
511            struct Guard<'a, T, S: Semaphore> {
512                list: &'a mut list::Rx<T>,
513                tx: &'a list::Tx<T>,
514                sem: &'a S,
515            }
516
517            impl<'a, T, S: Semaphore> Guard<'a, T, S> {
518                fn drain(&mut self) {
519                    // call T's destructor.
520                    while let Some(Value(_)) = self.list.pop(self.tx) {
521                        self.sem.add_permit();
522                    }
523                }
524            }
525
526            impl<'a, T, S: Semaphore> Drop for Guard<'a, T, S> {
527                fn drop(&mut self) {
528                    self.drain();
529                }
530            }
531
532            let mut guard = Guard {
533                list: &mut rx_fields.list,
534                tx: &self.inner.tx,
535                sem: &self.inner.semaphore,
536            };
537
538            // When Rx is dropped, there is nothing for a task to poll anymore.
539            // This means we can drop our waker to potentially free up resources.
540            // Do so before draining the channel where panics may occur.
541            self.inner.rx_waker.take_waker();
542
543            guard.drain();
544        });
545    }
546}
547
548// ===== impl Chan =====
549
550impl<T, S> Chan<T, S> {
551    fn send(&self, value: T) {
552        // Push the value
553        self.tx.push(value);
554
555        // Notify the rx task
556        self.rx_waker.wake();
557    }
558
559    pub(super) fn decrement_weak_count(&self) {
560        self.tx_weak_count.fetch_sub(1, Relaxed);
561    }
562
563    pub(super) fn increment_weak_count(&self) {
564        self.tx_weak_count.fetch_add(1, Relaxed);
565    }
566
567    pub(super) fn strong_count(&self) -> usize {
568        self.tx_count.load(Acquire)
569    }
570
571    pub(super) fn weak_count(&self) -> usize {
572        self.tx_weak_count.load(Relaxed)
573    }
574}
575
576impl<T, S> Drop for Chan<T, S> {
577    fn drop(&mut self) {
578        use super::block::Read::Value;
579
580        // Safety: the only owner of the rx fields is Chan, and being
581        // inside its own Drop means we're the last ones to touch it.
582        self.rx_fields.with_mut(|rx_fields_ptr| {
583            let rx_fields = unsafe { &mut *rx_fields_ptr };
584
585            while let Some(Value(_)) = rx_fields.list.pop(&self.tx) {}
586            unsafe { rx_fields.list.free_blocks() };
587        });
588    }
589}
590
591// ===== impl Semaphore for (::Semaphore, capacity) =====
592
593impl Semaphore for bounded::Semaphore {
594    fn add_permit(&self) {
595        self.semaphore.release(1);
596    }
597
598    fn add_permits(&self, n: usize) {
599        self.semaphore.release(n)
600    }
601
602    fn is_idle(&self) -> bool {
603        self.semaphore.available_permits() == self.bound
604    }
605
606    fn close(&self) {
607        self.semaphore.close();
608    }
609
610    fn is_closed(&self) -> bool {
611        self.semaphore.is_closed()
612    }
613}
614
615// ===== impl Semaphore for AtomicUsize =====
616
617impl Semaphore for unbounded::Semaphore {
618    fn add_permit(&self) {
619        let prev = self.0.fetch_sub(2, Release);
620
621        if prev >> 1 == 0 {
622            // Something went wrong
623            process::abort();
624        }
625    }
626
627    fn add_permits(&self, n: usize) {
628        let prev = self.0.fetch_sub(n << 1, Release);
629
630        if (prev >> 1) < n {
631            // Something went wrong
632            process::abort();
633        }
634    }
635
636    fn is_idle(&self) -> bool {
637        self.0.load(Acquire) >> 1 == 0
638    }
639
640    fn close(&self) {
641        self.0.fetch_or(1, Release);
642    }
643
644    fn is_closed(&self) -> bool {
645        self.0.load(Acquire) & 1 == 1
646    }
647}