use std::{
fmt::Debug,
hint::unreachable_unchecked,
mem::replace,
net::SocketAddr,
pin::Pin,
sync::{
Arc,
atomic::{AtomicBool, Ordering},
},
task::{Context, Poll},
time::{Duration, Instant},
};
use bytes::Bytes;
use futures::{Stream, StreamExt, TryStreamExt, stream};
use http::header::{CONTENT_LENGTH, HeaderMap};
use http_body_util::BodyStream;
use reqwest::{StatusCode, Url, Version};
use serde::de::DeserializeOwned;
use stream_shared::SharedStream;
use tokio::{io::AsyncWriteExt, sync::watch};
#[cfg(feature = "encoding")]
use web_faith_encoding::{Coding, decode_stream};
use crate::{
body::{Body, BodyHolder, DynStream, drain_body_inner},
error::{FaithError, FaithErrorKind},
stats::InnerAgentStats,
timing::TimingSlot,
};
use crate::integrity::{finish_integrity, integrity_checker, verify_integrity};
#[derive(Debug)]
pub struct PeerInformation {
pub address: Option<SocketAddr>,
pub certificate: Option<Vec<u8>>,
}
#[derive(Debug, Clone, Default)]
pub struct FileDestination {
pub overwrite: bool,
pub mode: Option<u32>,
}
pub const PROGRESS_INTERVAL: Duration = Duration::from_millis(50);
pub async fn open_destination(
path: &str,
options: &FileDestination,
) -> Result<tokio::fs::File, FaithError> {
let mut open = tokio::fs::OpenOptions::new();
open.write(true);
if options.overwrite {
open.create(true).truncate(true);
} else {
open.create_new(true);
}
#[cfg(unix)]
if let Some(mode) = options.mode {
open.mode(mode);
}
match open.open(path).await {
Ok(file) => Ok(file),
Err(err) => Err(classify_open_error(path, err).await),
}
}
pub async fn classify_open_error(path: &str, err: std::io::Error) -> FaithError {
let kind = if err.kind() == std::io::ErrorKind::AlreadyExists {
match tokio::fs::symlink_metadata(path).await {
Ok(meta) if meta.is_dir() => FaithErrorKind::FileWrite,
_ => FaithErrorKind::FileExists,
}
} else {
FaithErrorKind::FileWrite
};
FaithError::new(kind, Some(err.to_string()))
}
#[derive(Clone, Debug, Default)]
pub enum Trailers {
#[default]
NotYet,
None,
Some(HeaderMap),
}
#[derive(Debug)]
pub struct TrailersSlot(watch::Sender<Trailers>);
impl Default for TrailersSlot {
fn default() -> Self {
Self(watch::channel(Trailers::NotYet).0)
}
}
impl TrailersSlot {
pub fn arrived(&self, trailers: HeaderMap) {
self.0.send_replace(Trailers::Some(trailers));
}
pub fn ended(&self) {
self.0.send_if_modified(|state| {
if matches!(state, Trailers::NotYet) {
*state = Trailers::None;
true
} else {
false
}
});
}
pub async fn settled(&self) -> Trailers {
let mut rx = self.0.subscribe();
match rx
.wait_for(|state| !matches!(state, Trailers::NotYet))
.await
{
Ok(state) => state.clone(),
Err(_) => Trailers::None,
}
}
}
#[cfg(test)]
mod tests {
use std::{
future::Future,
pin::pin,
sync::atomic::{AtomicUsize, Ordering},
task::{Context, Poll, Wake, Waker},
};
use super::*;
struct CountingWaker(AtomicUsize);
impl CountingWaker {
fn wakes(&self) -> usize {
self.0.load(Ordering::SeqCst)
}
}
impl Wake for CountingWaker {
fn wake(self: Arc<Self>) {
self.wake_by_ref();
}
fn wake_by_ref(self: &Arc<Self>) {
self.0.fetch_add(1, Ordering::SeqCst);
}
}
#[test]
fn a_bodyless_response_converts_to_an_http_response() {
let mut headers = HeaderMap::new();
headers.insert("x-test", "yes".parse().expect("a valid header value"));
let response = Response {
body: BodyHolder::none(),
#[cfg(feature = "encoding")]
decode: None,
disturbed: Arc::new(AtomicBool::new(false)),
headers,
integrity: None,
peer: Arc::new(PeerInformation {
address: None,
certificate: None,
}),
redirected: false,
stats: Arc::new(InnerAgentStats::default()),
status_code: StatusCode::NO_CONTENT,
timing: Arc::new(TimingSlot::new(
Instant::now(),
crate::timing::RequestTiming::default(),
)),
trailers: Arc::new(TrailersSlot::default()),
url: Url::parse("https://example.com/").expect("a valid url"),
version: Version::HTTP_2,
};
let http = response.into_http().expect("an undisturbed body converts");
assert_eq!(http.status(), StatusCode::NO_CONTENT);
assert_eq!(http.version(), Version::HTTP_2);
assert_eq!(
http.headers().get("x-test").map(|v| v.as_bytes()),
Some(&b"yes"[..])
);
let collected =
futures::executor::block_on(http_body_util::BodyExt::collect(http.into_body()))
.expect("an empty body collects");
assert!(collected.to_bytes().is_empty());
}
#[test]
fn waiting_for_trailers_parks_rather_than_spinning() {
let slot = TrailersSlot::default();
let counter = Arc::new(CountingWaker(AtomicUsize::new(0)));
let waker = Waker::from(counter.clone());
let mut cx = Context::from_waker(&waker);
let mut settled = pin!(slot.settled());
assert!(matches!(settled.as_mut().poll(&mut cx), Poll::Pending));
assert_eq!(counter.wakes(), 0, "a parked wait asks for no wake-up");
assert!(matches!(settled.as_mut().poll(&mut cx), Poll::Pending));
assert_eq!(counter.wakes(), 0, "polling again does not arm a wake-up");
slot.ended();
assert!(counter.wakes() >= 1, "the body ending wakes the waiter");
assert!(matches!(
settled.as_mut().poll(&mut cx),
Poll::Ready(Trailers::None)
));
}
#[test]
fn trailers_already_there_resolve_on_the_first_poll() {
let slot = TrailersSlot::default();
let mut headers = HeaderMap::new();
headers.insert("x-checksum", "abc123".parse().unwrap());
slot.arrived(headers);
let counter = Arc::new(CountingWaker(AtomicUsize::new(0)));
let waker = Waker::from(counter.clone());
let mut cx = Context::from_waker(&waker);
let mut settled = pin!(slot.settled());
assert!(matches!(
settled.as_mut().poll(&mut cx),
Poll::Ready(Trailers::Some(_))
));
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct FileProgress {
pub bytes_written: u64,
pub content_length: Option<u64>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct FileWritten {
pub path: String,
pub bytes_written: u64,
}
#[derive(Debug, Clone)]
pub struct Response {
pub body: BodyHolder,
#[cfg(feature = "encoding")]
pub decode: Option<Coding>,
pub disturbed: Arc<AtomicBool>,
pub headers: HeaderMap,
pub integrity: Option<String>,
pub peer: Arc<PeerInformation>,
pub redirected: bool,
pub stats: Arc<InnerAgentStats>,
pub status_code: StatusCode,
pub timing: Arc<TimingSlot>,
pub trailers: Arc<TrailersSlot>,
pub url: Url,
pub version: Version,
}
impl Response {
pub fn status(&self) -> StatusCode {
self.status_code
}
pub fn status_text(&self) -> &'static str {
self.status_code.canonical_reason().unwrap_or_default()
}
pub fn ok(&self) -> bool {
self.status_code.is_success()
}
pub fn headers(&self) -> &HeaderMap {
&self.headers
}
pub fn url(&self) -> &Url {
&self.url
}
pub fn redirected(&self) -> bool {
self.redirected
}
pub fn version(&self) -> Version {
self.version
}
pub fn peer(&self) -> &PeerInformation {
&self.peer
}
pub fn body_used(&self) -> bool {
self.disturbed.load(Ordering::SeqCst)
}
pub async fn bytes(&self) -> Result<Vec<u8>, FaithError> {
self.check_stream_disturbed()?;
self.gather_contiguous().await
}
pub async fn text(&self) -> Result<String, FaithError> {
let bytes = self.bytes().await?;
Ok(String::from_utf8(bytes)
.unwrap_or_else(|err| String::from_utf8_lossy(err.as_bytes()).into_owned()))
}
pub async fn json<T: DeserializeOwned>(&self) -> Result<T, FaithError> {
let bytes = self.bytes().await?;
serde_json::from_slice(&bytes)
.map_err(|err| FaithError::new(FaithErrorKind::JsonParse, Some(err.to_string())))
}
pub fn body_stream(
&self,
) -> Result<Option<impl Stream<Item = Result<Bytes, FaithError>> + use<>>, FaithError> {
let _ = self.check_stream_disturbed();
let Some(lock) = &self.body.body else {
return Ok(None);
};
let mut body = lock
.try_lock()
.map_err(|_| FaithError::from(FaithErrorKind::ResponseAlreadyDisturbed))?;
let stream = self.ensure_stream(&mut body, self.body.drained.clone())?;
Ok(Some(stream.map_err(|err| {
FaithError::new(FaithErrorKind::BodyStream, Some(err))
})))
}
pub async fn discard(&self) {
if let Some(arc) = self.body.body.clone() {
if self.body.is_multiplexed() {
*arc.lock().await = Body::Consumed;
} else {
drain_body_inner(arc).await;
}
}
self.body.drained.store(true, Ordering::SeqCst);
self.trailers.ended();
self.timing.ended();
}
pub async fn timing(&self) -> crate::timing::RequestTiming {
self.timing.settled().await
}
pub async fn trailers(&self) -> Trailers {
self.trailers.settled().await
}
pub fn check_stream_disturbed(&self) -> Result<(), FaithError> {
if self.disturbed.swap(true, Ordering::SeqCst) {
Err(FaithErrorKind::ResponseAlreadyDisturbed.into())
} else {
Ok(())
}
}
pub fn ensure_stream(
&self,
body: &mut Body,
drained_flag: Arc<AtomicBool>,
) -> Result<SharedStream<Pin<Box<DynStream>>>, FaithError> {
match body {
Body::Consumed => Err(FaithErrorKind::ResponseAlreadyDisturbed.into()),
Body::Stream(stream) => Ok(stream.clone()),
lock @ Body::Inner(_) => {
let Body::Inner(inner) = replace(lock, Body::Consumed) else {
unsafe { unreachable_unchecked() }
};
self.stats.bodies_started.fetch_add(1, Ordering::Relaxed);
let trailers_stream = self.trailers.clone();
let trailers_finish = self.trailers.clone();
let stats_finish = self.stats.clone();
let timing_finish = self.timing.clone();
let drained_finish = drained_flag.clone();
let bytes = Box::pin(
BodyStream::new(inner)
.then(move |frame| {
let trailers_lock = trailers_stream.clone();
async move {
match frame {
Err(err) => Some(Err(err.to_string())),
Ok(frame) => match frame.into_trailers() {
Ok(trailers) => {
trailers_lock.arrived(trailers);
None
}
Err(frame) => Some(
frame
.into_data()
.map_err(|_| "unknown frame kind".to_string()),
),
},
}
}
})
.filter_map(async |item| item),
) as Pin<Box<DynStream>>;
#[cfg(feature = "encoding")]
let bytes = match self.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 bytes = Box::pin(
bytes.chain(
stream::once(async move {
trailers_finish.ended();
timing_finish.ended();
stats_finish.bodies_finished.fetch_add(1, Ordering::Relaxed);
drained_finish.store(true, Ordering::SeqCst);
})
.filter_map(async |()| None),
),
) as Pin<Box<DynStream>>;
let stream = SharedStream::new(bytes);
let _ = replace(lock, Body::Stream(stream.clone()));
Ok(stream)
}
}
}
pub async fn gather(&self) -> Result<Arc<[Bytes]>, FaithError> {
let Some(lock) = &self.body.body else {
return Ok(Default::default());
};
let mut body = lock.lock().await;
let stream = self.ensure_stream(&mut body, self.body.drained.clone())?;
drop(body);
let mut chunks = Vec::new();
futures::pin_mut!(stream);
while let Some(result) = stream.next().await {
let chunk =
result.map_err(|err| FaithError::new(FaithErrorKind::BodyStream, Some(err)))?;
chunks.push(chunk);
}
self.body.mark_drained();
Ok(Arc::from(chunks.into_boxed_slice()))
}
pub async fn gather_contiguous(&self) -> Result<Vec<u8>, FaithError> {
let body = self.gather().await?;
let length = body.iter().map(|chunk| chunk.len()).sum();
let mut bytes = Vec::with_capacity(length);
for chunk in body.into_iter() {
bytes.extend_from_slice(chunk);
}
if let Some(ref integrity) = self.integrity {
verify_integrity(&bytes, integrity)?;
}
Ok(bytes)
}
pub async fn write_to_file(
&self,
path: &str,
options: &FileDestination,
mut on_progress: impl FnMut(FileProgress),
) -> Result<FileWritten, FaithError> {
let Some(lock) = self.body.body.clone() else {
return Err(FaithErrorKind::ResponseBodyNull.into());
};
if self.disturbed.load(Ordering::SeqCst) {
return Err(FaithErrorKind::ResponseAlreadyDisturbed.into());
}
let mut checker = integrity_checker(self.integrity.as_deref())?;
let content_length = self
.headers
.get(CONTENT_LENGTH)
.and_then(|value| value.to_str().ok())
.and_then(|value| value.trim().parse::<u64>().ok());
let mut file = open_destination(path, options).await?;
self.check_stream_disturbed()?;
let stream = {
let mut body = lock.lock().await;
let stream = self.ensure_stream(&mut body, self.body.drained.clone())?;
drop(body); stream
};
let mut report = |written: u64| {
on_progress(FileProgress {
bytes_written: written,
content_length,
});
};
let mut written: u64 = 0;
let mut reported_at = Instant::now();
futures::pin_mut!(stream);
while let Some(result) = stream.next().await {
let chunk =
result.map_err(|err| FaithError::new(FaithErrorKind::BodyStream, Some(err)))?;
if let Some(checker) = checker.as_mut() {
checker.input(&chunk);
}
file.write_all(&chunk)
.await
.map_err(|err| FaithError::new(FaithErrorKind::FileWrite, Some(err.to_string())))?;
written += chunk.len() as u64;
if let Some(limit) = content_length {
if written > limit {
return Err(FaithErrorKind::ContentLengthOverrun.into());
}
}
if reported_at.elapsed() >= PROGRESS_INTERVAL {
reported_at = Instant::now();
report(written);
}
}
file.flush()
.await
.map_err(|err| FaithError::new(FaithErrorKind::FileWrite, Some(err.to_string())))?;
report(written);
if let Some(checker) = checker {
finish_integrity(checker)?;
}
self.body.mark_drained();
Ok(FileWritten {
path: std::path::absolute(path)
.map(|abs| abs.to_string_lossy().into_owned())
.unwrap_or_else(|_| path.to_owned()),
bytes_written: written,
})
}
}
pub struct ResponseBody {
chunks: Pin<Box<dyn Stream<Item = Result<Bytes, FaithError>> + Send>>,
}
impl Debug for ResponseBody {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ResponseBody").finish_non_exhaustive()
}
}
impl http_body::Body for ResponseBody {
type Data = Bytes;
type Error = FaithError;
fn poll_frame(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Option<Result<http_body::Frame<Self::Data>, Self::Error>>> {
self.chunks
.as_mut()
.poll_next(cx)
.map(|chunk| chunk.map(|chunk| chunk.map(http_body::Frame::data)))
}
}
impl Response {
pub fn into_http(self) -> Result<http::Response<ResponseBody>, FaithError> {
let chunks: Pin<Box<dyn Stream<Item = Result<Bytes, FaithError>> + Send>> =
match self.body_stream()? {
Some(stream) => Box::pin(stream),
None => Box::pin(stream::empty()),
};
let mut response = http::Response::new(ResponseBody { chunks });
*response.status_mut() = self.status_code;
*response.version_mut() = self.version;
*response.headers_mut() = self.headers.clone();
Ok(response)
}
}
impl TryFrom<Response> for http::Response<ResponseBody> {
type Error = FaithError;
fn try_from(response: Response) -> Result<Self, Self::Error> {
response.into_http()
}
}