Skip to main content

mpmc/
lib.rs

1// Copyright 2019 The Fuchsia Authors. All rights reserved.
2// Use of this source code is governed by a BSD-style license that can be
3// found in the LICENSE file.
4
5#![deny(missing_docs)]
6
7//! A library with futures-aware mpmc channels.
8
9use crossbeam::queue::SegQueue;
10use futures::Stream;
11use futures::channel::mpsc;
12use futures::lock::Mutex;
13use futures::sink::SinkExt;
14use futures::stream::FusedStream;
15use std::pin::Pin;
16use std::sync::{Arc, Weak};
17use std::task::{Context, Poll};
18
19/// The default number of messages that will be buffered per-receiver.
20pub const DEFAULT_CHANNEL_BUFFER_SIZE: usize = 100;
21
22/// An async sender end of an mpmc channel. Messages sent on this are received by
23/// _all_ receivers connected to it (they are duplicated).
24#[derive(Clone)]
25pub struct Sender<T> {
26    inner: Arc<Mutex<Vec<mpsc::Sender<T>>>>,
27    enqueued_senders: Arc<SegQueue<mpsc::Sender<T>>>,
28    buffer_size: usize,
29}
30
31impl<T> Default for Sender<T> {
32    fn default() -> Self {
33        Sender {
34            inner: Arc::default(),
35            enqueued_senders: Arc::default(),
36            buffer_size: DEFAULT_CHANNEL_BUFFER_SIZE,
37        }
38    }
39}
40
41impl<T: Clone> Sender<T> {
42    /// Construct a sender whose receivers will buffer the given number of messages.
43    pub fn with_buffer_size(buffer_size: usize) -> Self {
44        Self { buffer_size, ..Default::default() }
45    }
46
47    /// Sends `payload` to all receivers that exist at the time of send.
48    ///
49    /// Sending is never an error, even if there are no receivers.
50    pub async fn send(&self, payload: T) {
51        let mut inner = self.inner.lock().await;
52        while let Some(new_sender) = self.enqueued_senders.pop() {
53            inner.push(new_sender);
54        }
55
56        let mut living_senders = vec![];
57        for mut sender in inner.drain(0..) {
58            let should_live = match sender.try_send(payload.clone()).err() {
59                None => true,
60                Some(send_error) if send_error.is_disconnected() => false,
61                Some(e) => {
62                    // The receiver is full. Apply backpressure.
63                    let payload = e.into_inner();
64                    sender.send(payload).await.is_ok()
65                }
66            };
67
68            if should_live {
69                living_senders.push(sender);
70            }
71        }
72        inner.append(&mut living_senders);
73    }
74
75    /// Sends `payload` to all receivers that exist at the time of send.
76    ///
77    /// Receivers whose buffers are full will be disconnected instead of applying backpressure.
78    pub async fn send_or_disconnect(&self, payload: T) {
79        let mut inner = self.inner.lock().await;
80        while let Some(new_sender) = self.enqueued_senders.pop() {
81            inner.push(new_sender);
82        }
83
84        let mut living_senders = vec![];
85        for mut sender in inner.drain(0..) {
86            if sender.try_send(payload.clone()).is_ok() {
87                living_senders.push(sender);
88            }
89        }
90        inner.append(&mut living_senders);
91    }
92
93    /// Creates a new receiver who will receive a copy of all messages sent after its creation.
94    pub fn new_receiver(&self) -> Receiver<T> {
95        let (sender, receiver) = mpsc::channel(self.buffer_size);
96        self.enqueued_senders.push(sender);
97        Receiver {
98            sources: Arc::downgrade(&self.enqueued_senders),
99            inner: receiver,
100            buffer_size: self.buffer_size,
101        }
102    }
103}
104
105/// An async receiver end of an mpmc channel. All receivers connected to the same
106/// sender receive the same duplicated message sequence.
107///
108/// The message sequence is duplicated starting from the beginning of the
109/// instance's lifetime; messages sent before the receiver is added to the
110/// channel are not duplicated.
111pub struct Receiver<T> {
112    sources: Weak<SegQueue<mpsc::Sender<T>>>,
113    inner: mpsc::Receiver<T>,
114    buffer_size: usize,
115}
116
117impl<T: Clone> Clone for Receiver<T> {
118    fn clone(&self) -> Self {
119        if let Some(sender_set) = self.sources.upgrade() {
120            let (sender, receiver) = mpsc::channel(self.buffer_size);
121            let sources = sender_set;
122            sources.push(sender);
123            Self {
124                sources: Arc::downgrade(&sources),
125                inner: receiver,
126                buffer_size: self.buffer_size,
127            }
128        } else {
129            // The senders have all been dropped; clone to a dummy channel that just yields `None`
130            // to be consistent.
131            let (_, receiver) = mpsc::channel(1);
132            Self { sources: Weak::new(), inner: receiver, buffer_size: 1 }
133        }
134    }
135}
136
137impl<T: Clone> Stream for Receiver<T> {
138    type Item = T;
139    fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
140        Stream::poll_next(Pin::new(&mut self.inner), cx)
141    }
142}
143
144impl<T: Clone> FusedStream for Receiver<T> {
145    fn is_terminated(&self) -> bool {
146        self.inner.is_terminated()
147    }
148}