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#[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 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 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 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 #[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 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
313pub struct ProxyConnectorLayer {
318 socks_layer: Socks5ProxyConnectorLayer,
319 http_layer: HttpProxyConnectorLayer,
320 required: bool,
321}
322
323impl ProxyConnectorLayer {
324 #[must_use]
325 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 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}