Skip to main content

gurty_cli/
security.rs

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}