pub use crate::timing::RequestTiming;
use std::{
fmt::Debug,
net::SocketAddr,
path::{Path, PathBuf},
pin::Pin,
sync::{
Arc,
atomic::{AtomicBool, Ordering},
},
task::{Context, Poll},
time::{Duration, Instant},
};
use bytes::Bytes;
use futures::{Stream, StreamExt, stream};
use http::header::{CONTENT_LENGTH, HeaderMap};
use reqwest::{StatusCode, Url, Version};
use serde::de::DeserializeOwned;
use tokio::{io::AsyncWriteExt, sync::watch};
pub use crate::body::BodyReader;
use crate::{
body::Claim,
error::{FaithError, FaithErrorKind},
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(crate) const PROGRESS_INTERVAL: Duration = Duration::from_millis(50);
pub(crate) async fn open_destination(
path: &Path,
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(crate) async fn classify_open_error(path: &Path, 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, err.to_string())
}
#[derive(Clone, Debug, Default)]
pub enum Trailers {
#[default]
NotYet,
None,
Some(HeaderMap),
}
#[derive(Debug)]
pub(crate) 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 {
claim: None,
disturbed: Arc::new(AtomicBool::new(false)),
headers,
integrity: None,
peer: Arc::new(PeerInformation {
address: None,
certificate: None,
}),
redirected: false,
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)]
#[non_exhaustive]
pub struct FileProgress {
pub bytes_written: u64,
pub content_length: Option<u64>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub struct FileWritten {
pub path: PathBuf,
pub bytes_written: u64,
}
#[derive(Debug, Clone)]
pub struct Response {
pub(crate) claim: Option<Arc<Claim>>,
pub(crate) disturbed: Arc<AtomicBool>,
pub(crate) headers: HeaderMap,
pub(crate) integrity: Option<String>,
pub(crate) peer: Arc<PeerInformation>,
pub(crate) redirected: bool,
pub(crate) status_code: StatusCode,
pub(crate) timing: Arc<TimingSlot>,
pub(crate) trailers: Arc<TrailersSlot>,
pub(crate) url: Url,
pub(crate) 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 try_clone(&self) -> Result<Self, FaithError> {
if self.body_used() {
return Err(FaithErrorKind::ResponseAlreadyDisturbed.into());
}
let claim = match &self.claim {
None => None,
Some(claim) => Some(
claim
.duplicate()
.ok_or(FaithErrorKind::ResponseAlreadyDisturbed)?,
),
};
Ok(Self {
claim,
disturbed: Arc::new(AtomicBool::new(false)),
..Clone::clone(self)
})
}
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, err.to_string()))
}
pub fn body_stream(&self) -> Result<Option<BodyReader>, FaithError> {
let _ = self.check_stream_disturbed();
match &self.claim {
None => Ok(None),
Some(claim) => claim.reader().map(Some),
}
}
pub async fn discard(&self) {
let Some(claim) = &self.claim else {
return;
};
claim.give_up();
if claim.body().claims_left() == 0 {
claim.body().settled().await;
}
}
pub async fn timing(&self) -> crate::timing::RequestTiming {
self.timing.settled().await
}
pub async fn trailers(&self) -> Trailers {
self.trailers.settled().await
}
pub(crate) fn check_stream_disturbed(&self) -> Result<(), FaithError> {
if self.disturbed.swap(true, Ordering::SeqCst) {
Err(FaithErrorKind::ResponseAlreadyDisturbed.into())
} else {
Ok(())
}
}
pub(crate) async fn gather(&self) -> Result<Arc<[Bytes]>, FaithError> {
let Some(claim) = &self.claim else {
return Ok(Default::default());
};
let mut stream = claim.reader()?;
let mut chunks = Vec::new();
while let Some(chunk) = stream.next().await {
chunks.push(chunk?);
}
Ok(Arc::from(chunks.into_boxed_slice()))
}
pub(crate) 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: impl AsRef<Path>,
options: &FileDestination,
mut on_progress: impl FnMut(FileProgress),
) -> Result<FileWritten, FaithError> {
let path = path.as_ref();
let Some(claim) = self.claim.clone() else {
return Err(FaithErrorKind::ResponseBodyNull.into());
};
if self.disturbed.load(Ordering::SeqCst) || claim.is_given_up() {
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 = claim.reader()?;
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?;
if let Some(checker) = checker.as_mut() {
checker.input(&chunk);
}
file.write_all(&chunk)
.await
.map_err(|err| FaithError::new(FaithErrorKind::FileWrite, 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, err.to_string()))?;
report(written);
if let Some(checker) = checker {
finish_integrity(checker)?;
}
Ok(FileWritten {
path: std::path::absolute(path).unwrap_or_else(|_| path.to_path_buf()),
bytes_written: written,
})
}
}
#[cfg(feature = "unstable-internals")]
impl Response {
pub fn check_disturbed(&self) -> Result<(), FaithError> {
self.check_stream_disturbed()
}
pub fn abort_body(&self) {
if let Some(claim) = &self.claim {
claim.body().abort();
}
}
}
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()
}
}