salvo_proxy/
reqwest_client.rs1use futures_util::TryStreamExt;
2use hyper::upgrade::OnUpgrade;
3use reqwest::Client as InnerClient;
4use salvo_core::Error;
5use salvo_core::http::{ResBody, StatusCode};
6use salvo_core::rt::tokio::TokioIo;
7use tokio::io::copy_bidirectional;
8
9use crate::{BoxedError, Client, HyperRequest, HyperResponse, Proxy, Upstreams};
10
11#[derive(Clone, Debug)]
17pub struct ReqwestClient {
18 inner: InnerClient,
19}
20
21impl<U> Proxy<U, ReqwestClient>
22where
23 U: Upstreams,
24 U::Error: Into<BoxedError>,
25{
26 pub fn use_reqwest_client(upstreams: U) -> Self {
30 Self::new(upstreams, ReqwestClient::default())
31 }
32}
33
34impl Default for ReqwestClient {
35 fn default() -> Self {
36 #[cfg(feature = "ring")]
37 let _ = rustls::crypto::ring::default_provider().install_default();
38 Self::new(
39 InnerClient::builder()
40 .redirect(reqwest::redirect::Policy::none())
41 .build()
42 .expect("failed to build reqwest client"),
43 )
44 }
45}
46
47impl ReqwestClient {
48 #[must_use]
50 pub fn new(inner: InnerClient) -> Self {
51 Self { inner }
52 }
53}
54
55impl Client for ReqwestClient {
56 type Error = salvo_core::Error;
57
58 async fn execute(
59 &self,
60 proxied_request: HyperRequest,
61 request_upgraded: Option<OnUpgrade>,
62 ) -> Result<HyperResponse, Self::Error> {
63 let request_upgrade_type =
64 crate::get_upgrade_type(proxied_request.headers()).map(|s| s.to_owned());
65
66 let proxied_request = proxied_request.map(|body| {
67 reqwest::Body::wrap_stream(
72 body.try_filter_map(|frame| std::future::ready(Ok(frame.into_data().ok()))),
73 )
74 });
75 let mut response = self
76 .inner
77 .execute(proxied_request.try_into().map_err(Error::other)?)
78 .await
79 .map_err(Error::other)?;
80
81 let status = response.status();
82 let version = response.version();
83 let response_upgrade_type = if status == StatusCode::SWITCHING_PROTOCOLS {
84 crate::get_upgrade_type(response.headers()).map(str::to_owned)
85 } else {
86 None
87 };
88 let res_headers = std::mem::take(response.headers_mut());
89 let hyper_response = hyper::Response::builder()
90 .status(status)
91 .version(version);
92
93 let mut hyper_response = if status == StatusCode::SWITCHING_PROTOCOLS {
94 if crate::upgrade_types_match(
97 request_upgrade_type.as_deref(),
98 response_upgrade_type.as_deref(),
99 ) {
100 let mut response_upgraded = response.upgrade().await.map_err(|e| {
101 Error::other(format!("response does not have an upgrade extension. {e}"))
102 })?;
103 if let Some(request_upgraded) = request_upgraded {
104 tokio::spawn(async move {
105 match request_upgraded.await {
106 Ok(request_upgraded) => {
107 let mut request_upgraded = TokioIo::new(request_upgraded);
108 if let Err(e) = copy_bidirectional(
109 &mut response_upgraded,
110 &mut request_upgraded,
111 )
112 .await
113 {
114 tracing::error!(error = ?e, "copying between upgraded connections failed");
115 }
116 }
117 Err(e) => {
118 tracing::error!(error = ?e, "upgrade request failed");
119 }
120 }
121 });
122 } else {
123 return Err(Error::other("request does not have an upgrade extension"));
124 }
125 } else {
126 return Err(Error::other("upgrade type mismatch"));
127 }
128 hyper_response.body(ResBody::None).map_err(Error::other)?
129 } else {
130 hyper_response
131 .body(ResBody::stream(response.bytes_stream()))
132 .map_err(Error::other)?
133 };
134 *hyper_response.headers_mut() = res_headers;
135 Ok(hyper_response)
136 }
137}
138
139#[cfg(test)]
141mod tests {
142 use salvo_core::prelude::*;
143 use salvo_core::test::*;
144
145 use super::*;
146 use crate::{Proxy, Upstreams};
147
148 #[tokio::test]
149 async fn test_upstreams_elect() {
150 let upstreams = vec!["https://www.example.com", "https://www.example2.com"];
151 let proxy = Proxy::new(upstreams.clone(), ReqwestClient::default());
152 let request = Request::new();
153 let depot = Depot::new();
154 let elected_upstream = proxy.upstreams().elect(&request, &depot).await.unwrap();
155 assert!(upstreams.contains(&elected_upstream));
156 }
157
158 #[tokio::test]
159 async fn test_reqwest_client() {
160 let router = Router::new().push(Router::with_path("rust/{**rest}").goal(Proxy::new(
161 vec!["https://salvo.rs"],
162 ReqwestClient::default(),
163 )));
164
165 let content = TestClient::get("http://127.0.0.1:5801/rust/guide/index.html")
166 .send(router)
167 .await
168 .take_string()
169 .await
170 .unwrap();
171 assert!(content.contains("Salvo"));
172 }
173
174 #[test]
175 fn test_others() {
176 let mut handler = Proxy::new(["https://www.bing.com"], ReqwestClient::default());
177 assert_eq!(handler.upstreams().len(), 1);
178 assert_eq!(handler.upstreams_mut().len(), 1);
179 }
180}