use std::{
sync::{
Arc,
atomic::{AtomicBool, Ordering},
},
time::Instant,
};
use reqwest::{
Method, StatusCode,
header::{CONTENT_ENCODING, HeaderName, HeaderValue},
tls::TlsInfo,
};
use reqwest_middleware::ClientWithMiddleware;
#[cfg(feature = "connection-tracking")]
use hyper_util::client::legacy::connect::HttpInfo;
#[cfg(feature = "cache")]
use http_cache_reqwest::CacheMode;
#[cfg(feature = "encoding")]
use reqwest::header::ACCEPT_ENCODING;
use tokio::sync::Mutex;
#[cfg(feature = "encoding")]
use web_faith_encoding::{self as encoding, AcceptEncoding, Coding, DEFAULT_ACCEPT_ENCODING};
use crate::{
agent::Agent,
body::{Body, BodyHolder},
error::{FaithError, FaithErrorKind},
request::{Credentials, NORMALISED_METHODS, PRIORITY, RequestBody, RequestOptions},
response::{PeerInformation, Response},
timing::{HeadersStamp, RequestTiming, TimingSlot, alpn_protocol_id},
};
pub async fn send(
agent: &Agent,
client: ClientWithMiddleware,
url: &str,
options: RequestOptions,
body: RequestBody,
abort: Option<impl Future<Output = ()>>,
) -> Result<Response, FaithError> {
let method = options.method.as_deref().unwrap_or("GET");
let method = NORMALISED_METHODS
.into_iter()
.find(|normalised| normalised.eq_ignore_ascii_case(method))
.unwrap_or(method);
let method =
Method::from_bytes(method.as_bytes()).map_err(|_| FaithErrorKind::InvalidMethod)?;
let is_head = method == Method::HEAD;
let mut parsed_url = reqwest::Url::parse(&url).map_err(|_| FaithErrorKind::InvalidUrl)?;
#[cfg(feature = "encoding")]
let compress = options
.compress
.as_deref()
.map(|value| {
Coding::from_option(value).ok_or_else(|| {
FaithError::new(
FaithErrorKind::InvalidCompression,
Some(format!(
"compress: {value:?} names no coding; expected gzip, deflate, br, or zstd"
)),
)
})
})
.transpose()?;
if options.credentials == Credentials::Omit {
let _ = parsed_url.set_username("");
let _ = parsed_url.set_password(None);
}
let headers_stamp = HeadersStamp::default();
let mut request = client
.request(method, parsed_url.clone())
.with_extension(headers_stamp.clone());
#[cfg(feature = "cache")]
{
request = request.with_extension(CacheMode::from(options.cache));
}
if let Some(headers) = &options.headers {
for (key, value) in headers {
if options.credentials == Credentials::Omit && key.eq_ignore_ascii_case("cookie") {
continue;
}
let header_name = HeaderName::from_bytes(key.as_bytes()).map_err(|_| {
FaithError::new(
FaithErrorKind::InvalidHeader,
Some(format!("invalid header name: {key}")),
)
})?;
let header_value = HeaderValue::from_str(value).map_err(|_| {
FaithError::new(
FaithErrorKind::InvalidHeader,
Some(format!("invalid header value: {value}")),
)
})?;
#[cfg(feature = "encoding")]
if compress.is_some() && header_name == CONTENT_ENCODING {
continue;
}
request = request.header(header_name, header_value);
}
}
#[cfg(feature = "encoding")]
let declared_content_encoding = compress.and_then(|_| {
let from_request = options.headers.as_ref().and_then(|headers| {
let declared = headers
.iter()
.filter(|(name, _)| name.eq_ignore_ascii_case(CONTENT_ENCODING.as_str()))
.map(|(_, value)| value.as_str())
.collect::<Vec<_>>();
(!declared.is_empty()).then(|| declared.join(", "))
});
from_request.or_else(|| {
agent
.default_content_encoding
.as_ref()
.and_then(|value| value.to_str().ok().map(str::to_owned))
})
});
#[cfg(feature = "encoding")]
let request_accept_encoding = options.headers.as_ref().and_then(|headers| {
headers
.iter()
.find(|(name, _)| name.eq_ignore_ascii_case("accept-encoding"))
.map(|(_, value)| value.clone())
});
#[cfg(feature = "encoding")]
let accept_encoding = AcceptEncoding::parse(
&request_accept_encoding
.clone()
.or_else(|| {
agent
.default_accept_encoding
.as_ref()
.and_then(|value| value.to_str().ok().map(str::to_owned))
})
.unwrap_or_else(|| DEFAULT_ACCEPT_ENCODING.to_owned()),
);
#[cfg(feature = "encoding")]
if request_accept_encoding.is_none() && agent.default_accept_encoding.is_none() {
request = request.header(
ACCEPT_ENCODING,
HeaderValue::from_static(DEFAULT_ACCEPT_ENCODING),
);
}
if let Some(urgency) = options.priority
&& !agent.has_default_priority
&& !options.headers.as_ref().is_some_and(|headers| {
headers
.iter()
.any(|(name, _)| name.eq_ignore_ascii_case(PRIORITY))
}) {
request = request.header(
HeaderName::from_static(PRIORITY),
HeaderValue::from_static(urgency),
);
}
#[cfg(feature = "encoding")]
let mut applied_coding = None;
match body {
RequestBody::Stream(byte_stream) => {
if !agent.quirk_h1_request_streaming {
if parsed_url.scheme() != "https" {
return Err(FaithError::new(
FaithErrorKind::Network,
Some(format!(
"a streaming request body requires HTTP/2 or HTTP/3, and {} is served over HTTP/1.1; set the agent's quirks.h1RequestStreaming to send it anyway",
parsed_url.as_str()
)),
));
}
request = request.version(http::Version::HTTP_2);
}
#[cfg(feature = "encoding")]
let body = match compress {
Some(coding) => {
applied_coding = Some(coding);
reqwest::Body::wrap_stream(encoding::compress_stream(byte_stream, coding))
}
None => reqwest::Body::wrap_stream(byte_stream),
};
#[cfg(not(feature = "encoding"))]
let body = reqwest::Body::wrap_stream(byte_stream);
request = request.body(body);
}
RequestBody::Bytes(bytes) => {
#[cfg(feature = "encoding")]
let body = match compress {
Some(coding) => {
applied_coding = Some(coding);
encoding::compress_buffer(&bytes, coding)
.await
.map_err(|err| {
FaithError::new(
FaithErrorKind::Network,
Some(format!("could not compress the request body: {err}")),
)
})?
}
None => bytes.to_vec(),
};
#[cfg(not(feature = "encoding"))]
let body = bytes.to_vec();
request = request.body(body);
}
RequestBody::None => {}
}
#[cfg(feature = "encoding")]
if let Some(coding) = applied_coding {
let value = encoding::layer_content_encoding(declared_content_encoding.as_deref(), coding);
let value = HeaderValue::from_str(&value).map_err(|_| {
FaithError::new(
FaithErrorKind::InvalidHeader,
Some(format!("invalid header value: {value}")),
)
})?;
request = request.header(CONTENT_ENCODING, value);
}
if let Some(dur) = options.timeout {
request = request.timeout(dur);
}
agent.stats.requests_sent.fetch_add(1, Ordering::Relaxed);
let started = Instant::now();
let response = match abort {
Some(abort) => {
tokio::select! {
result = request.send() => result?,
_ = abort => {
return Err(FaithErrorKind::Aborted.into());
}
}
}
None => request.send().await?,
};
agent
.stats
.responses_received
.fetch_add(1, Ordering::Relaxed);
let status_code = response.status();
let empty = status_code == StatusCode::NO_CONTENT || is_head;
let response_url = response.url().clone();
let version = response.version();
let redirected = if agent.h3_follow_advertised_port && version == http::Version::HTTP_3 {
let without_port = |url: &reqwest::Url| {
let mut url = url.clone();
let _ = url.set_port(None);
url
};
without_port(&parsed_url) != without_port(&response_url)
} else {
parsed_url != response_url
};
#[cfg(feature = "connection-tracking")]
let reused = if let Some(http_info) = response.extensions().get::<HttpInfo>() {
let local_addr = http_info.local_addr();
let remote_addr = http_info.remote_addr();
agent.conn_tracker.track(local_addr, remote_addr)
} else {
false
};
#[cfg(not(feature = "connection-tracking"))]
let reused = false;
agent.mark_warm(&response_url);
let peer = PeerInformation {
address: response.remote_addr(),
certificate: response
.extensions()
.get::<TlsInfo>()
.and_then(|info| info.peer_certificate())
.map(|cert| cert.into()),
};
let mut headers = response.headers().clone();
if options.credentials == Credentials::Omit {
headers.remove("set-cookie");
}
let headers_at = headers_stamp.get().unwrap_or_else(Instant::now);
let timing = RequestTiming {
headers_ms: headers_at.duration_since(started).as_secs_f64() * 1000.0,
body_ms: None,
reused,
next_hop_protocol: alpn_protocol_id(version, &response_url),
content_encoding: headers
.get(CONTENT_ENCODING)
.and_then(|value| value.to_str().ok())
.map(str::to_owned),
from_cache: headers
.get("x-cache")
.and_then(|value| value.to_str().ok())
.is_some_and(|value| value.eq_ignore_ascii_case("HIT")),
};
#[cfg(feature = "encoding")]
let decode = if empty {
None
} else {
encoding::decision(&headers, &accept_encoding)
};
#[cfg(feature = "encoding")]
if decode.is_some() {
encoding::strip_decoded_headers(&mut headers);
}
let timing = Arc::new(TimingSlot::new(started, timing));
if empty {
timing.ended();
}
Ok(Response {
body: if empty {
BodyHolder::none()
} else {
let http_response: http::Response<_> = response.into();
BodyHolder::new(
Some(Arc::new(Mutex::new(Body::Inner(http_response.into_body())))),
version,
timing.clone(),
)
},
#[cfg(feature = "encoding")]
decode,
disturbed: Arc::new(AtomicBool::new(false)),
headers,
integrity: options.integrity,
peer: Arc::new(peer),
redirected,
stats: agent.stats.clone(),
status_code,
timing,
trailers: Default::default(),
url: response_url,
version,
})
}