1use crate::config::GurtConfig;
2use gurtlib::{prelude::*, GurtMethod, GurtStatusCode};
3use std::collections::HashMap;
4use std::net::IpAddr;
5use std::sync::{Arc, Mutex};
6use std::time::{Duration, Instant};
7use tracing::{warn, debug};
8
9#[derive(Debug)]
10pub struct RateLimitData {
11 requests: Vec<Instant>,
12 connections: u32,
13}
14
15impl RateLimitData {
16 fn new() -> Self {
17 Self {
18 requests: Vec::new(),
19 connections: 0,
20 }
21 }
22
23 fn cleanup_old_requests(&mut self, window: Duration) {
24 let cutoff = Instant::now() - window;
25 self.requests.retain(|&request_time| request_time > cutoff);
26 }
27
28 fn add_request(&mut self) {
29 self.requests.push(Instant::now());
30 }
31
32 fn request_count(&self) -> usize {
33 self.requests.len()
34 }
35
36 fn increment_connections(&mut self) {
37 self.connections += 1;
38 }
39
40 fn decrement_connections(&mut self) {
41 if self.connections > 0 {
42 self.connections -= 1;
43 }
44 }
45
46 fn connection_count(&self) -> u32 {
47 self.connections
48 }
49}
50
51pub struct SecurityMiddleware {
52 config: Arc<GurtConfig>,
53 rate_limit_data: Arc<Mutex<HashMap<IpAddr, RateLimitData>>>,
54}
55
56impl SecurityMiddleware {
57 pub fn new(config: Arc<GurtConfig>) -> Self {
58 Self {
59 config,
60 rate_limit_data: Arc::new(Mutex::new(HashMap::new())),
61 }
62 }
63
64 pub fn is_method_allowed(&self, method: &GurtMethod) -> bool {
65 if let Some(security) = &self.config.security {
66 let method_str = method.to_string();
67 security.allowed_methods.contains(&method_str)
68 } else {
69 true
70 }
71 }
72
73 pub fn check_rate_limit(&self, client_ip: IpAddr) -> bool {
74 if let Some(security) = &self.config.security {
75 let mut data = self.rate_limit_data.lock().unwrap();
76 let rate_data = data.entry(client_ip).or_insert_with(RateLimitData::new);
77
78 rate_data.cleanup_old_requests(Duration::from_secs(60));
79
80 if rate_data.request_count() >= security.rate_limit_requests as usize {
81 warn!("Rate limit exceeded for IP {}: {} requests in the last minute",
82 client_ip, rate_data.request_count());
83 return false;
84 }
85
86 rate_data.add_request();
87 debug!("Request from {}: {}/{} requests in the last minute",
88 client_ip, rate_data.request_count(), security.rate_limit_requests);
89 }
90
91 true
92 }
93
94 pub fn check_connection_limit(&self, client_ip: IpAddr) -> bool {
95 if let Some(security) = &self.config.security {
96 let mut data = self.rate_limit_data.lock().unwrap();
97 let rate_data = data.entry(client_ip).or_insert_with(RateLimitData::new);
98
99 if rate_data.connection_count() >= security.rate_limit_connections {
100 warn!("Connection limit exceeded for IP {}: {} concurrent connections",
101 client_ip, rate_data.connection_count());
102 return false;
103 }
104 }
105
106 true
107 }
108
109 pub fn register_connection(&self, client_ip: IpAddr) {
110 if self.config.security.is_some() {
111 let mut data = self.rate_limit_data.lock().unwrap();
112 let rate_data = data.entry(client_ip).or_insert_with(RateLimitData::new);
113 rate_data.increment_connections();
114 debug!("Connection registered for {}: {} concurrent connections",
115 client_ip, rate_data.connection_count());
116 }
117 }
118
119 pub fn unregister_connection(&self, client_ip: IpAddr) {
120 if self.config.security.is_some() {
121 let mut data = self.rate_limit_data.lock().unwrap();
122 if let Some(rate_data) = data.get_mut(&client_ip) {
123 rate_data.decrement_connections();
124 debug!("Connection unregistered for {}: {} concurrent connections remaining",
125 client_ip, rate_data.connection_count());
126 }
127 }
128 }
129
130 pub fn create_method_not_allowed_response(&self) -> std::result::Result<GurtResponse, GurtError> {
131 let response = GurtResponse::new(GurtStatusCode::MethodNotAllowed)
132 .with_header("Content-Type", "text/html");
133 Ok(response)
134 }
135
136 pub fn create_rate_limit_response(&self) -> std::result::Result<GurtResponse, GurtError> {
137 let response = GurtResponse::new(GurtStatusCode::TooManyRequests)
138 .with_header("Content-Type", "text/html")
139 .with_header("Retry-After", "60");
140 Ok(response)
141 }
142}
143
144#[cfg(test)]
145mod tests {
146 use super::*;
147 use std::net::IpAddr;
148 use std::sync::Arc;
149 use std::time::Duration;
150
151 fn create_test_config() -> Arc<GurtConfig> {
152 let mut config = crate::config::GurtConfig::default();
153 config.security = Some(crate::config::SecurityConfig {
154 deny_files: vec!["*.secret".to_string(), "private/*".to_string()],
155 allowed_methods: vec!["GET".to_string(), "POST".to_string()],
156 rate_limit_requests: 5,
157 rate_limit_connections: 2,
158 });
159 Arc::new(config)
160 }
161
162 #[test]
163 fn test_rate_limit_data_initialization() {
164 let data = RateLimitData::new();
165
166 assert_eq!(data.request_count(), 0);
167 assert_eq!(data.connection_count(), 0);
168 }
169
170 #[test]
171 fn test_rate_limit_data_request_tracking() {
172 let mut data = RateLimitData::new();
173
174 data.add_request();
175 data.add_request();
176 assert_eq!(data.request_count(), 2);
177
178 data.cleanup_old_requests(Duration::from_secs(0));
179 assert_eq!(data.request_count(), 0);
180 }
181
182 #[test]
183 fn test_rate_limit_data_connection_tracking() {
184 let mut data = RateLimitData::new();
185
186 data.increment_connections();
187 data.increment_connections();
188 assert_eq!(data.connection_count(), 2);
189
190 data.decrement_connections();
191 assert_eq!(data.connection_count(), 1);
192
193 data.decrement_connections();
194 data.decrement_connections();
195 assert_eq!(data.connection_count(), 0);
196 }
197
198 #[test]
199 fn test_security_middleware_initialization() {
200 let config = create_test_config();
201 let middleware = SecurityMiddleware::new(config.clone());
202
203 assert!(middleware.rate_limit_data.lock().unwrap().is_empty());
204 }
205
206 #[test]
207 fn test_connection_tracking() {
208 let config = create_test_config();
209 let middleware = SecurityMiddleware::new(config.clone());
210 let ip: IpAddr = "127.0.0.1".parse().unwrap();
211
212 middleware.register_connection(ip);
213 {
214 let data = middleware.rate_limit_data.lock().unwrap();
215 assert_eq!(data.get(&ip).unwrap().connection_count(), 1);
216 }
217
218 middleware.unregister_connection(ip);
219 {
220 let data = middleware.rate_limit_data.lock().unwrap();
221 assert_eq!(data.get(&ip).unwrap().connection_count(), 0);
222 }
223 }
224
225 #[test]
226 fn test_rate_limiting_requests() {
227 let config = create_test_config();
228 let middleware = SecurityMiddleware::new(config.clone());
229 let ip: IpAddr = "127.0.0.1".parse().unwrap();
230
231 for _ in 0..5 {
232 assert!(middleware.check_rate_limit(ip));
233 }
234
235 assert!(!middleware.check_rate_limit(ip));
236 }
237
238 #[test]
239 fn test_connection_limiting() {
240 let config = create_test_config();
241 let middleware = SecurityMiddleware::new(config.clone());
242 let ip: IpAddr = "127.0.0.1".parse().unwrap();
243
244 middleware.register_connection(ip);
245 middleware.register_connection(ip);
246
247 assert!(!middleware.check_connection_limit(ip));
248 }
249
250 #[test]
251 fn test_method_validation() {
252 let config = create_test_config();
253 let middleware = SecurityMiddleware::new(config.clone());
254
255 assert!(middleware.is_method_allowed(&GurtMethod::GET));
256 assert!(middleware.is_method_allowed(&GurtMethod::POST));
257
258 assert!(!middleware.is_method_allowed(&GurtMethod::PUT));
259 assert!(!middleware.is_method_allowed(&GurtMethod::DELETE));
260 }
261
262 #[test]
263 fn test_multiple_ips_isolation() {
264 let config = create_test_config();
265 let middleware = SecurityMiddleware::new(config.clone());
266 let ip1: IpAddr = "127.0.0.1".parse().unwrap();
267 let ip2: IpAddr = "127.0.0.2".parse().unwrap();
268
269 for _ in 0..6 {
270 middleware.check_rate_limit(ip1);
271 }
272
273 assert!(middleware.check_rate_limit(ip2));
274 assert!(!middleware.check_rate_limit(ip1));
275 }
276
277 #[test]
278 fn test_response_creation() {
279 let config = create_test_config();
280 let middleware = SecurityMiddleware::new(config.clone());
281
282 let response = middleware.create_method_not_allowed_response().unwrap();
283 assert_eq!(response.status_code, 405);
284
285 let response = middleware.create_rate_limit_response().unwrap();
286 assert_eq!(response.status_code, 429);
287 }
288}