reinhardt_server/server/
rate_limit_settings.rs1#![allow(deprecated)] use std::sync::Arc;
13use std::time::Duration;
14
15use reinhardt_core::macros::settings;
16use reinhardt_http::Handler;
17use serde::{Deserialize, Serialize};
18
19use super::rate_limit::{RateLimitConfig, RateLimitHandler, RateLimitStrategy};
20
21fn default_max_requests() -> usize {
24 60
25}
26
27fn default_window_secs() -> u64 {
28 60
29}
30
31fn default_strategy() -> RateLimitStrategyKind {
32 RateLimitStrategyKind::FixedWindow
33}
34
35#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize, Default)]
41#[serde(rename_all = "snake_case")]
42pub enum RateLimitStrategyKind {
43 #[default]
45 FixedWindow,
46 SlidingWindow,
48}
49
50impl From<RateLimitStrategyKind> for RateLimitStrategy {
51 fn from(kind: RateLimitStrategyKind) -> Self {
52 match kind {
53 RateLimitStrategyKind::FixedWindow => RateLimitStrategy::FixedWindow,
54 RateLimitStrategyKind::SlidingWindow => RateLimitStrategy::SlidingWindow,
55 }
56 }
57}
58
59impl From<RateLimitStrategy> for RateLimitStrategyKind {
60 fn from(strategy: RateLimitStrategy) -> Self {
61 match strategy {
62 RateLimitStrategy::FixedWindow => RateLimitStrategyKind::FixedWindow,
63 RateLimitStrategy::SlidingWindow => RateLimitStrategyKind::SlidingWindow,
64 }
65 }
66}
67
68#[settings(fragment = true, section = "server_rate_limit")]
72#[non_exhaustive]
73#[derive(Clone, Debug, Serialize, Deserialize)]
74pub struct RateLimitSettings {
75 #[serde(default = "default_max_requests")]
77 pub max_requests: usize,
78 #[serde(default = "default_window_secs")]
80 pub window_secs: u64,
81 #[serde(default = "default_strategy")]
83 pub strategy: RateLimitStrategyKind,
84 #[serde(default)]
89 pub trusted_proxies: Vec<String>,
90}
91
92impl Default for RateLimitSettings {
93 fn default() -> Self {
94 Self {
95 max_requests: default_max_requests(),
96 window_secs: default_window_secs(),
97 strategy: default_strategy(),
98 trusted_proxies: Vec::new(),
99 }
100 }
101}
102
103impl From<&RateLimitSettings> for RateLimitConfig {
104 fn from(settings: &RateLimitSettings) -> Self {
105 RateLimitConfig::new(
106 settings.max_requests,
107 Duration::from_secs(settings.window_secs),
108 settings.strategy.into(),
109 )
110 .with_trusted_proxies(settings.trusted_proxies.clone())
111 }
112}
113
114pub fn create_rate_limit_handler_from_settings(
117 inner: Arc<dyn Handler>,
118 settings: &RateLimitSettings,
119) -> RateLimitHandler {
120 RateLimitHandler::new(inner, RateLimitConfig::from(settings))
121}
122
123#[cfg(test)]
124mod tests {
125 use super::*;
126 use reinhardt_conf::settings::fragment::SettingsFragment;
127
128 #[rstest::rstest]
129 fn section_name_is_crate_prefixed() {
130 assert_eq!(RateLimitSettings::section(), "server_rate_limit");
132 }
133
134 #[rstest::rstest]
135 fn default_converts_to_config() {
136 let settings = RateLimitSettings::default();
138
139 let config = RateLimitConfig::from(&settings);
141
142 assert_eq!(config.max_requests, 60);
144 assert_eq!(config.window_duration, Duration::from_secs(60));
145 assert_eq!(config.strategy, RateLimitStrategy::FixedWindow);
146 assert!(config.trusted_proxies.is_empty());
147 }
148
149 #[rstest::rstest]
150 fn converts_window_seconds_strategy_and_proxies() {
151 let settings = RateLimitSettings {
153 max_requests: 10,
154 window_secs: 3600,
155 strategy: RateLimitStrategyKind::SlidingWindow,
156 trusted_proxies: vec!["10.0.0.0/8".to_string()],
157 };
158
159 let config = RateLimitConfig::from(&settings);
161
162 assert_eq!(config.max_requests, 10);
164 assert_eq!(config.window_duration, Duration::from_secs(3600));
165 assert_eq!(config.strategy, RateLimitStrategy::SlidingWindow);
166 assert_eq!(config.trusted_proxies, vec!["10.0.0.0/8".to_string()]);
167 }
168
169 #[rstest::rstest]
170 fn deserializes_with_defaults() {
171 let json = r#"{ "max_requests": 100 }"#;
173
174 let settings: RateLimitSettings = serde_json::from_str(json).unwrap();
176 let config = RateLimitConfig::from(&settings);
177
178 assert_eq!(config.max_requests, 100);
180 assert_eq!(config.window_duration, Duration::from_secs(60));
181 assert_eq!(config.strategy, RateLimitStrategy::FixedWindow);
182 assert!(config.trusted_proxies.is_empty());
183 }
184
185 #[rstest::rstest]
186 fn deserializes_strategy_in_snake_case() {
187 let json = r#"{ "strategy": "sliding_window" }"#;
189
190 let settings: RateLimitSettings = serde_json::from_str(json).unwrap();
192
193 assert_eq!(settings.strategy, RateLimitStrategyKind::SlidingWindow);
195 }
196}