http_body_util/combinators/
fuse.rs1use std::{
2 pin::Pin,
3 task::{Context, Poll},
4};
5
6use http_body::{Body, Frame, SizeHint};
7
8#[derive(Debug)]
21pub struct Fuse<B> {
22 inner: Option<B>,
23}
24
25impl<B> Fuse<B>
26where
27 B: Body,
28{
29 pub fn new(body: B) -> Self {
31 Self {
32 inner: if body.is_end_stream() {
33 None
34 } else {
35 Some(body)
36 },
37 }
38 }
39}
40
41impl<B> Body for Fuse<B>
42where
43 B: Body + Unpin,
44{
45 type Data = B::Data;
46 type Error = B::Error;
47
48 fn poll_frame(
49 self: Pin<&mut Self>,
50 cx: &mut Context<'_>,
51 ) -> Poll<Option<Result<Frame<B::Data>, B::Error>>> {
52 let Self { inner } = self.get_mut();
53
54 let poll = inner
55 .as_mut()
56 .map(|mut inner| match Pin::new(&mut inner).poll_frame(cx) {
57 frame @ Poll::Ready(Some(Ok(_))) => (frame, inner.is_end_stream()),
58 end @ Poll::Ready(Some(Err(_)) | None) => (end, true),
59 poll @ Poll::Pending => (poll, false),
60 });
61
62 if let Some((frame, eos)) = poll {
63 eos.then(|| inner.take());
64 frame
65 } else {
66 Poll::Ready(None)
67 }
68 }
69
70 fn is_end_stream(&self) -> bool {
71 self.inner.is_none()
72 }
73
74 fn size_hint(&self) -> SizeHint {
75 self.inner
76 .as_ref()
77 .map(B::size_hint)
78 .unwrap_or_else(|| SizeHint::with_exact(0))
79 }
80}
81
82#[cfg(test)]
83mod tests {
84 use super::*;
85 use bytes::Bytes;
86 use std::collections::VecDeque;
87
88 type PollFrame = Poll<Option<Result<Frame<Bytes>, Error>>>;
90
91 type Error = &'static str;
92
93 struct Mock<'count> {
94 poll_count: &'count mut u8,
95 polls: VecDeque<PollFrame>,
96 }
97
98 #[test]
99 fn empty_never_polls() {
100 let mut count = 0_u8;
101 let empty = Mock::new(&mut count, []);
102 debug_assert!(empty.is_end_stream());
103 let fused = Fuse::new(empty);
104 assert!(fused.inner.is_none());
105 drop(fused);
106 assert_eq!(count, 0);
107 }
108
109 #[test]
110 fn stops_polling_after_none() {
111 let mut count = 0_u8;
112 let empty = Mock::new(&mut count, [Poll::Ready(None)]);
113 debug_assert!(!empty.is_end_stream());
114 let mut fused = Fuse::new(empty);
115 assert!(fused.inner.is_some());
116
117 let waker = futures_util::task::noop_waker();
118 let mut cx = Context::from_waker(&waker);
119 match Pin::new(&mut fused).poll_frame(&mut cx) {
120 Poll::Ready(None) => {}
121 other => panic!("unexpected poll outcome: {:?}", other),
122 }
123
124 assert!(fused.inner.is_none());
125 match Pin::new(&mut fused).poll_frame(&mut cx) {
126 Poll::Ready(None) => {}
127 other => panic!("unexpected poll outcome: {:?}", other),
128 }
129
130 drop(fused);
131 assert_eq!(count, 1);
132 }
133
134 #[test]
135 fn stops_polling_after_some_eos() {
136 let mut count = 0_u8;
137 let body = Mock::new(
138 &mut count,
139 [Poll::Ready(Some(Ok(Frame::data(Bytes::from_static(
140 b"hello",
141 )))))],
142 );
143 debug_assert!(!body.is_end_stream());
144 let mut fused = Fuse::new(body);
145 assert!(fused.inner.is_some());
146
147 let waker = futures_util::task::noop_waker();
148 let mut cx = Context::from_waker(&waker);
149
150 match Pin::new(&mut fused).poll_frame(&mut cx) {
151 Poll::Ready(Some(Ok(bytes))) => assert_eq!(bytes.into_data().expect("data"), "hello"),
152 other => panic!("unexpected poll outcome: {:?}", other),
153 }
154
155 assert!(fused.inner.is_none());
156 match Pin::new(&mut fused).poll_frame(&mut cx) {
157 Poll::Ready(None) => {}
158 other => panic!("unexpected poll outcome: {:?}", other),
159 }
160
161 drop(fused);
162 assert_eq!(count, 1);
163 }
164
165 #[test]
166 fn stops_polling_after_some_error() {
167 let mut count = 0_u8;
168 let body = Mock::new(
169 &mut count,
170 [
171 Poll::Ready(Some(Ok(Frame::data(Bytes::from_static(b"hello"))))),
172 Poll::Ready(Some(Err("oh no"))),
173 Poll::Ready(Some(Ok(Frame::data(Bytes::from_static(b"world"))))),
174 ],
175 );
176 debug_assert!(!body.is_end_stream());
177 let mut fused = Fuse::new(body);
178 assert!(fused.inner.is_some());
179
180 let waker = futures_util::task::noop_waker();
181 let mut cx = Context::from_waker(&waker);
182
183 match Pin::new(&mut fused).poll_frame(&mut cx) {
184 Poll::Ready(Some(Ok(bytes))) => assert_eq!(bytes.into_data().expect("data"), "hello"),
185 other => panic!("unexpected poll outcome: {:?}", other),
186 }
187
188 assert!(fused.inner.is_some());
189 match Pin::new(&mut fused).poll_frame(&mut cx) {
190 Poll::Ready(Some(Err("oh no"))) => {}
191 other => panic!("unexpected poll outcome: {:?}", other),
192 }
193
194 assert!(fused.inner.is_none());
195 match Pin::new(&mut fused).poll_frame(&mut cx) {
196 Poll::Ready(None) => {}
197 other => panic!("unexpected poll outcome: {:?}", other),
198 }
199
200 drop(fused);
201 assert_eq!(count, 2);
202 }
203
204 impl<'count> Mock<'count> {
207 fn new(poll_count: &'count mut u8, polls: impl IntoIterator<Item = PollFrame>) -> Self {
208 Self {
209 poll_count,
210 polls: polls.into_iter().collect(),
211 }
212 }
213 }
214
215 impl Body for Mock<'_> {
216 type Data = Bytes;
217 type Error = &'static str;
218
219 fn poll_frame(
220 self: Pin<&mut Self>,
221 _cx: &mut Context<'_>,
222 ) -> Poll<Option<Result<Frame<Self::Data>, Self::Error>>> {
223 let Self { poll_count, polls } = self.get_mut();
224 **poll_count = poll_count.saturating_add(1);
225 polls.pop_front().unwrap_or(Poll::Ready(None))
226 }
227
228 fn is_end_stream(&self) -> bool {
229 self.polls.is_empty()
230 }
231 }
232}