marigold_impl/
multi_consumer_stream.rs1use 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 #[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 #[cfg(any(feature = "tokio", feature = "async-std"))]
133 mod runtime {
134 use super::super::MultiConsumerStream;
135 use futures::stream::StreamExt;
136
137 #[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 #[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 #[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 #[tokio::test]
187 async fn no_consumers_still_terminates() {
188 let source = futures::stream::iter(0u32..3);
189 let mcs = MultiConsumerStream::new(source);
190 mcs.run().await;
193 }
194
195 #[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 #[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 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}