httproxide-client-util 0.3.0

wrapper connector around hyper, hyper-rustls and hyperlocal
Documentation
use std::collections::HashMap;
use std::pin::Pin;
use std::sync::{LazyLock, Mutex};
use std::task::{Context, Poll};
use std::time::Duration;

use futures_util::future::BoxFuture;
use hyper::body::Body;
use hyper::rt::ReadBufCursor;
use hyper::Uri;
use hyper_util::client::legacy::connect::{Connection, HttpConnector};
use hyper_util::client::legacy::Client as HyperClient;
use hyper_util::rt::{TokioExecutor, TokioIo};
use serde::{Deserialize, Serialize};
use tower::{Service, ServiceExt};

pub type Client<B> = HyperClient<Connector, B>;

#[derive(Clone, Debug)]
pub struct Connector {
    #[cfg(feature = "https")]
    http: hyper_rustls::HttpsConnector<HttpConnector>,
    #[cfg(not(feature = "https"))]
    http: HttpConnector,
    #[cfg(feature = "unix")]
    unix: Option<hyper_unix_socket::UnixSocketConnector<String>>,
}

pub enum Stream {
    #[cfg(feature = "https")]
    Http(hyper_rustls::MaybeHttpsStream<TokioIo<tokio::net::TcpStream>>),
    #[cfg(not(feature = "https"))]
    Http(TokioIo<tokio::net::TcpStream>),
    #[cfg(feature = "unix")]
    Unix(hyper_unix_socket::UnixSocketConnection),
}

impl Service<Uri> for Connector {
    type Response = Stream;
    type Error = tower::BoxError;
    type Future = BoxFuture<'static, Result<Stream, Self::Error>>;

    fn poll_ready(&mut self, _cx: &mut Context) -> Poll<Result<(), Self::Error>> {
        Poll::Ready(Ok(()))
    }

    fn call(&mut self, dst: Uri) -> Self::Future {
        #[cfg(feature = "unix")]
        if dst.scheme_str() == Some("unix") {
            if let Some(unix_ref) = self.unix.as_mut() {
                let clone = unix_ref.clone();
                let unix = std::mem::replace(&mut *unix_ref, clone);
                return Box::pin(async move { Ok(Stream::Unix(unix.oneshot(dst).await?)) });
            }
        }
        let clone = self.http.clone();
        let http = std::mem::replace(&mut self.http, clone);
        Box::pin(async move { Ok(Stream::Http(http.oneshot(dst).await?)) })
    }
}

impl hyper::rt::Read for Stream {
    fn poll_read(
        self: Pin<&mut Self>,
        cx: &mut Context,
        buf: ReadBufCursor,
    ) -> Poll<std::io::Result<()>> {
        match self.get_mut() {
            Stream::Http(s) => Pin::new(s).poll_read(cx, buf),
            #[cfg(feature = "unix")]
            Stream::Unix(s) => Pin::new(s).poll_read(cx, buf),
        }
    }
}

impl hyper::rt::Write for Stream {
    fn poll_write(
        self: Pin<&mut Self>,
        cx: &mut Context,
        buf: &[u8],
    ) -> Poll<std::io::Result<usize>> {
        match self.get_mut() {
            Stream::Http(s) => Pin::new(s).poll_write(cx, buf),
            #[cfg(feature = "unix")]
            Stream::Unix(s) => Pin::new(s).poll_write(cx, buf),
        }
    }
    fn poll_flush(self: Pin<&mut Self>, cx: &mut Context) -> Poll<std::io::Result<()>> {
        match self.get_mut() {
            Stream::Http(s) => Pin::new(s).poll_flush(cx),
            #[cfg(feature = "unix")]
            Stream::Unix(s) => Pin::new(s).poll_flush(cx),
        }
    }
    fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context) -> Poll<std::io::Result<()>> {
        match self.get_mut() {
            Stream::Http(s) => Pin::new(s).poll_flush(cx),
            #[cfg(feature = "unix")]
            Stream::Unix(s) => Pin::new(s).poll_flush(cx),
        }
    }
}

impl Connection for Stream {
    fn connected(&self) -> hyper_util::client::legacy::connect::Connected {
        match self {
            Stream::Http(s) => s.connected(),
            #[cfg(feature = "unix")]
            Stream::Unix(s) => s.connected(),
        }
    }
}

#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq, Hash, Default)]
pub struct ClientConfig {
    #[serde(default)]
    #[cfg(feature = "https")]
    dangerous_skip_cert_check: bool,
    #[serde(default)]
    #[cfg(feature = "unix")]
    unix_socket_path: Option<String>,
}

static CACHE: LazyLock<Mutex<HashMap<ClientConfig, Box<dyn std::any::Any + Send>>>> =
    LazyLock::new(Default::default);

pub fn clear_cache() {
    *(*CACHE).lock().unwrap() = Default::default();
}

#[cfg(feature = "https")]
#[derive(Debug)]
struct NoCertVerifier {}

#[cfg(feature = "https")]
impl rustls::client::danger::ServerCertVerifier for NoCertVerifier {
    fn verify_server_cert(
        &self,
        _end_entity: &rustls::pki_types::CertificateDer,
        _intermediates: &[rustls::pki_types::CertificateDer],
        _server_name: &rustls::pki_types::ServerName,
        _ocsp_response: &[u8],
        _now: rustls::pki_types::UnixTime,
    ) -> Result<rustls::client::danger::ServerCertVerified, rustls::Error> {
        Ok(rustls::client::danger::ServerCertVerified::assertion())
    }

    fn verify_tls12_signature(
        &self,
        message: &[u8],
        cert: &rustls::pki_types::CertificateDer<'_>,
        dss: &rustls::DigitallySignedStruct,
    ) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
        rustls::crypto::verify_tls12_signature(
            message,
            cert,
            dss,
            &rustls::crypto::aws_lc_rs::default_provider().signature_verification_algorithms,
        )
    }
    fn verify_tls13_signature(
        &self,
        message: &[u8],
        cert: &rustls::pki_types::CertificateDer<'_>,
        dss: &rustls::DigitallySignedStruct,
    ) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
        rustls::crypto::verify_tls13_signature(
            message,
            cert,
            dss,
            &rustls::crypto::aws_lc_rs::default_provider().signature_verification_algorithms,
        )
    }

    fn supported_verify_schemes(&self) -> Vec<rustls::SignatureScheme> {
        rustls::crypto::aws_lc_rs::default_provider()
            .signature_verification_algorithms
            .supported_schemes()
    }
}

fn new_client<B>(cfg: ClientConfig) -> anyhow::Result<Client<B>>
where
    B: Body + Send + 'static,
    B::Data: Send,
{
    #[cfg(not(feature = "https"))]
    let http = HttpConnector::new();

    #[cfg(feature = "https")]
    let http = {
        use std::sync::Arc;
        use hyper_rustls::ConfigBuilderExt;

        let mut http = HttpConnector::new();
        http.enforce_http(false);

        let tls_config = { rustls::ClientConfig::builder() };

        let tls_config = if cfg.dangerous_skip_cert_check {
            tls_config
                .dangerous()
                .with_custom_certificate_verifier(Arc::new(NoCertVerifier {}))
                .with_no_client_auth()
        } else {
            tls_config.with_native_roots()?.with_no_client_auth()
        };

        let tls = hyper_rustls::HttpsConnectorBuilder::new()
            .with_tls_config(tls_config)
            .https_or_http()
            .enable_http1();

        #[cfg(feature = "http2")]
        {
            tls.enable_http2().wrap_connector(http)
        }

        #[cfg(not(feature = "http2"))]
        {
            tls.wrap_connector(http)
        }
    };

    let connector = Connector {
        http,
        #[cfg(feature = "unix")]
        unix: cfg
            .unix_socket_path
            .map(hyper_unix_socket::UnixSocketConnector::new),
    };

    let client = HyperClient::builder(TokioExecutor::new())
        .pool_idle_timeout(Duration::from_secs(30))
        .build(connector);

    Ok(client)
}

pub fn get_client<B>(cfg: ClientConfig) -> anyhow::Result<Client<B>>
where
    B: Body + Send + 'static,
    B::Data: Send,
{
    let mut cache = (*CACHE).lock().unwrap();
    if let Some(val) = cache.get(&cfg).and_then(|x| x.downcast_ref::<Client<B>>()) {
        Ok((val).clone())
    } else {
        let new_val = new_client(cfg.clone())?;
        cache.insert(cfg, Box::new(new_val.clone()));
        Ok(new_val)
    }
}