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)
}
}