use serde::{Deserialize, Serialize};
use crate::config::validation::require_nonzero;
use crate::errors::OrionError;
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct RateLimitConfig {
pub enabled: bool,
#[serde(default = "default_rps")]
pub default_rps: u32,
#[serde(default = "default_burst")]
pub default_burst: u32,
pub trusted_proxies: Vec<String>,
#[serde(default)]
pub endpoints: EndpointRateLimits,
}
fn default_rps() -> u32 {
100
}
fn default_burst() -> u32 {
50
}
impl Default for RateLimitConfig {
fn default() -> Self {
Self {
enabled: false,
default_rps: default_rps(),
default_burst: default_burst(),
trusted_proxies: Vec::new(),
endpoints: EndpointRateLimits::default(),
}
}
}
impl RateLimitConfig {
pub(crate) fn validate(&self) -> Result<(), OrionError> {
if self.enabled {
require_nonzero(
u64::from(self.default_rps),
"rate_limit.default_rps (required when rate limiting is enabled)",
)?;
require_nonzero(
u64::from(self.default_burst),
"rate_limit.default_burst (required when rate limiting is enabled)",
)?;
}
for entry in &self.trusted_proxies {
parse_proxy_entry(entry).map_err(|reason| OrionError::Config {
message: format!("rate_limit.trusted_proxies: invalid entry '{entry}': {reason}"),
})?;
}
Ok(())
}
pub fn parsed_trusted_proxies(&self) -> Vec<ipnet::IpNet> {
self.trusted_proxies
.iter()
.filter_map(|entry| parse_proxy_entry(entry).ok())
.collect()
}
}
fn parse_proxy_entry(entry: &str) -> Result<ipnet::IpNet, &'static str> {
let entry = entry.trim();
if let Ok(net) = entry.parse::<ipnet::IpNet>() {
return Ok(net);
}
entry
.parse::<std::net::IpAddr>()
.map(ipnet::IpNet::from)
.map_err(|_| "expected an IP address or CIDR block (e.g. \"10.0.0.0/8\")")
}
fn default_admin_rps() -> Option<u32> {
Some(20)
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct EndpointRateLimits {
#[serde(default = "default_admin_rps")]
pub admin_rps: Option<u32>,
pub data_rps: Option<u32>,
}
impl Default for EndpointRateLimits {
fn default() -> Self {
Self {
admin_rps: default_admin_rps(),
data_rps: None,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn config_with_proxies(entries: &[&str]) -> RateLimitConfig {
RateLimitConfig {
trusted_proxies: entries.iter().map(|s| s.to_string()).collect(),
..Default::default()
}
}
#[test]
fn test_validate_accepts_cidr_and_bare_ip() {
let config = config_with_proxies(&["10.0.0.0/8", "192.168.1.1", "fd00::/8", "::1"]);
assert!(config.validate().is_ok());
assert_eq!(config.parsed_trusted_proxies().len(), 4);
}
#[test]
fn test_validate_rejects_malformed_entry() {
for bad in ["not-an-ip", "10.0.0.0/33", "10.0.0/8", ""] {
let config = config_with_proxies(&[bad]);
let err = config.validate().expect_err("should reject");
assert!(
err.to_string().contains("trusted_proxies"),
"error for '{bad}' should mention trusted_proxies: {err}"
);
}
}
#[test]
fn test_bare_ip_parses_as_host_network() {
let config = config_with_proxies(&["192.168.1.1"]);
let nets = config.parsed_trusted_proxies();
assert_eq!(nets.len(), 1);
assert_eq!(nets[0].prefix_len(), 32);
}
#[test]
fn test_enabled_on_derived_default_passes_validation() {
let config = RateLimitConfig {
enabled: true,
..Default::default()
};
assert_eq!(config.default_rps, 100);
assert_eq!(config.default_burst, 50);
config
.validate()
.expect("enabling rate limiting without a config file must validate");
}
#[test]
fn test_derived_default_matches_empty_section() {
let from_file: RateLimitConfig =
toml::from_str("").expect("an empty rate_limit section must deserialize");
let derived = RateLimitConfig::default();
assert_eq!(from_file.default_rps, derived.default_rps);
assert_eq!(from_file.default_burst, derived.default_burst);
assert_eq!(from_file.enabled, derived.enabled);
}
#[test]
fn test_default_has_no_trusted_proxies() {
let config = RateLimitConfig::default();
assert!(config.trusted_proxies.is_empty());
assert!(config.parsed_trusted_proxies().is_empty());
assert!(config.validate().is_ok());
}
}