Skip to main content

salvo_proxy/
hyper_client.rs

1use std::io;
2
3use hyper::upgrade::OnUpgrade;
4use hyper_rustls::{HttpsConnector, HttpsConnectorBuilder};
5use hyper_util::client::legacy::Client as HyperUtilClient;
6use hyper_util::client::legacy::connect::{Connect, HttpConnector};
7use hyper_util::rt::TokioExecutor;
8use salvo_core::Error;
9use salvo_core::http::{ReqBody, ResBody, StatusCode};
10use salvo_core::rt::tokio::TokioIo;
11use tokio::io::copy_bidirectional;
12
13use crate::{BoxedError, Client, HyperRequest, HyperResponse, Proxy, Upstreams};
14
15/// A [`Client`] implementation based on [`hyper_util::client::legacy::Client`].
16///
17/// This client provides proxy capabilities using the Hyper HTTP client library.
18/// It's lightweight and tightly integrated with the Tokio runtime.
19#[derive(Clone, Debug)]
20pub struct HyperClient<C> {
21    inner: HyperUtilClient<C, ReqBody>,
22}
23
24fn build_default_https_connector_with(
25    native_roots: impl FnOnce() -> io::Result<HttpsConnector<HttpConnector>>,
26    webpki_roots: impl FnOnce() -> HttpsConnector<HttpConnector>,
27) -> HttpsConnector<HttpConnector> {
28    match native_roots() {
29        Ok(connector) => connector,
30        Err(error) => {
31            tracing::warn!(
32                error = ?error,
33                "failed to load native root certificates for proxy hyper client; falling back to webpki roots"
34            );
35            webpki_roots()
36        }
37    }
38}
39
40impl Default for HyperClient<HttpsConnector<HttpConnector>> {
41    fn default() -> Self {
42        #[cfg(feature = "ring")]
43        let _ = rustls::crypto::ring::default_provider().install_default();
44        let https = build_default_https_connector_with(
45            || {
46                Ok(HttpsConnectorBuilder::new()
47                    .with_native_roots()?
48                    .https_or_http()
49                    .enable_all_versions()
50                    .build())
51            },
52            || {
53                HttpsConnectorBuilder::new()
54                    .with_webpki_roots()
55                    .https_or_http()
56                    .enable_all_versions()
57                    .build()
58            },
59        );
60        Self {
61            inner: HyperUtilClient::builder(TokioExecutor::new()).build(https),
62        }
63    }
64}
65
66impl<U> Proxy<U, HyperClient<HttpsConnector<HttpConnector>>>
67where
68    U: Upstreams,
69    U::Error: Into<BoxedError>,
70{
71    /// Create a new `Proxy` using the default Hyper client.
72    ///
73    /// This is a convenient way to create a proxy with standard configuration.
74    pub fn use_hyper_client(upstreams: U) -> Self {
75        Self::new(upstreams, Default::default())
76    }
77}
78
79impl<C> HyperClient<C> {
80    /// Create a new `HyperClient` with the given `HyperClient`.
81    #[must_use]
82    pub fn new(inner: HyperUtilClient<C, ReqBody>) -> Self {
83        Self { inner }
84    }
85}
86
87impl<C> Client for HyperClient<C>
88where
89    C: Connect + Clone + Send + Sync + 'static,
90{
91    type Error = salvo_core::Error;
92
93    async fn execute(
94        &self,
95        proxied_request: HyperRequest,
96        request_upgraded: Option<OnUpgrade>,
97    ) -> Result<HyperResponse, Self::Error> {
98        let request_upgrade_type =
99            crate::get_upgrade_type(proxied_request.headers()).map(|s| s.to_owned());
100
101        let mut response = self
102            .inner
103            .request(proxied_request)
104            .await
105            .map_err(Error::other)?;
106
107        if response.status() == StatusCode::SWITCHING_PROTOCOLS {
108            let response_upgrade_type = crate::get_upgrade_type(response.headers());
109            if request_upgrade_type == response_upgrade_type.map(|s| s.to_lowercase()) {
110                let response_upgraded = hyper::upgrade::on(&mut response).await?;
111                if let Some(request_upgraded) = request_upgraded {
112                    tokio::spawn(async move {
113                        match request_upgraded.await {
114                            Ok(request_upgraded) => {
115                                let mut request_upgraded = TokioIo::new(request_upgraded);
116                                let mut response_upgraded = TokioIo::new(response_upgraded);
117                                if let Err(e) = copy_bidirectional(
118                                    &mut response_upgraded,
119                                    &mut request_upgraded,
120                                )
121                                .await
122                                {
123                                    tracing::error!(error = ?e, "copying between upgraded connections failed");
124                                }
125                            }
126                            Err(e) => {
127                                tracing::error!(error = ?e, "upgrade request failed");
128                            }
129                        }
130                    });
131                } else {
132                    return Err(Error::other("request does not have an upgrade extension"));
133                }
134            } else {
135                return Err(Error::other("upgrade type mismatch"));
136            }
137        }
138        Ok(response.map(ResBody::Hyper))
139    }
140}
141
142// Unit tests for Proxy
143#[cfg(test)]
144mod tests {
145    use std::io;
146
147    use salvo_core::prelude::*;
148    use salvo_core::test::*;
149
150    use super::*;
151    use crate::{Proxy, Upstreams};
152
153    #[test]
154    fn test_default_connector_falls_back_to_webpki_roots() {
155        let _ = rustls::crypto::aws_lc_rs::default_provider()
156            .install_default();
157
158        let connector = build_default_https_connector_with(
159            || Err(io::Error::other("missing native roots")),
160            || {
161                HttpsConnectorBuilder::new()
162                    .with_webpki_roots()
163                    .https_or_http()
164                    .enable_all_versions()
165                    .build()
166            },
167        );
168
169        let _client = HyperClient::new(HyperUtilClient::builder(TokioExecutor::new()).build(connector));
170    }
171
172    #[tokio::test]
173    async fn test_upstreams_elect() {
174        let _ = rustls::crypto::aws_lc_rs::default_provider()
175            .install_default();
176        let upstreams = vec!["https://www.example.com", "https://www.example2.com"];
177        let proxy = Proxy::new(upstreams.clone(), HyperClient::default());
178        let request = Request::new();
179        let depot = Depot::new();
180        let elected_upstream = proxy.upstreams().elect(&request, &depot).await.unwrap();
181        assert!(upstreams.contains(&elected_upstream));
182    }
183
184    #[tokio::test]
185    async fn test_hyper_client() {
186        let _ = rustls::crypto::aws_lc_rs::default_provider()
187            .install_default();
188        let router = Router::new().push(
189            Router::with_path("rust/{**rest}")
190                .goal(Proxy::new(vec!["https://salvo.rs"], HyperClient::default())),
191        );
192
193        let content = TestClient::get("http://127.0.0.1:5801/rust/guide/index.html")
194            .send(router)
195            .await
196            .take_string()
197            .await
198            .unwrap();
199        assert!(content.contains("Salvo"));
200    }
201
202    #[test]
203    fn test_others() {
204        let _ = rustls::crypto::aws_lc_rs::default_provider()
205            .install_default();
206        let mut handler = Proxy::new(["https://www.bing.com"], HyperClient::default());
207        assert_eq!(handler.upstreams().len(), 1);
208        assert_eq!(handler.upstreams_mut().len(), 1);
209    }
210}
211