1use std::collections::HashMap;
14use std::sync::{Arc, RwLock};
15use std::time::{Duration, Instant};
16
17pub const DEFAULT_MAX_KEYS: usize = 10_000;
22
23pub trait RateLimiter: Send + Sync {
24 fn acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError>;
25 fn try_acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError>;
26 fn reset(&self, key: &str) -> Result<(), RateLimitError>;
27}
28
29#[derive(Debug, Clone)]
30pub struct RateLimitResult {
31 pub allowed: bool,
32 pub remaining: u64,
33 pub reset_at: i64,
34}
35
36impl RateLimitResult {
37 pub fn allowed(remaining: u64, reset_at: i64) -> Self {
38 Self {
39 allowed: true,
40 remaining,
41 reset_at,
42 }
43 }
44
45 pub fn rejected(remaining: u64, reset_at: i64) -> Self {
46 Self {
47 allowed: false,
48 remaining,
49 reset_at,
50 }
51 }
52}
53
54pub struct SlidingWindowRateLimiter {
55 max_requests: u64,
56 window_size: Duration,
57 entries: Arc<RwLock<HashMap<String, SlidingWindowEntry>>>,
58 max_keys: usize,
60}
61
62#[derive(Clone)]
63struct SlidingWindowEntry {
64 requests: Vec<Instant>,
65}
66
67impl SlidingWindowRateLimiter {
68 pub fn new(max_requests: u64, window_size: Duration) -> Self {
69 Self {
70 max_requests,
71 window_size,
72 entries: Arc::new(RwLock::new(HashMap::new())),
73 max_keys: DEFAULT_MAX_KEYS,
74 }
75 }
76
77 pub fn with_max_keys(mut self, max_keys: usize) -> Self {
82 self.max_keys = max_keys;
83 self
84 }
85
86 fn cleanup_old_requests(&self, entry: &mut SlidingWindowEntry) {
87 let now = Instant::now();
88 entry
89 .requests
90 .retain(|&time| now.duration_since(time) < self.window_size);
91 }
92
93 fn enforce_max_keys(&self, entries: &mut HashMap<String, SlidingWindowEntry>) {
98 while entries.len() > self.max_keys {
99 let now = Instant::now();
101 let oldest_key = entries
102 .iter()
103 .min_by_key(|(_, e)| e.requests.first().copied().unwrap_or(now))
104 .map(|(k, _)| k.clone());
105 match oldest_key {
106 Some(k) => {
107 entries.remove(&k);
108 }
109 None => break,
110 }
111 }
112 }
113}
114
115impl RateLimiter for SlidingWindowRateLimiter {
116 fn acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
117 let mut entries = self
118 .entries
119 .write()
120 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
121
122 if entries.len() >= self.max_keys && !entries.contains_key(key) {
124 self.enforce_max_keys(&mut entries);
125 }
126
127 let entry = entries
128 .entry(key.to_string())
129 .or_insert_with(|| SlidingWindowEntry {
130 requests: Vec::new(),
131 });
132
133 self.cleanup_old_requests(entry);
134
135 if entry.requests.len() < self.max_requests as usize {
136 entry.requests.push(Instant::now());
137 let remaining = self.max_requests - entry.requests.len() as u64;
138 let reset_at = now_timestamp() + self.window_size.as_millis() as i64;
139 Ok(RateLimitResult::allowed(remaining, reset_at))
140 } else {
141 let oldest = entry
142 .requests
143 .first()
144 .map(|t| {
145 let elapsed = t.elapsed().as_millis() as i64;
146 let window_ms = self.window_size.as_millis() as i64;
147 now_timestamp() + (window_ms - elapsed)
148 })
149 .unwrap_or(now_timestamp());
150
151 let remaining = 0;
152 Ok(RateLimitResult::rejected(remaining, oldest))
153 }
154 }
155
156 fn try_acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
157 self.acquire(key)
158 }
159
160 fn reset(&self, key: &str) -> Result<(), RateLimitError> {
161 let mut entries = self
162 .entries
163 .write()
164 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
165 entries.remove(key);
166 Ok(())
167 }
168}
169
170pub struct TokenBucketRateLimiter {
171 capacity: f64,
172 refill_rate: f64,
173 entries: Arc<RwLock<HashMap<String, TokenBucketEntry>>>,
174 max_keys: usize,
176}
177
178#[derive(Clone)]
179struct TokenBucketEntry {
180 tokens: f64,
181 last_refill: Instant,
182}
183
184impl TokenBucketRateLimiter {
185 pub fn new(capacity: u64, refill_per_second: f64) -> Self {
186 Self {
187 capacity: capacity as f64,
188 refill_rate: refill_per_second,
189 entries: Arc::new(RwLock::new(HashMap::new())),
190 max_keys: DEFAULT_MAX_KEYS,
191 }
192 }
193
194 pub fn with_max_keys(mut self, max_keys: usize) -> Self {
196 self.max_keys = max_keys;
197 self
198 }
199
200 fn refill(&self, entry: &mut TokenBucketEntry) {
201 let now = Instant::now();
202 let elapsed = now.duration_since(entry.last_refill).as_secs_f64();
203 let tokens_to_add = if self.refill_rate > 0.0 {
206 elapsed * self.refill_rate
207 } else {
208 0.0
209 };
210
211 entry.tokens = (entry.tokens + tokens_to_add).min(self.capacity);
212 entry.last_refill = now;
213 }
214
215 fn enforce_max_keys(&self, entries: &mut HashMap<String, TokenBucketEntry>) {
220 while entries.len() > self.max_keys {
221 let oldest_key = entries
222 .iter()
223 .min_by_key(|(_, e)| e.last_refill)
224 .map(|(k, _)| k.clone());
225 match oldest_key {
226 Some(k) => {
227 entries.remove(&k);
228 }
229 None => break,
230 }
231 }
232 }
233}
234
235impl RateLimiter for TokenBucketRateLimiter {
236 fn acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
237 let mut entries = self
238 .entries
239 .write()
240 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
241
242 if entries.len() >= self.max_keys && !entries.contains_key(key) {
244 self.enforce_max_keys(&mut entries);
245 }
246
247 let entry = entries
248 .entry(key.to_string())
249 .or_insert_with(|| TokenBucketEntry {
250 tokens: self.capacity,
251 last_refill: Instant::now(),
252 });
253
254 self.refill(entry);
255
256 if entry.tokens >= 1.0 {
257 entry.tokens -= 1.0;
258 let remaining = entry.tokens.floor() as u64;
259 let reset_at = if self.refill_rate > 0.0 {
262 now_timestamp() + (1000.0 / self.refill_rate) as i64
263 } else {
264 i64::MAX
266 };
267 Ok(RateLimitResult::allowed(remaining, reset_at))
268 } else {
269 let reset_at = if self.refill_rate > 0.0 {
271 let wait_time = ((1.0 - entry.tokens) / self.refill_rate * 1000.0) as i64;
272 now_timestamp() + wait_time
273 } else {
274 i64::MAX
276 };
277 Ok(RateLimitResult::rejected(0, reset_at))
278 }
279 }
280
281 fn try_acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
282 self.acquire(key)
283 }
284
285 fn reset(&self, key: &str) -> Result<(), RateLimitError> {
286 let mut entries = self
287 .entries
288 .write()
289 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
290 entries.remove(key);
291 Ok(())
292 }
293}
294
295#[derive(Debug)]
296pub enum RateLimitError {
297 KeyNotFound(String),
298 Internal(String),
299}
300
301impl std::fmt::Display for RateLimitError {
302 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
303 match self {
304 RateLimitError::KeyNotFound(key) => write!(f, "Key not found: {}", key),
305 RateLimitError::Internal(msg) => write!(f, "Internal error: {}", msg),
306 }
307 }
308}
309
310impl std::error::Error for RateLimitError {}
311
312impl serde::Serialize for RateLimitError {
313 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
314 where
315 S: serde::Serializer,
316 {
317 serializer.serialize_str(&self.to_string())
318 }
319}
320
321fn now_timestamp() -> i64 {
322 use std::time::{SystemTime, UNIX_EPOCH};
323 SystemTime::now()
324 .duration_since(UNIX_EPOCH)
325 .unwrap_or_default()
326 .as_millis() as i64
327}
328
329fn now_secs() -> i64 {
330 use std::time::{SystemTime, UNIX_EPOCH};
331 SystemTime::now()
332 .duration_since(UNIX_EPOCH)
333 .unwrap_or_default()
334 .as_secs() as i64
335}
336
337pub struct FixedWindowRateLimiter {
360 max_requests: u64,
361 window_size: Duration,
362 entries: Arc<RwLock<HashMap<String, FixedWindowEntry>>>,
363 max_keys: usize,
364}
365
366#[derive(Clone)]
367struct FixedWindowEntry {
368 count: u64,
369 window_start: Instant,
370}
371
372impl FixedWindowRateLimiter {
373 pub fn new(max_requests: u64, window_size: Duration) -> Self {
378 Self {
379 max_requests,
380 window_size,
381 entries: Arc::new(RwLock::new(HashMap::new())),
382 max_keys: DEFAULT_MAX_KEYS,
383 }
384 }
385
386 pub fn with_max_keys(mut self, max_keys: usize) -> Self {
388 self.max_keys = max_keys;
389 self
390 }
391
392 fn enforce_max_keys(&self, entries: &mut HashMap<String, FixedWindowEntry>) {
394 while entries.len() > self.max_keys {
395 let oldest_key = entries
396 .iter()
397 .min_by_key(|(_, e)| e.window_start)
398 .map(|(k, _)| k.clone());
399 match oldest_key {
400 Some(k) => {
401 entries.remove(&k);
402 }
403 None => break,
404 }
405 }
406 }
407}
408
409impl RateLimiter for FixedWindowRateLimiter {
410 fn acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
411 let mut entries = self
412 .entries
413 .write()
414 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
415
416 if entries.len() >= self.max_keys && !entries.contains_key(key) {
417 self.enforce_max_keys(&mut entries);
418 }
419
420 let now = Instant::now();
421 let entry = entries
422 .entry(key.to_string())
423 .or_insert_with(|| FixedWindowEntry {
424 count: 0,
425 window_start: now,
426 });
427
428 if now.duration_since(entry.window_start) >= self.window_size {
430 entry.count = 0;
431 entry.window_start = now;
432 }
433
434 if entry.count < self.max_requests {
435 entry.count += 1;
436 let remaining = self.max_requests - entry.count;
437 let reset_at = now_timestamp() + self.window_size.as_millis() as i64;
438 Ok(RateLimitResult::allowed(remaining, reset_at))
439 } else {
440 let elapsed = now.duration_since(entry.window_start);
442 let remaining_window = self
443 .window_size
444 .checked_sub(elapsed)
445 .unwrap_or(Duration::ZERO);
446 let reset_at = now_timestamp() + remaining_window.as_millis() as i64;
447 Ok(RateLimitResult::rejected(0, reset_at))
448 }
449 }
450
451 fn try_acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
452 self.acquire(key)
453 }
454
455 fn reset(&self, key: &str) -> Result<(), RateLimitError> {
456 let mut entries = self
457 .entries
458 .write()
459 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
460 entries.remove(key);
461 Ok(())
462 }
463}
464
465pub trait DistributedBackend: Send + Sync {
482 fn incr_and_get(
497 &self,
498 key: &str,
499 window_secs: u64,
500 window_start: i64,
501 max_requests: u64,
502 ) -> Result<(u64, i64), RateLimitError>;
503
504 fn get(&self, key: &str) -> Result<Option<(u64, i64)>, RateLimitError>;
506
507 fn reset_key(&self, key: &str) -> Result<(), RateLimitError>;
509}
510
511pub struct InMemoryBackend {
516 entries: RwLock<HashMap<String, (u64, i64)>>, }
518
519impl InMemoryBackend {
520 pub fn new() -> Self {
521 Self {
522 entries: RwLock::new(HashMap::new()),
523 }
524 }
525}
526
527impl Default for InMemoryBackend {
528 fn default() -> Self {
529 Self::new()
530 }
531}
532
533impl DistributedBackend for InMemoryBackend {
534 fn incr_and_get(
535 &self,
536 key: &str,
537 window_secs: u64,
538 window_start: i64,
539 _max_requests: u64,
540 ) -> Result<(u64, i64), RateLimitError> {
541 let mut entries = self
542 .entries
543 .write()
544 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
545
546 let entry = entries
547 .entry(key.to_string())
548 .or_insert_with(|| (0, window_start));
549
550 if window_start - entry.1 >= window_secs as i64 {
552 *entry = (0, window_start);
554 }
555
556 entry.0 += 1;
557 let reset_at = entry.1 + window_secs as i64;
558 Ok((entry.0, reset_at))
559 }
560
561 fn get(&self, key: &str) -> Result<Option<(u64, i64)>, RateLimitError> {
562 let entries = self
563 .entries
564 .read()
565 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
566 Ok(entries.get(key).copied())
567 }
568
569 fn reset_key(&self, key: &str) -> Result<(), RateLimitError> {
570 let mut entries = self
571 .entries
572 .write()
573 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
574 entries.remove(key);
575 Ok(())
576 }
577}
578
579pub struct DistributedRateLimiter {
584 backend: Arc<dyn DistributedBackend>,
585 max_requests: u64,
586 window_secs: u64,
587}
588
589impl DistributedRateLimiter {
590 pub fn new(backend: Arc<dyn DistributedBackend>, max_requests: u64, window_secs: u64) -> Self {
596 Self {
597 backend,
598 max_requests,
599 window_secs,
600 }
601 }
602
603 pub fn in_memory(max_requests: u64, window_secs: u64) -> Self {
605 Self::new(Arc::new(InMemoryBackend::new()), max_requests, window_secs)
606 }
607}
608
609impl RateLimiter for DistributedRateLimiter {
610 fn acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
611 let window_start = now_secs();
612 let (count, reset_at) =
613 self.backend
614 .incr_and_get(key, self.window_secs, window_start, self.max_requests)?;
615
616 if count <= self.max_requests {
617 let remaining = self.max_requests - count;
618 Ok(RateLimitResult::allowed(remaining, reset_at * 1000))
619 } else {
620 Ok(RateLimitResult::rejected(0, reset_at * 1000))
621 }
622 }
623
624 fn try_acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
625 self.acquire(key)
626 }
627
628 fn reset(&self, key: &str) -> Result<(), RateLimitError> {
629 self.backend.reset_key(key)
630 }
631}
632
633#[derive(Debug, Clone)]
651pub struct RateLimitHeaders {
652 pub limit: u64,
654 pub remaining: u64,
656 pub reset: i64,
658 pub retry_after: Option<u64>,
660}
661
662impl RateLimitHeaders {
663 pub fn from_result(result: &RateLimitResult, limit: u64) -> Self {
668 let reset_secs = result.reset_at / 1000;
669 let now_secs_val = now_secs();
670 let retry_after = if !result.allowed {
671 let diff = reset_secs - now_secs_val;
672 if diff > 0 {
673 Some(diff as u64)
674 } else {
675 Some(1)
676 }
677 } else {
678 None
679 };
680
681 Self {
682 limit,
683 remaining: result.remaining,
684 reset: reset_secs,
685 retry_after,
686 }
687 }
688
689 pub fn to_headers(&self) -> Vec<(String, String)> {
691 let mut headers = vec![
692 ("X-RateLimit-Limit".to_string(), self.limit.to_string()),
693 (
694 "X-RateLimit-Remaining".to_string(),
695 self.remaining.to_string(),
696 ),
697 ("X-RateLimit-Reset".to_string(), self.reset.to_string()),
698 ];
699 if let Some(retry) = self.retry_after {
700 headers.push(("Retry-After".to_string(), retry.to_string()));
701 }
702 headers
703 }
704
705 pub fn to_json(&self) -> serde_json::Value {
707 let mut map = serde_json::json!({
708 "X-RateLimit-Limit": self.limit,
709 "X-RateLimit-Remaining": self.remaining,
710 "X-RateLimit-Reset": self.reset,
711 });
712 if let Some(retry) = self.retry_after {
713 map["Retry-After"] = serde_json::json!(retry);
714 }
715 map
716 }
717}
718
719#[derive(Debug, Clone)]
723pub enum RateLimitResponseStrategy {
724 TooManyRequests,
726 ServiceUnavailable,
728 Custom(u16),
730}
731
732impl RateLimitResponseStrategy {
733 pub fn status_code(&self) -> u16 {
735 match self {
736 RateLimitResponseStrategy::TooManyRequests => 429,
737 RateLimitResponseStrategy::ServiceUnavailable => 503,
738 RateLimitResponseStrategy::Custom(code) => *code,
739 }
740 }
741}
742
743#[derive(Debug, Clone)]
747pub struct RateLimitResponse {
748 pub status_code: u16,
750 pub headers: RateLimitHeaders,
752 pub body: serde_json::Value,
754}
755
756impl RateLimitResponse {
757 pub fn rejected(
763 result: &RateLimitResult,
764 limit: u64,
765 strategy: RateLimitResponseStrategy,
766 ) -> Self {
767 let headers = RateLimitHeaders::from_result(result, limit);
768 let status_code = strategy.status_code();
769 let body = serde_json::json!({
770 "error": "rate_limit_exceeded",
771 "message": "Rate limit exceeded. Please retry later.",
772 "retry_after": headers.retry_after.unwrap_or(1),
773 });
774
775 Self {
776 status_code,
777 headers,
778 body,
779 }
780 }
781
782 pub fn allowed(result: &RateLimitResult, limit: u64) -> Self {
784 let headers = RateLimitHeaders::from_result(result, limit);
785 Self {
786 status_code: 200,
787 headers,
788 body: serde_json::Value::Null,
789 }
790 }
791}
792
793pub struct MultiRateLimiter {
804 limiters: Vec<Arc<dyn RateLimiter>>,
805}
806
807impl MultiRateLimiter {
808 pub fn new(limiters: Vec<Arc<dyn RateLimiter>>) -> Self {
810 Self { limiters }
811 }
812
813 pub fn with_limiter(mut self, limiter: Arc<dyn RateLimiter>) -> Self {
815 self.limiters.push(limiter);
816 self
817 }
818
819 pub fn check_all(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
824 let mut best_result: Option<RateLimitResult> = None;
825 for limiter in &self.limiters {
826 let result = limiter.acquire(key)?;
827 match &best_result {
828 None => best_result = Some(result),
829 Some(current) => {
830 if !result.allowed {
832 if !current.allowed {
834 if result.remaining <= current.remaining {
836 best_result = Some(result);
837 }
838 } else {
839 best_result = Some(result);
841 }
842 } else if current.allowed && result.remaining < current.remaining {
843 best_result = Some(result);
845 }
846 }
847 }
848 }
849 best_result.ok_or_else(|| RateLimitError::Internal("No limiters configured".to_string()))
850 }
851}
852
853#[cfg(test)]
854mod tests {
855 use super::*;
856
857 #[test]
858 fn test_rate_limit_result_allowed() {
859 let result = RateLimitResult::allowed(5, 1000);
860 assert!(result.allowed);
861 assert_eq!(result.remaining, 5);
862 assert_eq!(result.reset_at, 1000);
863 }
864
865 #[test]
866 fn test_rate_limit_result_rejected() {
867 let result = RateLimitResult::rejected(0, 2000);
868 assert!(!result.allowed);
869 assert_eq!(result.remaining, 0);
870 }
871
872 #[test]
873 fn test_sliding_window_limiter_new() {
874 let limiter = SlidingWindowRateLimiter::new(10, Duration::from_secs(60));
875 let result = limiter.acquire("test-key");
876 assert!(result.is_ok());
877 assert!(result.unwrap().allowed);
878 }
879
880 #[test]
881 fn test_sliding_window_limiter_full() {
882 let limiter = SlidingWindowRateLimiter::new(2, Duration::from_secs(60));
883
884 let r1 = limiter.acquire("key1").unwrap();
885 assert!(r1.allowed);
886
887 let r2 = limiter.acquire("key1").unwrap();
888 assert!(r2.allowed);
889
890 let r3 = limiter.acquire("key1").unwrap();
891 assert!(!r3.allowed);
892 }
893
894 #[test]
895 fn test_sliding_window_different_keys() {
896 let limiter = SlidingWindowRateLimiter::new(1, Duration::from_secs(60));
897
898 let r1 = limiter.acquire("key-a").unwrap();
899 assert!(r1.allowed);
900
901 let r2 = limiter.acquire("key-b").unwrap();
902 assert!(r2.allowed);
903 }
904
905 #[test]
906 fn test_sliding_window_reset() {
907 let limiter = SlidingWindowRateLimiter::new(1, Duration::from_secs(60));
908
909 limiter.acquire("reset-key").unwrap();
910 limiter.acquire("reset-key").unwrap();
911
912 limiter.reset("reset-key").unwrap();
913
914 let result = limiter.acquire("reset-key").unwrap();
915 assert!(result.allowed);
916 }
917
918 #[test]
919 fn test_token_bucket_limiter_new() {
920 let limiter = TokenBucketRateLimiter::new(10, 1.0);
921 let result = limiter.acquire("test-key");
922 assert!(result.is_ok());
923 assert!(result.unwrap().allowed);
924 }
925
926 #[test]
927 fn test_token_bucket_limiter_depletes() {
928 let limiter = TokenBucketRateLimiter::new(2, 1.0);
929
930 let r1 = limiter.acquire("key1").unwrap();
931 assert!(r1.allowed);
932 assert_eq!(r1.remaining, 1);
933
934 let r2 = limiter.acquire("key1").unwrap();
935 assert!(r2.allowed);
936 assert_eq!(r2.remaining, 0);
937
938 let r3 = limiter.acquire("key1").unwrap();
939 assert!(!r3.allowed);
940 }
941
942 #[test]
943 fn test_token_bucket_different_keys() {
944 let limiter = TokenBucketRateLimiter::new(1, 1.0);
945
946 let r1 = limiter.acquire("key-a").unwrap();
947 assert!(r1.allowed);
948
949 let r2 = limiter.acquire("key-b").unwrap();
950 assert!(r2.allowed);
951 }
952
953 #[test]
954 fn test_token_bucket_reset() {
955 let limiter = TokenBucketRateLimiter::new(1, 1.0);
956
957 limiter.acquire("reset-key").unwrap();
958 limiter.acquire("reset-key").unwrap();
959
960 limiter.reset("reset-key").unwrap();
961
962 let result = limiter.acquire("reset-key").unwrap();
963 assert!(result.allowed);
964 }
965
966 #[test]
967 fn test_limiter_try_acquire() {
968 let limiter = SlidingWindowRateLimiter::new(1, Duration::from_secs(60));
969
970 let r1 = limiter.try_acquire("key").unwrap();
971 assert!(r1.allowed);
972
973 let r2 = limiter.try_acquire("key").unwrap();
974 assert!(!r2.allowed);
975 }
976
977 #[test]
980 fn test_token_bucket_zero_refill_rate_does_not_panic() {
981 let limiter = TokenBucketRateLimiter::new(1, 0.0);
984
985 let r1 = limiter.acquire("zero-refill").unwrap();
986 assert!(r1.allowed, "first acquire should be allowed");
987
988 let r2 = limiter.acquire("zero-refill").unwrap();
990 assert!(!r2.allowed, "second acquire should be rejected");
991 assert!(
993 r2.reset_at > 0,
994 "reset_at should be a valid timestamp, got: {}",
995 r2.reset_at
996 );
997 }
998
999 #[test]
1000 fn test_token_bucket_negative_refill_rate_does_not_panic() {
1001 let limiter = TokenBucketRateLimiter::new(1, -1.0);
1003
1004 let r1 = limiter.acquire("neg-refill").unwrap();
1005 assert!(r1.allowed, "first acquire should be allowed");
1006
1007 let r2 = limiter.acquire("neg-refill").unwrap();
1008 assert!(!r2.allowed, "second acquire should be rejected");
1009 assert!(
1010 r2.reset_at > 0,
1011 "reset_at should be a valid timestamp, got: {}",
1012 r2.reset_at
1013 );
1014 }
1015
1016 #[test]
1019 fn test_fixed_window_limiter_allows_within_limit() {
1020 let limiter = FixedWindowRateLimiter::new(5, Duration::from_secs(60));
1021 for i in 0..5 {
1022 let r = limiter.acquire("key").unwrap();
1023 assert!(r.allowed, "request {} should be allowed", i);
1024 }
1025 }
1026
1027 #[test]
1028 fn test_fixed_window_limiter_rejects_over_limit() {
1029 let limiter = FixedWindowRateLimiter::new(2, Duration::from_secs(60));
1030 assert!(limiter.acquire("key").unwrap().allowed);
1031 assert!(limiter.acquire("key").unwrap().allowed);
1032 let r3 = limiter.acquire("key").unwrap();
1033 assert!(!r3.allowed);
1034 assert_eq!(r3.remaining, 0);
1035 }
1036
1037 #[test]
1038 fn test_fixed_window_limiter_different_keys() {
1039 let limiter = FixedWindowRateLimiter::new(1, Duration::from_secs(60));
1040 assert!(limiter.acquire("key-a").unwrap().allowed);
1041 assert!(limiter.acquire("key-b").unwrap().allowed);
1042 }
1043
1044 #[test]
1045 fn test_fixed_window_limiter_reset() {
1046 let limiter = FixedWindowRateLimiter::new(1, Duration::from_secs(60));
1047 limiter.acquire("key").unwrap();
1048 assert!(!limiter.acquire("key").unwrap().allowed);
1049 limiter.reset("key").unwrap();
1050 assert!(limiter.acquire("key").unwrap().allowed);
1051 }
1052
1053 #[test]
1054 fn test_fixed_window_limiter_remaining_decreases() {
1055 let limiter = FixedWindowRateLimiter::new(3, Duration::from_secs(60));
1056 let r1 = limiter.acquire("key").unwrap();
1057 assert_eq!(r1.remaining, 2);
1058 let r2 = limiter.acquire("key").unwrap();
1059 assert_eq!(r2.remaining, 1);
1060 let r3 = limiter.acquire("key").unwrap();
1061 assert_eq!(r3.remaining, 0);
1062 }
1063
1064 #[test]
1065 fn test_fixed_window_limiter_reset_at_positive() {
1066 let limiter = FixedWindowRateLimiter::new(1, Duration::from_secs(60));
1067 let r = limiter.acquire("key").unwrap();
1068 assert!(r.reset_at > 0);
1069 }
1070
1071 #[test]
1072 fn test_fixed_window_limiter_try_acquire() {
1073 let limiter = FixedWindowRateLimiter::new(1, Duration::from_secs(60));
1074 assert!(limiter.try_acquire("key").unwrap().allowed);
1075 assert!(!limiter.try_acquire("key").unwrap().allowed);
1076 }
1077
1078 #[test]
1081 fn test_in_memory_backend_new() {
1082 let backend = InMemoryBackend::new();
1083 let result = backend.get("key").unwrap();
1084 assert!(result.is_none());
1085 }
1086
1087 #[test]
1088 fn test_in_memory_backend_incr_and_get() {
1089 let backend = InMemoryBackend::new();
1090 let now = now_secs();
1091 let (count1, reset1) = backend.incr_and_get("key", 60, now, 10).unwrap();
1092 assert_eq!(count1, 1);
1093 assert_eq!(reset1, now + 60);
1094
1095 let (count2, _) = backend.incr_and_get("key", 60, now, 10).unwrap();
1096 assert_eq!(count2, 2);
1097 }
1098
1099 #[test]
1100 fn test_in_memory_backend_get() {
1101 let backend = InMemoryBackend::new();
1102 let now = now_secs();
1103 backend.incr_and_get("key", 60, now, 10).unwrap();
1104 let result = backend.get("key").unwrap();
1105 assert!(result.is_some());
1106 assert_eq!(result.unwrap().0, 1);
1107 }
1108
1109 #[test]
1110 fn test_in_memory_backend_reset_key() {
1111 let backend = InMemoryBackend::new();
1112 let now = now_secs();
1113 backend.incr_and_get("key", 60, now, 10).unwrap();
1114 assert!(backend.get("key").unwrap().is_some());
1115 backend.reset_key("key").unwrap();
1116 assert!(backend.get("key").unwrap().is_none());
1117 }
1118
1119 #[test]
1120 fn test_in_memory_backend_window_expiry() {
1121 let backend = InMemoryBackend::new();
1122 let now = now_secs();
1123 backend.incr_and_get("key", 60, now, 10).unwrap();
1125 backend.incr_and_get("key", 60, now, 10).unwrap();
1126 let (count, _) = backend.incr_and_get("key", 60, now + 61, 10).unwrap();
1128 assert_eq!(count, 1);
1129 }
1130
1131 #[test]
1132 fn test_distributed_rate_limiter_allows() {
1133 let limiter = DistributedRateLimiter::in_memory(5, 60);
1134 for i in 0..5 {
1135 let r = limiter.acquire("key").unwrap();
1136 assert!(r.allowed, "request {} should be allowed", i);
1137 }
1138 }
1139
1140 #[test]
1141 fn test_distributed_rate_limiter_rejects() {
1142 let limiter = DistributedRateLimiter::in_memory(2, 60);
1143 assert!(limiter.acquire("key").unwrap().allowed);
1144 assert!(limiter.acquire("key").unwrap().allowed);
1145 assert!(!limiter.acquire("key").unwrap().allowed);
1146 }
1147
1148 #[test]
1149 fn test_distributed_rate_limiter_shared_backend() {
1150 let backend = Arc::new(InMemoryBackend::new());
1152 let limiter1 = DistributedRateLimiter::new(backend.clone(), 2, 60);
1153 let limiter2 = DistributedRateLimiter::new(backend.clone(), 2, 60);
1154
1155 assert!(limiter1.acquire("key").unwrap().allowed);
1157 assert!(limiter2.acquire("key").unwrap().allowed);
1159 assert!(!limiter1.acquire("key").unwrap().allowed);
1161 }
1162
1163 #[test]
1164 fn test_distributed_rate_limiter_reset() {
1165 let limiter = DistributedRateLimiter::in_memory(1, 60);
1166 limiter.acquire("key").unwrap();
1167 assert!(!limiter.acquire("key").unwrap().allowed);
1168 limiter.reset("key").unwrap();
1169 assert!(limiter.acquire("key").unwrap().allowed);
1170 }
1171
1172 #[test]
1173 fn test_distributed_rate_limiter_different_keys() {
1174 let limiter = DistributedRateLimiter::in_memory(1, 60);
1175 assert!(limiter.acquire("key-a").unwrap().allowed);
1176 assert!(limiter.acquire("key-b").unwrap().allowed);
1177 }
1178
1179 #[test]
1180 fn test_distributed_rate_limiter_remaining() {
1181 let limiter = DistributedRateLimiter::in_memory(3, 60);
1182 let r1 = limiter.acquire("key").unwrap();
1183 assert_eq!(r1.remaining, 2);
1184 let r2 = limiter.acquire("key").unwrap();
1185 assert_eq!(r2.remaining, 1);
1186 }
1187
1188 #[test]
1191 fn test_rate_limit_headers_allowed() {
1192 let result = RateLimitResult::allowed(5, now_secs() * 1000 + 60000);
1193 let headers = RateLimitHeaders::from_result(&result, 10);
1194 assert_eq!(headers.limit, 10);
1195 assert_eq!(headers.remaining, 5);
1196 assert!(headers.retry_after.is_none());
1197 }
1198
1199 #[test]
1200 fn test_rate_limit_headers_rejected() {
1201 let result = RateLimitResult::rejected(0, now_secs() * 1000 + 60000);
1202 let headers = RateLimitHeaders::from_result(&result, 10);
1203 assert_eq!(headers.limit, 10);
1204 assert_eq!(headers.remaining, 0);
1205 assert!(headers.retry_after.is_some());
1206 assert!(headers.retry_after.unwrap() > 0);
1207 }
1208
1209 #[test]
1210 fn test_rate_limit_headers_to_headers_allowed() {
1211 let result = RateLimitResult::allowed(5, now_secs() * 1000 + 60000);
1212 let headers = RateLimitHeaders::from_result(&result, 10);
1213 let hdrs = headers.to_headers();
1214 assert_eq!(hdrs.len(), 3); assert!(hdrs
1216 .iter()
1217 .any(|(k, v)| k == "X-RateLimit-Limit" && v == "10"));
1218 assert!(hdrs
1219 .iter()
1220 .any(|(k, v)| k == "X-RateLimit-Remaining" && v == "5"));
1221 }
1222
1223 #[test]
1224 fn test_rate_limit_headers_to_headers_rejected() {
1225 let result = RateLimitResult::rejected(0, now_secs() * 1000 + 60000);
1226 let headers = RateLimitHeaders::from_result(&result, 10);
1227 let hdrs = headers.to_headers();
1228 assert_eq!(hdrs.len(), 4); assert!(hdrs.iter().any(|(k, _)| k == "Retry-After"));
1230 }
1231
1232 #[test]
1233 fn test_rate_limit_headers_to_json() {
1234 let result = RateLimitResult::allowed(5, now_secs() * 1000 + 60000);
1235 let headers = RateLimitHeaders::from_result(&result, 10);
1236 let json = headers.to_json();
1237 assert_eq!(json["X-RateLimit-Limit"], 10);
1238 assert_eq!(json["X-RateLimit-Remaining"], 5);
1239 assert!(json.get("Retry-After").is_none());
1240 }
1241
1242 #[test]
1243 fn test_rate_limit_headers_to_json_rejected() {
1244 let result = RateLimitResult::rejected(0, now_secs() * 1000 + 60000);
1245 let headers = RateLimitHeaders::from_result(&result, 10);
1246 let json = headers.to_json();
1247 assert!(json.get("Retry-After").is_some());
1248 }
1249
1250 #[test]
1251 fn test_rate_limit_response_strategy_status_codes() {
1252 assert_eq!(
1253 RateLimitResponseStrategy::TooManyRequests.status_code(),
1254 429
1255 );
1256 assert_eq!(
1257 RateLimitResponseStrategy::ServiceUnavailable.status_code(),
1258 503
1259 );
1260 assert_eq!(RateLimitResponseStrategy::Custom(502).status_code(), 502);
1261 }
1262
1263 #[test]
1264 fn test_rate_limit_response_rejected() {
1265 let result = RateLimitResult::rejected(0, now_secs() * 1000 + 60000);
1266 let response =
1267 RateLimitResponse::rejected(&result, 10, RateLimitResponseStrategy::TooManyRequests);
1268 assert_eq!(response.status_code, 429);
1269 assert_eq!(response.body["error"], "rate_limit_exceeded");
1270 assert!(response.body["retry_after"].as_u64().unwrap() > 0);
1271 assert!(response.headers.retry_after.is_some());
1272 }
1273
1274 #[test]
1275 fn test_rate_limit_response_allowed() {
1276 let result = RateLimitResult::allowed(5, now_secs() * 1000 + 60000);
1277 let response = RateLimitResponse::allowed(&result, 10);
1278 assert_eq!(response.status_code, 200);
1279 assert!(response.body.is_null());
1280 assert!(response.headers.retry_after.is_none());
1281 }
1282
1283 #[test]
1284 fn test_rate_limit_response_custom_strategy() {
1285 let result = RateLimitResult::rejected(0, now_secs() * 1000 + 60000);
1286 let response =
1287 RateLimitResponse::rejected(&result, 10, RateLimitResponseStrategy::Custom(503));
1288 assert_eq!(response.status_code, 503);
1289 }
1290
1291 #[test]
1294 fn test_multi_rate_limiter_all_allowed() {
1295 let l1 = Arc::new(SlidingWindowRateLimiter::new(10, Duration::from_secs(60)));
1296 let l2 = Arc::new(TokenBucketRateLimiter::new(10, 1.0));
1297 let multi = MultiRateLimiter::new(vec![l1, l2]);
1298
1299 let result = multi.check_all("key").unwrap();
1300 assert!(result.allowed);
1301 }
1302
1303 #[test]
1304 fn test_multi_rate_limiter_one_rejects() {
1305 let l1 = Arc::new(SlidingWindowRateLimiter::new(10, Duration::from_secs(60)));
1306 let l2 = Arc::new(TokenBucketRateLimiter::new(1, 1.0));
1307 let multi = MultiRateLimiter::new(vec![l1, l2]);
1308
1309 assert!(multi.check_all("key").unwrap().allowed);
1311 let result = multi.check_all("key").unwrap();
1313 assert!(!result.allowed);
1314 }
1315
1316 #[test]
1317 fn test_multi_rate_limiter_takes_strictest() {
1318 let l1 = Arc::new(SlidingWindowRateLimiter::new(5, Duration::from_secs(60)));
1319 let l2 = Arc::new(SlidingWindowRateLimiter::new(2, Duration::from_secs(60)));
1320 let multi = MultiRateLimiter::new(vec![l1, l2]);
1321
1322 multi.check_all("key").unwrap();
1324 multi.check_all("key").unwrap();
1325 assert!(!multi.check_all("key").unwrap().allowed);
1327 }
1328
1329 #[test]
1330 fn test_multi_rate_limiter_with_limiter() {
1331 let l1 = Arc::new(SlidingWindowRateLimiter::new(10, Duration::from_secs(60)));
1332 let multi = MultiRateLimiter::new(vec![]).with_limiter(l1);
1333 assert!(multi.check_all("key").unwrap().allowed);
1334 }
1335
1336 #[test]
1337 fn test_multi_rate_limiter_empty_errors() {
1338 let multi = MultiRateLimiter::new(vec![]);
1339 let result = multi.check_all("key");
1340 assert!(result.is_err());
1341 }
1342}