Skip to main content

xmtp_api_grpc/streams/
multiplexed.rs

1//! Multiplexed Stream Type
2
3use std::{pin::Pin, task::Poll};
4
5use futures::{Stream, stream::FusedStream};
6use pin_project::pin_project;
7use std::task::Context;
8
9/// Attempts to pull items from both streams. S1 will always
10/// be polled before S2. if S1 finishes first, the stream is considered over.
11/// Attempts naive fairness by polling S2 before stream close if S1 finished and S2 still has Ready
12/// items.
13pub fn multiplexed<S1, S2>(s1: S1, s2: S2) -> MultiplexedStream<S1, S2> {
14    MultiplexedStream {
15        s1,
16        s2,
17        s1_ended: false,
18        terminated: false,
19    }
20}
21
22#[pin_project]
23/// Stream for the [multiplexed()] function
24pub struct MultiplexedStream<S1, S2> {
25    #[pin]
26    s1: S1,
27    #[pin]
28    s2: S2,
29    s1_ended: bool,
30    terminated: bool,
31}
32
33impl<S1, S2> MultiplexedStream<S1, S2>
34where
35    S1: Stream<Item = S2::Item>,
36    S2: Stream,
37{
38    fn poll_s2(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<S2::Item>> {
39        let mut this = self.as_mut().project();
40        if let Poll::Ready(Some(item)) = this.s2.as_mut().poll_next(cx) {
41            return Poll::Ready(Some(item));
42        }
43        *this.terminated = true;
44        Poll::Ready(None)
45    }
46}
47
48impl<S1, S2> Stream for MultiplexedStream<S1, S2>
49where
50    S1: Stream<Item = S2::Item>,
51    S2: Stream,
52{
53    type Item = S2::Item;
54
55    fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
56        let mut this = self.as_mut().project();
57        if *this.terminated {
58            return Poll::Ready(None);
59        }
60        if *this.s1_ended {
61            return self.poll_s2(cx);
62        }
63        match this.s1.as_mut().poll_next(cx) {
64            Poll::Ready(Some(item)) => {
65                return Poll::Ready(Some(item));
66            }
67            Poll::Ready(None) => {
68                *this.s1_ended = true;
69                return self.poll_s2(cx);
70            }
71            Poll::Pending => (),
72        };
73
74        match this.s2.as_mut().poll_next(cx) {
75            Poll::Ready(Some(item)) => Poll::Ready(Some(item)),
76            Poll::Ready(None) => Poll::Ready(None),
77            Poll::Pending => Poll::Pending,
78        }
79    }
80}
81
82impl<S1, S2> FusedStream for MultiplexedStream<S1, S2>
83where
84    S1: Stream<Item = S2::Item>,
85    S2: Stream,
86{
87    fn is_terminated(&self) -> bool {
88        self.terminated
89    }
90}
91
92#[cfg(test)]
93mod tests {
94    use super::*;
95    use futures::stream;
96    use futures_test::{
97        assert_stream_done, assert_stream_next, stream::StreamTestExt, task::noop_context,
98    };
99
100    #[xmtp_common::test]
101    fn does_not_starve_s2() {
102        let s1 = stream::iter(vec![1, 2, 3]);
103        let s2 = stream::iter(vec![4, 5, 6]);
104        let stream = multiplexed(s1, s2);
105        futures::pin_mut!(stream);
106        for i in 1..=6 {
107            assert_stream_next!(stream, i);
108        }
109        assert_stream_done!(stream)
110    }
111
112    #[xmtp_common::test]
113    fn polls_s2_in_between_s1() {
114        let s1 = stream::iter(vec![1, 2, 3]).interleave_pending();
115        let s2 = stream::iter(vec![4, 5, 6]);
116        let stream = multiplexed(s1, s2);
117        futures::pin_mut!(stream);
118        assert_stream_next!(stream, 4);
119        assert_stream_next!(stream, 1);
120        assert_stream_next!(stream, 5);
121        assert_stream_next!(stream, 2);
122        assert_stream_next!(stream, 6);
123        assert_stream_next!(stream, 3);
124        assert_stream_done!(stream)
125    }
126
127    #[xmtp_common::test]
128    fn ignores_items_after_s2_pending() {
129        let s1 = stream::iter(vec![1]);
130        let s2 = stream::iter(vec![4, 5, 6]).interleave_pending();
131        let stream = multiplexed(s1, s2);
132        futures::pin_mut!(stream);
133        assert_stream_next!(stream, 1);
134        assert_stream_done!(stream)
135    }
136
137    #[xmtp_common::test]
138    fn ends_when_s1_ends() {
139        let s1 = stream::iter(vec![1, 2, 3]);
140        let s2 = stream::iter(vec![]); // s2 ends immediately, but s1 should keep going
141        let stream = multiplexed(s1, s2);
142        futures::pin_mut!(stream);
143        for i in 1..=3 {
144            assert_stream_next!(stream, i);
145        }
146        assert_stream_done!(stream)
147    }
148
149    #[xmtp_common::test]
150    fn does_not_panic_on_polling_after_finish() {
151        let s1 = stream::iter(vec![1]);
152        let s2 = stream::iter(vec![]);
153        let stream = multiplexed(s1, s2);
154        futures::pin_mut!(stream);
155        assert_stream_next!(stream, 1);
156        assert_stream_done!(stream);
157        assert!(stream.is_terminated());
158        let mut cx = noop_context();
159        let res = stream.as_mut().poll_next(&mut cx);
160        assert_eq!(res, Poll::Ready(None));
161        let res = stream.as_mut().poll_next(&mut cx);
162        assert_eq!(res, Poll::Ready(None));
163    }
164}