salvo-proxy 0.93.0

HTTP proxy support for the Salvo web server framework. Provides flexible proxy middleware for forwarding requests to upstream servers.
Documentation
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};

/// A [`Client`] implementation based on [`hyper_util::client::legacy::Client`].
///
/// This client provides proxy capabilities using the Hyper HTTP client library.
/// It's lightweight and tightly integrated with the Tokio runtime.
#[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>,
{
    /// Create a new `Proxy` using the default Hyper client.
    ///
    /// This is a convenient way to create a proxy with standard configuration.
    pub fn use_hyper_client(upstreams: U) -> Self {
        Self::new(upstreams, Default::default())
    }
}

impl<C> HyperClient<C> {
    /// Create a new `HyperClient` with the given `HyperClient`.
    #[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 request_upgrade_type == response_upgrade_type.map(|s| s.to_lowercase()) {
                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))
    }
}

// Unit tests for Proxy
#[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);
    }
}