1use dashmap::DashMap;
4use serde::{Deserialize, Serialize};
5use std::{
6 net::IpAddr,
7 sync::{
8 Arc,
9 atomic::{AtomicBool, Ordering},
10 },
11 time::{Duration, Instant},
12};
13use thiserror::Error;
14
15pub const MAX_TRACKED_CLIENTS: usize = 100_000;
37
38pub const DEFAULT_CLEANUP_INTERVAL: Duration = Duration::from_secs(300);
41
42#[derive(Error, Debug, Clone)]
44pub enum RateLimitError {
45 #[error("Rate limit exceeded: {limit} requests per {window:?}")]
47 LimitExceeded {
48 limit: u32,
50 window: Duration,
52 },
53
54 #[error("Connection limit exceeded: {current}/{max} connections")]
56 ConnectionLimitExceeded {
57 current: usize,
59 max: usize,
61 },
62
63 #[error("Frame size limit exceeded: {size} bytes > {max} bytes")]
65 FrameSizeExceeded {
66 size: usize,
68 max: usize,
70 },
71
72 #[error("Rate limiter at capacity: {max} tracked clients")]
76 CapacityExceeded {
77 max: usize,
79 },
80}
81
82#[derive(Debug, Clone, Serialize, Deserialize)]
84pub struct RateLimitConfig {
85 pub max_requests_per_window: u32,
87 pub window_duration: Duration,
89 pub max_connections_per_ip: usize,
91 pub max_frame_size: usize,
93 pub max_messages_per_second: u32,
95 pub burst_allowance: u32,
97 pub write_timeout: Duration,
115}
116
117impl Default for RateLimitConfig {
118 fn default() -> Self {
119 Self {
120 max_requests_per_window: 100,
121 window_duration: Duration::from_secs(60),
122 max_connections_per_ip: 10,
123 max_frame_size: 1024 * 1024, max_messages_per_second: 30,
125 burst_allowance: 5,
126 write_timeout: Duration::from_secs(10),
127 }
128 }
129}
130
131impl RateLimitConfig {
132 pub fn high_traffic() -> Self {
134 Self {
135 max_requests_per_window: 1000,
136 max_connections_per_ip: 50,
137 max_messages_per_second: 100,
138 burst_allowance: 20,
139 ..Default::default()
140 }
141 }
142
143 pub fn low_resource() -> Self {
145 Self {
146 max_requests_per_window: 20,
147 max_connections_per_ip: 2,
148 max_frame_size: 256 * 1024, max_messages_per_second: 5,
150 burst_allowance: 2,
151 write_timeout: Duration::from_secs(3),
168 ..Default::default()
169 }
170 }
171}
172
173#[derive(Debug)]
175struct ClientRateLimit {
176 requests: Vec<Instant>,
178 connection_count: usize,
180 tokens: f64,
182 last_refill: Instant,
184}
185
186impl ClientRateLimit {
187 fn new(burst_allowance: u32) -> Self {
188 let now = Instant::now();
189 Self {
190 requests: Vec::new(),
191 connection_count: 0,
192 tokens: burst_allowance as f64, last_refill: now,
194 }
195 }
196
197 fn refill_tokens(&mut self, config: &RateLimitConfig) {
199 let now = Instant::now();
200 let time_passed = now.duration_since(self.last_refill).as_secs_f64();
201
202 let tokens_to_add = time_passed * config.max_messages_per_second as f64;
204 let max_tokens = (config.max_messages_per_second + config.burst_allowance) as f64;
205
206 self.tokens = (self.tokens + tokens_to_add).min(max_tokens);
207 self.last_refill = now;
208 }
209
210 fn check_message_rate(&mut self, config: &RateLimitConfig) -> Result<(), RateLimitError> {
212 self.refill_tokens(config);
213
214 if self.tokens >= 1.0 {
215 self.tokens -= 1.0;
216 Ok(())
217 } else {
218 Err(RateLimitError::LimitExceeded {
219 limit: config.max_messages_per_second,
220 window: Duration::from_secs(1),
221 })
222 }
223 }
224}
225
226#[derive(Debug)]
228pub struct WebSocketRateLimiter {
229 config: RateLimitConfig,
230 clients: Arc<DashMap<IpAddr, ClientRateLimit>>,
231 cleanup_spawned: AtomicBool,
243}
244
245impl Default for WebSocketRateLimiter {
246 fn default() -> Self {
247 Self::new(RateLimitConfig::default())
248 }
249}
250
251impl WebSocketRateLimiter {
252 pub fn new(config: RateLimitConfig) -> Self {
254 Self {
255 config,
256 clients: Arc::new(DashMap::new()),
257 cleanup_spawned: AtomicBool::new(false),
258 }
259 }
260
261 pub fn config(&self) -> &RateLimitConfig {
263 &self.config
264 }
265
266 pub fn remaining_for(&self, ip: IpAddr) -> u32 {
284 let Some(client) = self.clients.get(&ip) else {
285 return self.config.max_requests_per_window;
286 };
287
288 let now = Instant::now();
289 let window_start = now.checked_sub(self.config.window_duration);
290 let used = client
291 .requests
292 .iter()
293 .filter(|&&t| window_start.is_none_or(|start| t > start))
294 .count();
295
296 self.config
297 .max_requests_per_window
298 .saturating_sub(used as u32)
299 }
300
301 pub fn reset_after(&self, ip: IpAddr) -> Duration {
321 let Some(client) = self.clients.get(&ip) else {
322 return Duration::ZERO;
323 };
324
325 let now = Instant::now();
326 let Some(window_start) = now.checked_sub(self.config.window_duration) else {
327 return self.config.window_duration;
328 };
329 let earliest_active = client.requests.iter().find(|&&t| t > window_start);
330
331 match earliest_active {
332 Some(&earliest) => self
333 .config
334 .window_duration
335 .saturating_sub(now.saturating_duration_since(earliest)),
336 None => Duration::ZERO,
337 }
338 }
339
340 pub fn spawn_cleanup_task(self: &Arc<Self>, period: Duration) {
354 if self.cleanup_spawned.swap(true, Ordering::AcqRel) {
366 return;
367 }
368
369 let Ok(handle) = tokio::runtime::Handle::try_current() else {
370 self.cleanup_spawned.store(false, Ordering::Release);
374 tracing::warn!(
375 "WebSocketRateLimiter::spawn_cleanup_task: no Tokio runtime available; \
376 periodic cleanup not started"
377 );
378 return;
379 };
380
381 let weak = Arc::downgrade(self);
382 handle.spawn(async move {
383 let mut interval = tokio::time::interval(period);
384 loop {
385 interval.tick().await;
386 let Some(limiter) = weak.upgrade() else {
387 break;
388 };
389 limiter.cleanup_expired();
390 tracing::debug!("WebSocketRateLimiter: cleanup pass completed");
391 }
392 });
393 }
394
395 #[cfg(test)]
401 pub(crate) fn is_cleanup_task_spawned(&self) -> bool {
402 self.cleanup_spawned.load(Ordering::Acquire)
403 }
404
405 pub fn check_request(&self, ip: IpAddr) -> Result<(), RateLimitError> {
407 if !self.clients.contains_key(&ip) && self.clients.len() >= MAX_TRACKED_CLIENTS {
408 return Err(RateLimitError::CapacityExceeded {
409 max: MAX_TRACKED_CLIENTS,
410 });
411 }
412
413 let now = Instant::now();
414 let burst = self.config.burst_allowance;
415 let mut client = self
416 .clients
417 .entry(ip)
418 .or_insert_with(|| ClientRateLimit::new(burst));
419
420 if let Some(window_start) = now.checked_sub(self.config.window_duration) {
446 client.requests.retain(|&time| time > window_start);
447 }
448
449 if client.requests.len() >= self.config.max_requests_per_window as usize {
451 return Err(RateLimitError::LimitExceeded {
452 limit: self.config.max_requests_per_window,
453 window: self.config.window_duration,
454 });
455 }
456
457 client.requests.push(now);
459 Ok(())
460 }
461
462 pub fn check_connection(&self, ip: IpAddr) -> Result<(), RateLimitError> {
464 if !self.clients.contains_key(&ip) && self.clients.len() >= MAX_TRACKED_CLIENTS {
465 return Err(RateLimitError::CapacityExceeded {
466 max: MAX_TRACKED_CLIENTS,
467 });
468 }
469
470 let burst = self.config.burst_allowance;
471 let mut client = self
472 .clients
473 .entry(ip)
474 .or_insert_with(|| ClientRateLimit::new(burst));
475
476 if client.connection_count >= self.config.max_connections_per_ip {
477 return Err(RateLimitError::ConnectionLimitExceeded {
478 current: client.connection_count,
479 max: self.config.max_connections_per_ip,
480 });
481 }
482
483 client.connection_count += 1;
484 Ok(())
485 }
486
487 pub fn close_connection(&self, ip: IpAddr) {
489 if let Some(mut client) = self.clients.get_mut(&ip) {
490 client.connection_count = client.connection_count.saturating_sub(1);
491 }
492 }
493
494 pub fn check_message(&self, ip: IpAddr, frame_size: usize) -> Result<(), RateLimitError> {
496 if frame_size > self.config.max_frame_size {
498 return Err(RateLimitError::FrameSizeExceeded {
499 size: frame_size,
500 max: self.config.max_frame_size,
501 });
502 }
503
504 if let Some(mut client) = self.clients.get_mut(&ip) {
506 client.check_message_rate(&self.config)?;
507 }
508
509 Ok(())
510 }
511
512 pub fn stats(&self) -> RateLimitStats {
514 let mut stats = RateLimitStats::default();
515
516 for entry in self.clients.iter() {
517 stats.total_clients += 1;
518 stats.total_connections += entry.value().connection_count;
519
520 if entry.value().connection_count > 0 {
521 stats.active_clients += 1;
522 }
523 }
524
525 stats
526 }
527
528 pub fn cleanup_expired(&self) {
530 let now = Instant::now();
531 let Some(cutoff) = now.checked_sub(self.config.window_duration * 2) else {
538 return;
539 };
540
541 self.clients.retain(|_, client| {
542 !(client.connection_count == 0
544 && client.requests.last().is_none_or(|&time| time < cutoff))
545 });
546 }
547}
548
549#[derive(Debug, Default, Clone)]
551pub struct RateLimitStats {
552 pub total_clients: usize,
554 pub active_clients: usize,
556 pub total_connections: usize,
558}
559
560#[derive(Debug)]
567pub struct RateLimitGuard {
568 rate_limiter: Arc<WebSocketRateLimiter>,
569 client_ip: IpAddr,
570}
571
572impl RateLimitGuard {
573 pub fn new(
575 rate_limiter: Arc<WebSocketRateLimiter>,
576 client_ip: IpAddr,
577 ) -> Result<Self, RateLimitError> {
578 rate_limiter.check_connection(client_ip)?;
579
580 Ok(Self {
581 rate_limiter,
582 client_ip,
583 })
584 }
585
586 pub fn check_message(&self, frame_size: usize) -> Result<(), RateLimitError> {
588 self.rate_limiter.check_message(self.client_ip, frame_size)
589 }
590}
591
592impl Drop for RateLimitGuard {
593 fn drop(&mut self) {
594 self.rate_limiter.close_connection(self.client_ip);
595 }
596}
597
598#[cfg(test)]
599mod tests {
600 use super::*;
601 use std::net::Ipv4Addr;
602 use std::thread;
603 use std::time::Duration;
604
605 #[test]
606 fn test_rate_limit_requests() {
607 let config = RateLimitConfig {
608 max_requests_per_window: 2,
609 window_duration: Duration::from_millis(100),
610 ..Default::default()
611 };
612
613 let limiter = WebSocketRateLimiter::new(config);
614 let ip = IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1));
615
616 assert!(limiter.check_request(ip).is_ok());
618 assert!(limiter.check_request(ip).is_ok());
619
620 assert!(limiter.check_request(ip).is_err());
622
623 thread::sleep(Duration::from_millis(110));
625
626 assert!(limiter.check_request(ip).is_ok());
628 }
629
630 #[test]
631 fn test_connection_limits() {
632 let config = RateLimitConfig {
633 max_connections_per_ip: 2,
634 ..Default::default()
635 };
636
637 let limiter = WebSocketRateLimiter::new(config);
638 let ip = IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1));
639
640 assert!(limiter.check_connection(ip).is_ok());
642 assert!(limiter.check_connection(ip).is_ok());
643
644 assert!(limiter.check_connection(ip).is_err());
646
647 limiter.close_connection(ip);
649
650 assert!(limiter.check_connection(ip).is_ok());
652 }
653
654 #[test]
655 fn test_message_rate_limiting() {
656 let config = RateLimitConfig {
657 max_messages_per_second: 2,
658 burst_allowance: 2, ..Default::default()
660 };
661
662 let limiter = WebSocketRateLimiter::new(config.clone());
663 let ip = IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1));
664
665 let client = limiter
667 .clients
668 .entry(ip)
669 .or_insert_with(|| ClientRateLimit::new(config.burst_allowance));
670 drop(client);
672
673 assert!(limiter.check_message(ip, 1024).is_ok());
675 assert!(limiter.check_message(ip, 1024).is_ok());
676
677 assert!(limiter.check_message(ip, 1024).is_err());
679 }
680
681 #[test]
682 fn test_frame_size_limits() {
683 let config = RateLimitConfig {
684 max_frame_size: 1024,
685 ..Default::default()
686 };
687
688 let limiter = WebSocketRateLimiter::new(config);
689 let ip = IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1));
690
691 assert!(limiter.check_message(ip, 512).is_ok());
693
694 assert!(limiter.check_message(ip, 2048).is_err());
696 }
697
698 #[test]
699 fn test_rate_limit_guard() {
700 let config = RateLimitConfig {
701 max_connections_per_ip: 1,
702 ..Default::default()
703 };
704
705 let limiter = Arc::new(WebSocketRateLimiter::new(config));
706 let ip = IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1));
707
708 let guard = RateLimitGuard::new(limiter.clone(), ip).unwrap();
710
711 assert!(RateLimitGuard::new(limiter.clone(), ip).is_err());
713
714 drop(guard);
716
717 assert!(RateLimitGuard::new(limiter, ip).is_ok());
719 }
720
721 #[test]
722 fn test_token_refill_over_time() {
723 let config = RateLimitConfig {
724 max_messages_per_second: 1,
725 burst_allowance: 0,
726 window_duration: Duration::from_millis(100),
727 ..Default::default()
728 };
729
730 let limiter = WebSocketRateLimiter::new(config.clone());
731 let ip = IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1));
732
733 {
735 let mut client = limiter
736 .clients
737 .entry(ip)
738 .or_insert_with(|| ClientRateLimit::new(config.burst_allowance));
739 client.tokens = 0.5; }
741
742 assert!(limiter.check_message(ip, 512).is_err());
744
745 thread::sleep(Duration::from_millis(1100));
747
748 let result = limiter.check_message(ip, 512);
750 assert!(result.is_ok(), "Expected refilled tokens to allow message");
752 }
753
754 #[test]
755 fn test_cleanup_expired_entries() {
756 let config = RateLimitConfig {
757 window_duration: Duration::from_millis(100),
758 ..Default::default()
759 };
760
761 let limiter = WebSocketRateLimiter::new(config);
762 let ip1 = IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1));
763 let ip2 = IpAddr::V4(Ipv4Addr::new(127, 0, 0, 2));
764
765 assert!(limiter.check_connection(ip1).is_ok());
767 assert!(limiter.check_connection(ip2).is_ok());
768
769 assert_eq!(limiter.stats().total_clients, 2);
771
772 limiter.close_connection(ip1);
774
775 thread::sleep(Duration::from_millis(250));
777
778 limiter.cleanup_expired();
780
781 let stats = limiter.stats();
783 assert!(stats.total_clients <= 2);
785 }
786
787 #[test]
788 fn test_multiple_ips_isolation() {
789 let config = RateLimitConfig {
790 max_requests_per_window: 1,
791 window_duration: Duration::from_millis(100),
792 ..Default::default()
793 };
794
795 let limiter = WebSocketRateLimiter::new(config);
796 let ip1 = IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1));
797 let ip2 = IpAddr::V4(Ipv4Addr::new(127, 0, 0, 2));
798
799 assert!(limiter.check_request(ip1).is_ok());
801 assert!(limiter.check_request(ip1).is_err());
802
803 assert!(limiter.check_request(ip2).is_ok());
805 assert!(limiter.check_request(ip2).is_err());
806 }
807
808 #[test]
809 fn test_burst_allowance_boundary() {
810 let config = RateLimitConfig {
811 max_messages_per_second: 1,
812 burst_allowance: 0,
813 ..Default::default()
814 };
815
816 let limiter = WebSocketRateLimiter::new(config.clone());
817 let ip = IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1));
818
819 let mut client = limiter
822 .clients
823 .entry(ip)
824 .or_insert_with(|| ClientRateLimit::new(config.burst_allowance));
825 client.tokens = 0.0;
826 drop(client);
827
828 assert!(limiter.check_message(ip, 512).is_err());
830 }
831
832 #[test]
833 fn test_rate_limit_config_high_traffic() {
834 let config = RateLimitConfig::high_traffic();
835
836 assert_eq!(config.max_requests_per_window, 1000);
837 assert_eq!(config.max_connections_per_ip, 50);
838 assert_eq!(config.max_messages_per_second, 100);
839 assert_eq!(config.burst_allowance, 20);
840 assert!(config.max_frame_size >= 1024 * 1024);
841 }
842
843 #[test]
844 fn test_rate_limit_config_low_resource() {
845 let config = RateLimitConfig::low_resource();
846
847 assert_eq!(config.max_requests_per_window, 20);
848 assert_eq!(config.max_connections_per_ip, 2);
849 assert_eq!(config.max_messages_per_second, 5);
850 assert_eq!(config.burst_allowance, 2);
851 assert_eq!(config.max_frame_size, 256 * 1024);
852 assert_eq!(config.write_timeout, Duration::from_secs(3));
853 assert!(config.write_timeout < RateLimitConfig::default().write_timeout);
854 }
855
856 #[test]
857 fn test_frame_size_boundary_exact() {
858 let config = RateLimitConfig {
859 max_frame_size: 1024,
860 ..Default::default()
861 };
862
863 let limiter = WebSocketRateLimiter::new(config);
864 let ip = IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1));
865
866 assert!(limiter.check_message(ip, 1024).is_ok());
868
869 assert!(limiter.check_message(ip, 1025).is_err());
871
872 assert!(limiter.check_message(ip, 0).is_ok());
874 }
875
876 #[test]
877 fn test_stats_accuracy() {
878 let config = RateLimitConfig {
879 max_connections_per_ip: 5,
880 ..Default::default()
881 };
882
883 let limiter = WebSocketRateLimiter::new(config);
884 let ip1 = IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1));
885 let ip2 = IpAddr::V4(Ipv4Addr::new(127, 0, 0, 2));
886
887 assert!(limiter.check_connection(ip1).is_ok());
889 assert!(limiter.check_connection(ip1).is_ok());
890 assert!(limiter.check_connection(ip2).is_ok());
891
892 let stats = limiter.stats();
893 assert_eq!(stats.total_clients, 2);
894 assert_eq!(stats.total_connections, 3);
895 assert_eq!(stats.active_clients, 2);
896
897 limiter.close_connection(ip1);
899
900 let stats = limiter.stats();
901 assert_eq!(stats.total_connections, 2);
902 }
903
904 #[test]
905 fn test_window_duration_respected() {
906 let config = RateLimitConfig {
907 max_requests_per_window: 1,
908 window_duration: Duration::from_millis(50),
909 ..Default::default()
910 };
911
912 let limiter = WebSocketRateLimiter::new(config);
913 let ip = IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1));
914
915 assert!(limiter.check_request(ip).is_ok());
917
918 assert!(limiter.check_request(ip).is_err());
920
921 thread::sleep(Duration::from_millis(60));
923
924 assert!(limiter.check_request(ip).is_ok());
926 }
927
928 #[test]
929 fn test_default_limiter() {
930 let limiter = WebSocketRateLimiter::default();
932 let ip = IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1));
933
934 assert!(limiter.check_request(ip).is_ok());
936 assert!(limiter.check_connection(ip).is_ok());
937
938 let stats = limiter.stats();
940 assert_eq!(stats.total_clients, 1);
941 assert_eq!(stats.total_connections, 1);
942 }
943
944 #[test]
945 fn test_cleanup_expired_removes_inactive_clients() {
946 let config = RateLimitConfig {
947 window_duration: Duration::from_millis(50),
948 ..Default::default()
949 };
950
951 let limiter = WebSocketRateLimiter::new(config);
952 let ip1 = IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1));
953 let ip2 = IpAddr::V4(Ipv4Addr::new(127, 0, 0, 2));
954 let ip3 = IpAddr::V4(Ipv4Addr::new(127, 0, 0, 3));
955
956 assert!(limiter.check_request(ip1).is_ok());
958 assert!(limiter.check_request(ip2).is_ok());
959 assert!(limiter.check_connection(ip3).is_ok());
960
961 let initial_stats = limiter.stats();
962 assert_eq!(initial_stats.total_clients, 3);
963
964 thread::sleep(Duration::from_millis(150));
966
967 limiter.cleanup_expired();
969
970 let after_cleanup = limiter.stats();
971 assert!(after_cleanup.total_clients <= initial_stats.total_clients);
973 }
974
975 #[test]
976 fn test_client_with_zero_connections_and_no_recent_requests_cleaned() {
977 let config = RateLimitConfig {
978 window_duration: Duration::from_millis(100),
979 ..Default::default()
980 };
981
982 let limiter = WebSocketRateLimiter::new(config);
983 let ip = IpAddr::V4(Ipv4Addr::new(192, 168, 1, 100));
984
985 assert!(limiter.check_request(ip).is_ok());
987
988 let initial_stats = limiter.stats();
990 assert_eq!(initial_stats.total_clients, 1);
991
992 thread::sleep(Duration::from_millis(250));
994
995 limiter.cleanup_expired();
997
998 let final_stats = limiter.stats();
999 assert_eq!(final_stats.total_clients, 0);
1001 }
1002
1003 #[test]
1004 fn test_cleanup_preserves_active_clients() {
1005 let config = RateLimitConfig {
1006 window_duration: Duration::from_millis(100),
1007 ..Default::default()
1008 };
1009
1010 let limiter = WebSocketRateLimiter::new(config);
1011 let ip1 = IpAddr::V4(Ipv4Addr::new(10, 0, 0, 1));
1012 let ip2 = IpAddr::V4(Ipv4Addr::new(10, 0, 0, 2));
1013
1014 assert!(limiter.check_connection(ip1).is_ok());
1016
1017 assert!(limiter.check_request(ip2).is_ok());
1019
1020 let initial_stats = limiter.stats();
1021 assert_eq!(initial_stats.total_clients, 2);
1022
1023 thread::sleep(Duration::from_millis(80));
1025
1026 let _ = limiter.check_request(ip2);
1028
1029 limiter.cleanup_expired();
1031
1032 let final_stats = limiter.stats();
1033 assert!(final_stats.total_clients >= 1);
1035 }
1036
1037 #[test]
1038 fn test_close_connection_on_nonexistent_ip() {
1039 let limiter = WebSocketRateLimiter::default();
1040 let ip = IpAddr::V4(Ipv4Addr::new(127, 0, 0, 99));
1041
1042 limiter.close_connection(ip);
1044
1045 let stats = limiter.stats();
1047 assert_eq!(stats.total_clients, 0);
1048 }
1049
1050 #[test]
1051 fn test_check_message_on_nonexistent_client() {
1052 let limiter = WebSocketRateLimiter::default();
1053 let ip = IpAddr::V4(Ipv4Addr::new(127, 0, 0, 88));
1054
1055 assert!(limiter.check_message(ip, 512).is_ok());
1058 }
1059
1060 #[test]
1061 fn test_rate_limit_guard_check_message() {
1062 let config = RateLimitConfig {
1063 max_connections_per_ip: 5,
1064 max_frame_size: 1024,
1065 max_messages_per_second: 10,
1066 burst_allowance: 5,
1067 ..Default::default()
1068 };
1069
1070 let limiter = Arc::new(WebSocketRateLimiter::new(config));
1071 let ip = IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1));
1072
1073 let guard = RateLimitGuard::new(limiter.clone(), ip).unwrap();
1074
1075 assert!(guard.check_message(512).is_ok());
1076 assert!(guard.check_message(512).is_ok());
1077 assert!(guard.check_message(2048).is_err());
1078 }
1079
1080 #[test]
1081 fn test_rate_limit_guard_check_message_rate_limit() {
1082 let config = RateLimitConfig {
1083 max_connections_per_ip: 5,
1084 max_frame_size: 10_000,
1085 max_messages_per_second: 2,
1086 burst_allowance: 2,
1087 ..Default::default()
1088 };
1089
1090 let limiter = Arc::new(WebSocketRateLimiter::new(config));
1091 let ip = IpAddr::V4(Ipv4Addr::new(127, 0, 0, 2));
1092
1093 let guard = RateLimitGuard::new(limiter.clone(), ip).unwrap();
1094
1095 assert!(guard.check_message(512).is_ok());
1096 assert!(guard.check_message(512).is_ok());
1097 assert!(guard.check_message(512).is_err());
1098 }
1099
1100 #[test]
1101 fn test_capacity_cap_rejects_new_clients_when_full() {
1102 let limiter = WebSocketRateLimiter::default();
1103
1104 for i in 0..MAX_TRACKED_CLIENTS as u32 {
1105 let ip = IpAddr::V4(Ipv4Addr::from(i));
1106 limiter.check_request(ip).unwrap();
1107 }
1108 assert_eq!(limiter.stats().total_clients, MAX_TRACKED_CLIENTS);
1109
1110 let overflow_ip = IpAddr::V4(Ipv4Addr::from(MAX_TRACKED_CLIENTS as u32));
1114 let result = limiter.check_request(overflow_ip);
1115 assert!(matches!(
1116 result,
1117 Err(RateLimitError::CapacityExceeded { max }) if max == MAX_TRACKED_CLIENTS
1118 ));
1119 assert_eq!(limiter.stats().total_clients, MAX_TRACKED_CLIENTS);
1120
1121 let existing_ip = IpAddr::V4(Ipv4Addr::from(0u32));
1123 assert!(limiter.check_request(existing_ip).is_ok());
1124 }
1125
1126 #[test]
1127 fn test_cleanup_expired_never_panics_regardless_of_window_duration() {
1128 for window_secs in [1, 60, 3600, u64::MAX / 8] {
1141 let config = RateLimitConfig {
1142 window_duration: Duration::from_secs(window_secs),
1143 ..Default::default()
1144 };
1145 let limiter = WebSocketRateLimiter::new(config);
1146 let ip = IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1));
1147 limiter.check_request(ip).unwrap();
1148
1149 limiter.cleanup_expired(); assert_eq!(
1152 limiter.stats().total_clients,
1153 1,
1154 "a client with a just-now request must survive cleanup regardless \
1155 of window_secs={window_secs}"
1156 );
1157 }
1158 }
1159
1160 #[test]
1161 fn test_check_request_never_panics_regardless_of_window_duration() {
1162 for window_secs in [1, 60, 3600, u64::MAX / 8] {
1173 let config = RateLimitConfig {
1174 window_duration: Duration::from_secs(window_secs),
1175 ..Default::default()
1176 };
1177 let limiter = WebSocketRateLimiter::new(config);
1178 let ip = IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1));
1179
1180 let _ = limiter.check_request(ip); let _ = limiter.remaining_for(ip);
1182 let _ = limiter.reset_after(ip);
1183 }
1184 }
1185
1186 #[test]
1187 fn test_remaining_for_fresh_ip_returns_full_quota() {
1188 let config = RateLimitConfig {
1189 max_requests_per_window: 10,
1190 ..Default::default()
1191 };
1192 let limiter = WebSocketRateLimiter::new(config);
1193 let ip = IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1));
1194
1195 assert_eq!(limiter.remaining_for(ip), 10);
1197 }
1198
1199 #[test]
1200 fn test_remaining_for_decreases_with_consumed_requests() {
1201 let config = RateLimitConfig {
1202 max_requests_per_window: 5,
1203 window_duration: Duration::from_secs(60),
1204 ..Default::default()
1205 };
1206 let limiter = WebSocketRateLimiter::new(config);
1207 let ip = IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1));
1208
1209 assert_eq!(limiter.remaining_for(ip), 5);
1210
1211 limiter.check_request(ip).unwrap();
1212 assert_eq!(limiter.remaining_for(ip), 4);
1213
1214 limiter.check_request(ip).unwrap();
1215 limiter.check_request(ip).unwrap();
1216 assert_eq!(limiter.remaining_for(ip), 2);
1217 }
1218
1219 #[test]
1220 fn test_remaining_for_saturates_at_zero_when_quota_exhausted() {
1221 let config = RateLimitConfig {
1222 max_requests_per_window: 2,
1223 window_duration: Duration::from_secs(60),
1224 ..Default::default()
1225 };
1226 let limiter = WebSocketRateLimiter::new(config);
1227 let ip = IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1));
1228
1229 assert!(limiter.check_request(ip).is_ok());
1232 assert!(limiter.check_request(ip).is_ok());
1233 assert!(limiter.check_request(ip).is_err());
1234
1235 assert_eq!(limiter.remaining_for(ip), 0);
1238 }
1239
1240 #[test]
1241 fn test_remaining_for_isolated_per_ip() {
1242 let config = RateLimitConfig {
1243 max_requests_per_window: 3,
1244 window_duration: Duration::from_secs(60),
1245 ..Default::default()
1246 };
1247 let limiter = WebSocketRateLimiter::new(config);
1248 let ip1 = IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1));
1249 let ip2 = IpAddr::V4(Ipv4Addr::new(127, 0, 0, 2));
1250
1251 limiter.check_request(ip1).unwrap();
1252 limiter.check_request(ip1).unwrap();
1253
1254 assert_eq!(limiter.remaining_for(ip1), 1);
1255 assert_eq!(limiter.remaining_for(ip2), 3);
1256 }
1257
1258 #[test]
1259 fn test_remaining_for_prunes_expired_requests_like_check_request() {
1260 let config = RateLimitConfig {
1267 max_requests_per_window: 5,
1268 window_duration: Duration::from_millis(500),
1269 ..Default::default()
1270 };
1271 let limiter = WebSocketRateLimiter::new(config);
1272 let ip = IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1));
1273
1274 for _ in 0..5 {
1275 limiter.check_request(ip).unwrap();
1276 }
1277 assert_eq!(limiter.remaining_for(ip), 0);
1278
1279 thread::sleep(Duration::from_millis(1000));
1280
1281 assert_eq!(limiter.remaining_for(ip), 5);
1282 assert!(limiter.check_request(ip).is_ok());
1283 }
1284
1285 #[test]
1286 fn test_reset_after_untracked_ip_is_zero() {
1287 let limiter = WebSocketRateLimiter::default();
1288 let ip = IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1));
1289
1290 assert_eq!(limiter.reset_after(ip), Duration::ZERO);
1291 }
1292
1293 #[test]
1294 fn test_reset_after_reflects_oldest_active_request() {
1295 let config = RateLimitConfig {
1298 max_requests_per_window: 5,
1299 window_duration: Duration::from_millis(1000),
1300 ..Default::default()
1301 };
1302 let limiter = WebSocketRateLimiter::new(config);
1303 let ip = IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1));
1304
1305 limiter.check_request(ip).unwrap();
1306 let just_after = limiter.reset_after(ip);
1307 assert!(just_after > Duration::from_millis(750));
1309 assert!(just_after <= Duration::from_millis(1000));
1310
1311 thread::sleep(Duration::from_millis(600));
1312 let later = limiter.reset_after(ip);
1313 assert!(later < just_after);
1315 assert!(later <= Duration::from_millis(400));
1316
1317 thread::sleep(Duration::from_millis(1000));
1318 assert_eq!(limiter.reset_after(ip), Duration::ZERO);
1320 }
1321
1322 #[tokio::test]
1323 async fn test_spawn_cleanup_task_is_idempotent() {
1324 let limiter = Arc::new(WebSocketRateLimiter::new(RateLimitConfig {
1325 window_duration: Duration::from_millis(1),
1326 ..Default::default()
1327 }));
1328
1329 limiter.spawn_cleanup_task(Duration::from_millis(10));
1333 limiter.spawn_cleanup_task(Duration::from_millis(10));
1334
1335 let ip = IpAddr::V4(Ipv4Addr::new(10, 0, 0, 1));
1336 limiter.check_request(ip).unwrap();
1337
1338 tokio::time::sleep(Duration::from_millis(100)).await;
1339
1340 assert_eq!(limiter.stats().total_clients, 0);
1341 }
1342}