Skip to main content

rama/http/client/
proxy_connector.rs

1use crate::{
2    Layer, Service,
3    error::{BoxError, BoxErrorExt},
4    extensions::{Extensions, ExtensionsRef},
5    http::client::proxy::layer::{
6        HttpProxyConnector, HttpProxyConnectorLayer, MaybeHttpProxiedConnection,
7    },
8    io::Io,
9    net::{
10        AuthorityInputExt, HttpVersionInputExt, Protocol, ProtocolInputExt,
11        TargetHttpVersionInputExt,
12        client::{
13            ConnectionError, ConnectionErrorKind, ConnectorService, EstablishedClientConnection,
14            EstablishedProxyRoute, ProxyRoute,
15        },
16    },
17    proxy::socks5::{Socks5ProxyConnector, Socks5ProxyConnectorLayer},
18    telemetry::tracing,
19};
20#[cfg(feature = "tls")]
21use crate::{net::client::ConnectionPolicyScope, tls::TlsTunnel};
22use pin_project_lite::pin_project;
23use std::{
24    fmt::Debug,
25    pin::Pin,
26    task::{self, Poll},
27};
28use tokio::io::{AsyncRead, AsyncWrite};
29
30/// Proxy connector which supports http(s) and socks5(h) proxy address
31///
32/// Connector will look at [`ProxyRoute`] to determine which proxy
33/// connector to use if one is configured
34#[derive(Debug, Clone)]
35pub struct ProxyConnector<S> {
36    inner: S,
37    socks: Socks5ProxyConnector<S>,
38    http: HttpProxyConnector<S>,
39    required: bool,
40}
41
42impl<S: Clone> ProxyConnector<S> {
43    /// Creates a new [`ProxyConnector`].
44    fn new(
45        inner: S,
46        socks_proxy_layer: Socks5ProxyConnectorLayer,
47        http_proxy_layer: HttpProxyConnectorLayer,
48        required: bool,
49    ) -> Self {
50        Self {
51            socks: socks_proxy_layer.into_layer(inner.clone()),
52            http: http_proxy_layer.into_layer(inner.clone()),
53            inner,
54            required,
55        }
56    }
57
58    #[inline]
59    /// Creates a new required [`ProxyConnector`].
60    ///
61    /// This connector will fail unless a proxied [`ProxyRoute`] is configured.
62    pub fn required(
63        inner: S,
64        socks_proxy_layer: Socks5ProxyConnectorLayer,
65        http_proxy_layer: HttpProxyConnectorLayer,
66    ) -> Self {
67        Self::new(inner, socks_proxy_layer, http_proxy_layer, true)
68    }
69
70    #[inline]
71    /// Creates a new optional [`ProxyConnector`].
72    ///
73    /// This connector will forward to the inner connector for a direct or missing [`ProxyRoute`].
74    pub fn optional(
75        inner: S,
76        socks_proxy_layer: Socks5ProxyConnectorLayer,
77        http_proxy_layer: HttpProxyConnectorLayer,
78    ) -> Self {
79        Self::new(inner, socks_proxy_layer, http_proxy_layer, false)
80    }
81}
82
83impl<Input, S> Service<Input> for ProxyConnector<S>
84where
85    S: ConnectorService<Input, Connection: Io + Unpin>,
86    Input: AuthorityInputExt
87        + ProtocolInputExt
88        + HttpVersionInputExt
89        + TargetHttpVersionInputExt
90        + Send
91        + ExtensionsRef
92        + 'static,
93{
94    type Output = EstablishedClientConnection<MaybeProxiedConnection<S::Connection>, Input>;
95    type Error = ConnectionError;
96
97    async fn serve(&self, input: Input) -> Result<Self::Output, Self::Error> {
98        let route = input.extensions().get_ref::<ProxyRoute>();
99        let route_requested = route.is_some();
100
101        match route {
102            None | Some(ProxyRoute::Direct) => {
103                if self.required {
104                    return Err(ConnectionError::local(
105                        BoxError::from_static_str("proxy required but none is defined"),
106                        ConnectionErrorKind::InvalidInput,
107                    ));
108                }
109                tracing::trace!("no proxy detected in ctx, using inner connector");
110                let EstablishedClientConnection { input, conn } = self.inner.connect(input).await?;
111
112                if route_requested {
113                    conn.extensions().insert(EstablishedProxyRoute::Direct);
114                }
115
116                let conn = MaybeProxiedConnection::direct(conn);
117                Ok(EstablishedClientConnection { input, conn })
118            }
119            Some(ProxyRoute::Proxy(proxy)) => {
120                let protocol = proxy.protocol.as_ref();
121                tracing::trace!(?protocol, "proxy detected in ctx");
122
123                let protocol = protocol.unwrap_or_else(|| {
124                    tracing::trace!("no protocol detected, using http as protocol");
125                    &Protocol::HTTP
126                });
127
128                if protocol.is_socks5() {
129                    tracing::trace!(
130                        target = %&proxy.address,
131                        "using socks proxy connector",
132                    );
133
134                    // SOCKS does not supply a tunnel TLS policy of its own.
135                    // An explicit request policy must not poison shared routes.
136                    #[cfg(feature = "tls")]
137                    let request_tunnel = input.extensions().contains::<TlsTunnel>();
138                    let established = self.socks.connect(input).await;
139                    #[cfg(feature = "tls")]
140                    let established = established.map_err(|error| {
141                        if request_tunnel {
142                            error.with_policy_scope(ConnectionPolicyScope::Request)
143                        } else {
144                            error
145                        }
146                    });
147                    let EstablishedClientConnection { input, conn } = established?;
148
149                    let conn = MaybeProxiedConnection::socks(conn);
150                    Ok(EstablishedClientConnection { input, conn })
151                } else if protocol.is_http() {
152                    tracing::trace!(
153                        target = %&proxy.address,
154                        "using http proxy connector"
155                    );
156
157                    let EstablishedClientConnection { input, conn } =
158                        self.http.connect(input).await?;
159
160                    let conn = MaybeProxiedConnection::http(conn);
161                    Ok(EstablishedClientConnection { input, conn })
162                } else {
163                    Err(ConnectionError::transport(
164                        BoxError::from_static_str("received unsupported proxy protocol"),
165                        ConnectionErrorKind::Protocol,
166                    )
167                    .context_debug_field("protocol", protocol.clone()))
168                }
169            }
170        }
171    }
172}
173
174pin_project! {
175    /// A connection which will be proxied if a proxied [`ProxyRoute`] was configured.
176    pub struct MaybeProxiedConnection<S> {
177        #[pin]
178        inner: Connection<S>,
179    }
180}
181
182impl<S: ExtensionsRef> MaybeProxiedConnection<S> {
183    pub fn direct(conn: S) -> Self {
184        Self {
185            inner: Connection::Direct { conn },
186        }
187    }
188
189    pub fn socks(conn: S) -> Self {
190        Self {
191            inner: Connection::Socks { conn },
192        }
193    }
194
195    pub fn http(conn: MaybeHttpProxiedConnection<S>) -> Self {
196        Self {
197            inner: Connection::Http { conn },
198        }
199    }
200}
201
202impl<S: Debug> Debug for MaybeProxiedConnection<S> {
203    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
204        f.debug_struct("MaybeProxiedConnection")
205            .field("inner", &self.inner)
206            .finish()
207    }
208}
209
210impl<S: ExtensionsRef> ExtensionsRef for MaybeProxiedConnection<S> {
211    fn extensions(&self) -> &Extensions {
212        match &self.inner {
213            Connection::Direct { conn } | Connection::Socks { conn } => conn.extensions(),
214            Connection::Http { conn } => conn.extensions(),
215        }
216    }
217}
218
219pin_project! {
220    #[project = ConnectionProj]
221    enum Connection<S> {
222        Direct{ #[pin] conn: S },
223        Socks{ #[pin] conn: S },
224        Http{ #[pin] conn: MaybeHttpProxiedConnection<S> },
225
226    }
227}
228
229impl<S: Debug> Debug for Connection<S> {
230    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
231        match self {
232            Self::Direct { conn } => f.debug_struct("Direct").field("conn", conn).finish(),
233            Self::Socks { conn } => f.debug_struct("Socks").field("conn", conn).finish(),
234            Self::Http { conn } => f.debug_struct("Http").field("conn", conn).finish(),
235        }
236    }
237}
238
239#[warn(clippy::missing_trait_methods)]
240impl<Conn: AsyncWrite> AsyncWrite for MaybeProxiedConnection<Conn> {
241    fn poll_write(
242        self: Pin<&mut Self>,
243        cx: &mut task::Context<'_>,
244        buf: &[u8],
245    ) -> Poll<Result<usize, std::io::Error>> {
246        match self.project().inner.project() {
247            ConnectionProj::Direct { conn } | ConnectionProj::Socks { conn } => {
248                conn.poll_write(cx, buf)
249            }
250            ConnectionProj::Http { conn } => conn.poll_write(cx, buf),
251        }
252    }
253
254    fn poll_flush(
255        self: Pin<&mut Self>,
256        cx: &mut task::Context<'_>,
257    ) -> Poll<Result<(), std::io::Error>> {
258        match self.project().inner.project() {
259            ConnectionProj::Direct { conn } | ConnectionProj::Socks { conn } => conn.poll_flush(cx),
260            ConnectionProj::Http { conn } => conn.poll_flush(cx),
261        }
262    }
263
264    fn poll_shutdown(
265        self: Pin<&mut Self>,
266        cx: &mut task::Context<'_>,
267    ) -> Poll<Result<(), std::io::Error>> {
268        match self.project().inner.project() {
269            ConnectionProj::Direct { conn } | ConnectionProj::Socks { conn } => {
270                conn.poll_shutdown(cx)
271            }
272            ConnectionProj::Http { conn } => conn.poll_shutdown(cx),
273        }
274    }
275
276    fn is_write_vectored(&self) -> bool {
277        match &self.inner {
278            Connection::Direct { conn } | Connection::Socks { conn } => conn.is_write_vectored(),
279            Connection::Http { conn } => conn.is_write_vectored(),
280        }
281    }
282
283    fn poll_write_vectored(
284        self: Pin<&mut Self>,
285        cx: &mut task::Context<'_>,
286        bufs: &[std::io::IoSlice<'_>],
287    ) -> Poll<Result<usize, std::io::Error>> {
288        match self.project().inner.project() {
289            ConnectionProj::Direct { conn } | ConnectionProj::Socks { conn } => {
290                conn.poll_write_vectored(cx, bufs)
291            }
292            ConnectionProj::Http { conn } => conn.poll_write_vectored(cx, bufs),
293        }
294    }
295}
296
297#[warn(clippy::missing_trait_methods)]
298impl<Conn: AsyncRead> AsyncRead for MaybeProxiedConnection<Conn> {
299    fn poll_read(
300        self: Pin<&mut Self>,
301        cx: &mut task::Context<'_>,
302        buf: &mut tokio::io::ReadBuf<'_>,
303    ) -> Poll<std::io::Result<()>> {
304        match self.project().inner.project() {
305            ConnectionProj::Direct { conn } | ConnectionProj::Socks { conn } => {
306                conn.poll_read(cx, buf)
307            }
308            ConnectionProj::Http { conn } => conn.poll_read(cx, buf),
309        }
310    }
311}
312
313/// Proxy connector layer which supports http(s) and socks5(h) proxy address
314///
315/// Connector will look at [`ProxyRoute`] to determine which proxy
316/// connector to use if one is configured
317pub struct ProxyConnectorLayer {
318    socks_layer: Socks5ProxyConnectorLayer,
319    http_layer: HttpProxyConnectorLayer,
320    required: bool,
321}
322
323impl ProxyConnectorLayer {
324    #[must_use]
325    /// Creates a new required [`ProxyConnectorLayer`].
326    ///
327    /// This connector will fail unless a proxied [`ProxyRoute`] is configured.
328    pub fn required(
329        socks_proxy_layer: Socks5ProxyConnectorLayer,
330        http_proxy_layer: HttpProxyConnectorLayer,
331    ) -> Self {
332        Self {
333            socks_layer: socks_proxy_layer,
334            http_layer: http_proxy_layer,
335            required: true,
336        }
337    }
338
339    #[must_use]
340    /// Creates a new optional [`ProxyConnectorLayer`].
341    ///
342    /// This connector will forward to the inner connector for a direct or missing [`ProxyRoute`].
343    pub fn optional(
344        socks_proxy_layer: Socks5ProxyConnectorLayer,
345        http_proxy_layer: HttpProxyConnectorLayer,
346    ) -> Self {
347        Self {
348            socks_layer: socks_proxy_layer,
349            http_layer: http_proxy_layer,
350            required: false,
351        }
352    }
353}
354
355impl<S: Clone> Layer<S> for ProxyConnectorLayer {
356    type Service = ProxyConnector<S>;
357
358    fn layer(&self, inner: S) -> Self::Service {
359        ProxyConnector::new(
360            inner,
361            self.socks_layer.clone(),
362            self.http_layer.clone(),
363            self.required,
364        )
365    }
366
367    fn into_layer(self, inner: S) -> Self::Service {
368        ProxyConnector::new(inner, self.socks_layer, self.http_layer, self.required)
369    }
370}
371
372#[cfg(test)]
373mod tests {
374    use super::*;
375    use crate::{
376        net::{proxy::IoForwardService, test_utils::client::MockSocket},
377        proxy::socks5::{
378            Socks5ProxyConnectorLayer,
379            server::{Connector as EagerSocks5Connector, Socks5Connector},
380        },
381        tcp::client::service::TcpConnector,
382    };
383
384    #[cfg(feature = "tls")]
385    use {
386        crate::{
387            net::{
388                address::HostWithPort,
389                client::{
390                    ConnectRequest, ProxyRouteFailureCache, ProxyRouteFailureCacheConfig,
391                    ProxyRouteFailureCacheConnector, ProxyRouteFailureCacheScope,
392                },
393            },
394            service::service_fn,
395        },
396        std::sync::{
397            Arc,
398            atomic::{AtomicUsize, Ordering},
399        },
400    };
401
402    fn assert_socks5_connector<S, C: Socks5Connector<S>>(_: &C) {}
403
404    #[cfg(feature = "tls")]
405    #[tokio::test]
406    async fn request_tunnel_failure_does_not_poison_socks_proxy_route() {
407        for scope in [
408            ProxyRouteFailureCacheScope::PerDestination,
409            ProxyRouteFailureCacheScope::PerProxy,
410        ] {
411            let attempts = Arc::new(AtomicUsize::new(0));
412            let transport = service_fn({
413                let attempts = attempts.clone();
414                move |input: ConnectRequest| {
415                    attempts.fetch_add(1, Ordering::Relaxed);
416                    async move {
417                        Err::<EstablishedClientConnection<MockSocket, ConnectRequest>, _>(
418                            ConnectionError::transport(
419                                BoxError::from_static_str("scripted proxy failure"),
420                                if input.extensions().contains::<TlsTunnel>() {
421                                    ConnectionErrorKind::Protocol
422                                } else {
423                                    ConnectionErrorKind::Unavailable
424                                },
425                            ),
426                        )
427                    }
428                }
429            });
430            let mut config = ProxyRouteFailureCacheConfig::default();
431            config.scope = scope;
432            let connector = ProxyRouteFailureCacheConnector::new(
433                ProxyConnector::required(
434                    transport,
435                    Socks5ProxyConnectorLayer::required(),
436                    HttpProxyConnectorLayer::required(),
437                ),
438                ProxyRouteFailureCache::try_new(config).unwrap(),
439            );
440            let request = || {
441                let input = ConnectRequest::new(HostWithPort::example_domain_http())
442                    .with_application_protocol(Protocol::HTTP);
443                input.extensions().insert(ProxyRoute::Proxy(
444                    "socks5://127.0.0.1:1080".parse().unwrap(),
445                ));
446                input
447            };
448            let customized = request();
449            customized.extensions().insert(TlsTunnel {
450                server_identity: Some("proxy.example".parse().unwrap()),
451                application_protocol: None,
452                alpn: None,
453            });
454            let error = connector.serve(customized).await.unwrap_err();
455            assert_eq!(error.kind(), ConnectionErrorKind::Protocol);
456            assert_eq!(error.policy_scope(), ConnectionPolicyScope::Request);
457
458            let error = connector.serve(request()).await.unwrap_err();
459            assert_eq!(error.kind(), ConnectionErrorKind::Unavailable);
460            assert_eq!(error.policy_scope(), ConnectionPolicyScope::Unknown);
461            assert_eq!(attempts.load(Ordering::Relaxed), 2);
462            _ = connector.serve(request()).await.unwrap_err();
463            assert_eq!(
464                attempts.load(Ordering::Relaxed),
465                2,
466                "ordinary failures remain cached"
467            );
468        }
469    }
470
471    #[test]
472    fn eager_socks5_accepts_combined_proxy_connection() {
473        let proxy_connector = ProxyConnectorLayer::optional(
474            Socks5ProxyConnectorLayer::optional(),
475            HttpProxyConnectorLayer::optional(),
476        )
477        .into_layer(TcpConnector::new());
478        let connector = EagerSocks5Connector::new(proxy_connector, IoForwardService::default());
479
480        assert_socks5_connector::<MockSocket, _>(&connector);
481    }
482}