use std::{
fmt::Debug,
mem::replace,
pin::Pin,
sync::{
Arc,
atomic::{AtomicBool, Ordering},
},
};
use bytes::Bytes;
use futures::{Stream, StreamExt};
use http_body_util::BodyExt;
use reqwest::Version;
use stream_shared::SharedStream;
use tokio::sync::Mutex;
use crate::timing::TimingSlot;
pub type DynStream = dyn Stream<Item = std::result::Result<Bytes, String>> + Send + Sync;
pub enum Body {
Inner(reqwest::Body),
Consumed,
Stream(SharedStream<Pin<Box<DynStream>>>),
}
pub struct BodyHolder {
pub body: Option<Arc<Mutex<Body>>>,
pub drained: Arc<AtomicBool>,
pub version: Version,
pub timing: Option<Arc<TimingSlot>>,
}
impl BodyHolder {
pub fn new(body: Option<Arc<Mutex<Body>>>, version: Version, timing: Arc<TimingSlot>) -> Self {
Self {
body,
version,
drained: Arc::new(AtomicBool::new(false)),
timing: Some(timing),
}
}
pub fn none() -> Self {
Self {
body: None,
version: Version::HTTP_11,
drained: Arc::new(AtomicBool::new(true)),
timing: None,
}
}
pub fn is_multiplexed(&self) -> bool {
matches!(self.version, Version::HTTP_2 | Version::HTTP_3)
}
pub fn mark_drained(&self) {
self.drained.store(true, Ordering::SeqCst);
}
}
impl Clone for BodyHolder {
fn clone(&self) -> Self {
Self {
body: self.body.clone(),
drained: self.drained.clone(),
version: self.version,
timing: self.timing.clone(),
}
}
}
impl Debug for BodyHolder {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("BodyHolder")
.field("body", &self.body)
.field("drained", &self.drained.load(Ordering::SeqCst))
.field("version", &self.version)
.field("timing", &self.timing)
.finish()
}
}
impl Drop for BodyHolder {
fn drop(&mut self) {
if self.drained.load(Ordering::SeqCst) {
return;
}
if self
.body
.as_ref()
.is_some_and(|arc| Arc::strong_count(arc) > 1)
{
return;
}
let timing = self.timing.take();
if self.is_multiplexed() {
if let Some(timing) = timing {
timing.ended();
}
return;
}
if let Some(arc) = self.body.take() {
if let Ok(handle) = tokio::runtime::Handle::try_current() {
handle.spawn(async move {
drain_body_inner(arc).await;
if let Some(timing) = timing {
timing.ended();
}
});
} else if let Some(timing) = timing {
timing.ended();
}
} else if let Some(timing) = timing {
timing.ended();
}
}
}
pub async fn drain_body_inner(arc: Arc<Mutex<Body>>) {
let mut guard = arc.lock().await;
match replace(&mut *guard, Body::Consumed) {
Body::Inner(body) => {
let mut body = body;
while body.frame().await.is_some() {}
}
Body::Stream(shared) => {
futures::pin_mut!(shared);
while shared.next().await.is_some() {}
}
Body::Consumed => {}
}
}
impl Debug for Body {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Inner(body) => write!(f, "{body:?}"),
Self::Consumed => write!(f, "Consumed"),
Self::Stream(stream) => {
let field = f
.debug_struct("SharedStream")
.field("stats", &stream.stats())
.finish_non_exhaustive();
f.debug_tuple("Stream").field(&field).finish()
}
}
}
}