1#![allow(deprecated)] use reinhardt_http::Handler;
4use reinhardt_http::{Request, Response};
5use std::collections::HashMap;
6use std::net::IpAddr;
7use std::sync::Arc;
8use std::time::{Duration, Instant};
9use tokio::sync::RwLock;
10
11#[derive(Debug, Clone, Copy, PartialEq, Eq)]
13pub enum RateLimitStrategy {
14 FixedWindow,
16 SlidingWindow,
18}
19
20#[deprecated(
22 since = "0.2.0",
23 note = "Use `RateLimitSettings` with the `#[settings]` macro instead."
24)]
25#[derive(Debug, Clone)]
26pub struct RateLimitConfig {
27 pub max_requests: usize,
29 pub window_duration: Duration,
31 pub strategy: RateLimitStrategy,
33 pub trusted_proxies: Vec<String>,
36}
37
38impl RateLimitConfig {
39 pub fn new(
54 max_requests: usize,
55 window_duration: Duration,
56 strategy: RateLimitStrategy,
57 ) -> Self {
58 Self {
59 max_requests,
60 window_duration,
61 strategy,
62 trusted_proxies: Vec::new(),
63 }
64 }
65
66 pub fn per_minute(max_requests: usize) -> Self {
76 Self::new(
77 max_requests,
78 Duration::from_secs(60),
79 RateLimitStrategy::FixedWindow,
80 )
81 }
82
83 pub fn per_hour(max_requests: usize) -> Self {
93 Self::new(
94 max_requests,
95 Duration::from_secs(3600),
96 RateLimitStrategy::FixedWindow,
97 )
98 }
99
100 pub fn with_trusted_proxies(mut self, proxies: Vec<String>) -> Self {
114 self.trusted_proxies = proxies;
115 self
116 }
117}
118
119#[derive(Debug, Clone)]
121struct RateLimitEntry {
122 count: usize,
123 window_start: Instant,
124}
125
126#[derive(Debug, Clone)]
131struct SlidingWindowEntry {
132 timestamps: Vec<Instant>,
133}
134
135pub struct RateLimitHandler {
164 inner: Arc<dyn Handler>,
165 config: RateLimitConfig,
166 limits: Arc<RwLock<HashMap<IpAddr, RateLimitEntry>>>,
167 sliding_limits: Arc<RwLock<HashMap<IpAddr, SlidingWindowEntry>>>,
168}
169
170impl RateLimitHandler {
171 pub fn new(inner: Arc<dyn Handler>, config: RateLimitConfig) -> Self {
200 Self {
201 inner,
202 config,
203 limits: Arc::new(RwLock::new(HashMap::new())),
204 sliding_limits: Arc::new(RwLock::new(HashMap::new())),
205 }
206 }
207
208 async fn is_allowed(&self, ip: IpAddr) -> bool {
213 match self.config.strategy {
214 RateLimitStrategy::FixedWindow => self.is_allowed_fixed_window(ip).await,
215 RateLimitStrategy::SlidingWindow => self.is_allowed_sliding_window(ip).await,
216 }
217 }
218
219 async fn is_allowed_fixed_window(&self, ip: IpAddr) -> bool {
224 let now = Instant::now();
225 let mut limits = self.limits.write().await;
226
227 if limits.len() > 1024 {
230 limits.retain(|_, entry| {
231 now.duration_since(entry.window_start) < self.config.window_duration * 2
232 });
233 }
234
235 let entry = limits.entry(ip).or_insert(RateLimitEntry {
236 count: 0,
237 window_start: now,
238 });
239
240 if now.duration_since(entry.window_start) >= self.config.window_duration {
242 entry.count = 0;
244 entry.window_start = now;
245 }
246
247 if entry.count < self.config.max_requests {
249 entry.count += 1;
250 true
251 } else {
252 false
253 }
254 }
255
256 async fn is_allowed_sliding_window(&self, ip: IpAddr) -> bool {
262 let now = Instant::now();
263 let window = self.config.window_duration;
264 let mut limits = self.sliding_limits.write().await;
265
266 if limits.len() > 1024 {
268 limits.retain(|_, entry| {
269 entry
270 .timestamps
271 .last()
272 .is_some_and(|&last| now.duration_since(last) < window * 2)
273 });
274 }
275
276 let entry = limits.entry(ip).or_insert(SlidingWindowEntry {
277 timestamps: Vec::new(),
278 });
279
280 entry
282 .timestamps
283 .retain(|&ts| now.duration_since(ts) < window);
284
285 if entry.timestamps.len() < self.config.max_requests {
287 entry.timestamps.push(now);
288 true
289 } else {
290 false
291 }
292 }
293
294 fn extract_client_ip(&self, request: &Request) -> IpAddr {
300 let peer_ip = request.remote_addr.map(|addr| addr.ip());
301
302 let from_trusted_proxy = peer_ip.map(|ip| self.is_trusted_proxy(ip)).unwrap_or(false);
304
305 if from_trusted_proxy {
306 if let Some(xff) = request.headers.get("X-Forwarded-For")
308 && let Ok(xff_str) = xff.to_str()
309 && let Some(first_ip) = xff_str.split(',').next()
310 && let Ok(ip) = first_ip.trim().parse()
311 {
312 return ip;
313 }
314
315 if let Some(xri) = request.headers.get("X-Real-IP")
317 && let Ok(ip_str) = xri.to_str()
318 && let Ok(ip) = ip_str.parse()
319 {
320 return ip;
321 }
322 }
323
324 if let Some(ip) = peer_ip {
326 return ip;
327 }
328
329 "127.0.0.1".parse().unwrap()
331 }
332
333 fn is_trusted_proxy(&self, ip: IpAddr) -> bool {
335 self.config.trusted_proxies.iter().any(|proxy| {
336 if let Ok(network) = proxy.parse::<ipnet::IpNet>() {
338 return network.contains(&ip);
339 }
340 if let Ok(proxy_ip) = proxy.parse::<IpAddr>() {
342 return proxy_ip == ip;
343 }
344 false
345 })
346 }
347}
348
349#[async_trait::async_trait]
350impl Handler for RateLimitHandler {
351 async fn handle(&self, request: Request) -> reinhardt_core::exception::Result<Response> {
352 let client_ip = self.extract_client_ip(&request);
353
354 if self.is_allowed(client_ip).await {
355 self.inner.handle(request).await
356 } else {
357 Ok(Response::new(http::StatusCode::TOO_MANY_REQUESTS).with_body("Rate limit exceeded"))
358 }
359 }
360}
361
362#[cfg(test)]
363mod tests {
364 use super::*;
365 use std::time::Duration;
366
367 async fn poll_until<F, Fut>(
369 timeout: std::time::Duration,
370 interval: std::time::Duration,
371 mut condition: F,
372 ) -> Result<(), String>
373 where
374 F: FnMut() -> Fut,
375 Fut: std::future::Future<Output = bool>,
376 {
377 let start = std::time::Instant::now();
378 while start.elapsed() < timeout {
379 if condition().await {
380 return Ok(());
381 }
382 tokio::time::sleep(interval).await;
383 }
384 Err(format!("Timeout after {:?} waiting for condition", timeout))
385 }
386
387 struct TestHandler;
388
389 #[async_trait::async_trait]
390 impl Handler for TestHandler {
391 async fn handle(&self, _request: Request) -> reinhardt_core::exception::Result<Response> {
392 Ok(Response::ok().with_body("Success"))
393 }
394 }
395
396 #[tokio::test]
397 async fn test_rate_limit_config_creation() {
398 let config = RateLimitConfig::per_minute(60);
400
401 assert_eq!(config.max_requests, 60);
403 assert_eq!(config.window_duration, Duration::from_secs(60));
404
405 let config = RateLimitConfig::per_hour(1000);
406 assert_eq!(config.max_requests, 1000);
407 assert_eq!(config.window_duration, Duration::from_secs(3600));
408 }
409
410 #[tokio::test]
411 async fn test_rate_limit_handler_creation() {
412 let handler = Arc::new(TestHandler);
414 let config = RateLimitConfig::per_minute(10);
415 let _rate_limit_handler = RateLimitHandler::new(handler, config);
416 }
417
418 #[tokio::test]
419 async fn test_requests_within_limit() {
420 let handler = Arc::new(TestHandler);
422 let config = RateLimitConfig::per_minute(5);
423 let rate_limit_handler = RateLimitHandler::new(handler, config);
424
425 for _ in 0..5 {
427 let request = Request::builder()
428 .method(http::Method::GET)
429 .uri("/")
430 .version(http::Version::HTTP_11)
431 .headers(http::HeaderMap::new())
432 .body(bytes::Bytes::new())
433 .build()
434 .unwrap();
435
436 let response = rate_limit_handler.handle(request).await.unwrap();
437 assert_eq!(response.status, http::StatusCode::OK);
438 }
439 }
440
441 #[tokio::test]
442 async fn test_requests_exceed_limit() {
443 let handler = Arc::new(TestHandler);
445 let config = RateLimitConfig::per_minute(3);
446 let rate_limit_handler = RateLimitHandler::new(handler, config);
447
448 for _ in 0..3 {
450 let request = Request::builder()
451 .method(http::Method::GET)
452 .uri("/")
453 .version(http::Version::HTTP_11)
454 .headers(http::HeaderMap::new())
455 .body(bytes::Bytes::new())
456 .build()
457 .unwrap();
458
459 let response = rate_limit_handler.handle(request).await.unwrap();
460 assert_eq!(response.status, http::StatusCode::OK);
461 }
462
463 let request = Request::builder()
465 .method(http::Method::GET)
466 .uri("/")
467 .version(http::Version::HTTP_11)
468 .headers(http::HeaderMap::new())
469 .body(bytes::Bytes::new())
470 .build()
471 .unwrap();
472
473 let response = rate_limit_handler.handle(request).await.unwrap();
474
475 assert_eq!(response.status, http::StatusCode::TOO_MANY_REQUESTS);
477 }
478
479 #[tokio::test]
480 async fn test_rate_limit_window_reset() {
481 let handler = Arc::new(TestHandler);
483 let config = RateLimitConfig::new(
484 2,
485 Duration::from_millis(100),
486 RateLimitStrategy::FixedWindow,
487 );
488 let rate_limit_handler = RateLimitHandler::new(handler, config);
489
490 for _ in 0..2 {
492 let request = Request::builder()
493 .method(http::Method::GET)
494 .uri("/")
495 .version(http::Version::HTTP_11)
496 .headers(http::HeaderMap::new())
497 .body(bytes::Bytes::new())
498 .build()
499 .unwrap();
500 let response = rate_limit_handler.handle(request).await.unwrap();
501 assert_eq!(response.status, http::StatusCode::OK);
502 }
503
504 poll_until(
506 Duration::from_millis(200),
507 Duration::from_millis(10),
508 || async {
509 let test_request = Request::builder()
510 .method(http::Method::GET)
511 .uri("/")
512 .version(http::Version::HTTP_11)
513 .headers(http::HeaderMap::new())
514 .body(bytes::Bytes::new())
515 .build()
516 .unwrap();
517 let test_response = rate_limit_handler.handle(test_request).await.unwrap();
518 test_response.status == http::StatusCode::OK
519 },
520 )
521 .await
522 .expect("Window should reset within 200ms");
523 }
524
525 #[tokio::test]
528 async fn test_sliding_window_requests_within_limit() {
529 let handler = Arc::new(TestHandler);
531 let config = RateLimitConfig::new(
532 3,
533 Duration::from_millis(200),
534 RateLimitStrategy::SlidingWindow,
535 );
536 let rate_limit_handler = RateLimitHandler::new(handler, config);
537
538 for _ in 0..3 {
540 let request = Request::builder()
541 .method(http::Method::GET)
542 .uri("/")
543 .version(http::Version::HTTP_11)
544 .headers(http::HeaderMap::new())
545 .body(bytes::Bytes::new())
546 .build()
547 .unwrap();
548
549 let response = rate_limit_handler.handle(request).await.unwrap();
550 assert_eq!(response.status, http::StatusCode::OK);
551 }
552 }
553
554 #[tokio::test]
555 async fn test_sliding_window_requests_exceed_limit() {
556 let handler = Arc::new(TestHandler);
558 let config = RateLimitConfig::new(
559 2,
560 Duration::from_millis(200),
561 RateLimitStrategy::SlidingWindow,
562 );
563 let rate_limit_handler = RateLimitHandler::new(handler, config);
564
565 for _ in 0..2 {
567 let request = Request::builder()
568 .method(http::Method::GET)
569 .uri("/")
570 .version(http::Version::HTTP_11)
571 .headers(http::HeaderMap::new())
572 .body(bytes::Bytes::new())
573 .build()
574 .unwrap();
575 let response = rate_limit_handler.handle(request).await.unwrap();
576 assert_eq!(response.status, http::StatusCode::OK);
577 }
578
579 let request = Request::builder()
581 .method(http::Method::GET)
582 .uri("/")
583 .version(http::Version::HTTP_11)
584 .headers(http::HeaderMap::new())
585 .body(bytes::Bytes::new())
586 .build()
587 .unwrap();
588 let response = rate_limit_handler.handle(request).await.unwrap();
589
590 assert_eq!(response.status, http::StatusCode::TOO_MANY_REQUESTS);
592 }
593
594 #[tokio::test]
595 async fn test_sliding_window_expires_old_requests() {
596 let handler = Arc::new(TestHandler);
598 let config = RateLimitConfig::new(
599 2,
600 Duration::from_millis(100),
601 RateLimitStrategy::SlidingWindow,
602 );
603 let rate_limit_handler = RateLimitHandler::new(handler, config);
604
605 for _ in 0..2 {
607 let request = Request::builder()
608 .method(http::Method::GET)
609 .uri("/")
610 .version(http::Version::HTTP_11)
611 .headers(http::HeaderMap::new())
612 .body(bytes::Bytes::new())
613 .build()
614 .unwrap();
615 let response = rate_limit_handler.handle(request).await.unwrap();
616 assert_eq!(response.status, http::StatusCode::OK);
617 }
618
619 poll_until(
621 Duration::from_millis(200),
622 Duration::from_millis(10),
623 || async {
624 let test_request = Request::builder()
625 .method(http::Method::GET)
626 .uri("/")
627 .version(http::Version::HTTP_11)
628 .headers(http::HeaderMap::new())
629 .body(bytes::Bytes::new())
630 .build()
631 .unwrap();
632 let test_response = rate_limit_handler.handle(test_request).await.unwrap();
633 test_response.status == http::StatusCode::OK
634 },
635 )
636 .await
637 .expect("Sliding window should allow requests after old timestamps expire");
638 }
639
640 #[test]
641 fn test_extract_client_ip_from_trusted_xff() {
642 let handler = Arc::new(TestHandler);
644 let config =
645 RateLimitConfig::per_minute(10).with_trusted_proxies(vec!["10.0.0.1".to_string()]);
646 let rate_limit_handler = RateLimitHandler::new(handler, config);
647
648 let mut headers = http::HeaderMap::new();
649 headers.insert(
650 "X-Forwarded-For",
651 "192.168.1.100, 10.0.0.1, 172.16.0.1".parse().unwrap(),
652 );
653
654 let mut request = Request::builder()
655 .method(http::Method::GET)
656 .uri("/")
657 .version(http::Version::HTTP_11)
658 .headers(headers)
659 .body(bytes::Bytes::new())
660 .build()
661 .unwrap();
662 request.remote_addr = Some("10.0.0.1:12345".parse().unwrap());
663
664 let ip = rate_limit_handler.extract_client_ip(&request);
666
667 assert_eq!(ip, "192.168.1.100".parse::<IpAddr>().unwrap());
669 }
670
671 #[test]
672 fn test_extract_client_ip_ignores_untrusted_xff() {
673 let handler = Arc::new(TestHandler);
675 let config =
676 RateLimitConfig::per_minute(10).with_trusted_proxies(vec!["10.0.0.1".to_string()]);
677 let rate_limit_handler = RateLimitHandler::new(handler, config);
678
679 let mut headers = http::HeaderMap::new();
680 headers.insert("X-Forwarded-For", "192.168.1.100".parse().unwrap());
681
682 let mut request = Request::builder()
683 .method(http::Method::GET)
684 .uri("/")
685 .version(http::Version::HTTP_11)
686 .headers(headers)
687 .body(bytes::Bytes::new())
688 .build()
689 .unwrap();
690 request.remote_addr = Some("203.0.113.42:54321".parse().unwrap());
692
693 let ip = rate_limit_handler.extract_client_ip(&request);
695
696 assert_eq!(ip, "203.0.113.42".parse::<IpAddr>().unwrap());
698 }
699
700 #[test]
701 fn test_extract_client_ip_from_trusted_x_real_ip() {
702 let handler = Arc::new(TestHandler);
704 let config =
705 RateLimitConfig::per_minute(10).with_trusted_proxies(vec!["10.0.0.0/8".to_string()]);
706 let rate_limit_handler = RateLimitHandler::new(handler, config);
707
708 let mut headers = http::HeaderMap::new();
709 headers.insert("X-Real-IP", "203.0.113.42".parse().unwrap());
710
711 let mut request = Request::builder()
712 .method(http::Method::GET)
713 .uri("/")
714 .version(http::Version::HTTP_11)
715 .headers(headers)
716 .body(bytes::Bytes::new())
717 .build()
718 .unwrap();
719 request.remote_addr = Some("10.0.0.5:8080".parse().unwrap());
720
721 let ip = rate_limit_handler.extract_client_ip(&request);
723
724 assert_eq!(ip, "203.0.113.42".parse::<IpAddr>().unwrap());
726 }
727
728 #[test]
729 fn test_extract_client_ip_fallback_to_localhost() {
730 let handler = Arc::new(TestHandler);
732 let config = RateLimitConfig::per_minute(10);
733 let rate_limit_handler = RateLimitHandler::new(handler, config);
734
735 let headers = http::HeaderMap::new();
736 let request = Request::builder()
737 .method(http::Method::GET)
738 .uri("/")
739 .version(http::Version::HTTP_11)
740 .headers(headers)
741 .body(bytes::Bytes::new())
742 .build()
743 .unwrap();
744
745 let ip = rate_limit_handler.extract_client_ip(&request);
747
748 assert_eq!(ip, "127.0.0.1".parse::<IpAddr>().unwrap());
750 }
751
752 #[test]
753 fn test_extract_client_ip_no_trusted_proxies() {
754 let handler = Arc::new(TestHandler);
756 let config = RateLimitConfig::per_minute(10);
757 let rate_limit_handler = RateLimitHandler::new(handler, config);
758
759 let mut headers = http::HeaderMap::new();
760 headers.insert("X-Forwarded-For", "192.168.1.100".parse().unwrap());
761
762 let mut request = Request::builder()
763 .method(http::Method::GET)
764 .uri("/")
765 .version(http::Version::HTTP_11)
766 .headers(headers)
767 .body(bytes::Bytes::new())
768 .build()
769 .unwrap();
770 request.remote_addr = Some("203.0.113.1:8080".parse().unwrap());
771
772 let ip = rate_limit_handler.extract_client_ip(&request);
774
775 assert_eq!(ip, "203.0.113.1".parse::<IpAddr>().unwrap());
777 }
778
779 #[test]
780 fn test_extract_client_ip_with_invalid_header() {
781 let handler = Arc::new(TestHandler);
783 let config =
784 RateLimitConfig::per_minute(10).with_trusted_proxies(vec!["10.0.0.1".to_string()]);
785 let rate_limit_handler = RateLimitHandler::new(handler, config);
786
787 let mut headers = http::HeaderMap::new();
788 headers.insert("X-Forwarded-For", "invalid-ip".parse().unwrap());
789
790 let mut request = Request::builder()
791 .method(http::Method::GET)
792 .uri("/")
793 .version(http::Version::HTTP_11)
794 .headers(headers)
795 .body(bytes::Bytes::new())
796 .build()
797 .unwrap();
798 request.remote_addr = Some("10.0.0.1:8080".parse().unwrap());
799
800 let ip = rate_limit_handler.extract_client_ip(&request);
802
803 assert_eq!(ip, "10.0.0.1".parse::<IpAddr>().unwrap());
805 }
806}