Skip to main content

marigold_impl/
multi_consumer_stream.rs

1use core::marker::PhantomData;
2use core::pin::Pin;
3use futures::channel::mpsc::Receiver;
4use futures::channel::mpsc::Sender;
5use futures::future::Future;
6use futures::sink::SinkExt;
7use futures::stream::FuturesUnordered;
8use futures::stream::Stream;
9use futures::stream::StreamExt;
10use futures::task::Context;
11use futures::task::Poll;
12
13const BUFFER_SIZE: usize = 1;
14
15pub struct MultiConsumerStream<
16    T: std::marker::Send + 'static,
17    S: Stream<Item = T> + std::marker::Unpin + std::marker::Send + 'static,
18> {
19    inner_stream: S,
20    senders: Vec<Sender<T>>,
21}
22
23impl<
24        T: std::marker::Send + Copy + 'static,
25        S: Stream<Item = T> + std::marker::Unpin + std::marker::Send + 'static,
26    > MultiConsumerStream<T, S>
27{
28    pub fn new(s: S) -> Self {
29        MultiConsumerStream {
30            inner_stream: s,
31            senders: Vec::new(),
32        }
33    }
34
35    pub fn get(&mut self) -> Receiver<T> {
36        let (sender, receiver) = futures::channel::mpsc::channel(BUFFER_SIZE);
37        self.senders.push(sender);
38        receiver
39    }
40
41    pub async fn run(mut self) {
42        self.senders.shrink_to_fit();
43
44        #[cfg(any(feature = "async-std", feature = "tokio"))]
45        crate::async_runtime::spawn(async move {
46            while let Some(v) = self.inner_stream.next().await {
47                let mut futures = self
48                    .senders
49                    .iter_mut()
50                    .map(|sender| sender.feed(v))
51                    .collect::<FuturesUnordered<_>>();
52                while let Some(_result) = futures.next().await {}
53            }
54            self.senders.iter_mut().for_each(|s| s.disconnect());
55        });
56
57        #[cfg(not(any(feature = "async-std", feature = "tokio")))]
58        {
59            while let Some(v) = self.inner_stream.next().await {
60                let mut futures = self
61                    .senders
62                    .iter_mut()
63                    .map(|sender| sender.feed(v))
64                    .collect::<FuturesUnordered<_>>();
65                while let Some(_result) = futures.next().await {}
66            }
67            self.senders.iter_mut().for_each(|s| s.disconnect());
68        }
69    }
70}
71
72pub struct RunFutureAsStream<T: Unpin, O, F: Future<Output = O>> {
73    future: Pin<Box<F>>,
74    t: PhantomData<T>,
75}
76
77impl<T: Unpin, O, F: Future<Output = O>> RunFutureAsStream<T, O, F> {
78    pub fn new(f: Pin<Box<F>>) -> RunFutureAsStream<T, O, F> {
79        RunFutureAsStream {
80            future: f,
81            t: PhantomData,
82        }
83    }
84}
85
86impl<T: std::marker::Send + Unpin + 'static, O, F: Future<Output = O>> Stream
87    for RunFutureAsStream<T, O, F>
88{
89    type Item = T;
90
91    fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
92        let future = &mut self.future;
93        match Pin::new(future).poll(cx) {
94            Poll::Pending => Poll::Pending,
95            Poll::Ready(_) => Poll::Ready(None),
96        }
97    }
98
99    fn size_hint(&self) -> (usize, Option<usize>) {
100        (0, None)
101    }
102}
103
104#[cfg(test)]
105mod tests {
106    use super::*;
107    use futures::stream::StreamExt;
108
109    // `RunFutureAsStream` should always yield `None` once the wrapped future
110    // resolves and otherwise propagates `Pending`. These invariants matter
111    // because in generated marigold code the future drives a
112    // `MultiConsumerStream::run`, and downstream consumers poll the wrapper
113    // alongside their receivers.
114    #[tokio::test]
115    async fn run_future_as_stream_completes_with_none() {
116        let fut = Box::pin(async {});
117        let mut s: RunFutureAsStream<u8, (), _> = RunFutureAsStream::new(fut);
118        assert_eq!(s.next().await, None);
119    }
120
121    #[tokio::test]
122    async fn run_future_as_stream_size_hint_is_unbounded() {
123        let fut = Box::pin(async {});
124        let s: RunFutureAsStream<u8, (), _> = RunFutureAsStream::new(fut);
125        assert_eq!(s.size_hint(), (0, None));
126    }
127
128    // The remaining tests exercise concurrent fan-out, which only happens
129    // when an async runtime is enabled (otherwise `MultiConsumerStream::run`
130    // is sequential and would deadlock against the BUFFER_SIZE-1 channel
131    // when consumers haven't been awaited yet).
132    #[cfg(any(feature = "tokio", feature = "async-std"))]
133    mod runtime {
134        use super::super::MultiConsumerStream;
135        use futures::stream::StreamExt;
136
137        // Single-consumer fan-out should be a faithful relay of the source
138        // stream, in order, with end-of-stream signalled by the receiver
139        // closing.
140        #[tokio::test]
141        async fn single_consumer_receives_all_items_in_order() {
142            let source = futures::stream::iter(0u32..10);
143            let mut mcs = MultiConsumerStream::new(source);
144            let recv = mcs.get();
145            mcs.run().await;
146            let collected: Vec<u32> = recv.collect().await;
147            assert_eq!(collected, (0u32..10).collect::<Vec<_>>());
148        }
149
150        // Every registered consumer must observe the entire source stream.
151        #[tokio::test]
152        async fn multiple_consumers_each_receive_full_stream() {
153            let source = futures::stream::iter(vec![10i32, 20, 30, 40, 50]);
154            let mut mcs = MultiConsumerStream::new(source);
155            let r1 = mcs.get();
156            let r2 = mcs.get();
157            let r3 = mcs.get();
158            mcs.run().await;
159            let (a, b, c) = futures::join!(
160                r1.collect::<Vec<_>>(),
161                r2.collect::<Vec<_>>(),
162                r3.collect::<Vec<_>>()
163            );
164            let expected = vec![10, 20, 30, 40, 50];
165            assert_eq!(a, expected);
166            assert_eq!(b, expected);
167            assert_eq!(c, expected);
168        }
169
170        // An empty source yields empty consumers but must still terminate
171        // (no hang waiting on items that never arrive).
172        #[tokio::test]
173        async fn empty_source_produces_empty_consumers() {
174            let source = futures::stream::iter(Vec::<u8>::new());
175            let mut mcs = MultiConsumerStream::new(source);
176            let r1 = mcs.get();
177            let r2 = mcs.get();
178            mcs.run().await;
179            let (a, b) = futures::join!(r1.collect::<Vec<_>>(), r2.collect::<Vec<_>>());
180            assert!(a.is_empty());
181            assert!(b.is_empty());
182        }
183
184        // Registering zero consumers is allowed; `run` must still drain the
185        // source and return.
186        #[tokio::test]
187        async fn no_consumers_still_terminates() {
188            let source = futures::stream::iter(0u32..3);
189            let mcs = MultiConsumerStream::new(source);
190            // No `.get()` calls. If run() ever blocked without consumers,
191            // tokio::test's default timeout would fail this test.
192            mcs.run().await;
193        }
194
195        // Each consumer must terminate (receiver closes) once the source
196        // stream is exhausted, so downstream `.collect()` does not hang.
197        #[tokio::test]
198        async fn consumers_terminate_after_source_exhaustion() {
199            let source = futures::stream::iter(vec![1u8]);
200            let mut mcs = MultiConsumerStream::new(source);
201            let mut r = mcs.get();
202            mcs.run().await;
203            assert_eq!(r.next().await, Some(1));
204            assert_eq!(r.next().await, None);
205        }
206
207        // Consumers must be able to poll concurrently without one blocking
208        // another beyond the BUFFER_SIZE backpressure window. This guards
209        // against a regression where `feed` is replaced with a serial
210        // `send_all`-style call that would prevent fast consumers from
211        // making progress when a slow consumer exists.
212        #[tokio::test]
213        async fn slow_consumer_does_not_lose_items() {
214            let source = futures::stream::iter(0u32..50);
215            let mut mcs = MultiConsumerStream::new(source);
216            let fast = mcs.get();
217            let slow = mcs.get();
218            mcs.run().await;
219
220            // Slow consumer yields between items but must still get all of them.
221            let slow_task = async {
222                let mut out = Vec::new();
223                let mut s = slow;
224                while let Some(v) = s.next().await {
225                    tokio::task::yield_now().await;
226                    out.push(v);
227                }
228                out
229            };
230            let (fast_out, slow_out) = futures::join!(fast.collect::<Vec<_>>(), slow_task);
231            let expected: Vec<u32> = (0u32..50).collect();
232            assert_eq!(fast_out, expected);
233            assert_eq!(slow_out, expected);
234        }
235    }
236}