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 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#[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