Skip to main content

rama_net/client/proxy/
address.rs

1use std::{
2    fmt,
3    sync::{Arc, OnceLock},
4};
5
6use crate::address::ProxyAddress;
7use rama_core::{
8    Layer, Service, error::BoxError, error_sink::ErrorSink, extensions::ExtensionsRef,
9    telemetry::tracing,
10};
11
12use super::{
13    ProxyRoute, ProxyRoutes,
14    env::proxy_address_from_env,
15    load::{CachedLoadError, LoadErrorPolicy},
16};
17
18#[derive(Debug, Clone, Default)]
19/// Apply one fixed proxy address to any service input with extensions.
20///
21/// This layer reads application environment variables only when constructed
22/// with [`try_from_env`][Self::try_from_env]. It does not inspect operating
23/// system proxy settings; use [`SystemProxyLayer`] for those. When the layers
24/// are chained, existing route decisions are preserved by default. Opt into
25/// overwriting only when this address is authoritative. Use
26/// [`LazyProxyAddressLayer`] when environment lookup should happen only if
27/// no higher-priority route has already been selected.
28///
29/// See [`ProxyAddressService`] for more information.
30///
31/// [`Extensions`]: rama_core::extensions::Extensions
32/// [`SystemProxyLayer`]: crate::client::SystemProxyLayer
33pub struct ProxyAddressLayer {
34    address: Option<ProxyAddress>,
35    overwrite: bool,
36}
37
38impl ProxyAddressLayer {
39    /// Create a new [`ProxyAddressLayer`] that will create
40    /// a service to set the given [`ProxyAddress`] as a proxied [`ProxyRoute`].
41    #[must_use]
42    pub fn new(address: ProxyAddress) -> Self {
43        Self::maybe(Some(address))
44    }
45
46    /// Create a new [`ProxyAddressLayer`] which will create
47    /// a service that will set the given [`ProxyAddress`] as a proxied [`ProxyRoute`] if it is not
48    /// `None`.
49    #[must_use]
50    pub fn maybe(address: Option<ProxyAddress>) -> Self {
51        Self {
52            address,
53            ..Default::default()
54        }
55    }
56
57    /// Return the configured proxy address, when this layer has one.
58    #[must_use]
59    pub const fn proxy_address(&self) -> Option<&ProxyAddress> {
60        self.address.as_ref()
61    }
62
63    /// Try to create a new [`ProxyAddressLayer`] which will establish
64    /// a proxy connection over the environment variable `http_proxy`.
65    ///
66    /// Uppercase `HTTP_PROXY` is deliberately not accepted by default because
67    /// CGI derives it from an incoming `Proxy` header. Use
68    /// [`ProxyEnvLayer`] for curl-compatible HTTP, HTTPS, and all-protocol
69    /// environment selection.
70    ///
71    /// [`ProxyEnvLayer`]: crate::client::ProxyEnvLayer
72    pub fn try_from_env_default() -> Result<Self, BoxError> {
73        Self::try_from_env("http_proxy")
74    }
75
76    /// Try to create a new [`ProxyAddressLayer`] which will establish
77    /// a proxy connection over the given environment variable.
78    pub fn try_from_env(key: impl AsRef<str>) -> Result<Self, BoxError> {
79        proxy_address_from_env(key.as_ref()).map(Self::maybe)
80    }
81
82    rama_utils::macros::generate_set_and_with! {
83        /// Replace an existing [`ProxyRoute`] or [`ProxyRoutes`] decision.
84        /// Existing routes are preserved by default.
85        pub fn overwrite(mut self, overwrite: bool) -> Self {
86            self.overwrite = overwrite;
87            self
88        }
89    }
90}
91
92impl<S> Layer<S> for ProxyAddressLayer {
93    type Service = ProxyAddressService<S>;
94
95    fn layer(&self, inner: S) -> Self::Service {
96        ProxyAddressService::maybe(inner, self.address.clone()).with_overwrite(self.overwrite)
97    }
98
99    fn into_layer(self, inner: S) -> Self::Service {
100        ProxyAddressService::maybe(inner, self.address).with_overwrite(self.overwrite)
101    }
102}
103
104/// Service produced by [`ProxyAddressLayer`].
105///
106/// [`Extensions`]: rama_core::extensions::Extensions
107#[derive(Debug, Clone)]
108pub struct ProxyAddressService<S> {
109    inner: S,
110    proxy_info: Option<ProxyAddress>,
111    overwrite: bool,
112}
113
114impl<S> ProxyAddressService<S> {
115    /// Create a new [`ProxyAddressService`] that will create
116    /// a service to set the given [`ProxyAddress`] as a proxied [`ProxyRoute`].
117    pub const fn new(inner: S, address: ProxyAddress) -> Self {
118        Self::maybe(inner, Some(address))
119    }
120
121    /// Create a new [`ProxyAddressService`] which will create
122    /// a service that will set the given [`ProxyAddress`] as a proxied [`ProxyRoute`] if it is not
123    /// `None`.
124    pub const fn maybe(inner: S, address: Option<ProxyAddress>) -> Self {
125        Self {
126            inner,
127            proxy_info: address,
128            overwrite: false,
129        }
130    }
131
132    /// Try to create a new [`ProxyAddressService`] which will establish
133    /// a proxy connection over the environment variable `http_proxy`.
134    ///
135    /// Uppercase `HTTP_PROXY` is deliberately not accepted by default because
136    /// CGI derives it from an incoming `Proxy` header. Use
137    /// [`ProxyEnvLayer`] for curl-compatible HTTP, HTTPS, and all-protocol
138    /// environment selection.
139    ///
140    /// [`ProxyEnvLayer`]: crate::client::ProxyEnvLayer
141    pub fn try_from_env_default(inner: S) -> Result<Self, BoxError> {
142        Self::try_from_env(inner, "http_proxy")
143    }
144
145    /// Try to create a new [`ProxyAddressService`] which will establish
146    /// a proxy connection over the given environment variable.
147    pub fn try_from_env(inner: S, key: impl AsRef<str>) -> Result<Self, BoxError> {
148        proxy_address_from_env(key.as_ref()).map(|address| Self::maybe(inner, address))
149    }
150
151    rama_utils::macros::generate_set_and_with! {
152        /// Replace an existing [`ProxyRoute`] or [`ProxyRoutes`] decision.
153        /// Existing routes are preserved by default.
154        pub fn overwrite(mut self, overwrite: bool) -> Self {
155            self.overwrite = overwrite;
156            self
157        }
158    }
159}
160
161type ProxyAddressLoader =
162    dyn Fn() -> Result<Option<ProxyAddress>, BoxError> + Send + Sync + 'static;
163
164type CachedProxyAddress = Result<Option<ProxyAddress>, CachedLoadError>;
165
166/// Lazily resolve and apply a proxy address to any input with extensions.
167///
168/// Existing routes are preserved by default, in which case the loader is not
169/// consulted if a [`ProxyRoute`] or [`ProxyRoutes`] decision exists. Otherwise,
170/// its result is cached and shared by every clone of the layer and service.
171/// This is useful for environment configuration that may never be needed.
172#[derive(Clone)]
173pub struct LazyProxyAddressLayer {
174    loader: Arc<ProxyAddressLoader>,
175    cached: Arc<OnceLock<CachedProxyAddress>>,
176    load_error_policy: LoadErrorPolicy,
177    overwrite: bool,
178}
179
180impl fmt::Debug for LazyProxyAddressLayer {
181    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
182        f.debug_struct("LazyProxyAddressLayer")
183            .field("cached", &self.cached.get())
184            .field("load_error_policy", &self.load_error_policy)
185            .field("overwrite", &self.overwrite)
186            .finish_non_exhaustive()
187    }
188}
189
190impl LazyProxyAddressLayer {
191    /// Create a lazy layer backed by a synchronous, non-blocking loader.
192    ///
193    /// The loader runs at most once, on the first request that does not already
194    /// have a preserved route decision. Both success and failure are cached.
195    #[must_use]
196    pub fn new<F>(loader: F) -> Self
197    where
198        F: Fn() -> Result<Option<ProxyAddress>, BoxError> + Send + Sync + 'static,
199    {
200        Self {
201            loader: Arc::new(loader),
202            cached: Arc::new(OnceLock::new()),
203            load_error_policy: LoadErrorPolicy::Reject,
204            overwrite: false,
205        }
206    }
207
208    /// Lazily read and parse the `http_proxy` environment variable on the
209    /// first request without a preserved route.
210    ///
211    /// Uppercase `HTTP_PROXY` is deliberately not accepted by default because
212    /// CGI derives it from an incoming `Proxy` header. Use
213    /// [`ProxyEnvLayer`] for curl-compatible HTTP, HTTPS, and all-protocol
214    /// environment selection.
215    ///
216    /// [`ProxyEnvLayer`]: crate::client::ProxyEnvLayer
217    #[must_use]
218    pub fn from_env_default() -> Self {
219        Self::from_env("http_proxy")
220    }
221
222    /// Lazily read and parse a proxy address from the named environment
223    /// variable on the first request without a preserved route.
224    #[must_use]
225    pub fn from_env(key: impl Into<String>) -> Self {
226        let key = key.into();
227        Self::new(move || proxy_address_from_env(&key))
228    }
229
230    rama_utils::macros::generate_set_and_with! {
231        /// Handle a loader error through an [`ErrorSink`] and continue without
232        /// selecting a proxy. By default loader errors reject the request.
233        ///
234        /// The sink is invoked at most once because the handled result is
235        /// cached and shared by every clone of this layer and its service.
236        pub fn load_error_sink(
237            mut self,
238            sink: impl ErrorSink,
239        ) -> Self {
240            self.load_error_policy = LoadErrorPolicy::Handle(Arc::new(sink));
241            self.cached = Arc::new(OnceLock::new());
242            self
243        }
244    }
245
246    rama_utils::macros::generate_set_and_with! {
247        /// Replace an existing [`ProxyRoute`] or [`ProxyRoutes`] decision.
248        /// Existing routes are preserved by default.
249        pub fn overwrite(mut self, overwrite: bool) -> Self {
250            self.overwrite = overwrite;
251            self
252        }
253    }
254}
255
256impl<S> Layer<S> for LazyProxyAddressLayer {
257    type Service = LazyProxyAddressService<S>;
258
259    fn layer(&self, inner: S) -> Self::Service {
260        LazyProxyAddressService {
261            inner,
262            loader: self.loader.clone(),
263            cached: self.cached.clone(),
264            load_error_policy: self.load_error_policy.clone(),
265            overwrite: self.overwrite,
266        }
267    }
268
269    fn into_layer(self, inner: S) -> Self::Service {
270        LazyProxyAddressService {
271            inner,
272            loader: self.loader,
273            cached: self.cached,
274            load_error_policy: self.load_error_policy,
275            overwrite: self.overwrite,
276        }
277    }
278}
279
280/// Service produced by [`LazyProxyAddressLayer`].
281#[derive(Clone)]
282pub struct LazyProxyAddressService<S> {
283    inner: S,
284    loader: Arc<ProxyAddressLoader>,
285    cached: Arc<OnceLock<CachedProxyAddress>>,
286    load_error_policy: LoadErrorPolicy,
287    overwrite: bool,
288}
289
290impl<S: fmt::Debug> fmt::Debug for LazyProxyAddressService<S> {
291    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
292        f.debug_struct("LazyProxyAddressService")
293            .field("inner", &self.inner)
294            .field("cached", &self.cached.get())
295            .field("load_error_policy", &self.load_error_policy)
296            .field("overwrite", &self.overwrite)
297            .finish_non_exhaustive()
298    }
299}
300
301impl<S, Input> Service<Input> for LazyProxyAddressService<S>
302where
303    S: Service<Input, Error: Into<BoxError>>,
304    Input: ExtensionsRef + Send + 'static,
305{
306    type Output = S::Output;
307    type Error = BoxError;
308
309    async fn serve(&self, input: Input) -> Result<Self::Output, Self::Error> {
310        if !self.overwrite
311            && (input.extensions().contains::<ProxyRoute>()
312                || input.extensions().contains::<ProxyRoutes>())
313        {
314            return self.inner.serve(input).await.map_err(Into::into);
315        }
316
317        let proxy_info = self.cached.get_or_init(|| match (self.loader)() {
318            Ok(proxy_info) => Ok(proxy_info),
319            Err(error) => self.load_error_policy.handle_cached(error, None),
320        });
321        let proxy_info = match proxy_info {
322            Ok(proxy_info) => proxy_info,
323            Err(error) => return Err(Box::new(error.clone())),
324        };
325
326        if let Some(proxy_info) = proxy_info {
327            tracing::trace!(
328                server.address = %proxy_info.address.host,
329                server.port = proxy_info.address.port,
330                "setting lazily resolved proxy address",
331            );
332            input
333                .extensions()
334                .insert(ProxyRoute::Proxy(proxy_info.clone()));
335        }
336
337        self.inner.serve(input).await.map_err(Into::into)
338    }
339}
340
341impl<S, Input> Service<Input> for ProxyAddressService<S>
342where
343    S: Service<Input>,
344    Input: ExtensionsRef + Send + 'static,
345{
346    type Output = S::Output;
347    type Error = S::Error;
348
349    fn serve(
350        &self,
351        input: Input,
352    ) -> impl Future<Output = Result<Self::Output, Self::Error>> + Send + '_ {
353        if let Some(ref proxy_info) = self.proxy_info
354            && (self.overwrite
355                || (!input.extensions().contains::<ProxyRoute>()
356                    && !input.extensions().contains::<ProxyRoutes>()))
357        {
358            tracing::trace!(
359                server.address = %proxy_info.address.host,
360                server.port = proxy_info.address.port,
361                "setting proxy address",
362            );
363            input
364                .extensions()
365                .insert(ProxyRoute::Proxy(proxy_info.clone()));
366        }
367        self.inner.serve(input)
368    }
369}
370
371#[cfg(test)]
372mod tests {
373    use std::{
374        convert::Infallible,
375        sync::{
376            Arc,
377            atomic::{AtomicUsize, Ordering},
378        },
379    };
380
381    use parking_lot::Mutex;
382    use rama_core::{Layer as _, Service as _, extensions::Extensions, service::service_fn};
383
384    use super::*;
385
386    #[derive(Debug, Clone)]
387    struct TestInput {
388        extensions: Extensions,
389    }
390
391    impl TestInput {
392        fn new() -> Self {
393            Self {
394                extensions: Extensions::new(),
395            }
396        }
397    }
398
399    impl ExtensionsRef for TestInput {
400        fn extensions(&self) -> &Extensions {
401            &self.extensions
402        }
403    }
404
405    #[tokio::test]
406    async fn preserve_respects_singular_and_collected_route_decisions() {
407        let seen = Arc::new(Mutex::new(Vec::new()));
408        let inner = service_fn({
409            let seen = seen.clone();
410            move |request: TestInput| {
411                seen.lock().push((
412                    request.extensions().contains::<ProxyRoute>(),
413                    request.extensions().contains::<ProxyRoutes>(),
414                ));
415                async { Ok::<_, Infallible>(()) }
416            }
417        });
418        let layer = ProxyAddressLayer::new("http://proxy.example:8080".parse().unwrap())
419            .with_overwrite(false);
420        let service = layer.into_layer(inner);
421
422        let singular = TestInput::new();
423        singular.extensions().insert(ProxyRoute::Direct);
424        service.serve(singular).await.unwrap();
425
426        let collected = TestInput::new();
427        collected
428            .extensions()
429            .insert(ProxyRoutes::from(ProxyRoute::Direct));
430        service.serve(collected).await.unwrap();
431
432        let undecided = TestInput::new();
433        service.serve(undecided).await.unwrap();
434
435        assert_eq!(
436            seen.lock().as_slice(),
437            [(true, false), (false, true), (true, false)]
438        );
439    }
440
441    #[tokio::test]
442    async fn overwrite_replaces_an_authoritative_plural_plan() {
443        let proxy: ProxyAddress = "http://new.proxy:8080".parse().unwrap();
444        let service = ProxyAddressLayer::new(proxy.clone())
445            .with_overwrite(true)
446            .into_layer(
447                crate::client::ProxyRoutesLayer::new().into_layer(service_fn(
448                    |request: TestInput| async move {
449                        let route = request.extensions().get_ref::<ProxyRoute>().cloned();
450                        Ok::<_, Infallible>(route)
451                    },
452                )),
453            );
454        let request = TestInput::new();
455        request
456            .extensions()
457            .insert(ProxyRoutes::new([ProxyRoute::Direct, ProxyRoute::Direct]));
458
459        assert_eq!(
460            service.serve(request).await.unwrap(),
461            Some(ProxyRoute::Proxy(proxy))
462        );
463    }
464
465    #[tokio::test]
466    async fn lazy_loader_skips_preserved_routes_and_shares_cached_result() {
467        let calls = Arc::new(AtomicUsize::new(0));
468        let proxy: ProxyAddress = "http://proxy.example:8080".parse().unwrap();
469        let layer = LazyProxyAddressLayer::new({
470            let calls = calls.clone();
471            let proxy = proxy.clone();
472            move || {
473                calls.fetch_add(1, Ordering::AcqRel);
474                Ok(Some(proxy.clone()))
475            }
476        })
477        .with_overwrite(false);
478
479        let seen = Arc::new(Mutex::new(Vec::new()));
480        let service = layer.into_layer(service_fn({
481            let seen = seen.clone();
482            move |request: TestInput| {
483                seen.lock().push((
484                    request.extensions().get_ref::<ProxyRoute>().cloned(),
485                    request.extensions().contains::<ProxyRoutes>(),
486                ));
487                async { Ok::<_, Infallible>(()) }
488            }
489        }));
490        let cloned_service = service.clone();
491
492        let singular = TestInput::new();
493        singular.extensions().insert(ProxyRoute::Direct);
494        service.serve(singular).await.unwrap();
495
496        let collected = TestInput::new();
497        collected
498            .extensions()
499            .insert(ProxyRoutes::from(ProxyRoute::Direct));
500        service.serve(collected).await.unwrap();
501        assert_eq!(calls.load(Ordering::Acquire), 0);
502
503        service.serve(TestInput::new()).await.unwrap();
504        cloned_service.serve(TestInput::new()).await.unwrap();
505
506        assert_eq!(calls.load(Ordering::Acquire), 1);
507        assert_eq!(
508            seen.lock().as_slice(),
509            [
510                (Some(ProxyRoute::Direct), false),
511                (None, true),
512                (Some(ProxyRoute::Proxy(proxy.clone())), false),
513                (Some(ProxyRoute::Proxy(proxy)), false),
514            ]
515        );
516    }
517
518    #[tokio::test]
519    async fn lazy_loader_caches_absence_and_failure() {
520        let absent_calls = Arc::new(AtomicUsize::new(0));
521        let absent_service = LazyProxyAddressLayer::new({
522            let absent_calls = absent_calls.clone();
523            move || {
524                absent_calls.fetch_add(1, Ordering::AcqRel);
525                Ok(None)
526            }
527        })
528        .into_layer(service_fn(|request: TestInput| async move {
529            Ok::<_, Infallible>(request.extensions().contains::<ProxyRoute>())
530        }));
531
532        assert!(!absent_service.serve(TestInput::new()).await.unwrap());
533        assert!(!absent_service.serve(TestInput::new()).await.unwrap());
534        assert_eq!(absent_calls.load(Ordering::Acquire), 1);
535
536        let error_calls = Arc::new(AtomicUsize::new(0));
537        let error_service = LazyProxyAddressLayer::new({
538            let error_calls = error_calls.clone();
539            move || {
540                error_calls.fetch_add(1, Ordering::AcqRel);
541                Err(std::io::Error::other("invalid proxy environment").into())
542            }
543        })
544        .into_layer(service_fn(|_request: TestInput| async move {
545            Ok::<_, Infallible>(())
546        }));
547
548        for _ in 0..2 {
549            let error = error_service.serve(TestInput::new()).await.unwrap_err();
550            assert_eq!(error.to_string(), "invalid proxy environment");
551        }
552        assert_eq!(error_calls.load(Ordering::Acquire), 1);
553    }
554
555    #[tokio::test]
556    async fn handled_lazy_loader_error_is_sunk_once_and_treated_as_absent() {
557        let loader_calls = Arc::new(AtomicUsize::new(0));
558        let sink_calls = Arc::new(AtomicUsize::new(0));
559        let service = LazyProxyAddressLayer::new({
560            let loader_calls = loader_calls.clone();
561            move || {
562                loader_calls.fetch_add(1, Ordering::AcqRel);
563                Err(std::io::Error::other("invalid proxy environment").into())
564            }
565        })
566        .with_load_error_sink({
567            let sink_calls = sink_calls.clone();
568            move |error: BoxError| {
569                assert_eq!(error.to_string(), "invalid proxy environment");
570                sink_calls.fetch_add(1, Ordering::AcqRel);
571            }
572        })
573        .into_layer(service_fn(|request: TestInput| async move {
574            Ok::<_, Infallible>(request.extensions().contains::<ProxyRoute>())
575        }));
576
577        for _ in 0..2 {
578            assert!(!service.serve(TestInput::new()).await.unwrap());
579        }
580        assert_eq!(loader_calls.load(Ordering::Acquire), 1);
581        assert_eq!(sink_calls.load(Ordering::Acquire), 1);
582    }
583}