use std::{
fmt::Debug,
mem::replace,
pin::Pin,
sync::{
Arc, Mutex, MutexGuard, PoisonError,
atomic::{AtomicBool, AtomicUsize, Ordering},
},
task::{Context, Poll},
time::Duration,
};
use bytes::Bytes;
use futures::{Stream, StreamExt, stream, task::AtomicWaker};
use http_body::{Body as _, Frame};
use http_body_util::BodyExt;
use reqwest::Version;
use stream_shared::SharedStream;
use tokio::{runtime::Handle, sync::watch};
#[cfg(feature = "encoding")]
use web_faith_encoding::{Coding, response::decode_stream};
use crate::{
error::{FaithError, FaithErrorKind},
response::TrailersSlot,
stats::InnerAgentStats,
timing::TimingSlot,
};
pub type DynStream = dyn Stream<Item = std::result::Result<Bytes, String>> + Send + Sync;
type Chain = SharedStream<Pin<Box<DynStream>>>;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct DrainPolicy {
pub limit: u64,
pub timeout: Duration,
}
impl Default for DrainPolicy {
fn default() -> Self {
Self {
limit: 128 * 1024,
timeout: Duration::from_secs(1),
}
}
}
enum Upstream {
Live(reqwest::Body),
Stopped,
Ended,
}
pub struct BodyShared {
upstream: Mutex<Upstream>,
upstream_waker: AtomicWaker,
claims: AtomicUsize,
aborted: AtomicBool,
started: AtomicBool,
finished: AtomicBool,
settled: watch::Sender<bool>,
version: Version,
drain: DrainPolicy,
runtime: Option<Handle>,
trailers: Arc<TrailersSlot>,
timing: Arc<TimingSlot>,
stats: Arc<InnerAgentStats>,
}
impl Debug for BodyShared {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("BodyShared")
.field("claims", &self.claims.load(Ordering::SeqCst))
.field("aborted", &self.aborted.load(Ordering::SeqCst))
.field("version", &self.version)
.field("drain", &self.drain)
.finish_non_exhaustive()
}
}
pub(crate) struct BodyParts {
pub body: reqwest::Body,
pub version: Version,
pub drain: DrainPolicy,
#[cfg(feature = "encoding")]
pub decode: Option<Coding>,
pub trailers: Arc<TrailersSlot>,
pub timing: Arc<TimingSlot>,
pub stats: Arc<InnerAgentStats>,
}
impl BodyShared {
pub(crate) fn first_claim(parts: BodyParts) -> Arc<Claim> {
let shared = Arc::new(Self {
upstream: Mutex::new(Upstream::Live(parts.body)),
upstream_waker: AtomicWaker::new(),
claims: AtomicUsize::new(1),
aborted: AtomicBool::new(false),
started: AtomicBool::new(false),
finished: AtomicBool::new(false),
settled: watch::channel(false).0,
version: parts.version,
drain: parts.drain,
runtime: Handle::try_current().ok(),
trailers: parts.trailers,
timing: parts.timing,
stats: parts.stats,
});
let chain = SharedStream::new(Self::pipeline(
&shared,
#[cfg(feature = "encoding")]
parts.decode,
));
Arc::new(Claim {
body: shared,
cursor: Mutex::new(Some(chain)),
given_up: AtomicBool::new(false),
ended: AtomicBool::new(false),
waker: AtomicWaker::new(),
})
}
fn pipeline(
shared: &Arc<Self>,
#[cfg(feature = "encoding")] decode: Option<Coding>,
) -> Pin<Box<DynStream>> {
let trailers = shared.trailers.clone();
let bytes = Box::pin(
UpstreamFrames {
shared: shared.clone(),
}
.filter_map(move |frame| {
let item = match frame {
Err(err) => Some(Err(err)),
Ok(frame) => match frame.into_trailers() {
Ok(headers) => {
trailers.arrived(headers);
None
}
Err(frame) => Some(
frame
.into_data()
.map_err(|_| "unknown frame kind".to_string()),
),
},
};
async move { item }
}),
) as Pin<Box<DynStream>>;
#[cfg(feature = "encoding")]
let bytes = match decode {
Some(coding) => decode_stream(bytes, coding),
None => bytes,
};
let bytes = Box::pin(bytes.filter(|item| {
let empty = matches!(item, Ok(chunk) if chunk.is_empty());
async move { !empty }
})) as Pin<Box<DynStream>>;
let finish = shared.clone();
Box::pin(
bytes.chain(
stream::once(async move {
finish.finish();
})
.filter_map(async |()| None),
),
)
}
fn upstream(&self) -> MutexGuard<'_, Upstream> {
self.upstream.lock().unwrap_or_else(PoisonError::into_inner)
}
fn finish(&self) {
if self.finished.swap(true, Ordering::SeqCst) {
return;
}
self.trailers.ended();
self.timing.ended();
if self.started.load(Ordering::SeqCst) {
self.stats.bodies_finished.fetch_add(1, Ordering::Relaxed);
}
}
fn opened(&self) {
if !self.started.swap(true, Ordering::SeqCst) {
self.stats.bodies_started.fetch_add(1, Ordering::Relaxed);
}
}
fn stop(self: &Arc<Self>) {
let taken = {
let mut upstream = self.upstream();
match replace(&mut *upstream, Upstream::Stopped) {
Upstream::Live(body) => Some(body),
other => {
*upstream = other;
None
}
}
};
self.upstream_waker.wake();
self.finish();
let Some(body) = taken else {
self.settled.send_replace(true);
return;
};
let http1 = matches!(
self.version,
Version::HTTP_09 | Version::HTTP_10 | Version::HTTP_11
);
match (&self.runtime, http1) {
(Some(runtime), true) => {
let shared = self.clone();
runtime.spawn(async move {
drain(body, shared.drain).await;
shared.settled.send_replace(true);
});
}
_ => {
drop(body);
self.settled.send_replace(true);
}
}
}
#[cfg_attr(not(feature = "unstable-internals"), allow(dead_code))]
pub(crate) fn abort(self: &Arc<Self>) {
if !self.aborted.swap(true, Ordering::SeqCst) {
self.stop();
}
}
pub(crate) fn claims_left(&self) -> usize {
self.claims.load(Ordering::SeqCst)
}
pub(crate) async fn settled(&self) {
let mut rx = self.settled.subscribe();
let _ = rx.wait_for(|settled| *settled).await;
}
}
async fn drain(mut body: reqwest::Body, policy: DrainPolicy) {
if policy.limit == 0 || body.size_hint().lower() > policy.limit {
return;
}
let _ = tokio::time::timeout(policy.timeout, async {
let mut read: u64 = 0;
while let Some(frame) = body.frame().await {
let Ok(frame) = frame else {
return;
};
if let Some(data) = frame.data_ref() {
read += data.len() as u64;
if read > policy.limit {
return;
}
}
}
})
.await;
}
struct UpstreamFrames {
shared: Arc<BodyShared>,
}
impl Stream for UpstreamFrames {
type Item = Result<Frame<Bytes>, String>;
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
self.shared.upstream_waker.register(cx.waker());
let mut upstream = self.shared.upstream();
match &mut *upstream {
Upstream::Live(body) => match Pin::new(body).poll_frame(cx) {
Poll::Pending => Poll::Pending,
Poll::Ready(Some(Ok(frame))) => Poll::Ready(Some(Ok(frame))),
Poll::Ready(Some(Err(err))) => {
*upstream = Upstream::Ended;
drop(upstream);
self.shared.settled.send_replace(true);
Poll::Ready(Some(Err(err.to_string())))
}
Poll::Ready(None) => {
*upstream = Upstream::Ended;
drop(upstream);
self.shared.settled.send_replace(true);
Poll::Ready(None)
}
},
Upstream::Stopped => {
*upstream = Upstream::Ended;
Poll::Ready(Some(Err("the transfer was stopped".to_string())))
}
Upstream::Ended => Poll::Ready(None),
}
}
}
pub struct Claim {
body: Arc<BodyShared>,
cursor: Mutex<Option<Chain>>,
given_up: AtomicBool,
ended: AtomicBool,
waker: AtomicWaker,
}
impl Debug for Claim {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Claim")
.field("body", &self.body)
.field("given_up", &self.given_up.load(Ordering::SeqCst))
.finish_non_exhaustive()
}
}
impl Claim {
fn cursor(&self) -> MutexGuard<'_, Option<Chain>> {
self.cursor.lock().unwrap_or_else(PoisonError::into_inner)
}
pub(crate) fn duplicate(&self) -> Option<Arc<Self>> {
let cursor = self.cursor();
let chain = cursor.as_ref()?.clone();
self.body.claims.fetch_add(1, Ordering::SeqCst);
Some(Arc::new(Self {
body: self.body.clone(),
cursor: Mutex::new(Some(chain)),
given_up: AtomicBool::new(false),
ended: AtomicBool::new(false),
waker: AtomicWaker::new(),
}))
}
pub(crate) fn is_given_up(&self) -> bool {
self.given_up.load(Ordering::SeqCst)
}
pub(crate) fn body(&self) -> &Arc<BodyShared> {
&self.body
}
pub(crate) fn give_up(&self) -> bool {
if self.given_up.swap(true, Ordering::SeqCst) {
return false;
}
let cursor = self.cursor().take();
drop(cursor);
self.waker.wake();
if self.body.claims.fetch_sub(1, Ordering::SeqCst) == 1 {
self.body.stop();
true
} else {
false
}
}
pub(crate) fn reader(self: &Arc<Self>) -> Result<BodyReader, FaithError> {
if self.is_given_up() {
return Err(FaithErrorKind::ResponseAlreadyDisturbed.into());
}
self.body.opened();
Ok(BodyReader {
claim: self.clone(),
done: false,
})
}
}
impl Drop for Claim {
fn drop(&mut self) {
self.give_up();
}
}
pub struct BodyReader {
claim: Arc<Claim>,
done: bool,
}
impl Debug for BodyReader {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("BodyReader")
.field("claim", &self.claim)
.field("done", &self.done)
.finish()
}
}
impl BodyReader {
#[cfg(feature = "unstable-internals")]
pub fn canceller(&self) -> BodyCanceller {
BodyCanceller(self.claim.clone())
}
fn fail(&mut self, kind: FaithErrorKind) -> Poll<Option<Result<Bytes, FaithError>>> {
self.done = true;
Poll::Ready(Some(Err(kind.into())))
}
}
impl Stream for BodyReader {
type Item = Result<Bytes, FaithError>;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
if self.done {
return Poll::Ready(None);
}
self.claim.waker.register(cx.waker());
if self.claim.body.aborted.load(Ordering::SeqCst) {
return self.fail(FaithErrorKind::Aborted);
}
let polled = self
.claim
.cursor()
.as_mut()
.map(|chain| Pin::new(chain).poll_next(cx));
match polled {
None if self.claim.ended.load(Ordering::SeqCst) => {
self.done = true;
Poll::Ready(None)
}
None => self.fail(FaithErrorKind::ResponseAlreadyDisturbed),
Some(Poll::Pending) => Poll::Pending,
Some(Poll::Ready(None)) => {
self.claim.ended.store(true, Ordering::SeqCst);
self.done = true;
Poll::Ready(None)
}
Some(Poll::Ready(Some(Ok(chunk)))) => Poll::Ready(Some(Ok(chunk))),
Some(Poll::Ready(Some(Err(err)))) => {
if self.claim.body.aborted.load(Ordering::SeqCst) {
return self.fail(FaithErrorKind::Aborted);
}
self.done = true;
Poll::Ready(Some(Err(FaithError::new(FaithErrorKind::BodyStream, err))))
}
}
}
}
impl Drop for BodyReader {
fn drop(&mut self) {
self.claim.give_up();
}
}
#[cfg(feature = "unstable-internals")]
#[derive(Debug, Clone)]
pub struct BodyCanceller(Arc<Claim>);
#[cfg(feature = "unstable-internals")]
impl BodyCanceller {
pub fn cancel(&self) {
self.0.give_up();
}
}
#[cfg(test)]
mod tests {
use std::{sync::atomic::AtomicU64, time::Instant};
use http_body::SizeHint;
use super::*;
use crate::timing::RequestTiming;
#[derive(Clone, Default)]
struct Probe {
read: Arc<AtomicU64>,
dropped: Arc<AtomicBool>,
}
impl Probe {
fn read(&self) -> u64 {
self.read.load(Ordering::SeqCst)
}
fn dropped(&self) -> bool {
self.dropped.load(Ordering::SeqCst)
}
}
struct TestBody {
probe: Probe,
length: Option<u64>,
stall_after: Option<u64>,
sent: u64,
}
const CHUNK: u64 = 1024;
impl http_body::Body for TestBody {
type Data = Bytes;
type Error = std::io::Error;
fn poll_frame(
mut self: Pin<&mut Self>,
_cx: &mut Context<'_>,
) -> Poll<Option<Result<Frame<Bytes>, Self::Error>>> {
if self.length.is_some_and(|length| self.sent >= length) {
return Poll::Ready(None);
}
if self.stall_after.is_some_and(|stall| self.sent >= stall) {
return Poll::Pending;
}
self.sent += 1;
self.probe.read.fetch_add(CHUNK, Ordering::SeqCst);
Poll::Ready(Some(Ok(Frame::data(Bytes::from(vec![7; CHUNK as usize])))))
}
fn size_hint(&self) -> SizeHint {
match self.length {
Some(length) => SizeHint::with_exact((length - self.sent) * CHUNK),
None => SizeHint::default(),
}
}
}
impl Drop for TestBody {
fn drop(&mut self) {
self.probe.dropped.store(true, Ordering::SeqCst);
}
}
struct Built {
claim: Arc<Claim>,
probe: Probe,
stats: Arc<InnerAgentStats>,
trailers: Arc<TrailersSlot>,
}
fn build(version: Version, length: Option<u64>, stall_after: Option<u64>) -> Built {
build_with(version, length, stall_after, DrainPolicy::default())
}
fn build_with(
version: Version,
length: Option<u64>,
stall_after: Option<u64>,
drain: DrainPolicy,
) -> Built {
let probe = Probe::default();
let stats = Arc::new(InnerAgentStats::default());
let trailers = Arc::new(TrailersSlot::default());
let claim = BodyShared::first_claim(BodyParts {
body: reqwest::Body::wrap(TestBody {
probe: probe.clone(),
length,
stall_after,
sent: 0,
}),
version,
drain,
#[cfg(feature = "encoding")]
decode: None,
trailers: trailers.clone(),
timing: Arc::new(TimingSlot::new(Instant::now(), RequestTiming::default())),
stats: stats.clone(),
});
Built {
claim,
probe,
stats,
trailers,
}
}
#[tokio::test]
async fn the_last_claim_going_drops_a_multiplexed_body() {
let built = build(Version::HTTP_2, None, None);
assert!(built.claim.give_up(), "the only claim is the last");
assert!(built.probe.dropped(), "the body is dropped");
assert_eq!(built.probe.read(), 0, "without reading any of it");
built.claim.body().settled().await;
}
#[tokio::test]
async fn a_clone_keeps_the_transfer_going() {
let built = build(Version::HTTP_2, None, None);
let clone = built.claim.duplicate().expect("an unread claim duplicates");
assert!(!built.claim.give_up(), "the original is not the last claim");
assert!(!built.probe.dropped(), "the body stays for the clone");
let mut reader = clone.reader().expect("the clone reads");
assert!(
reader.next().await.is_some_and(|chunk| chunk.is_ok()),
"the clone reads on"
);
drop(reader);
assert!(
built.probe.dropped(),
"dropping the clone's reader stops the transfer"
);
}
#[tokio::test]
async fn a_given_up_claim_has_nothing_to_give() {
let built = build(Version::HTTP_2, None, None);
let _clone = built.claim.duplicate().expect("an unread claim duplicates");
built.claim.give_up();
assert!(
built.claim.duplicate().is_none(),
"no copy of a given-up claim"
);
assert!(
matches!(
built.claim.reader().map(|_| ()).map_err(|err| err.kind()),
Err(FaithErrorKind::ResponseAlreadyDisturbed)
),
"and no reader either"
);
}
#[tokio::test]
async fn a_small_http1_remainder_is_drained() {
let built = build(Version::HTTP_11, Some(10), None);
built.claim.give_up();
built.claim.body().settled().await;
assert_eq!(
built.probe.read(),
10 * CHUNK,
"the whole remainder is read"
);
}
#[tokio::test]
async fn an_http1_remainder_over_the_limit_closes_at_once() {
let built = build(Version::HTTP_11, Some(1024), None);
built.claim.give_up();
built.claim.body().settled().await;
assert_eq!(built.probe.read(), 0, "none of it is read");
assert!(
built.probe.dropped(),
"the body is dropped, closing the connection"
);
}
#[tokio::test]
async fn an_endless_http1_body_is_read_to_the_limit_then_dropped() {
let built = build(Version::HTTP_11, None, None);
built.claim.give_up();
built.claim.body().settled().await;
let limit = DrainPolicy::default().limit;
assert!(built.probe.read() > limit, "the drain reads past the limit");
assert!(
built.probe.read() <= limit + CHUNK,
"by no more than a chunk"
);
assert!(built.probe.dropped(), "then drops the body");
}
#[tokio::test]
async fn a_zero_drain_limit_always_closes() {
let built = build_with(
Version::HTTP_11,
Some(1),
None,
DrainPolicy {
limit: 0,
..Default::default()
},
);
built.claim.give_up();
built.claim.body().settled().await;
assert_eq!(built.probe.read(), 0, "nothing is read");
assert!(built.probe.dropped(), "the body is dropped");
}
#[tokio::test]
async fn a_stalled_drain_is_bounded_by_its_timeout() {
let built = build_with(
Version::HTTP_11,
Some(20),
Some(5),
DrainPolicy {
timeout: Duration::from_millis(50),
..Default::default()
},
);
built.claim.give_up();
tokio::time::timeout(Duration::from_secs(5), built.claim.body().settled())
.await
.expect("the drain settles");
assert!(built.probe.dropped(), "the stalled body is dropped");
}
#[tokio::test]
async fn an_abort_errors_readers_ahead_of_buffered_chunks() {
let built = build(Version::HTTP_2, None, None);
let clone = built.claim.duplicate().expect("an unread claim duplicates");
let mut ahead = built.claim.reader().expect("a reader");
for _ in 0..3 {
ahead.next().await.expect("a chunk").expect("that reads");
}
built.claim.body().abort();
assert!(built.probe.dropped(), "the abort drops the body");
let mut behind = clone.reader().expect("a reader");
let first = behind.next().await.expect("an item");
assert!(
matches!(
first.map_err(|err| err.kind()),
Err(FaithErrorKind::Aborted)
),
"the clone's first read is the abort, not a buffered chunk"
);
assert!(behind.next().await.is_none(), "and nothing after it");
}
#[tokio::test]
async fn a_body_read_to_its_end_finishes_once() {
let built = build(Version::HTTP_2, Some(3), None);
let mut reader = built.claim.reader().expect("a reader");
let mut second = built.claim.reader().expect("a second reader");
let mut bytes = 0;
while let Some(chunk) = reader.next().await {
bytes += chunk.expect("the chunk reads").len() as u64;
}
assert_eq!(bytes, 3 * CHUNK);
drop(reader);
assert!(
second.next().await.is_none(),
"the second reader sees the end"
);
assert_eq!(built.stats.bodies_started.load(Ordering::SeqCst), 1);
assert_eq!(built.stats.bodies_finished.load(Ordering::SeqCst), 1);
assert!(matches!(
built.trailers.settled().await,
crate::response::Trailers::None
));
}
#[tokio::test]
async fn a_body_given_up_early_settles_its_bookkeeping() {
let built = build(Version::HTTP_2, None, None);
let mut reader = built.claim.reader().expect("a reader");
reader.next().await.expect("a chunk").expect("that reads");
drop(reader);
assert!(matches!(
built.trailers.settled().await,
crate::response::Trailers::None
));
assert_eq!(built.stats.bodies_finished.load(Ordering::SeqCst), 1);
}
}