1use std::{
16 fmt::Debug,
17 mem::replace,
18 pin::Pin,
19 sync::{
20 Arc, Mutex, MutexGuard, PoisonError,
21 atomic::{AtomicBool, AtomicUsize, Ordering},
22 },
23 task::{Context, Poll},
24 time::Duration,
25};
26
27use bytes::Bytes;
28use futures::{Stream, StreamExt, stream, task::AtomicWaker};
29use http_body::{Body as _, Frame};
30use http_body_util::BodyExt;
31use reqwest::Version;
32use stream_shared::SharedStream;
33use tokio::{runtime::Handle, sync::watch};
34
35#[cfg(feature = "encoding")]
36use web_faith_encoding::{Coding, response::decode_stream};
37
38use crate::{
39 error::{FaithError, FaithErrorKind},
40 response::TrailersSlot,
41 stats::InnerAgentStats,
42 timing::TimingSlot,
43};
44
45pub type DynStream = dyn Stream<Item = std::result::Result<Bytes, String>> + Send + Sync;
47
48type Chain = SharedStream<Pin<Box<DynStream>>>;
50
51#[derive(Debug, Clone, Copy, PartialEq, Eq)]
54pub struct DrainPolicy {
55 pub limit: u64,
57 pub timeout: Duration,
59}
60
61impl Default for DrainPolicy {
62 fn default() -> Self {
63 Self {
64 limit: 128 * 1024,
65 timeout: Duration::from_secs(1),
66 }
67 }
68}
69
70enum Upstream {
72 Live(reqwest::Body),
74 Stopped,
76 Ended,
78}
79
80pub struct BodyShared {
82 upstream: Mutex<Upstream>,
83 upstream_waker: AtomicWaker,
85 claims: AtomicUsize,
87 aborted: AtomicBool,
89 started: AtomicBool,
91 finished: AtomicBool,
93 settled: watch::Sender<bool>,
95 version: Version,
96 drain: DrainPolicy,
97 runtime: Option<Handle>,
100 trailers: Arc<TrailersSlot>,
101 timing: Arc<TimingSlot>,
102 stats: Arc<InnerAgentStats>,
103}
104
105impl Debug for BodyShared {
106 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
107 f.debug_struct("BodyShared")
108 .field("claims", &self.claims.load(Ordering::SeqCst))
109 .field("aborted", &self.aborted.load(Ordering::SeqCst))
110 .field("version", &self.version)
111 .field("drain", &self.drain)
112 .finish_non_exhaustive()
113 }
114}
115
116pub(crate) struct BodyParts {
118 pub body: reqwest::Body,
119 pub version: Version,
120 pub drain: DrainPolicy,
121 #[cfg(feature = "encoding")]
122 pub decode: Option<Coding>,
123 pub trailers: Arc<TrailersSlot>,
124 pub timing: Arc<TimingSlot>,
125 pub stats: Arc<InnerAgentStats>,
126}
127
128impl BodyShared {
129 pub(crate) fn first_claim(parts: BodyParts) -> Arc<Claim> {
131 let shared = Arc::new(Self {
132 upstream: Mutex::new(Upstream::Live(parts.body)),
133 upstream_waker: AtomicWaker::new(),
134 claims: AtomicUsize::new(1),
135 aborted: AtomicBool::new(false),
136 started: AtomicBool::new(false),
137 finished: AtomicBool::new(false),
138 settled: watch::channel(false).0,
139 version: parts.version,
140 drain: parts.drain,
141 runtime: Handle::try_current().ok(),
142 trailers: parts.trailers,
143 timing: parts.timing,
144 stats: parts.stats,
145 });
146
147 let chain = SharedStream::new(Self::pipeline(
148 &shared,
149 #[cfg(feature = "encoding")]
150 parts.decode,
151 ));
152
153 Arc::new(Claim {
154 body: shared,
155 cursor: Mutex::new(Some(chain)),
156 given_up: AtomicBool::new(false),
157 ended: AtomicBool::new(false),
158 waker: AtomicWaker::new(),
159 })
160 }
161
162 fn pipeline(
164 shared: &Arc<Self>,
165 #[cfg(feature = "encoding")] decode: Option<Coding>,
166 ) -> Pin<Box<DynStream>> {
167 let trailers = shared.trailers.clone();
170 let bytes = Box::pin(
171 UpstreamFrames {
172 shared: shared.clone(),
173 }
174 .filter_map(move |frame| {
175 let item = match frame {
176 Err(err) => Some(Err(err)),
177 Ok(frame) => match frame.into_trailers() {
178 Ok(headers) => {
179 trailers.arrived(headers);
180 None
181 }
182 Err(frame) => Some(
183 frame
184 .into_data()
185 .map_err(|_| "unknown frame kind".to_string()),
186 ),
187 },
188 };
189 async move { item }
190 }),
191 ) as Pin<Box<DynStream>>;
192
193 #[cfg(feature = "encoding")]
194 let bytes = match decode {
195 Some(coding) => decode_stream(bytes, coding),
196 None => bytes,
197 };
198
199 let bytes = Box::pin(bytes.filter(|item| {
204 let empty = matches!(item, Ok(chunk) if chunk.is_empty());
205 async move { !empty }
206 })) as Pin<Box<DynStream>>;
207
208 let finish = shared.clone();
212 Box::pin(
213 bytes.chain(
214 stream::once(async move {
215 finish.finish();
217 })
218 .filter_map(async |()| None),
219 ),
220 )
221 }
222
223 fn upstream(&self) -> MutexGuard<'_, Upstream> {
224 self.upstream.lock().unwrap_or_else(PoisonError::into_inner)
225 }
226
227 fn finish(&self) {
232 if self.finished.swap(true, Ordering::SeqCst) {
233 return;
234 }
235 self.trailers.ended();
236 self.timing.ended();
237 if self.started.load(Ordering::SeqCst) {
238 self.stats.bodies_finished.fetch_add(1, Ordering::Relaxed);
239 }
240 }
241
242 fn opened(&self) {
244 if !self.started.swap(true, Ordering::SeqCst) {
245 self.stats.bodies_started.fetch_add(1, Ordering::Relaxed);
246 }
247 }
248
249 fn stop(self: &Arc<Self>) {
255 let taken = {
256 let mut upstream = self.upstream();
257 match replace(&mut *upstream, Upstream::Stopped) {
258 Upstream::Live(body) => Some(body),
259 other => {
260 *upstream = other;
261 None
262 }
263 }
264 };
265 self.upstream_waker.wake();
267 self.finish();
268
269 let Some(body) = taken else {
270 self.settled.send_replace(true);
271 return;
272 };
273
274 let http1 = matches!(
278 self.version,
279 Version::HTTP_09 | Version::HTTP_10 | Version::HTTP_11
280 );
281 match (&self.runtime, http1) {
282 (Some(runtime), true) => {
283 let shared = self.clone();
284 runtime.spawn(async move {
285 drain(body, shared.drain).await;
286 shared.settled.send_replace(true);
287 });
288 }
289 _ => {
290 drop(body);
291 self.settled.send_replace(true);
292 }
293 }
294 }
295
296 #[cfg_attr(not(feature = "unstable-internals"), allow(dead_code))]
299 pub(crate) fn abort(self: &Arc<Self>) {
300 if !self.aborted.swap(true, Ordering::SeqCst) {
301 self.stop();
302 }
303 }
304
305 pub(crate) fn claims_left(&self) -> usize {
307 self.claims.load(Ordering::SeqCst)
308 }
309
310 pub(crate) async fn settled(&self) {
312 let mut rx = self.settled.subscribe();
313 let _ = rx.wait_for(|settled| *settled).await;
314 }
315}
316
317async fn drain(mut body: reqwest::Body, policy: DrainPolicy) {
321 if policy.limit == 0 || body.size_hint().lower() > policy.limit {
324 return;
325 }
326
327 let _ = tokio::time::timeout(policy.timeout, async {
328 let mut read: u64 = 0;
329 while let Some(frame) = body.frame().await {
330 let Ok(frame) = frame else {
331 return;
332 };
333 if let Some(data) = frame.data_ref() {
334 read += data.len() as u64;
335 if read > policy.limit {
336 return;
337 }
338 }
339 }
340 })
341 .await;
342}
343
344struct UpstreamFrames {
346 shared: Arc<BodyShared>,
347}
348
349impl Stream for UpstreamFrames {
350 type Item = Result<Frame<Bytes>, String>;
351
352 fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
353 self.shared.upstream_waker.register(cx.waker());
354 let mut upstream = self.shared.upstream();
355 match &mut *upstream {
356 Upstream::Live(body) => match Pin::new(body).poll_frame(cx) {
357 Poll::Pending => Poll::Pending,
358 Poll::Ready(Some(Ok(frame))) => Poll::Ready(Some(Ok(frame))),
359 Poll::Ready(Some(Err(err))) => {
360 *upstream = Upstream::Ended;
363 drop(upstream);
364 self.shared.settled.send_replace(true);
365 Poll::Ready(Some(Err(err.to_string())))
366 }
367 Poll::Ready(None) => {
368 *upstream = Upstream::Ended;
369 drop(upstream);
370 self.shared.settled.send_replace(true);
371 Poll::Ready(None)
372 }
373 },
374 Upstream::Stopped => {
375 *upstream = Upstream::Ended;
378 Poll::Ready(Some(Err("the transfer was stopped".to_string())))
379 }
380 Upstream::Ended => Poll::Ready(None),
381 }
382 }
383}
384
385pub struct Claim {
391 body: Arc<BodyShared>,
392 cursor: Mutex<Option<Chain>>,
394 given_up: AtomicBool,
395 ended: AtomicBool,
398 waker: AtomicWaker,
400}
401
402impl Debug for Claim {
403 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
404 f.debug_struct("Claim")
405 .field("body", &self.body)
406 .field("given_up", &self.given_up.load(Ordering::SeqCst))
407 .finish_non_exhaustive()
408 }
409}
410
411impl Claim {
412 fn cursor(&self) -> MutexGuard<'_, Option<Chain>> {
413 self.cursor.lock().unwrap_or_else(PoisonError::into_inner)
414 }
415
416 pub(crate) fn duplicate(&self) -> Option<Arc<Self>> {
420 let cursor = self.cursor();
421 let chain = cursor.as_ref()?.clone();
422 self.body.claims.fetch_add(1, Ordering::SeqCst);
423 Some(Arc::new(Self {
424 body: self.body.clone(),
425 cursor: Mutex::new(Some(chain)),
426 given_up: AtomicBool::new(false),
427 ended: AtomicBool::new(false),
428 waker: AtomicWaker::new(),
429 }))
430 }
431
432 pub(crate) fn is_given_up(&self) -> bool {
434 self.given_up.load(Ordering::SeqCst)
435 }
436
437 pub(crate) fn body(&self) -> &Arc<BodyShared> {
439 &self.body
440 }
441
442 pub(crate) fn give_up(&self) -> bool {
446 if self.given_up.swap(true, Ordering::SeqCst) {
447 return false;
448 }
449
450 let cursor = self.cursor().take();
453 drop(cursor);
454 self.waker.wake();
455
456 if self.body.claims.fetch_sub(1, Ordering::SeqCst) == 1 {
457 self.body.stop();
458 true
459 } else {
460 false
461 }
462 }
463
464 pub(crate) fn reader(self: &Arc<Self>) -> Result<BodyReader, FaithError> {
466 if self.is_given_up() {
467 return Err(FaithErrorKind::ResponseAlreadyDisturbed.into());
468 }
469 self.body.opened();
470 Ok(BodyReader {
471 claim: self.clone(),
472 done: false,
473 })
474 }
475}
476
477impl Drop for Claim {
478 fn drop(&mut self) {
479 self.give_up();
480 }
481}
482
483pub struct BodyReader {
490 claim: Arc<Claim>,
491 done: bool,
492}
493
494impl Debug for BodyReader {
495 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
496 f.debug_struct("BodyReader")
497 .field("claim", &self.claim)
498 .field("done", &self.done)
499 .finish()
500 }
501}
502
503impl BodyReader {
504 #[cfg(feature = "unstable-internals")]
507 pub fn canceller(&self) -> BodyCanceller {
508 BodyCanceller(self.claim.clone())
509 }
510
511 fn fail(&mut self, kind: FaithErrorKind) -> Poll<Option<Result<Bytes, FaithError>>> {
512 self.done = true;
513 Poll::Ready(Some(Err(kind.into())))
514 }
515}
516
517impl Stream for BodyReader {
518 type Item = Result<Bytes, FaithError>;
519
520 fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
521 if self.done {
522 return Poll::Ready(None);
523 }
524
525 self.claim.waker.register(cx.waker());
526 if self.claim.body.aborted.load(Ordering::SeqCst) {
529 return self.fail(FaithErrorKind::Aborted);
530 }
531
532 let polled = self
535 .claim
536 .cursor()
537 .as_mut()
538 .map(|chain| Pin::new(chain).poll_next(cx));
539
540 match polled {
541 None if self.claim.ended.load(Ordering::SeqCst) => {
543 self.done = true;
544 Poll::Ready(None)
545 }
546 None => self.fail(FaithErrorKind::ResponseAlreadyDisturbed),
547 Some(Poll::Pending) => Poll::Pending,
548 Some(Poll::Ready(None)) => {
549 self.claim.ended.store(true, Ordering::SeqCst);
550 self.done = true;
551 Poll::Ready(None)
552 }
553 Some(Poll::Ready(Some(Ok(chunk)))) => Poll::Ready(Some(Ok(chunk))),
554 Some(Poll::Ready(Some(Err(err)))) => {
555 if self.claim.body.aborted.load(Ordering::SeqCst) {
557 return self.fail(FaithErrorKind::Aborted);
558 }
559 self.done = true;
560 Poll::Ready(Some(Err(FaithError::new(FaithErrorKind::BodyStream, err))))
561 }
562 }
563 }
564}
565
566impl Drop for BodyReader {
567 fn drop(&mut self) {
568 self.claim.give_up();
569 }
570}
571
572#[cfg(feature = "unstable-internals")]
574#[derive(Debug, Clone)]
575pub struct BodyCanceller(Arc<Claim>);
576
577#[cfg(feature = "unstable-internals")]
578impl BodyCanceller {
579 pub fn cancel(&self) {
581 self.0.give_up();
582 }
583}
584
585#[cfg(test)]
586mod tests {
587 use std::{sync::atomic::AtomicU64, time::Instant};
588
589 use http_body::SizeHint;
590
591 use super::*;
592 use crate::timing::RequestTiming;
593
594 #[derive(Clone, Default)]
596 struct Probe {
597 read: Arc<AtomicU64>,
598 dropped: Arc<AtomicBool>,
599 }
600
601 impl Probe {
602 fn read(&self) -> u64 {
603 self.read.load(Ordering::SeqCst)
604 }
605
606 fn dropped(&self) -> bool {
607 self.dropped.load(Ordering::SeqCst)
608 }
609 }
610
611 struct TestBody {
614 probe: Probe,
615 length: Option<u64>,
616 stall_after: Option<u64>,
617 sent: u64,
618 }
619
620 const CHUNK: u64 = 1024;
621
622 impl http_body::Body for TestBody {
623 type Data = Bytes;
624 type Error = std::io::Error;
625
626 fn poll_frame(
627 mut self: Pin<&mut Self>,
628 _cx: &mut Context<'_>,
629 ) -> Poll<Option<Result<Frame<Bytes>, Self::Error>>> {
630 if self.length.is_some_and(|length| self.sent >= length) {
631 return Poll::Ready(None);
632 }
633 if self.stall_after.is_some_and(|stall| self.sent >= stall) {
634 return Poll::Pending;
635 }
636 self.sent += 1;
637 self.probe.read.fetch_add(CHUNK, Ordering::SeqCst);
638 Poll::Ready(Some(Ok(Frame::data(Bytes::from(vec![7; CHUNK as usize])))))
639 }
640
641 fn size_hint(&self) -> SizeHint {
642 match self.length {
643 Some(length) => SizeHint::with_exact((length - self.sent) * CHUNK),
644 None => SizeHint::default(),
645 }
646 }
647 }
648
649 impl Drop for TestBody {
650 fn drop(&mut self) {
651 self.probe.dropped.store(true, Ordering::SeqCst);
652 }
653 }
654
655 struct Built {
656 claim: Arc<Claim>,
657 probe: Probe,
658 stats: Arc<InnerAgentStats>,
659 trailers: Arc<TrailersSlot>,
660 }
661
662 fn build(version: Version, length: Option<u64>, stall_after: Option<u64>) -> Built {
663 build_with(version, length, stall_after, DrainPolicy::default())
664 }
665
666 fn build_with(
667 version: Version,
668 length: Option<u64>,
669 stall_after: Option<u64>,
670 drain: DrainPolicy,
671 ) -> Built {
672 let probe = Probe::default();
673 let stats = Arc::new(InnerAgentStats::default());
674 let trailers = Arc::new(TrailersSlot::default());
675 let claim = BodyShared::first_claim(BodyParts {
676 body: reqwest::Body::wrap(TestBody {
677 probe: probe.clone(),
678 length,
679 stall_after,
680 sent: 0,
681 }),
682 version,
683 drain,
684 #[cfg(feature = "encoding")]
685 decode: None,
686 trailers: trailers.clone(),
687 timing: Arc::new(TimingSlot::new(Instant::now(), RequestTiming::default())),
688 stats: stats.clone(),
689 });
690 Built {
691 claim,
692 probe,
693 stats,
694 trailers,
695 }
696 }
697
698 #[tokio::test]
701 async fn the_last_claim_going_drops_a_multiplexed_body() {
702 let built = build(Version::HTTP_2, None, None);
703 assert!(built.claim.give_up(), "the only claim is the last");
704 assert!(built.probe.dropped(), "the body is dropped");
705 assert_eq!(built.probe.read(), 0, "without reading any of it");
706 built.claim.body().settled().await;
707 }
708
709 #[tokio::test]
712 async fn a_clone_keeps_the_transfer_going() {
713 let built = build(Version::HTTP_2, None, None);
714 let clone = built.claim.duplicate().expect("an unread claim duplicates");
715
716 assert!(!built.claim.give_up(), "the original is not the last claim");
717 assert!(!built.probe.dropped(), "the body stays for the clone");
718
719 let mut reader = clone.reader().expect("the clone reads");
720 assert!(
721 reader.next().await.is_some_and(|chunk| chunk.is_ok()),
722 "the clone reads on"
723 );
724
725 drop(reader);
726 assert!(
727 built.probe.dropped(),
728 "dropping the clone's reader stops the transfer"
729 );
730 }
731
732 #[tokio::test]
734 async fn a_given_up_claim_has_nothing_to_give() {
735 let built = build(Version::HTTP_2, None, None);
736 let _clone = built.claim.duplicate().expect("an unread claim duplicates");
737 built.claim.give_up();
738
739 assert!(
740 built.claim.duplicate().is_none(),
741 "no copy of a given-up claim"
742 );
743 assert!(
744 matches!(
745 built.claim.reader().map(|_| ()).map_err(|err| err.kind()),
746 Err(FaithErrorKind::ResponseAlreadyDisturbed)
747 ),
748 "and no reader either"
749 );
750 }
751
752 #[tokio::test]
754 async fn a_small_http1_remainder_is_drained() {
755 let built = build(Version::HTTP_11, Some(10), None);
756 built.claim.give_up();
757 built.claim.body().settled().await;
758 assert_eq!(
759 built.probe.read(),
760 10 * CHUNK,
761 "the whole remainder is read"
762 );
763 }
764
765 #[tokio::test]
767 async fn an_http1_remainder_over_the_limit_closes_at_once() {
768 let built = build(Version::HTTP_11, Some(1024), None);
769 built.claim.give_up();
770 built.claim.body().settled().await;
771 assert_eq!(built.probe.read(), 0, "none of it is read");
772 assert!(
773 built.probe.dropped(),
774 "the body is dropped, closing the connection"
775 );
776 }
777
778 #[tokio::test]
780 async fn an_endless_http1_body_is_read_to_the_limit_then_dropped() {
781 let built = build(Version::HTTP_11, None, None);
782 built.claim.give_up();
783 built.claim.body().settled().await;
784 let limit = DrainPolicy::default().limit;
785 assert!(built.probe.read() > limit, "the drain reads past the limit");
786 assert!(
787 built.probe.read() <= limit + CHUNK,
788 "by no more than a chunk"
789 );
790 assert!(built.probe.dropped(), "then drops the body");
791 }
792
793 #[tokio::test]
795 async fn a_zero_drain_limit_always_closes() {
796 let built = build_with(
797 Version::HTTP_11,
798 Some(1),
799 None,
800 DrainPolicy {
801 limit: 0,
802 ..Default::default()
803 },
804 );
805 built.claim.give_up();
806 built.claim.body().settled().await;
807 assert_eq!(built.probe.read(), 0, "nothing is read");
808 assert!(built.probe.dropped(), "the body is dropped");
809 }
810
811 #[tokio::test]
813 async fn a_stalled_drain_is_bounded_by_its_timeout() {
814 let built = build_with(
815 Version::HTTP_11,
816 Some(20),
817 Some(5),
818 DrainPolicy {
819 timeout: Duration::from_millis(50),
820 ..Default::default()
821 },
822 );
823 built.claim.give_up();
824 tokio::time::timeout(Duration::from_secs(5), built.claim.body().settled())
825 .await
826 .expect("the drain settles");
827 assert!(built.probe.dropped(), "the stalled body is dropped");
828 }
829
830 #[tokio::test]
832 async fn an_abort_errors_readers_ahead_of_buffered_chunks() {
833 let built = build(Version::HTTP_2, None, None);
834 let clone = built.claim.duplicate().expect("an unread claim duplicates");
835
836 let mut ahead = built.claim.reader().expect("a reader");
838 for _ in 0..3 {
839 ahead.next().await.expect("a chunk").expect("that reads");
840 }
841
842 built.claim.body().abort();
843 assert!(built.probe.dropped(), "the abort drops the body");
844
845 let mut behind = clone.reader().expect("a reader");
846 let first = behind.next().await.expect("an item");
847 assert!(
848 matches!(
849 first.map_err(|err| err.kind()),
850 Err(FaithErrorKind::Aborted)
851 ),
852 "the clone's first read is the abort, not a buffered chunk"
853 );
854 assert!(behind.next().await.is_none(), "and nothing after it");
855 }
856
857 #[tokio::test]
860 async fn a_body_read_to_its_end_finishes_once() {
861 let built = build(Version::HTTP_2, Some(3), None);
862
863 let mut reader = built.claim.reader().expect("a reader");
864 let mut second = built.claim.reader().expect("a second reader");
865 let mut bytes = 0;
866 while let Some(chunk) = reader.next().await {
867 bytes += chunk.expect("the chunk reads").len() as u64;
868 }
869 assert_eq!(bytes, 3 * CHUNK);
870 drop(reader);
871
872 assert!(
873 second.next().await.is_none(),
874 "the second reader sees the end"
875 );
876 assert_eq!(built.stats.bodies_started.load(Ordering::SeqCst), 1);
877 assert_eq!(built.stats.bodies_finished.load(Ordering::SeqCst), 1);
878 assert!(matches!(
879 built.trailers.settled().await,
880 crate::response::Trailers::None
881 ));
882 }
883
884 #[tokio::test]
886 async fn a_body_given_up_early_settles_its_bookkeeping() {
887 let built = build(Version::HTTP_2, None, None);
888 let mut reader = built.claim.reader().expect("a reader");
889 reader.next().await.expect("a chunk").expect("that reads");
890 drop(reader);
891
892 assert!(matches!(
893 built.trailers.settled().await,
894 crate::response::Trailers::None
895 ));
896 assert_eq!(built.stats.bodies_finished.load(Ordering::SeqCst), 1);
897 }
898}