salvo_proxy/
hyper_client.rs1use 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#[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 pub fn use_hyper_client(upstreams: U) -> Self {
75 Self::new(upstreams, Default::default())
76 }
77}
78
79impl<C> HyperClient<C> {
80 #[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 crate::upgrade_types_match(request_upgrade_type.as_deref(), response_upgrade_type) {
112 let response_upgraded = hyper::upgrade::on(&mut response).await?;
113 if let Some(request_upgraded) = request_upgraded {
114 tokio::spawn(async move {
115 match request_upgraded.await {
116 Ok(request_upgraded) => {
117 let mut request_upgraded = TokioIo::new(request_upgraded);
118 let mut response_upgraded = TokioIo::new(response_upgraded);
119 if let Err(e) = copy_bidirectional(
120 &mut response_upgraded,
121 &mut request_upgraded,
122 )
123 .await
124 {
125 tracing::error!(error = ?e, "copying between upgraded connections failed");
126 }
127 }
128 Err(e) => {
129 tracing::error!(error = ?e, "upgrade request failed");
130 }
131 }
132 });
133 } else {
134 return Err(Error::other("request does not have an upgrade extension"));
135 }
136 } else {
137 return Err(Error::other("upgrade type mismatch"));
138 }
139 }
140 Ok(response.map(ResBody::Hyper))
141 }
142}
143
144#[cfg(test)]
146mod tests {
147 use std::io;
148
149 use salvo_core::prelude::*;
150 use salvo_core::test::*;
151
152 use super::*;
153 use crate::{Proxy, Upstreams};
154
155 #[test]
156 fn test_default_connector_falls_back_to_webpki_roots() {
157 let _ = rustls::crypto::aws_lc_rs::default_provider()
158 .install_default();
159
160 let connector = build_default_https_connector_with(
161 || Err(io::Error::other("missing native roots")),
162 || {
163 HttpsConnectorBuilder::new()
164 .with_webpki_roots()
165 .https_or_http()
166 .enable_all_versions()
167 .build()
168 },
169 );
170
171 let _client = HyperClient::new(HyperUtilClient::builder(TokioExecutor::new()).build(connector));
172 }
173
174 #[tokio::test]
175 async fn test_upstreams_elect() {
176 let _ = rustls::crypto::aws_lc_rs::default_provider()
177 .install_default();
178 let upstreams = vec!["https://www.example.com", "https://www.example2.com"];
179 let proxy = Proxy::new(upstreams.clone(), HyperClient::default());
180 let request = Request::new();
181 let depot = Depot::new();
182 let elected_upstream = proxy.upstreams().elect(&request, &depot).await.unwrap();
183 assert!(upstreams.contains(&elected_upstream));
184 }
185
186 #[tokio::test]
187 async fn test_hyper_client() {
188 let _ = rustls::crypto::aws_lc_rs::default_provider()
189 .install_default();
190 let router = Router::new().push(
191 Router::with_path("rust/{**rest}")
192 .goal(Proxy::new(vec!["https://salvo.rs"], HyperClient::default())),
193 );
194
195 let content = TestClient::get("http://127.0.0.1:5801/rust/guide/index.html")
196 .send(router)
197 .await
198 .take_string()
199 .await
200 .unwrap();
201 assert!(content.contains("Salvo"));
202 }
203
204 #[test]
205 fn test_others() {
206 let _ = rustls::crypto::aws_lc_rs::default_provider()
207 .install_default();
208 let mut handler = Proxy::new(["https://www.bing.com"], HyperClient::default());
209 assert_eq!(handler.upstreams().len(), 1);
210 assert_eq!(handler.upstreams_mut().len(), 1);
211 }
212}
213