xmtp_api_grpc/streams/
multiplexed.rs1use std::{pin::Pin, task::Poll};
4
5use futures::{Stream, stream::FusedStream};
6use pin_project::pin_project;
7use std::task::Context;
8
9pub 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]
23pub 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![]); 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}