Skip to main content

salvo_proxy/
reqwest_client.rs

1use 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/// A [`Client`] implementation based on [`reqwest::Client`].
12///
13/// This client provides proxy capabilities using the Reqwest HTTP client.
14/// Redirect following is disabled by default for proxy safety. If you need to
15/// follow upstream redirects, pass a custom [`reqwest::Client`] to [`ReqwestClient::new`].
16#[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    /// Create a new `Proxy` using the default Reqwest client.
27    ///
28    /// This is a convenient way to create a proxy with standard configuration.
29    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    /// Create a new `ReqwestClient` with the given [`reqwest::Client`].
49    #[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            // Forward only data frames; drop non-data frames (e.g. trailers, which
68            // reqwest cannot send anyway) instead of turning them into empty bytes,
69            // and propagate stream errors so a failed/aborted upstream body surfaces
70            // as an error rather than being silently truncated into a "successful" body.
71            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            // RFC 7230 ยง6.7 makes Upgrade tokens case-insensitive (`websocket`,
95            // `WebSocket`, ... are all the same protocol). Compare without allocating.
96            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// Unit tests for Proxy
140#[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}