use std::{
fmt,
sync::{Arc, OnceLock},
};
use crate::address::ProxyAddress;
use rama_core::{
Layer, Service, error::BoxError, error_sink::ErrorSink, extensions::ExtensionsRef,
telemetry::tracing,
};
use super::{
ProxyRoute, ProxyRoutes,
env::proxy_address_from_env,
load::{CachedLoadError, LoadErrorPolicy},
};
#[derive(Debug, Clone, Default)]
pub struct ProxyAddressLayer {
address: Option<ProxyAddress>,
overwrite: bool,
}
impl ProxyAddressLayer {
#[must_use]
pub fn new(address: ProxyAddress) -> Self {
Self::maybe(Some(address))
}
#[must_use]
pub fn maybe(address: Option<ProxyAddress>) -> Self {
Self {
address,
..Default::default()
}
}
#[must_use]
pub const fn proxy_address(&self) -> Option<&ProxyAddress> {
self.address.as_ref()
}
pub fn try_from_env_default() -> Result<Self, BoxError> {
Self::try_from_env("http_proxy")
}
pub fn try_from_env(key: impl AsRef<str>) -> Result<Self, BoxError> {
proxy_address_from_env(key.as_ref()).map(Self::maybe)
}
rama_utils::macros::generate_set_and_with! {
pub fn overwrite(mut self, overwrite: bool) -> Self {
self.overwrite = overwrite;
self
}
}
}
impl<S> Layer<S> for ProxyAddressLayer {
type Service = ProxyAddressService<S>;
fn layer(&self, inner: S) -> Self::Service {
ProxyAddressService::maybe(inner, self.address.clone()).with_overwrite(self.overwrite)
}
fn into_layer(self, inner: S) -> Self::Service {
ProxyAddressService::maybe(inner, self.address).with_overwrite(self.overwrite)
}
}
#[derive(Debug, Clone)]
pub struct ProxyAddressService<S> {
inner: S,
proxy_info: Option<ProxyAddress>,
overwrite: bool,
}
impl<S> ProxyAddressService<S> {
pub const fn new(inner: S, address: ProxyAddress) -> Self {
Self::maybe(inner, Some(address))
}
pub const fn maybe(inner: S, address: Option<ProxyAddress>) -> Self {
Self {
inner,
proxy_info: address,
overwrite: false,
}
}
pub fn try_from_env_default(inner: S) -> Result<Self, BoxError> {
Self::try_from_env(inner, "http_proxy")
}
pub fn try_from_env(inner: S, key: impl AsRef<str>) -> Result<Self, BoxError> {
proxy_address_from_env(key.as_ref()).map(|address| Self::maybe(inner, address))
}
rama_utils::macros::generate_set_and_with! {
pub fn overwrite(mut self, overwrite: bool) -> Self {
self.overwrite = overwrite;
self
}
}
}
type ProxyAddressLoader =
dyn Fn() -> Result<Option<ProxyAddress>, BoxError> + Send + Sync + 'static;
type CachedProxyAddress = Result<Option<ProxyAddress>, CachedLoadError>;
#[derive(Clone)]
pub struct LazyProxyAddressLayer {
loader: Arc<ProxyAddressLoader>,
cached: Arc<OnceLock<CachedProxyAddress>>,
load_error_policy: LoadErrorPolicy,
overwrite: bool,
}
impl fmt::Debug for LazyProxyAddressLayer {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("LazyProxyAddressLayer")
.field("cached", &self.cached.get())
.field("load_error_policy", &self.load_error_policy)
.field("overwrite", &self.overwrite)
.finish_non_exhaustive()
}
}
impl LazyProxyAddressLayer {
#[must_use]
pub fn new<F>(loader: F) -> Self
where
F: Fn() -> Result<Option<ProxyAddress>, BoxError> + Send + Sync + 'static,
{
Self {
loader: Arc::new(loader),
cached: Arc::new(OnceLock::new()),
load_error_policy: LoadErrorPolicy::Reject,
overwrite: false,
}
}
#[must_use]
pub fn from_env_default() -> Self {
Self::from_env("http_proxy")
}
#[must_use]
pub fn from_env(key: impl Into<String>) -> Self {
let key = key.into();
Self::new(move || proxy_address_from_env(&key))
}
rama_utils::macros::generate_set_and_with! {
pub fn load_error_sink(
mut self,
sink: impl ErrorSink,
) -> Self {
self.load_error_policy = LoadErrorPolicy::Handle(Arc::new(sink));
self.cached = Arc::new(OnceLock::new());
self
}
}
rama_utils::macros::generate_set_and_with! {
pub fn overwrite(mut self, overwrite: bool) -> Self {
self.overwrite = overwrite;
self
}
}
}
impl<S> Layer<S> for LazyProxyAddressLayer {
type Service = LazyProxyAddressService<S>;
fn layer(&self, inner: S) -> Self::Service {
LazyProxyAddressService {
inner,
loader: self.loader.clone(),
cached: self.cached.clone(),
load_error_policy: self.load_error_policy.clone(),
overwrite: self.overwrite,
}
}
fn into_layer(self, inner: S) -> Self::Service {
LazyProxyAddressService {
inner,
loader: self.loader,
cached: self.cached,
load_error_policy: self.load_error_policy,
overwrite: self.overwrite,
}
}
}
#[derive(Clone)]
pub struct LazyProxyAddressService<S> {
inner: S,
loader: Arc<ProxyAddressLoader>,
cached: Arc<OnceLock<CachedProxyAddress>>,
load_error_policy: LoadErrorPolicy,
overwrite: bool,
}
impl<S: fmt::Debug> fmt::Debug for LazyProxyAddressService<S> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("LazyProxyAddressService")
.field("inner", &self.inner)
.field("cached", &self.cached.get())
.field("load_error_policy", &self.load_error_policy)
.field("overwrite", &self.overwrite)
.finish_non_exhaustive()
}
}
impl<S, Input> Service<Input> for LazyProxyAddressService<S>
where
S: Service<Input, Error: Into<BoxError>>,
Input: ExtensionsRef + Send + 'static,
{
type Output = S::Output;
type Error = BoxError;
async fn serve(&self, input: Input) -> Result<Self::Output, Self::Error> {
if !self.overwrite
&& (input.extensions().contains::<ProxyRoute>()
|| input.extensions().contains::<ProxyRoutes>())
{
return self.inner.serve(input).await.map_err(Into::into);
}
let proxy_info = self.cached.get_or_init(|| match (self.loader)() {
Ok(proxy_info) => Ok(proxy_info),
Err(error) => self.load_error_policy.handle_cached(error, None),
});
let proxy_info = match proxy_info {
Ok(proxy_info) => proxy_info,
Err(error) => return Err(Box::new(error.clone())),
};
if let Some(proxy_info) = proxy_info {
tracing::trace!(
server.address = %proxy_info.address.host,
server.port = proxy_info.address.port,
"setting lazily resolved proxy address",
);
input
.extensions()
.insert(ProxyRoute::Proxy(proxy_info.clone()));
}
self.inner.serve(input).await.map_err(Into::into)
}
}
impl<S, Input> Service<Input> for ProxyAddressService<S>
where
S: Service<Input>,
Input: ExtensionsRef + Send + 'static,
{
type Output = S::Output;
type Error = S::Error;
fn serve(
&self,
input: Input,
) -> impl Future<Output = Result<Self::Output, Self::Error>> + Send + '_ {
if let Some(ref proxy_info) = self.proxy_info
&& (self.overwrite
|| (!input.extensions().contains::<ProxyRoute>()
&& !input.extensions().contains::<ProxyRoutes>()))
{
tracing::trace!(
server.address = %proxy_info.address.host,
server.port = proxy_info.address.port,
"setting proxy address",
);
input
.extensions()
.insert(ProxyRoute::Proxy(proxy_info.clone()));
}
self.inner.serve(input)
}
}
#[cfg(test)]
mod tests {
use std::{
convert::Infallible,
sync::{
Arc,
atomic::{AtomicUsize, Ordering},
},
};
use parking_lot::Mutex;
use rama_core::{Layer as _, Service as _, extensions::Extensions, service::service_fn};
use super::*;
#[derive(Debug, Clone)]
struct TestInput {
extensions: Extensions,
}
impl TestInput {
fn new() -> Self {
Self {
extensions: Extensions::new(),
}
}
}
impl ExtensionsRef for TestInput {
fn extensions(&self) -> &Extensions {
&self.extensions
}
}
#[tokio::test]
async fn preserve_respects_singular_and_collected_route_decisions() {
let seen = Arc::new(Mutex::new(Vec::new()));
let inner = service_fn({
let seen = seen.clone();
move |request: TestInput| {
seen.lock().push((
request.extensions().contains::<ProxyRoute>(),
request.extensions().contains::<ProxyRoutes>(),
));
async { Ok::<_, Infallible>(()) }
}
});
let layer = ProxyAddressLayer::new("http://proxy.example:8080".parse().unwrap())
.with_overwrite(false);
let service = layer.into_layer(inner);
let singular = TestInput::new();
singular.extensions().insert(ProxyRoute::Direct);
service.serve(singular).await.unwrap();
let collected = TestInput::new();
collected
.extensions()
.insert(ProxyRoutes::from(ProxyRoute::Direct));
service.serve(collected).await.unwrap();
let undecided = TestInput::new();
service.serve(undecided).await.unwrap();
assert_eq!(
seen.lock().as_slice(),
[(true, false), (false, true), (true, false)]
);
}
#[tokio::test]
async fn overwrite_replaces_an_authoritative_plural_plan() {
let proxy: ProxyAddress = "http://new.proxy:8080".parse().unwrap();
let service = ProxyAddressLayer::new(proxy.clone())
.with_overwrite(true)
.into_layer(
crate::client::ProxyRoutesLayer::new().into_layer(service_fn(
|request: TestInput| async move {
let route = request.extensions().get_ref::<ProxyRoute>().cloned();
Ok::<_, Infallible>(route)
},
)),
);
let request = TestInput::new();
request
.extensions()
.insert(ProxyRoutes::new([ProxyRoute::Direct, ProxyRoute::Direct]));
assert_eq!(
service.serve(request).await.unwrap(),
Some(ProxyRoute::Proxy(proxy))
);
}
#[tokio::test]
async fn lazy_loader_skips_preserved_routes_and_shares_cached_result() {
let calls = Arc::new(AtomicUsize::new(0));
let proxy: ProxyAddress = "http://proxy.example:8080".parse().unwrap();
let layer = LazyProxyAddressLayer::new({
let calls = calls.clone();
let proxy = proxy.clone();
move || {
calls.fetch_add(1, Ordering::AcqRel);
Ok(Some(proxy.clone()))
}
})
.with_overwrite(false);
let seen = Arc::new(Mutex::new(Vec::new()));
let service = layer.into_layer(service_fn({
let seen = seen.clone();
move |request: TestInput| {
seen.lock().push((
request.extensions().get_ref::<ProxyRoute>().cloned(),
request.extensions().contains::<ProxyRoutes>(),
));
async { Ok::<_, Infallible>(()) }
}
}));
let cloned_service = service.clone();
let singular = TestInput::new();
singular.extensions().insert(ProxyRoute::Direct);
service.serve(singular).await.unwrap();
let collected = TestInput::new();
collected
.extensions()
.insert(ProxyRoutes::from(ProxyRoute::Direct));
service.serve(collected).await.unwrap();
assert_eq!(calls.load(Ordering::Acquire), 0);
service.serve(TestInput::new()).await.unwrap();
cloned_service.serve(TestInput::new()).await.unwrap();
assert_eq!(calls.load(Ordering::Acquire), 1);
assert_eq!(
seen.lock().as_slice(),
[
(Some(ProxyRoute::Direct), false),
(None, true),
(Some(ProxyRoute::Proxy(proxy.clone())), false),
(Some(ProxyRoute::Proxy(proxy)), false),
]
);
}
#[tokio::test]
async fn lazy_loader_caches_absence_and_failure() {
let absent_calls = Arc::new(AtomicUsize::new(0));
let absent_service = LazyProxyAddressLayer::new({
let absent_calls = absent_calls.clone();
move || {
absent_calls.fetch_add(1, Ordering::AcqRel);
Ok(None)
}
})
.into_layer(service_fn(|request: TestInput| async move {
Ok::<_, Infallible>(request.extensions().contains::<ProxyRoute>())
}));
assert!(!absent_service.serve(TestInput::new()).await.unwrap());
assert!(!absent_service.serve(TestInput::new()).await.unwrap());
assert_eq!(absent_calls.load(Ordering::Acquire), 1);
let error_calls = Arc::new(AtomicUsize::new(0));
let error_service = LazyProxyAddressLayer::new({
let error_calls = error_calls.clone();
move || {
error_calls.fetch_add(1, Ordering::AcqRel);
Err(std::io::Error::other("invalid proxy environment").into())
}
})
.into_layer(service_fn(|_request: TestInput| async move {
Ok::<_, Infallible>(())
}));
for _ in 0..2 {
let error = error_service.serve(TestInput::new()).await.unwrap_err();
assert_eq!(error.to_string(), "invalid proxy environment");
}
assert_eq!(error_calls.load(Ordering::Acquire), 1);
}
#[tokio::test]
async fn handled_lazy_loader_error_is_sunk_once_and_treated_as_absent() {
let loader_calls = Arc::new(AtomicUsize::new(0));
let sink_calls = Arc::new(AtomicUsize::new(0));
let service = LazyProxyAddressLayer::new({
let loader_calls = loader_calls.clone();
move || {
loader_calls.fetch_add(1, Ordering::AcqRel);
Err(std::io::Error::other("invalid proxy environment").into())
}
})
.with_load_error_sink({
let sink_calls = sink_calls.clone();
move |error: BoxError| {
assert_eq!(error.to_string(), "invalid proxy environment");
sink_calls.fetch_add(1, Ordering::AcqRel);
}
})
.into_layer(service_fn(|request: TestInput| async move {
Ok::<_, Infallible>(request.extensions().contains::<ProxyRoute>())
}));
for _ in 0..2 {
assert!(!service.serve(TestInput::new()).await.unwrap());
}
assert_eq!(loader_calls.load(Ordering::Acquire), 1);
assert_eq!(sink_calls.load(Ordering::Acquire), 1);
}
}