use std::fmt;
use std::sync::Arc;
use http::{HeaderMap, Response};
use tower_layer::Layer;
use crate::event::LimitEvent;
use crate::gcra::RateLimited;
use crate::identifier::{ClosureIdentifier, TierIdentifier, TierIdentity};
use crate::on_storage_error::OnStorageError;
use crate::service::TierLimitService;
use crate::tier::RateTier;
pub type OnLimitedFn = dyn Fn(&str, &str, &RateLimited) + Send + Sync;
pub type RateLimitedResponseFn = dyn Fn(&str, &str, &RateLimited) -> Response<String> + Send + Sync;
pub type OnEventFn = dyn for<'a> Fn(&LimitEvent<'a>) + Send + Sync;
pub type CostFn = dyn Fn(&http::request::Parts) -> u32 + Send + Sync;
#[derive(Clone, Debug)]
pub struct TierLimitLayer {
pub(crate) rate_tier: Arc<RateTier>,
pub(crate) settings: Settings,
}
#[derive(Clone)]
pub(crate) struct Settings {
pub(crate) identifier: Arc<dyn TierIdentifier>,
pub(crate) on_storage_error: OnStorageError,
pub(crate) on_limited: Option<Arc<OnLimitedFn>>,
pub(crate) rate_limited_response: Option<Arc<RateLimitedResponseFn>>,
pub(crate) on_event: Option<Arc<OnEventFn>>,
pub(crate) cost_fn: Option<Arc<CostFn>>,
}
impl fmt::Debug for Settings {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Settings")
.field("on_storage_error", &self.on_storage_error)
.field("on_limited", &self.on_limited.is_some())
.field(
"rate_limited_response",
&self.rate_limited_response.is_some(),
)
.field("on_event", &self.on_event.is_some())
.field("cost_fn", &self.cost_fn.is_some())
.finish_non_exhaustive()
}
}
struct NoopIdentifier;
impl TierIdentifier for NoopIdentifier {
fn identify(
&self,
_headers: &HeaderMap,
) -> std::pin::Pin<Box<dyn std::future::Future<Output = Option<TierIdentity>> + Send + '_>>
{
Box::pin(std::future::ready(None))
}
}
impl TierLimitLayer {
pub fn new(rate_tier: impl Into<Arc<RateTier>>) -> Self {
Self {
rate_tier: rate_tier.into(),
settings: Settings {
identifier: Arc::new(NoopIdentifier),
on_storage_error: OnStorageError::default(),
on_limited: None,
rate_limited_response: None,
on_event: None,
cost_fn: None,
},
}
}
pub fn identifier(mut self, identifier: impl TierIdentifier) -> Self {
self.settings.identifier = Arc::new(identifier);
self
}
pub fn identifier_fn<F>(mut self, f: F) -> Self
where
F: Fn(&HeaderMap) -> Option<TierIdentity> + Send + Sync + 'static,
{
self.settings.identifier = Arc::new(ClosureIdentifier(f));
self
}
pub fn on_storage_error(mut self, policy: OnStorageError) -> Self {
self.settings.on_storage_error = policy;
self
}
pub fn on_limited(
mut self,
f: impl Fn(&str, &str, &RateLimited) + Send + Sync + 'static,
) -> Self {
self.settings.on_limited = Some(Arc::new(f));
self
}
pub fn rate_limited_response(
mut self,
f: impl Fn(&str, &str, &RateLimited) -> Response<String> + Send + Sync + 'static,
) -> Self {
self.settings.rate_limited_response = Some(Arc::new(f));
self
}
pub fn cost_fn(
mut self,
f: impl Fn(&http::request::Parts) -> u32 + Send + Sync + 'static,
) -> Self {
self.settings.cost_fn = Some(Arc::new(f));
self
}
pub fn on_event(mut self, f: impl Fn(&LimitEvent<'_>) + Send + Sync + 'static) -> Self {
self.settings.on_event = Some(Arc::new(f));
self
}
#[cfg(feature = "buffered-body")]
pub fn buffer_body(self) -> crate::buffered::BufferedTierLimitLayer {
crate::buffered::BufferedTierLimitLayer {
rate_tier: self.rate_tier,
settings: self.settings,
max_body_size: 64 * 1024,
}
}
}
impl<S> Layer<S> for TierLimitLayer {
type Service = TierLimitService<S>;
fn layer(&self, inner: S) -> Self::Service {
TierLimitService {
inner,
rate_tier: self.rate_tier.clone(),
settings: Arc::new(self.settings.clone()),
}
}
}