1use std::fmt;
18use std::pin::Pin;
19use std::task::{Context, Poll, ready};
20
21use bytes::Bytes;
22use futures_core::Stream;
23use memchr::memmem;
24
25use crate::Error;
26use crate::buffer::StreamBuffer;
27use crate::delimiter::{DataSearch, search_data};
28
29pub struct FinalPartDataStream<S>
39where
40 S: Stream<Item = Result<Bytes, Error>> + Send + Unpin,
41{
42 buffer: StreamBuffer<S>,
43 delimiter_finder: Box<memmem::Finder<'static>>,
44 state: DataState,
45 multipart_consumed: u64,
46 aborted: bool,
50 terminated: bool,
51}
52
53#[derive(Debug, Clone, Copy, PartialEq, Eq)]
54enum DataState {
55 Data,
58 AfterBoundary,
61 FinalCRLF,
64 Eof,
66 Done,
67}
68
69impl<S> FinalPartDataStream<S>
70where
71 S: Stream<Item = Result<Bytes, Error>> + Send + Unpin,
72{
73 pub(super) fn new(
81 buffer: StreamBuffer<S>,
82 delimiter_finder: Box<memmem::Finder<'static>>,
83 multipart_consumed: u64,
84 data_done: bool,
85 aborted: bool,
86 ) -> Self {
87 Self {
88 buffer,
89 delimiter_finder,
90 state: if data_done {
91 DataState::AfterBoundary
92 } else {
93 DataState::Data
94 },
95 multipart_consumed,
96 aborted,
97 terminated: false,
98 }
99 }
100
101 #[must_use]
109 pub fn multipart_consumed(&self) -> u64 {
110 self.multipart_consumed
111 }
112
113 fn poll_data(&mut self, cx: &mut Context<'_>) -> Poll<Result<Option<Bytes>, Error>> {
114 if self.aborted {
115 self.terminated = true;
119 return Poll::Ready(Err(Error::IncompleteStreamPart));
120 }
121
122 let delimiter_finder = &self.delimiter_finder;
123 let delimiter_len = delimiter_finder.needle().len();
124
125 loop {
126 if !self.buffer.buf.is_empty() {
127 match search_data(&self.buffer.buf, delimiter_finder) {
128 DataSearch::Found { index } => {
129 let data = self.buffer.buf.split_to(index).freeze();
130 let _ = self.buffer.buf.split_to(delimiter_len);
131 self.state = DataState::AfterBoundary;
132 return Poll::Ready(Ok((!data.is_empty()).then_some(data)));
133 }
134 DataSearch::Emit { end } => {
135 return Poll::Ready(Ok(Some(self.buffer.buf.split_to(end).freeze())));
136 }
137 DataSearch::KeepAll => {}
138 }
139 }
140
141 match ready!(self.buffer.poll_stream(cx)) {
142 Some(Ok(chunk)) => {
143 if !self.buffer.buf.is_empty() {
144 self.buffer.buf.extend_from_slice(&chunk);
145 continue;
146 }
147
148 match search_data(&chunk, delimiter_finder) {
149 DataSearch::Found { index } => {
150 let data = chunk.slice(..index);
151 self.buffer.buf.clear();
152 self.buffer
153 .buf
154 .extend_from_slice(&chunk[index.saturating_add(delimiter_len)..]);
155 self.state = DataState::AfterBoundary;
156 return Poll::Ready(Ok((!data.is_empty()).then_some(data)));
157 }
158 DataSearch::Emit { end } => {
159 let data = chunk.slice(..end);
160 self.buffer.buf.extend_from_slice(&chunk[end..]);
161 return Poll::Ready(Ok(Some(data)));
162 }
163 DataSearch::KeepAll => {
164 self.buffer.buf.extend_from_slice(&chunk);
165 }
166 }
167 }
168 Some(Err(err)) => return Poll::Ready(Err(err)),
169 None => return Poll::Ready(Err(Error::IncompleteStreamPart)),
170 }
171 }
172 }
173
174 fn poll_trailer(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Error>> {
175 loop {
176 match self.state {
177 DataState::AfterBoundary => {
178 ready!(self.fill_buf(2, cx))?;
179 if !self.buffer.buf.starts_with(b"--") {
180 let err = if self.buffer.buf.starts_with(b"\r")
181 || self.buffer.buf.first().is_some_and(|b| matches!(b, b' ' | b'\t'))
182 {
183 Error::StreamPartNotLast
184 } else {
185 Error::InvalidFormat
186 };
187 return Poll::Ready(Err(err));
188 }
189 let _ = self.buffer.buf.split_to(2);
190 self.state = DataState::FinalCRLF;
191 }
192 DataState::FinalCRLF => {
193 loop {
194 ready!(self.fill_buf(1, cx))?;
195
196 let first = self.buffer.buf[0];
197 match first {
198 b' ' | b'\t' => {
199 let _ = self.buffer.buf.split_to(1);
200 }
201 b'\r' => break,
202 _ => return Poll::Ready(Err(Error::InvalidFormat)),
203 }
204 }
205
206 ready!(self.fill_buf(2, cx))?;
207 if !self.buffer.buf.starts_with(b"\r\n") {
208 return Poll::Ready(Err(Error::InvalidFormat));
209 }
210 let _ = self.buffer.buf.split_to(2);
211 self.state = DataState::Eof;
212 }
213 DataState::Eof => {
214 if !self.buffer.buf.is_empty() {
215 return Poll::Ready(Err(Error::StreamPartNotLast));
216 }
217
218 match ready!(self.buffer.poll_stream(cx)) {
219 None => {
220 self.state = DataState::Done;
221 return Poll::Ready(Ok(()));
222 }
223 Some(Ok(_)) => return Poll::Ready(Err(Error::StreamPartNotLast)),
226 Some(Err(err)) => return Poll::Ready(Err(err)),
227 }
228 }
229 DataState::Data | DataState::Done => return Poll::Ready(Ok(())),
233 }
234 }
235 }
236
237 fn fill_buf(&mut self, len: usize, cx: &mut Context<'_>) -> Poll<Result<(), Error>> {
238 while self.buffer.buf.len() < len {
239 match ready!(self.buffer.poll_stream(cx)) {
240 Some(Ok(chunk)) => self.buffer.buf.extend_from_slice(&chunk),
241 Some(Err(err)) => return Poll::Ready(Err(err)),
242 None => return Poll::Ready(Err(Error::IncompleteStreamPart)),
243 }
244 }
245 Poll::Ready(Ok(()))
246 }
247}
248
249impl<S> fmt::Debug for FinalPartDataStream<S>
250where
251 S: Stream<Item = Result<Bytes, Error>> + Send + Unpin,
252{
253 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
254 f.debug_struct("FinalPartDataStream")
255 .field("state", &self.state)
256 .field("multipart_consumed", &self.multipart_consumed)
257 .finish_non_exhaustive()
258 }
259}
260
261impl<S> Stream for FinalPartDataStream<S>
262where
263 S: Stream<Item = Result<Bytes, Error>> + Send + Unpin,
264{
265 type Item = Result<Bytes, Error>;
266
267 fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
268 let this = self.get_mut();
269 loop {
270 if this.terminated {
271 return Poll::Ready(None);
272 }
273
274 match this.state {
275 DataState::Data => match ready!(this.poll_data(cx)) {
276 Ok(Some(item)) => return Poll::Ready(Some(Ok(item))),
277 Ok(None) => {}
281 Err(err) => {
282 this.terminated = true;
283 return Poll::Ready(Some(Err(err)));
284 }
285 },
286 DataState::AfterBoundary | DataState::FinalCRLF | DataState::Eof => {
287 if let Err(err) = ready!(this.poll_trailer(cx)) {
288 this.terminated = true;
289 return Poll::Ready(Some(Err(err)));
290 }
291 }
292 DataState::Done => return Poll::Ready(None),
293 }
294 }
295 }
296
297 fn size_hint(&self) -> (usize, Option<usize>) {
298 (0, None)
299 }
300}
301
302#[cfg(test)]
303#[allow(
304 clippy::expect_used,
305 clippy::indexing_slicing,
306 clippy::panic,
307 clippy::unreachable,
308 clippy::unwrap_used
309)]
310mod tests {
311 use super::*;
312
313 use std::pin::Pin as StdPin;
314 use std::task::{Context as TaskContext, Poll as TaskPoll};
315
316 use futures_util::StreamExt;
317 use futures_util::stream;
318 use futures_util::task::noop_waker;
319
320 use crate::delimiter::make_delimiter_finder;
321 use crate::part_data_stream::PartDataStream;
322
323 struct PendingStream;
324
325 impl Stream for PendingStream {
326 type Item = Result<Bytes, Error>;
327
328 fn poll_next(self: StdPin<&mut Self>, _cx: &mut TaskContext<'_>) -> TaskPoll<Option<Self::Item>> {
329 TaskPoll::Pending
330 }
331 }
332
333 fn with_cx<R>(f: impl FnOnce(&mut TaskContext<'_>) -> R) -> R {
334 let waker = noop_waker();
335 let mut cx = TaskContext::from_waker(&waker);
336 f(&mut cx)
337 }
338
339 fn final_with_buffer(
340 items: Vec<Result<Bytes, Error>>,
341 state: DataState,
342 ) -> FinalPartDataStream<impl Stream<Item = Result<Bytes, Error>> + Send + Sync> {
343 FinalPartDataStream {
344 buffer: StreamBuffer::new(stream::iter(items)),
345 delimiter_finder: make_delimiter_finder(b"boundary"),
346 state,
347 multipart_consumed: 0,
348 aborted: false,
349 terminated: false,
350 }
351 }
352
353 fn final_with_error(
354 err: Error,
355 state: DataState,
356 ) -> FinalPartDataStream<impl Stream<Item = Result<Bytes, Error>> + Send + Sync> {
357 final_with_buffer(vec![Err(err)], state)
358 }
359
360 fn final_with_prefix(
361 prefix: &[u8],
362 state: DataState,
363 ) -> FinalPartDataStream<impl Stream<Item = Result<Bytes, Error>> + Send + Sync> {
364 let mut buffer = StreamBuffer::new(PendingStream);
365 buffer.buf.extend_from_slice(prefix);
366 FinalPartDataStream {
367 buffer,
368 delimiter_finder: make_delimiter_finder(b"boundary"),
369 state,
370 multipart_consumed: 0,
371 aborted: false,
372 terminated: false,
373 }
374 }
375
376 #[test]
377 fn part_data_stream_pending_polls() {
378 with_cx(|cx| {
379 let buffer = StreamBuffer::new(PendingStream);
380 let mut stream = FinalPartDataStream {
381 buffer,
382 delimiter_finder: make_delimiter_finder(b"boundary"),
383 state: DataState::Data,
384 multipart_consumed: 0,
385 aborted: false,
386 terminated: false,
387 };
388 assert!(matches!(stream.poll_next_unpin(cx), TaskPoll::Pending));
389 });
390 }
391
392 #[test]
393 fn poll_data_buffered_empty_and_prefix_paths() {
394 with_cx(|cx| {
395 let mut ds = final_with_buffer(Vec::new(), DataState::Data);
396 ds.buffer.buf.extend_from_slice(b"\r\n--boundary--\r\n");
397 assert!(matches!(ds.poll_next_unpin(cx), TaskPoll::Ready(None)));
399
400 let mut ds = final_with_buffer(Vec::new(), DataState::Data);
401 ds.buffer.buf.extend_from_slice(b"0123456789abcdefghij");
402 assert!(matches!(ds.poll_next_unpin(cx), TaskPoll::Ready(Some(Ok(_)))));
403
404 let mut ds = final_with_buffer(vec![Ok(Bytes::from_static(b"tiny"))], DataState::Data);
405 let poll = ds.poll_next_unpin(cx);
406 assert!(matches!(poll, TaskPoll::Ready(Some(Err(Error::IncompleteStreamPart)))));
407 });
408 }
409
410 #[test]
411 fn fill_buf_error_and_pending_paths() {
412 with_cx(|cx| {
413 let mut ds = final_with_error(Error::InvalidFormat, DataState::AfterBoundary);
414 assert!(matches!(ds.poll_next_unpin(cx), TaskPoll::Ready(Some(Err(Error::InvalidFormat)))));
415
416 let mut ds = final_with_buffer(Vec::new(), DataState::AfterBoundary);
417 ds.buffer.buf.extend_from_slice(b"xx");
418 assert!(matches!(ds.poll_next_unpin(cx), TaskPoll::Ready(Some(Err(Error::InvalidFormat)))));
419
420 let mut ds = final_with_buffer(vec![Ok(Bytes::from_static(b"--\r\n"))], DataState::AfterBoundary);
421 assert!(matches!(ds.poll_next_unpin(cx), TaskPoll::Ready(None)));
422 });
423 }
424
425 #[test]
426 fn trailer_padding_error_and_pending_paths() {
427 with_cx(|cx| {
428 let mut ds = final_with_error(Error::InvalidFormat, DataState::FinalCRLF);
429 ds.buffer.buf.extend_from_slice(b" ");
430 assert!(matches!(ds.poll_next_unpin(cx), TaskPoll::Ready(Some(Err(Error::InvalidFormat)))));
431
432 let mut ds = final_with_prefix(b" ", DataState::FinalCRLF);
433 assert!(matches!(ds.poll_next_unpin(cx), TaskPoll::Pending));
434
435 let mut ds = final_with_buffer(Vec::new(), DataState::FinalCRLF);
436 ds.buffer.buf.extend_from_slice(b" ");
437 assert!(matches!(ds.poll_next_unpin(cx), TaskPoll::Ready(Some(Err(Error::IncompleteStreamPart)))));
438
439 let mut ds = final_with_error(Error::InvalidFormat, DataState::FinalCRLF);
440 ds.buffer.buf.extend_from_slice(b"\r");
441 assert!(matches!(ds.poll_next_unpin(cx), TaskPoll::Ready(Some(Err(Error::InvalidFormat)))));
442
443 let mut ds = final_with_prefix(b"\r", DataState::FinalCRLF);
444 assert!(matches!(ds.poll_next_unpin(cx), TaskPoll::Pending));
445
446 let mut ds = final_with_buffer(Vec::new(), DataState::FinalCRLF);
447 ds.buffer.buf.extend_from_slice(b"xx");
448 assert!(matches!(ds.poll_next_unpin(cx), TaskPoll::Ready(Some(Err(Error::InvalidFormat)))));
449 });
450 }
451
452 #[test]
453 fn trailer_final_crlf_requires_exact_bytes() {
454 with_cx(|cx| {
455 let mut ds = final_with_buffer(Vec::new(), DataState::FinalCRLF);
456 ds.buffer.buf.extend_from_slice(b"\rX");
457 assert!(matches!(ds.poll_next_unpin(cx), TaskPoll::Ready(Some(Err(Error::InvalidFormat)))));
458 });
459 }
460
461 #[test]
462 fn terminated_stream_returns_none_on_second_poll() {
463 with_cx(|cx| {
464 let ds = final_with_error(Error::InvalidFormat, DataState::Data);
465 let mut ds = Box::pin(ds);
466 assert!(matches!(
467 ds.as_mut().poll_next_unpin(cx),
468 TaskPoll::Ready(Some(Err(Error::InvalidFormat)))
469 ));
470 assert!(matches!(ds.as_mut().poll_next_unpin(cx), TaskPoll::Ready(None)));
471 });
472 }
473
474 #[test]
475 fn yields_remaining_data_then_rejects_another_part() {
476 with_cx(|cx| {
477 let mut ds = final_with_buffer(
481 vec![Ok(Bytes::from_static(
482 b"hello\r\n--boundary\r\nX: y\r\n\r\nworld\r\n--boundary--\r\n",
483 ))],
484 DataState::Data,
485 );
486 assert!(matches!(
487 ds.poll_next_unpin(cx),
488 TaskPoll::Ready(Some(Ok(bytes))) if bytes.as_ref() == b"hello"
489 ));
490 assert!(matches!(ds.poll_next_unpin(cx), TaskPoll::Ready(Some(Err(Error::StreamPartNotLast)))));
491 assert!(matches!(ds.poll_next_unpin(cx), TaskPoll::Ready(None)));
492 });
493 }
494
495 #[test]
498 fn empty_chunks_after_the_trailer_are_not_an_epilogue() {
499 with_cx(|cx| {
500 let mut ds = final_with_buffer(
501 vec![
502 Ok(Bytes::from_static(b"")),
503 Ok(Bytes::from_static(b"")),
504 Ok(Bytes::from_static(b"")),
505 ],
506 DataState::Eof,
507 );
508 assert!(matches!(ds.poll_next_unpin(cx), TaskPoll::Ready(None)));
509 });
510 }
511
512 #[test]
518 fn accessors_and_debug_are_stable() {
519 let ds = final_with_buffer(Vec::new(), DataState::Data);
520 assert_eq!(ds.size_hint(), (0, None));
521 assert_eq!(ds.multipart_consumed(), 0);
522
523 let rendered = format!("{ds:?}");
524 assert!(rendered.starts_with("FinalPartDataStream"), "{rendered}");
525 assert!(rendered.contains("state"), "{rendered}");
526 assert!(rendered.contains("multipart_consumed"), "{rendered}");
527
528 let stream = PartDataStream::new(
531 StreamBuffer::new(stream::iter(Vec::<Result<Bytes, Error>>::new())),
532 make_delimiter_finder(b"boundary"),
533 42,
534 );
535 assert_eq!(stream.multipart_consumed(), 42);
536 assert_eq!(stream.into_final().multipart_consumed(), 42);
537 }
538
539 #[test]
544 fn a_failed_taken_stream_stays_failed() {
545 with_cx(|cx| {
546 let mut stream = PartDataStream::new(
547 StreamBuffer::new(stream::iter(vec![
548 Ok(Bytes::from_static(b"0123456789abcdefghij")),
549 Err(Error::InvalidFormat),
550 Ok(Bytes::from_static(b"world\r\n--boundary--\r\n")),
551 ])),
552 make_delimiter_finder(b"boundary"),
553 5,
554 );
555 assert!(matches!(
556 stream.poll_next_unpin(cx),
557 TaskPoll::Ready(Some(Ok(bytes))) if bytes.as_ref() == b"012345678"
558 ));
559 assert!(matches!(stream.poll_next_unpin(cx), TaskPoll::Ready(Some(Err(Error::InvalidFormat)))));
560
561 let mut ds = stream.into_final();
562 assert!(matches!(ds.poll_next_unpin(cx), TaskPoll::Ready(Some(Err(Error::IncompleteStreamPart)))));
563 assert!(matches!(ds.poll_next_unpin(cx), TaskPoll::Ready(None)));
564 });
565 }
566
567 #[test]
570 fn a_truncated_taken_stream_stays_failed() {
571 with_cx(|cx| {
572 let mut stream = PartDataStream::new(
573 StreamBuffer::new(stream::iter(vec![Ok(Bytes::from_static(b"0123456789abcdefghij"))])),
574 make_delimiter_finder(b"boundary"),
575 0,
576 );
577 assert!(matches!(stream.poll_next_unpin(cx), TaskPoll::Ready(Some(Ok(_)))));
578 assert!(matches!(
579 stream.poll_next_unpin(cx),
580 TaskPoll::Ready(Some(Err(Error::IncompleteStreamPart)))
581 ));
582
583 let mut ds = stream.into_final();
584 assert!(matches!(ds.poll_next_unpin(cx), TaskPoll::Ready(Some(Err(Error::IncompleteStreamPart)))));
585 assert!(matches!(ds.poll_next_unpin(cx), TaskPoll::Ready(None)));
586 });
587 }
588}