use std::{
error::Error as StdError,
fmt,
future::Future,
io,
pin::Pin,
sync::Arc,
task::{Context, Poll},
time::Duration,
};
use ::hyper::body::Incoming;
use bytes::Bytes;
use http::{Request, Response};
use http_body::{Frame, SizeHint};
use hyper_rustls::{HttpsConnector, HttpsConnectorBuilder};
use hyper_util::{
client::legacy::{self, connect::HttpConnector},
rt::{TokioExecutor, TokioTimer},
};
use rustls::{ClientConfig, pki_types::CertificateDer};
use rustls_platform_verifier::{BuilderVerifierExt as _, Verifier};
use tower_service::Service;
use super::{Body, BoxError};
use crate::{error::Error, text};
const POOL_IDLE_TIMEOUT: Duration = Duration::from_secs(90);
const KEEP_ALIVE_INTERVAL: Duration = Duration::from_secs(30);
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[non_exhaustive]
pub enum HttpVersion {
Http2Only,
Auto,
}
pub(crate) struct TransportSettings {
pub(crate) version: HttpVersion,
pub(crate) extra_roots: Vec<Vec<u8>>,
pub(crate) connect_timeout: Option<Duration>,
}
#[derive(Clone)]
pub struct HyperTransport {
client: legacy::Client<HttpsConnector<HttpConnector>, Body>,
version: HttpVersion,
extra_roots: usize,
connect_timeout: Option<Duration>,
}
impl HyperTransport {
pub(crate) fn new(settings: TransportSettings) -> Result<Self, Error> {
let TransportSettings { version, extra_roots, connect_timeout } = settings;
let root_count = extra_roots.len();
let tls = tls_config(extra_roots.into_iter().map(CertificateDer::from).collect())?;
let mut http = HttpConnector::new();
http.enforce_http(false);
http.set_nodelay(true);
http.set_connect_timeout(connect_timeout);
let https = HttpsConnectorBuilder::new().with_tls_config(tls).https_or_http();
let connector = match version {
HttpVersion::Http2Only => https.enable_http2().wrap_connector(http),
HttpVersion::Auto => https.enable_http1().enable_http2().wrap_connector(http),
};
let mut builder = legacy::Client::builder(TokioExecutor::new());
builder
.timer(TokioTimer::new())
.pool_timer(TokioTimer::new())
.pool_idle_timeout(POOL_IDLE_TIMEOUT)
.http2_keep_alive_interval(KEEP_ALIVE_INTERVAL)
.http2_keep_alive_while_idle(true)
.http2_only(version == HttpVersion::Http2Only);
Ok(Self {
client: builder.build(connector),
version,
extra_roots: root_count,
connect_timeout,
})
}
}
impl fmt::Debug for HyperTransport {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("HyperTransport")
.field("http_version", &self.version)
.field("extra_roots", &self.extra_roots)
.field("connect_timeout", &self.connect_timeout)
.finish()
}
}
impl Service<Request<Body>> for HyperTransport {
type Response = Response<ResponseBody>;
type Error = BoxError;
type Future = HyperResponseFuture;
fn poll_ready(&mut self, _: &mut Context<'_>) -> Poll<Result<(), BoxError>> {
Poll::Ready(Ok(()))
}
fn call(&mut self, request: Request<Body>) -> HyperResponseFuture {
HyperResponseFuture {
inner: self.client.request(request),
connect_timeout: self.connect_timeout,
}
}
}
#[must_use = "futures do nothing unless polled"]
pub struct HyperResponseFuture {
inner: legacy::ResponseFuture,
connect_timeout: Option<Duration>,
}
impl fmt::Debug for HyperResponseFuture {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.debug_struct("HyperResponseFuture").finish_non_exhaustive()
}
}
impl Future for HyperResponseFuture {
type Output = Result<Response<ResponseBody>, BoxError>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let this = self.get_mut();
match Pin::new(&mut this.inner).poll(cx) {
Poll::Pending => Poll::Pending,
Poll::Ready(Ok(response)) => Poll::Ready(Ok(response.map(ResponseBody))),
Poll::Ready(Err(error)) => Poll::Ready(Err(failure(error, this.connect_timeout))),
}
}
}
pub struct ResponseBody(Incoming);
impl fmt::Debug for ResponseBody {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.debug_struct("ResponseBody").finish_non_exhaustive()
}
}
impl http_body::Body for ResponseBody {
type Data = Bytes;
type Error = BoxError;
fn poll_frame(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Option<Result<Frame<Bytes>, BoxError>>> {
Pin::new(&mut self.get_mut().0).poll_frame(cx).map_err(Into::into)
}
fn is_end_stream(&self) -> bool {
self.0.is_end_stream()
}
fn size_hint(&self) -> SizeHint {
self.0.size_hint()
}
}
fn failure(error: legacy::Error, connect_timeout: Option<Duration>) -> BoxError {
match connect_timeout {
Some(timeout) if error.is_connect() && timed_out(&error) => {
Box::new(Error::timeout(timeout))
}
_ => Box::new(error),
}
}
fn timed_out(error: &(dyn StdError + 'static)) -> bool {
let mut link = Some(error);
while let Some(current) = link {
if current
.downcast_ref::<io::Error>()
.is_some_and(|io| io.kind() == io::ErrorKind::TimedOut)
{
return true;
}
link = current.source();
}
false
}
fn tls_config(extra_roots: Vec<CertificateDer<'static>>) -> Result<ClientConfig, Error> {
let provider = Arc::new(rustls::crypto::aws_lc_rs::default_provider());
let builder = ClientConfig::builder_with_provider(Arc::clone(&provider))
.with_safe_default_protocol_versions()
.map_err(verifier_error)?;
let config = if extra_roots.is_empty() {
builder.with_platform_verifier().map_err(verifier_error)?.with_no_client_auth()
} else {
let verifier =
Verifier::new_with_extra_roots(extra_roots, provider).map_err(verifier_error)?;
builder
.dangerous()
.with_custom_certificate_verifier(Arc::new(verifier))
.with_no_client_auth()
};
Ok(config)
}
fn verifier_error(error: rustls::Error) -> Error {
Error::config(format!(
"The TLS certificate verifier could not be built: {}.",
text::bounded(&error, text::MAX_MESSAGE_CHARS)
))
}
#[cfg(test)]
#[path = "hyper_tests.rs"]
mod tests;