s3s_multipart/
part_data_stream.rs1use std::fmt;
14use std::pin::Pin;
15use std::task::{Context, Poll, ready};
16
17use bytes::Bytes;
18use futures_core::Stream;
19use memchr::memmem;
20
21use crate::Error;
22use crate::buffer::StreamBuffer;
23use crate::delimiter::{DataSearch, search_data};
24use crate::final_part_data_stream::FinalPartDataStream;
25
26pub struct PartDataStream<S>
35where
36 S: Stream<Item = Result<Bytes, Error>> + Send + Unpin,
37{
38 buffer: StreamBuffer<S>,
39 delimiter_finder: Box<memmem::Finder<'static>>,
40 multipart_consumed: u64,
41 done: bool,
42 terminated: bool,
43}
44
45impl<S> PartDataStream<S>
46where
47 S: Stream<Item = Result<Bytes, Error>> + Send + Unpin,
48{
49 pub(super) fn new(buffer: StreamBuffer<S>, delimiter_finder: Box<memmem::Finder<'static>>, multipart_consumed: u64) -> Self {
51 Self {
52 buffer,
53 delimiter_finder,
54 multipart_consumed,
55 done: false,
56 terminated: false,
57 }
58 }
59
60 #[must_use]
68 pub fn multipart_consumed(&self) -> u64 {
69 self.multipart_consumed
70 }
71
72 #[must_use]
84 pub fn into_final(self) -> FinalPartDataStream<S> {
85 FinalPartDataStream::new(self.buffer, self.delimiter_finder, self.multipart_consumed, self.done, self.terminated)
86 }
87
88 fn poll_data(&mut self, cx: &mut Context<'_>) -> Poll<Option<Result<Bytes, Error>>> {
89 if self.done || self.terminated {
90 return Poll::Ready(None);
91 }
92
93 let delimiter_finder = &self.delimiter_finder;
94 let delimiter_len = delimiter_finder.needle().len();
95
96 loop {
97 if !self.buffer.buf.is_empty() {
98 match search_data(&self.buffer.buf, delimiter_finder) {
99 DataSearch::Found { index } => {
100 let data = self.buffer.buf.split_to(index).freeze();
101 let _ = self.buffer.buf.split_to(delimiter_len);
102 self.done = true;
103 return if data.is_empty() {
104 Poll::Ready(None)
105 } else {
106 Poll::Ready(Some(Ok(data)))
107 };
108 }
109 DataSearch::Emit { end } => {
110 return Poll::Ready(Some(Ok(self.buffer.buf.split_to(end).freeze())));
111 }
112 DataSearch::KeepAll => {}
113 }
114 }
115
116 match ready!(self.buffer.poll_stream(cx)) {
117 Some(Ok(chunk)) => {
118 if !self.buffer.buf.is_empty() {
119 self.buffer.buf.extend_from_slice(&chunk);
120 continue;
121 }
122
123 match search_data(&chunk, delimiter_finder) {
124 DataSearch::Found { index } => {
125 let data = chunk.slice(..index);
126 self.buffer.buf.clear();
127 self.buffer
128 .buf
129 .extend_from_slice(&chunk[index.saturating_add(delimiter_len)..]);
130 self.done = true;
131 return if data.is_empty() {
132 Poll::Ready(None)
133 } else {
134 Poll::Ready(Some(Ok(data)))
135 };
136 }
137 DataSearch::Emit { end } => {
138 let data = chunk.slice(..end);
139 self.buffer.buf.extend_from_slice(&chunk[end..]);
140 return Poll::Ready(Some(Ok(data)));
141 }
142 DataSearch::KeepAll => {
143 self.buffer.buf.extend_from_slice(&chunk);
144 }
145 }
146 }
147 Some(Err(err)) => {
148 self.terminated = true;
149 return Poll::Ready(Some(Err(err)));
150 }
151 None => {
152 self.terminated = true;
153 return Poll::Ready(Some(Err(Error::IncompleteStreamPart)));
154 }
155 }
156 }
157 }
158}
159
160impl<S> fmt::Debug for PartDataStream<S>
161where
162 S: Stream<Item = Result<Bytes, Error>> + Send + Unpin,
163{
164 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
165 f.debug_struct("PartDataStream")
166 .field("multipart_consumed", &self.multipart_consumed)
167 .field("done", &self.done)
168 .finish_non_exhaustive()
169 }
170}
171
172impl<S> Stream for PartDataStream<S>
173where
174 S: Stream<Item = Result<Bytes, Error>> + Send + Unpin,
175{
176 type Item = Result<Bytes, Error>;
177
178 fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
179 self.get_mut().poll_data(cx)
180 }
181
182 fn size_hint(&self) -> (usize, Option<usize>) {
183 (0, None)
184 }
185}
186
187#[cfg(test)]
188#[allow(
189 clippy::expect_used,
190 clippy::indexing_slicing,
191 clippy::panic,
192 clippy::unreachable,
193 clippy::unwrap_used
194)]
195mod tests {
196 use super::*;
197
198 use std::pin::Pin as StdPin;
199 use std::task::{Context as TaskContext, Poll as TaskPoll};
200
201 use futures_util::StreamExt;
202 use futures_util::stream;
203 use futures_util::task::noop_waker;
204
205 use crate::delimiter::make_delimiter_finder;
206
207 struct PendingStream;
208
209 impl Stream for PendingStream {
210 type Item = Result<Bytes, Error>;
211
212 fn poll_next(self: StdPin<&mut Self>, _cx: &mut TaskContext<'_>) -> TaskPoll<Option<Self::Item>> {
213 TaskPoll::Pending
214 }
215 }
216
217 fn with_cx<R>(f: impl FnOnce(&mut TaskContext<'_>) -> R) -> R {
218 let waker = noop_waker();
219 let mut cx = TaskContext::from_waker(&waker);
220 f(&mut cx)
221 }
222
223 fn pending_part_data_stream() -> PartDataStream<impl Stream<Item = Result<Bytes, Error>> + Send + Sync> {
224 PartDataStream::new(StreamBuffer::new(PendingStream), make_delimiter_finder(b"boundary"), 0)
225 }
226
227 fn ds_with_buffer(
228 items: Vec<Result<Bytes, Error>>,
229 ) -> PartDataStream<impl Stream<Item = Result<Bytes, Error>> + Send + Sync> {
230 PartDataStream::new(StreamBuffer::new(stream::iter(items)), make_delimiter_finder(b"boundary"), 0)
231 }
232
233 #[test]
234 fn part_data_stream_pending_polls() {
235 with_cx(|cx| {
236 let mut stream = pending_part_data_stream();
237 assert!(matches!(stream.poll_next_unpin(cx), TaskPoll::Pending));
238 });
239 }
240
241 #[test]
242 fn poll_data_buffered_empty_and_prefix_paths() {
243 with_cx(|cx| {
244 let mut ds = ds_with_buffer(Vec::new());
245 ds.buffer.buf.extend_from_slice(b"\r\n--boundary--\r\n");
246 assert!(matches!(ds.poll_next_unpin(cx), TaskPoll::Ready(None)));
248 assert!(matches!(ds.poll_next_unpin(cx), TaskPoll::Ready(None)));
249
250 let mut ds = ds_with_buffer(Vec::new());
251 ds.buffer.buf.extend_from_slice(b"0123456789abcdefghij");
252 assert!(matches!(ds.poll_next_unpin(cx), TaskPoll::Ready(Some(Ok(_)))));
253
254 let mut ds = ds_with_buffer(vec![Ok(Bytes::from_static(b"tiny"))]);
255 let poll = ds.poll_next_unpin(cx);
256 assert!(matches!(poll, TaskPoll::Ready(Some(Err(Error::IncompleteStreamPart)))));
257 assert!(matches!(ds.poll_next_unpin(cx), TaskPoll::Ready(None)));
258 });
259 }
260
261 #[test]
262 fn stream_ends_at_delimiter_without_trailer_checks() {
263 with_cx(|cx| {
264 let mut ds = ds_with_buffer(vec![Ok(Bytes::from_static(
267 b"hello\r\n--boundary\r\nX: y\r\n\r\nworld\r\n--boundary--\r\n",
268 ))]);
269 assert!(matches!(
270 ds.poll_next_unpin(cx),
271 TaskPoll::Ready(Some(Ok(bytes))) if bytes.as_ref() == b"hello"
272 ));
273 assert!(matches!(ds.poll_next_unpin(cx), TaskPoll::Ready(None)));
274 });
275 }
276
277 #[test]
278 fn into_final_after_delimiter_validates_trailer() {
279 with_cx(|cx| {
280 let mut ds = ds_with_buffer(vec![Ok(Bytes::from_static(b"hello\r\n--boundary--\r\n"))]);
281 assert!(matches!(ds.poll_next_unpin(cx), TaskPoll::Ready(Some(Ok(_)))));
282 let mut final_stream = ds.into_final();
283 assert!(matches!(final_stream.poll_next_unpin(cx), TaskPoll::Ready(None)));
284 });
285 }
286
287 #[test]
292 fn stream_error_is_reported_then_the_stream_ends() {
293 with_cx(|cx| {
294 let mut ds = ds_with_buffer(vec![
295 Ok(Bytes::from_static(b"0123456789abcdefghij")),
299 Err(Error::InvalidFormat),
300 Ok(Bytes::from_static(b"world")),
301 ]);
302 assert!(matches!(
303 ds.poll_next_unpin(cx),
304 TaskPoll::Ready(Some(Ok(bytes))) if bytes.as_ref() == b"012345678"
305 ));
306 assert!(matches!(ds.poll_next_unpin(cx), TaskPoll::Ready(Some(Err(Error::InvalidFormat)))));
307 assert!(matches!(ds.poll_next_unpin(cx), TaskPoll::Ready(None)));
308 });
309 }
310}