#![allow(deprecated)]
use std::sync::Arc;
use std::time::Duration;
use reinhardt_core::macros::settings;
use reinhardt_http::Handler;
use serde::{Deserialize, Serialize};
use super::rate_limit::{RateLimitConfig, RateLimitHandler, RateLimitStrategy};
fn default_max_requests() -> usize {
60
}
fn default_window_secs() -> u64 {
60
}
fn default_strategy() -> RateLimitStrategyKind {
RateLimitStrategyKind::FixedWindow
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize, Default)]
#[serde(rename_all = "snake_case")]
pub enum RateLimitStrategyKind {
#[default]
FixedWindow,
SlidingWindow,
}
impl From<RateLimitStrategyKind> for RateLimitStrategy {
fn from(kind: RateLimitStrategyKind) -> Self {
match kind {
RateLimitStrategyKind::FixedWindow => RateLimitStrategy::FixedWindow,
RateLimitStrategyKind::SlidingWindow => RateLimitStrategy::SlidingWindow,
}
}
}
impl From<RateLimitStrategy> for RateLimitStrategyKind {
fn from(strategy: RateLimitStrategy) -> Self {
match strategy {
RateLimitStrategy::FixedWindow => RateLimitStrategyKind::FixedWindow,
RateLimitStrategy::SlidingWindow => RateLimitStrategyKind::SlidingWindow,
}
}
}
#[settings(fragment = true, section = "server_rate_limit")]
#[non_exhaustive]
#[derive(Clone, Debug, Serialize, Deserialize)]
pub struct RateLimitSettings {
#[serde(default = "default_max_requests")]
pub max_requests: usize,
#[serde(default = "default_window_secs")]
pub window_secs: u64,
#[serde(default = "default_strategy")]
pub strategy: RateLimitStrategyKind,
#[serde(default)]
pub trusted_proxies: Vec<String>,
}
impl Default for RateLimitSettings {
fn default() -> Self {
Self {
max_requests: default_max_requests(),
window_secs: default_window_secs(),
strategy: default_strategy(),
trusted_proxies: Vec::new(),
}
}
}
impl From<&RateLimitSettings> for RateLimitConfig {
fn from(settings: &RateLimitSettings) -> Self {
RateLimitConfig::new(
settings.max_requests,
Duration::from_secs(settings.window_secs),
settings.strategy.into(),
)
.with_trusted_proxies(settings.trusted_proxies.clone())
}
}
pub fn create_rate_limit_handler_from_settings(
inner: Arc<dyn Handler>,
settings: &RateLimitSettings,
) -> RateLimitHandler {
RateLimitHandler::new(inner, RateLimitConfig::from(settings))
}
#[cfg(test)]
mod tests {
use super::*;
use reinhardt_conf::settings::fragment::SettingsFragment;
#[rstest::rstest]
fn section_name_is_crate_prefixed() {
assert_eq!(RateLimitSettings::section(), "server_rate_limit");
}
#[rstest::rstest]
fn default_converts_to_config() {
let settings = RateLimitSettings::default();
let config = RateLimitConfig::from(&settings);
assert_eq!(config.max_requests, 60);
assert_eq!(config.window_duration, Duration::from_secs(60));
assert_eq!(config.strategy, RateLimitStrategy::FixedWindow);
assert!(config.trusted_proxies.is_empty());
}
#[rstest::rstest]
fn converts_window_seconds_strategy_and_proxies() {
let settings = RateLimitSettings {
max_requests: 10,
window_secs: 3600,
strategy: RateLimitStrategyKind::SlidingWindow,
trusted_proxies: vec!["10.0.0.0/8".to_string()],
};
let config = RateLimitConfig::from(&settings);
assert_eq!(config.max_requests, 10);
assert_eq!(config.window_duration, Duration::from_secs(3600));
assert_eq!(config.strategy, RateLimitStrategy::SlidingWindow);
assert_eq!(config.trusted_proxies, vec!["10.0.0.0/8".to_string()]);
}
#[rstest::rstest]
fn deserializes_with_defaults() {
let json = r#"{ "max_requests": 100 }"#;
let settings: RateLimitSettings = serde_json::from_str(json).unwrap();
let config = RateLimitConfig::from(&settings);
assert_eq!(config.max_requests, 100);
assert_eq!(config.window_duration, Duration::from_secs(60));
assert_eq!(config.strategy, RateLimitStrategy::FixedWindow);
assert!(config.trusted_proxies.is_empty());
}
#[rstest::rstest]
fn deserializes_strategy_in_snake_case() {
let json = r#"{ "strategy": "sliding_window" }"#;
let settings: RateLimitSettings = serde_json::from_str(json).unwrap();
assert_eq!(settings.strategy, RateLimitStrategyKind::SlidingWindow);
}
}