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)]
19pub struct ProxyAddressLayer {
34 address: Option<ProxyAddress>,
35 overwrite: bool,
36}
37
38impl ProxyAddressLayer {
39 #[must_use]
42 pub fn new(address: ProxyAddress) -> Self {
43 Self::maybe(Some(address))
44 }
45
46 #[must_use]
50 pub fn maybe(address: Option<ProxyAddress>) -> Self {
51 Self {
52 address,
53 ..Default::default()
54 }
55 }
56
57 #[must_use]
59 pub const fn proxy_address(&self) -> Option<&ProxyAddress> {
60 self.address.as_ref()
61 }
62
63 pub fn try_from_env_default() -> Result<Self, BoxError> {
73 Self::try_from_env("http_proxy")
74 }
75
76 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 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#[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 pub const fn new(inner: S, address: ProxyAddress) -> Self {
118 Self::maybe(inner, Some(address))
119 }
120
121 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 pub fn try_from_env_default(inner: S) -> Result<Self, BoxError> {
142 Self::try_from_env(inner, "http_proxy")
143 }
144
145 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 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#[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 #[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 #[must_use]
218 pub fn from_env_default() -> Self {
219 Self::from_env("http_proxy")
220 }
221
222 #[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 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 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#[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}