use serde::{Deserialize, Serialize};
pub const DEFAULT_FAILED_LOGIN_MAX_ATTEMPTS: u32 = 10;
pub const DEFAULT_FAILED_LOGIN_LOCKOUT_SECS: u64 = 900;
pub use fraiseql_core::schema::RateLimitingSecurityConfig;
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct RateLimitConfig {
pub enabled: bool,
pub rps_per_ip: u32,
pub rps_per_user: u32,
pub burst_size: u32,
pub cleanup_interval_secs: u64,
pub trust_proxy_headers: bool,
pub trusted_proxy_cidrs: Vec<ipnet::IpNet>,
pub max_buckets: usize,
}
impl Default for RateLimitConfig {
fn default() -> Self {
Self {
enabled: true,
rps_per_ip: 100, rps_per_user: 1000, burst_size: 500, cleanup_interval_secs: 300, trust_proxy_headers: false,
trusted_proxy_cidrs: Vec::new(),
max_buckets: fraiseql_core::schema::DEFAULT_RATE_LIMIT_MAX_BUCKETS,
}
}
}
impl RateLimitConfig {
#[must_use]
pub fn from_security_config(sec: &RateLimitingSecurityConfig) -> Self {
let trusted_proxy_cidrs = parse_trusted_proxy_cidrs(sec).unwrap_or_else(|_| Vec::new());
Self::assemble(sec, trusted_proxy_cidrs)
}
pub fn try_from_security_config(
sec: &RateLimitingSecurityConfig,
) -> std::result::Result<Self, String> {
Ok(Self::assemble(sec, parse_trusted_proxy_cidrs(sec)?))
}
fn assemble(sec: &RateLimitingSecurityConfig, trusted_proxy_cidrs: Vec<ipnet::IpNet>) -> Self {
Self {
enabled: sec.enabled,
rps_per_ip: sec.requests_per_second,
rps_per_user: sec
.requests_per_second_per_user
.unwrap_or_else(|| sec.requests_per_second.saturating_mul(10)),
burst_size: sec.burst_size,
cleanup_interval_secs: 300,
trust_proxy_headers: sec.trust_proxy_headers,
trusted_proxy_cidrs,
max_buckets: sec.max_buckets,
}
}
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct RateLimitOverrides {
pub enabled: Option<bool>,
pub rps_per_ip: Option<u32>,
pub rps_per_user: Option<u32>,
pub burst_size: Option<u32>,
pub max_buckets: Option<usize>,
}
impl RateLimitOverrides {
#[must_use]
pub const fn is_empty(&self) -> bool {
self.enabled.is_none()
&& self.rps_per_ip.is_none()
&& self.rps_per_user.is_none()
&& self.burst_size.is_none()
&& self.max_buckets.is_none()
}
#[must_use]
pub fn enables(&self) -> bool {
self.enabled == Some(true) || (self.enabled.is_none() && !self.is_empty())
}
pub const fn apply_to(&self, config: &mut RateLimitConfig) {
if let Some(v) = self.enabled {
config.enabled = v;
}
if let Some(v) = self.rps_per_ip {
config.rps_per_ip = v;
}
if let Some(v) = self.rps_per_user {
config.rps_per_user = v;
}
if let Some(v) = self.burst_size {
config.burst_size = v;
}
if let Some(v) = self.max_buckets {
config.max_buckets = v;
}
}
}
fn parse_trusted_proxy_cidrs(
sec: &RateLimitingSecurityConfig,
) -> std::result::Result<Vec<ipnet::IpNet>, String> {
sec.trusted_proxy_cidrs
.as_deref()
.unwrap_or(&[])
.iter()
.map(|s| {
s.parse::<ipnet::IpNet>().map_err(|e| {
format!(
"[security.rate_limiting] trusted_proxy_cidrs contains {s:?}, which is not a \
CIDR range ({e}). Every entry must parse, because an entry that is dropped \
leaves the list shorter than it looks — and an empty list means every peer \
is trusted to set X-Forwarded-For."
)
})
})
.collect()
}
#[derive(Debug, Clone)]
pub struct CheckResult {
pub allowed: bool,
pub remaining: f64,
pub retry_after_secs: u32,
}
impl CheckResult {
pub(super) const fn allow(remaining: f64) -> Self {
Self {
allowed: true,
remaining,
retry_after_secs: 0,
}
}
pub(super) const fn deny(retry_after_secs: u32) -> Self {
Self {
allowed: false,
remaining: 0.0,
retry_after_secs,
}
}
}