use std::io;
use hyper::upgrade::OnUpgrade;
use hyper_rustls::{HttpsConnector, HttpsConnectorBuilder};
use hyper_util::client::legacy::Client as HyperUtilClient;
use hyper_util::client::legacy::connect::{Connect, HttpConnector};
use hyper_util::rt::TokioExecutor;
use salvo_core::Error;
use salvo_core::http::{ReqBody, ResBody, StatusCode};
use salvo_core::rt::tokio::TokioIo;
use tokio::io::copy_bidirectional;
use crate::{BoxedError, Client, HyperRequest, HyperResponse, Proxy, Upstreams};
#[derive(Clone, Debug)]
pub struct HyperClient<C> {
inner: HyperUtilClient<C, ReqBody>,
}
fn build_default_https_connector_with(
native_roots: impl FnOnce() -> io::Result<HttpsConnector<HttpConnector>>,
webpki_roots: impl FnOnce() -> HttpsConnector<HttpConnector>,
) -> HttpsConnector<HttpConnector> {
match native_roots() {
Ok(connector) => connector,
Err(error) => {
tracing::warn!(
error = ?error,
"failed to load native root certificates for proxy hyper client; falling back to webpki roots"
);
webpki_roots()
}
}
}
impl Default for HyperClient<HttpsConnector<HttpConnector>> {
fn default() -> Self {
#[cfg(feature = "ring")]
let _ = rustls::crypto::ring::default_provider().install_default();
let https = build_default_https_connector_with(
|| {
Ok(HttpsConnectorBuilder::new()
.with_native_roots()?
.https_or_http()
.enable_all_versions()
.build())
},
|| {
HttpsConnectorBuilder::new()
.with_webpki_roots()
.https_or_http()
.enable_all_versions()
.build()
},
);
Self {
inner: HyperUtilClient::builder(TokioExecutor::new()).build(https),
}
}
}
impl<U> Proxy<U, HyperClient<HttpsConnector<HttpConnector>>>
where
U: Upstreams,
U::Error: Into<BoxedError>,
{
pub fn use_hyper_client(upstreams: U) -> Self {
Self::new(upstreams, Default::default())
}
}
impl<C> HyperClient<C> {
#[must_use]
pub fn new(inner: HyperUtilClient<C, ReqBody>) -> Self {
Self { inner }
}
}
impl<C> Client for HyperClient<C>
where
C: Connect + Clone + Send + Sync + 'static,
{
type Error = salvo_core::Error;
async fn execute(
&self,
proxied_request: HyperRequest,
request_upgraded: Option<OnUpgrade>,
) -> Result<HyperResponse, Self::Error> {
let request_upgrade_type =
crate::get_upgrade_type(proxied_request.headers()).map(|s| s.to_owned());
let mut response = self
.inner
.request(proxied_request)
.await
.map_err(Error::other)?;
if response.status() == StatusCode::SWITCHING_PROTOCOLS {
let response_upgrade_type = crate::get_upgrade_type(response.headers());
if crate::upgrade_types_match(request_upgrade_type.as_deref(), response_upgrade_type) {
let response_upgraded = hyper::upgrade::on(&mut response).await?;
if let Some(request_upgraded) = request_upgraded {
tokio::spawn(async move {
match request_upgraded.await {
Ok(request_upgraded) => {
let mut request_upgraded = TokioIo::new(request_upgraded);
let mut response_upgraded = TokioIo::new(response_upgraded);
if let Err(e) = copy_bidirectional(
&mut response_upgraded,
&mut request_upgraded,
)
.await
{
tracing::error!(error = ?e, "copying between upgraded connections failed");
}
}
Err(e) => {
tracing::error!(error = ?e, "upgrade request failed");
}
}
});
} else {
return Err(Error::other("request does not have an upgrade extension"));
}
} else {
return Err(Error::other("upgrade type mismatch"));
}
}
Ok(response.map(ResBody::Hyper))
}
}
#[cfg(test)]
mod tests {
use std::io;
use salvo_core::prelude::*;
use salvo_core::test::*;
use super::*;
use crate::{Proxy, Upstreams};
#[test]
fn test_default_connector_falls_back_to_webpki_roots() {
let _ = rustls::crypto::aws_lc_rs::default_provider()
.install_default();
let connector = build_default_https_connector_with(
|| Err(io::Error::other("missing native roots")),
|| {
HttpsConnectorBuilder::new()
.with_webpki_roots()
.https_or_http()
.enable_all_versions()
.build()
},
);
let _client = HyperClient::new(HyperUtilClient::builder(TokioExecutor::new()).build(connector));
}
#[tokio::test]
async fn test_upstreams_elect() {
let _ = rustls::crypto::aws_lc_rs::default_provider()
.install_default();
let upstreams = vec!["https://www.example.com", "https://www.example2.com"];
let proxy = Proxy::new(upstreams.clone(), HyperClient::default());
let request = Request::new();
let depot = Depot::new();
let elected_upstream = proxy.upstreams().elect(&request, &depot).await.unwrap();
assert!(upstreams.contains(&elected_upstream));
}
#[tokio::test]
async fn test_hyper_client() {
let _ = rustls::crypto::aws_lc_rs::default_provider()
.install_default();
let router = Router::new().push(
Router::with_path("rust/{**rest}")
.goal(Proxy::new(vec!["https://salvo.rs"], HyperClient::default())),
);
let content = TestClient::get("http://127.0.0.1:5801/rust/guide/index.html")
.send(router)
.await
.take_string()
.await
.unwrap();
assert!(content.contains("Salvo"));
}
#[test]
fn test_others() {
let _ = rustls::crypto::aws_lc_rs::default_provider()
.install_default();
let mut handler = Proxy::new(["https://www.bing.com"], HyperClient::default());
assert_eq!(handler.upstreams().len(), 1);
assert_eq!(handler.upstreams_mut().len(), 1);
}
}