use std::borrow::Cow;
use std::io;
use std::time::Duration;
use moirai_async::timer::timeout;
use moirai_tls::TlsConnector;
use crate::codec::{DEFAULT_MAX_RESPONSE_BYTES, read_response, write_request};
use crate::conn::Conn;
use crate::pool::IdlePool;
use crate::redirect::{
forwarded_headers, is_redirect, parse_url, redirects_to_get, resolve_redirect,
};
use crate::{Origin, Response};
const DEFAULT_MAX_IDLE_PER_ORIGIN: usize = 8;
const DEFAULT_IDLE_TIMEOUT: Duration = Duration::from_secs(300);
const DEFAULT_MAX_REDIRECTS: usize = 10;
const DEFAULT_REQUEST_TIMEOUT: Duration = Duration::from_secs(30);
pub struct HttpClient {
tls: TlsConnector,
pool: IdlePool<Origin, Conn>,
max_idle_per_host: usize,
idle_timeout: Duration,
max_redirects: usize,
request_timeout: Duration,
max_response_bytes: usize,
}
impl Default for HttpClient {
fn default() -> Self {
Self::new()
}
}
impl HttpClient {
#[must_use]
pub fn new() -> Self {
Self::with_tls(TlsConnector::with_webpki_roots())
}
#[must_use]
pub fn with_tls(tls: TlsConnector) -> Self {
Self {
tls,
pool: IdlePool::default(),
max_idle_per_host: DEFAULT_MAX_IDLE_PER_ORIGIN,
idle_timeout: DEFAULT_IDLE_TIMEOUT,
max_redirects: DEFAULT_MAX_REDIRECTS,
request_timeout: DEFAULT_REQUEST_TIMEOUT,
max_response_bytes: DEFAULT_MAX_RESPONSE_BYTES,
}
}
pub fn set_timeout(&mut self, duration: Duration) {
self.request_timeout = duration;
}
pub fn set_max_response_bytes(&mut self, bytes: usize) {
self.max_response_bytes = bytes;
}
pub fn set_max_idle_per_host(&mut self, connections: usize) {
self.max_idle_per_host = connections;
}
pub fn set_idle_timeout(&mut self, duration: Duration) {
self.idle_timeout = duration;
}
pub fn set_max_redirects(&mut self, redirects: usize) {
self.max_redirects = redirects;
}
pub async fn get(&self, url: &str, headers: &[(&str, &str)]) -> io::Result<Response> {
self.request("GET", url, headers, None).await
}
pub async fn head(&self, url: &str, headers: &[(&str, &str)]) -> io::Result<Response> {
self.request("HEAD", url, headers, None).await
}
pub async fn request(
&self,
method: &str,
url: &str,
headers: &[(&str, &str)],
body: Option<&[u8]>,
) -> io::Result<Response> {
match timeout(
self.request_timeout,
self.request_following_redirects(method, url, headers, body),
)
.await
{
Ok(result) => result,
Err(_) => Err(io::Error::new(
io::ErrorKind::TimedOut,
"logical HTTP request timed out",
)),
}
}
async fn request_following_redirects<'a>(
&self,
method: &'a str,
url: &'a str,
headers: &'a [(&'a str, &'a str)],
body: Option<&'a [u8]>,
) -> io::Result<Response> {
let mut current_method = Cow::Borrowed(method);
let mut current_url = Cow::Borrowed(url);
let mut current_headers: Option<Vec<(&str, &str)>> = None;
let mut current_body = body;
let mut redirects_followed = 0usize;
loop {
let request_headers = current_headers.as_deref().unwrap_or(headers);
let response = self
.request_once_url(
current_method.as_ref(),
current_url.as_ref(),
request_headers,
current_body,
)
.await?;
if !is_redirect(response.status) {
return Ok(response);
}
let Some(location) = response.header("location") else {
return Ok(response);
};
if redirects_followed >= self.max_redirects {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
format!("redirect limit of {} exceeded", self.max_redirects),
));
}
let (origin, _) = parse_url(current_url.as_ref())?;
let next_url = resolve_redirect(current_url.as_ref(), location)?;
let (next_origin, _) = parse_url(&next_url)?;
let body_dropped = redirects_to_get(response.status, current_method.as_ref());
let next_headers =
forwarded_headers(request_headers, next_origin != origin, body_dropped);
if body_dropped {
current_method = Cow::Borrowed("GET");
current_body = None;
}
current_headers = Some(next_headers);
current_url = Cow::Owned(next_url);
redirects_followed = redirects_followed.checked_add(1).ok_or_else(|| {
io::Error::new(io::ErrorKind::InvalidData, "redirect counter overflow")
})?;
}
}
async fn request_once_url(
&self,
method: &str,
url: &str,
headers: &[(&str, &str)],
body: Option<&[u8]>,
) -> io::Result<Response> {
let (origin, path) = parse_url(url)?;
if let Some(connection) = self.pool.take(&origin, self.idle_timeout) {
match self
.try_once(connection, method, &origin, &path, headers, body)
.await
{
Ok(response) => return Ok(response),
Err(pooled_error) if !is_idempotent(method) => return Err(pooled_error),
Err(pooled_error) => {
let connection = Conn::connect(&origin, &self.tls)
.await
.map_err(|error| with_retry_context(error, &pooled_error))?;
return self
.try_once(connection, method, &origin, &path, headers, body)
.await
.map_err(|error| with_retry_context(error, &pooled_error));
}
}
}
let connection = Conn::connect(&origin, &self.tls).await?;
self.try_once(connection, method, &origin, &path, headers, body)
.await
}
async fn try_once(
&self,
mut connection: Conn,
method: &str,
origin: &Origin,
path: &str,
headers: &[(&str, &str)],
body: Option<&[u8]>,
) -> io::Result<Response> {
let host = origin.host_header();
write_request(&mut connection, method, &host, path, headers, body).await?;
let response = read_response(
&mut connection,
method.eq_ignore_ascii_case("HEAD"),
self.max_response_bytes,
)
.await?;
if response.keep_alive {
self.pool.put(origin, connection, self.max_idle_per_host);
}
Ok(response)
}
}
fn is_idempotent(method: &str) -> bool {
["GET", "HEAD", "OPTIONS", "TRACE", "PUT", "DELETE"]
.iter()
.any(|candidate| method.eq_ignore_ascii_case(candidate))
}
fn with_retry_context(error: io::Error, pooled_error: &io::Error) -> io::Error {
io::Error::new(
error.kind(),
format!("{error} (retry after stale pooled connection failed with: {pooled_error})"),
)
}
#[cfg(test)]
mod tests;