#![forbid(unsafe_code)]
#![deny(missing_docs)]
#![deny(rustdoc::broken_intra_doc_links)]
#![deny(clippy::unwrap_used)]
#![deny(clippy::expect_used)]
#![deny(clippy::panic)]
#![deny(clippy::unreachable)]
#![warn(clippy::pedantic)]
#![cfg_attr(docsrs, feature(doc_cfg))]
pub mod mtls;
use bytes::Bytes;
use http::Request;
use huskarl_core::{
Error, RetryAdvice,
http::{HttpClient, HttpResponse, Idempotency},
platform::MaybeSendBoxFuture,
};
use snafu::ResultExt as _;
#[derive(Debug, snafu::Snafu, huskarl_macros::Classify)]
#[non_exhaustive]
pub(crate) enum ReqwestSetupError {
#[snafu(display("building HTTP client"))]
#[classify(no)]
BuildingClient {
source: reqwest::Error,
},
#[snafu(display("building HTTP request"))]
#[classify(no)]
BuildingRequest {
source: reqwest::Error,
},
#[snafu(display("response body exceeded {max} byte limit"))]
#[classify(no)]
OversizeBody {
max: usize,
},
}
pub const DEFAULT_MAX_RESPONSE_BYTES: usize = 1024 * 1024;
#[derive(Clone)]
pub struct ReqwestClient {
client: reqwest::Client,
uses_mtls: bool,
max_response_bytes: Option<usize>,
#[cfg(all(
not(target_arch = "wasm32"),
any(feature = "rustls-tls", feature = "native-tls")
))]
identity: Option<reqwest::Identity>,
}
impl From<reqwest::Client> for ReqwestClient {
fn from(client: reqwest::Client) -> Self {
Self {
client,
uses_mtls: false,
max_response_bytes: Some(DEFAULT_MAX_RESPONSE_BYTES),
#[cfg(all(
not(target_arch = "wasm32"),
any(feature = "rustls-tls", feature = "native-tls")
))]
identity: None,
}
}
}
#[bon::bon]
impl ReqwestClient {
#[builder]
pub async fn new(
#[builder(required, into, default = Some(concat!("huskarl/", env!("CARGO_PKG_VERSION")).to_string()))]
user_agent: Option<String>,
#[builder(
with = |provider: impl mtls::MtlsProvider + 'static| Box::new(provider) as Box<dyn mtls::MtlsProvider>,
default = Box::new(mtls::NoMtls),
)]
mtls: Box<dyn mtls::MtlsProvider>,
#[cfg(all(
not(target_arch = "wasm32"),
any(feature = "rustls-tls", feature = "native-tls")
))]
root_certificates: Option<Vec<reqwest::Certificate>>,
#[cfg(not(target_arch = "wasm32"))]
#[builder(default = false)]
follow_redirects: bool,
#[cfg(not(target_arch = "wasm32"))]
#[builder(required, default = Some(std::time::Duration::from_secs(30)))]
timeout: Option<std::time::Duration>,
#[builder(required, default = Some(DEFAULT_MAX_RESPONSE_BYTES))]
max_response_bytes: Option<usize>,
configure_builder: Option<
Box<dyn FnOnce(reqwest::ClientBuilder) -> reqwest::ClientBuilder>,
>,
) -> Result<Self, Error> {
let mut reqwest_builder = reqwest::Client::builder();
#[cfg(not(target_arch = "wasm32"))]
if !follow_redirects {
reqwest_builder = reqwest_builder.redirect(reqwest::redirect::Policy::none());
}
#[cfg(not(target_arch = "wasm32"))]
if let Some(timeout) = timeout {
reqwest_builder = reqwest_builder.timeout(timeout);
}
if let Some(user_agent) = user_agent {
reqwest_builder = reqwest_builder.user_agent(user_agent);
}
#[cfg(all(
not(target_arch = "wasm32"),
any(feature = "rustls-tls", feature = "native-tls")
))]
if let Some(root_certificates) = root_certificates {
reqwest_builder = reqwest_builder.tls_certs_only(root_certificates);
}
if let Some(configure_builder) = configure_builder {
reqwest_builder = configure_builder(reqwest_builder);
}
let uses_mtls = mtls.uses_mtls();
let mtls_output = mtls.apply(reqwest_builder).await?;
Ok(Self {
client: mtls_output.builder.build().context(BuildingClientSnafu)?,
uses_mtls,
max_response_bytes,
#[cfg(all(
not(target_arch = "wasm32"),
any(feature = "rustls-tls", feature = "native-tls")
))]
identity: mtls_output.identity,
})
}
}
impl ReqwestClient {
#[must_use]
pub fn with_uses_mtls(mut self, uses_mtls: bool) -> Self {
self.uses_mtls = uses_mtls;
self
}
#[cfg(all(
not(target_arch = "wasm32"),
any(feature = "rustls-tls", feature = "native-tls")
))]
#[must_use]
pub fn identity(&self) -> Option<&reqwest::Identity> {
self.identity.as_ref()
}
}
#[track_caller]
fn oversize_error(max: usize) -> Error {
Error::from(ReqwestSetupError::OversizeBody { max })
}
async fn read_body(
response: reqwest::Response,
max_response_bytes: Option<usize>,
idempotency: Idempotency,
) -> Result<Bytes, Error> {
let Some(max) = max_response_bytes else {
return response
.bytes()
.await
.map_err(|e| transport_error(e, idempotency));
};
if response
.content_length()
.is_some_and(|len| len > max as u64)
{
return Err(oversize_error(max));
}
#[cfg(not(target_arch = "wasm32"))]
{
let mut response = response;
let mut body = bytes::BytesMut::new();
while let Some(chunk) = response
.chunk()
.await
.map_err(|e| transport_error(e, idempotency))?
{
if body.len() + chunk.len() > max {
return Err(oversize_error(max));
}
body.extend_from_slice(&chunk);
}
Ok(body.freeze())
}
#[cfg(target_arch = "wasm32")]
{
let body = response
.bytes()
.await
.map_err(|e| transport_error(e, idempotency))?;
if body.len() > max {
return Err(oversize_error(max));
}
Ok(body)
}
}
#[track_caller]
fn transport_error(source: reqwest::Error, idempotency: Idempotency) -> Error {
#[cfg(not(target_arch = "wasm32"))]
let retryable = source.is_connect()
|| (matches!(idempotency, Idempotency::Idempotent)
&& (source.is_timeout() || source.is_body()));
#[cfg(target_arch = "wasm32")]
let retryable = {
let _ = idempotency;
false
};
Error::new(RetryAdvice::retry_if(retryable), source)
}
impl HttpClient for ReqwestClient {
fn uses_mtls(&self) -> bool {
self.uses_mtls
}
fn execute(
&self,
request: Request<Bytes>,
idempotency: Idempotency,
) -> MaybeSendBoxFuture<'_, Result<HttpResponse, Error>> {
Box::pin(async move {
let (parts, body) = request.into_parts();
let reqwest_request = self
.client
.request(parts.method, parts.uri.to_string())
.headers(parts.headers)
.body(body)
.build()
.context(BuildingRequestSnafu)?;
let response = self
.client
.execute(reqwest_request)
.await
.map_err(|e| transport_error(e, idempotency))?;
let status = response.status();
let headers = response.headers().clone();
let body = read_body(response, self.max_response_bytes, idempotency).await?;
Ok(HttpResponse {
status,
headers,
body,
})
})
}
}
#[cfg(test)]
#[cfg(not(target_arch = "wasm32"))]
mod tests {
use std::time::Duration;
use bytes::Bytes;
use http::Request;
use huskarl_core::{
RetryAdvice,
http::{HttpClient, Idempotency},
};
use super::{DEFAULT_MAX_RESPONSE_BYTES, ReqwestClient, transport_error};
#[test]
fn with_uses_mtls_overrides_the_from_default() {
use huskarl_core::http::HttpClient as _;
let client = ReqwestClient::from(reqwest::Client::new());
assert!(!client.uses_mtls(), "From defaults to no mTLS");
assert!(client.with_uses_mtls(true).uses_mtls());
}
fn serve_once(body: Vec<u8>, content_length: bool) -> String {
use std::io::{Read, Write};
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
std::thread::spawn(move || {
let (mut stream, _) = listener.accept().unwrap();
let mut buf = [0u8; 1024];
let _ = stream.read(&mut buf);
let header = if content_length {
format!(
"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\n\r\n",
body.len()
)
} else {
"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nConnection: close\r\n\r\n"
.to_string()
};
let _ = stream.write_all(header.as_bytes());
let _ = stream.write_all(&body);
});
format!("http://{addr}/")
}
async fn get(
url: &str,
max_response_bytes: Option<usize>,
) -> Result<Bytes, huskarl_core::Error> {
let client = ReqwestClient::builder()
.max_response_bytes(max_response_bytes)
.build()
.await
.unwrap();
let request = Request::builder().uri(url).body(Bytes::new()).unwrap();
client
.execute(request, Idempotency::Idempotent)
.await
.map(|response| response.body)
}
#[tokio::test]
async fn body_within_limit_is_returned() {
let url = serve_once(b"{\"ok\":true}".to_vec(), true);
let body = get(&url, Some(64)).await.unwrap();
assert_eq!(&body[..], b"{\"ok\":true}");
}
#[tokio::test]
async fn oversized_content_length_is_rejected_early() {
let url = serve_once(vec![b'x'; 5000], true);
let error = get(&url, Some(1000)).await.unwrap_err();
assert_eq!(error.retry_advice(), RetryAdvice::No);
}
#[tokio::test]
async fn oversized_streamed_body_without_content_length_is_rejected() {
let url = serve_once(vec![b'x'; 5000], false);
let error = get(&url, Some(1000)).await.unwrap_err();
assert_eq!(error.retry_advice(), RetryAdvice::No);
}
#[tokio::test]
async fn streamed_body_is_reassembled_without_corruption() {
let expected: Vec<u8> = (0usize..256 * 1024)
.map(|i| u8::try_from(i % 251).unwrap())
.collect();
let url = serve_once(expected.clone(), false);
let body = get(&url, Some(expected.len() + 1)).await.unwrap();
assert_eq!(body.len(), expected.len());
assert!(
body[..] == expected[..],
"streamed body must match byte-for-byte"
);
}
#[tokio::test]
async fn body_exactly_at_limit_is_returned() {
let expected = vec![b'x'; 1000];
let url = serve_once(expected.clone(), false);
let body = get(&url, Some(1000)).await.unwrap();
assert_eq!(&body[..], &expected[..]);
}
#[tokio::test]
async fn no_limit_reads_a_large_body() {
let url = serve_once(vec![b'x'; 5000], true);
let body = get(&url, None).await.unwrap();
assert_eq!(body.len(), 5000);
}
#[tokio::test]
async fn default_limit_is_one_mib() {
assert_eq!(DEFAULT_MAX_RESPONSE_BYTES, 1024 * 1024);
}
async fn connect_error() -> reqwest::Error {
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
let port = listener.local_addr().unwrap().port();
drop(listener);
reqwest::Client::new()
.get(format!("http://127.0.0.1:{port}/"))
.send()
.await
.unwrap_err()
}
async fn timeout_error() -> (reqwest::Error, std::net::TcpListener) {
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
let port = listener.local_addr().unwrap().port();
let error = reqwest::Client::builder()
.timeout(Duration::from_millis(100))
.build()
.unwrap()
.get(format!("http://127.0.0.1:{port}/"))
.send()
.await
.unwrap_err();
(error, listener)
}
#[tokio::test]
async fn connect_failure_is_retryable_regardless_of_idempotency() {
for idempotency in [Idempotency::Idempotent, Idempotency::Unknown] {
let error = transport_error(connect_error().await, idempotency);
assert_eq!(
error.retry_advice(),
RetryAdvice::RETRY,
"connect failure with {idempotency:?} should be retryable"
);
}
}
#[tokio::test]
async fn timeout_is_retryable_only_when_idempotent() {
let (error, _listener) = timeout_error().await;
let error = transport_error(error, Idempotency::Idempotent);
assert_eq!(error.retry_advice(), RetryAdvice::RETRY);
let (error, _listener) = timeout_error().await;
let error = transport_error(error, Idempotency::Unknown);
assert_eq!(error.retry_advice(), RetryAdvice::No);
}
}