1pub mod composite;
15pub mod concurrency;
16pub mod config_builder;
17pub mod leaky_bucket;
18pub mod metrics;
19
20pub use composite::{
21 CompositeLimiter, CompositeStrategy, FallbackLimiter, LimitKeyBuilder, RateLimitRule, RuleSet,
22};
23pub use concurrency::{
24 ConcurrencyGuard, ConcurrencyLimiter, ConcurrencyStats, TimedConcurrencyLimiter,
25};
26pub use config_builder::{
27 ConfigError, LimitAlgorithm, RateLimitConfig, RateLimitConfigBuilder, TieredConfigBuilder,
28 TieredRateLimitConfig,
29};
30pub use leaky_bucket::{LeakyBucketLimiter, LeakyBucketStats, SlidingWindowLogLimiter};
31pub use metrics::{Alert, MetricsSnapshot, RateLimitMetrics, RateLimitMonitor};
32
33use std::collections::HashMap;
34use std::sync::atomic::{AtomicU64, Ordering};
35use std::sync::{Arc, RwLock};
36use std::time::{Duration, Instant};
37
38pub const DEFAULT_MAX_KEYS: usize = 10_000;
43
44pub trait RateLimiter: Send + Sync {
45 fn acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError>;
46 fn try_acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError>;
47 fn reset(&self, key: &str) -> Result<(), RateLimitError>;
48}
49
50#[derive(Debug, Clone)]
51pub struct RateLimitResult {
52 pub allowed: bool,
53 pub remaining: u64,
54 pub reset_at: i64,
55}
56
57impl RateLimitResult {
58 pub fn allowed(remaining: u64, reset_at: i64) -> Self {
59 Self {
60 allowed: true,
61 remaining,
62 reset_at,
63 }
64 }
65
66 pub fn rejected(remaining: u64, reset_at: i64) -> Self {
67 Self {
68 allowed: false,
69 remaining,
70 reset_at,
71 }
72 }
73
74 pub fn is_allowed(&self) -> bool {
75 self.allowed
76 }
77
78 pub fn is_rejected(&self) -> bool {
79 !self.allowed
80 }
81}
82
83pub struct SlidingWindowRateLimiter {
84 max_requests: Arc<AtomicU64>,
85 window_size: Duration,
86 entries: Arc<RwLock<HashMap<String, SlidingWindowEntry>>>,
87 max_keys: usize,
89 allowed_count: AtomicU64,
91 rejected_count: AtomicU64,
93}
94
95#[derive(Clone)]
96struct SlidingWindowEntry {
97 requests: Vec<Instant>,
98}
99
100impl SlidingWindowRateLimiter {
101 pub fn new(max_requests: u64, window_size: Duration) -> Self {
102 Self {
103 max_requests: Arc::new(AtomicU64::new(max_requests)),
104 window_size,
105 entries: Arc::new(RwLock::new(HashMap::new())),
106 max_keys: DEFAULT_MAX_KEYS,
107 allowed_count: AtomicU64::new(0),
108 rejected_count: AtomicU64::new(0),
109 }
110 }
111
112 pub fn with_max_keys(mut self, max_keys: usize) -> Self {
117 self.max_keys = max_keys;
118 self
119 }
120
121 pub fn max_requests(&self) -> u64 {
122 self.max_requests.load(Ordering::Relaxed)
123 }
124
125 pub fn window_size(&self) -> Duration {
126 self.window_size
127 }
128
129 pub fn max_keys(&self) -> usize {
130 self.max_keys
131 }
132
133 pub fn key_count(&self) -> usize {
134 self.entries
135 .read()
136 .map(|entries| entries.len())
137 .unwrap_or(0)
138 }
139
140 fn cleanup_old_requests(&self, entry: &mut SlidingWindowEntry) {
141 let now = Instant::now();
142 entry
143 .requests
144 .retain(|&time| now.duration_since(time) < self.window_size);
145 }
146
147 fn enforce_max_keys(&self, entries: &mut HashMap<String, SlidingWindowEntry>) {
153 while entries.len() > self.max_keys {
154 let now = Instant::now();
156 let oldest_key = entries
157 .iter()
158 .min_by_key(|(_, e)| e.requests.first().copied().unwrap_or(now))
159 .map(|(k, _)| k.clone());
160 match oldest_key {
161 Some(k) => {
162 entries.remove(&k);
163 }
164 None => break,
165 }
166 }
167 }
168
169 #[cfg(feature = "prod-rate-limit-tuning")]
171 pub fn set_capacity(&self, capacity: u64) {
172 self.max_requests.store(capacity, Ordering::Relaxed);
173 }
174
175 #[cfg(feature = "prod-rate-limit-tuning")]
177 pub fn set_rate(&self, rate: u64) {
178 let window_secs = self.window_size.as_secs().max(1);
179 self.max_requests
180 .store(rate * window_secs, Ordering::Relaxed);
181 }
182
183 #[cfg(feature = "prod-rate-limit-tuning")]
185 pub fn capacity(&self) -> u64 {
186 self.max_requests.load(Ordering::Relaxed)
187 }
188
189 #[cfg(feature = "prod-rate-limit-tuning")]
191 pub fn stats(&self) -> RateLimitStats {
192 RateLimitStats {
193 capacity: self.max_requests.load(Ordering::Relaxed),
194 allowed_count: self.allowed_count.load(Ordering::Relaxed),
195 rejected_count: self.rejected_count.load(Ordering::Relaxed),
196 }
197 }
198}
199
200impl RateLimiter for SlidingWindowRateLimiter {
201 fn acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
202 let mut entries = self
203 .entries
204 .write()
205 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
206
207 if entries.len() >= self.max_keys && !entries.contains_key(key) {
209 self.enforce_max_keys(&mut entries);
210 }
211
212 let entry = entries
213 .entry(key.to_string())
214 .or_insert_with(|| SlidingWindowEntry {
215 requests: Vec::new(),
216 });
217
218 self.cleanup_old_requests(entry);
219
220 let max_req = self.max_requests.load(Ordering::Relaxed);
221 if entry.requests.len() < max_req as usize {
222 entry.requests.push(Instant::now());
223 let remaining = max_req - entry.requests.len() as u64;
224 let reset_at = now_timestamp() + self.window_size.as_millis() as i64;
225 self.allowed_count.fetch_add(1, Ordering::Relaxed);
226 Ok(RateLimitResult::allowed(remaining, reset_at))
227 } else {
228 let oldest = entry
229 .requests
230 .first()
231 .map(|t| {
232 let elapsed = t.elapsed().as_millis() as i64;
233 let window_ms = self.window_size.as_millis() as i64;
234 now_timestamp() + (window_ms - elapsed)
235 })
236 .unwrap_or(now_timestamp());
237
238 let remaining = 0;
239 self.rejected_count.fetch_add(1, Ordering::Relaxed);
240 Ok(RateLimitResult::rejected(remaining, oldest))
241 }
242 }
243
244 fn try_acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
245 self.acquire(key)
246 }
247
248 fn reset(&self, key: &str) -> Result<(), RateLimitError> {
249 let mut entries = self
250 .entries
251 .write()
252 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
253 entries.remove(key);
254 Ok(())
255 }
256}
257
258pub struct TokenBucketRateLimiter {
259 capacity: f64,
260 refill_rate: f64,
261 entries: Arc<RwLock<HashMap<String, TokenBucketEntry>>>,
262 max_keys: usize,
264}
265
266#[derive(Clone)]
267struct TokenBucketEntry {
268 tokens: f64,
269 last_refill: Instant,
270}
271
272impl TokenBucketRateLimiter {
273 pub fn new(capacity: u64, refill_per_second: f64) -> Self {
274 Self {
275 capacity: capacity as f64,
276 refill_rate: refill_per_second,
277 entries: Arc::new(RwLock::new(HashMap::new())),
278 max_keys: DEFAULT_MAX_KEYS,
279 }
280 }
281
282 pub fn with_max_keys(mut self, max_keys: usize) -> Self {
284 self.max_keys = max_keys;
285 self
286 }
287
288 pub fn capacity(&self) -> u64 {
289 self.capacity as u64
290 }
291
292 pub fn refill_rate(&self) -> f64 {
293 self.refill_rate
294 }
295
296 pub fn max_keys(&self) -> usize {
297 self.max_keys
298 }
299
300 pub fn key_count(&self) -> usize {
301 self.entries
302 .read()
303 .map(|entries| entries.len())
304 .unwrap_or(0)
305 }
306
307 fn refill(&self, entry: &mut TokenBucketEntry) {
308 let now = Instant::now();
309 let elapsed = now.duration_since(entry.last_refill).as_secs_f64();
310 let tokens_to_add = if self.refill_rate > 0.0 {
313 elapsed * self.refill_rate
314 } else {
315 0.0
316 };
317
318 entry.tokens = (entry.tokens + tokens_to_add).min(self.capacity);
319 entry.last_refill = now;
320 }
321
322 fn enforce_max_keys(&self, entries: &mut HashMap<String, TokenBucketEntry>) {
327 while entries.len() > self.max_keys {
328 let oldest_key = entries
329 .iter()
330 .min_by_key(|(_, e)| e.last_refill)
331 .map(|(k, _)| k.clone());
332 match oldest_key {
333 Some(k) => {
334 entries.remove(&k);
335 }
336 None => break,
337 }
338 }
339 }
340}
341
342impl RateLimiter for TokenBucketRateLimiter {
343 fn acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
344 let mut entries = self
345 .entries
346 .write()
347 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
348
349 if entries.len() >= self.max_keys && !entries.contains_key(key) {
351 self.enforce_max_keys(&mut entries);
352 }
353
354 let entry = entries
355 .entry(key.to_string())
356 .or_insert_with(|| TokenBucketEntry {
357 tokens: self.capacity,
358 last_refill: Instant::now(),
359 });
360
361 self.refill(entry);
362
363 if entry.tokens >= 1.0 {
364 entry.tokens -= 1.0;
365 let remaining = entry.tokens.floor() as u64;
366 let reset_at = if self.refill_rate > 0.0 {
369 now_timestamp() + (1000.0 / self.refill_rate) as i64
370 } else {
371 i64::MAX
373 };
374 Ok(RateLimitResult::allowed(remaining, reset_at))
375 } else {
376 let reset_at = if self.refill_rate > 0.0 {
378 let wait_time = ((1.0 - entry.tokens) / self.refill_rate * 1000.0) as i64;
379 now_timestamp() + wait_time
380 } else {
381 i64::MAX
383 };
384 Ok(RateLimitResult::rejected(0, reset_at))
385 }
386 }
387
388 fn try_acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
389 self.acquire(key)
390 }
391
392 fn reset(&self, key: &str) -> Result<(), RateLimitError> {
393 let mut entries = self
394 .entries
395 .write()
396 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
397 entries.remove(key);
398 Ok(())
399 }
400}
401
402#[derive(Debug)]
403pub enum RateLimitError {
404 KeyNotFound(String),
405 Internal(String),
406}
407
408impl std::fmt::Display for RateLimitError {
409 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
410 match self {
411 RateLimitError::KeyNotFound(key) => write!(f, "Key not found: {}", key),
412 RateLimitError::Internal(msg) => write!(f, "Internal error: {}", msg),
413 }
414 }
415}
416
417impl std::error::Error for RateLimitError {}
418
419impl serde::Serialize for RateLimitError {
420 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
421 where
422 S: serde::Serializer,
423 {
424 serializer.serialize_str(&self.to_string())
425 }
426}
427
428fn now_timestamp() -> i64 {
429 use std::time::{SystemTime, UNIX_EPOCH};
430 SystemTime::now()
431 .duration_since(UNIX_EPOCH)
432 .unwrap_or_default()
433 .as_millis() as i64
434}
435
436fn now_secs() -> i64 {
437 use std::time::{SystemTime, UNIX_EPOCH};
438 SystemTime::now()
439 .duration_since(UNIX_EPOCH)
440 .unwrap_or_default()
441 .as_secs() as i64
442}
443
444pub struct FixedWindowRateLimiter {
467 max_requests: u64,
468 window_size: Duration,
469 entries: Arc<RwLock<HashMap<String, FixedWindowEntry>>>,
470 max_keys: usize,
471}
472
473#[derive(Clone)]
474struct FixedWindowEntry {
475 count: u64,
476 window_start: Instant,
477}
478
479impl FixedWindowRateLimiter {
480 pub fn new(max_requests: u64, window_size: Duration) -> Self {
485 Self {
486 max_requests,
487 window_size,
488 entries: Arc::new(RwLock::new(HashMap::new())),
489 max_keys: DEFAULT_MAX_KEYS,
490 }
491 }
492
493 pub fn with_max_keys(mut self, max_keys: usize) -> Self {
495 self.max_keys = max_keys;
496 self
497 }
498
499 pub fn max_requests(&self) -> u64 {
500 self.max_requests
501 }
502
503 pub fn window_size(&self) -> Duration {
504 self.window_size
505 }
506
507 pub fn max_keys(&self) -> usize {
508 self.max_keys
509 }
510
511 pub fn key_count(&self) -> usize {
512 self.entries
513 .read()
514 .map(|entries| entries.len())
515 .unwrap_or(0)
516 }
517
518 fn enforce_max_keys(&self, entries: &mut HashMap<String, FixedWindowEntry>) {
520 while entries.len() > self.max_keys {
521 let oldest_key = entries
522 .iter()
523 .min_by_key(|(_, e)| e.window_start)
524 .map(|(k, _)| k.clone());
525 match oldest_key {
526 Some(k) => {
527 entries.remove(&k);
528 }
529 None => break,
530 }
531 }
532 }
533}
534
535impl RateLimiter for FixedWindowRateLimiter {
536 fn acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
537 let mut entries = self
538 .entries
539 .write()
540 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
541
542 if entries.len() >= self.max_keys && !entries.contains_key(key) {
543 self.enforce_max_keys(&mut entries);
544 }
545
546 let now = Instant::now();
547 let entry = entries
548 .entry(key.to_string())
549 .or_insert_with(|| FixedWindowEntry {
550 count: 0,
551 window_start: now,
552 });
553
554 if now.duration_since(entry.window_start) >= self.window_size {
556 entry.count = 0;
557 entry.window_start = now;
558 }
559
560 if entry.count < self.max_requests {
561 entry.count += 1;
562 let remaining = self.max_requests - entry.count;
563 let reset_at = now_timestamp() + self.window_size.as_millis() as i64;
564 Ok(RateLimitResult::allowed(remaining, reset_at))
565 } else {
566 let elapsed = now.duration_since(entry.window_start);
568 let remaining_window = self
569 .window_size
570 .checked_sub(elapsed)
571 .unwrap_or(Duration::ZERO);
572 let reset_at = now_timestamp() + remaining_window.as_millis() as i64;
573 Ok(RateLimitResult::rejected(0, reset_at))
574 }
575 }
576
577 fn try_acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
578 self.acquire(key)
579 }
580
581 fn reset(&self, key: &str) -> Result<(), RateLimitError> {
582 let mut entries = self
583 .entries
584 .write()
585 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
586 entries.remove(key);
587 Ok(())
588 }
589}
590
591pub trait DistributedBackend: Send + Sync {
608 fn incr_and_get(
623 &self,
624 key: &str,
625 window_secs: u64,
626 window_start: i64,
627 max_requests: u64,
628 ) -> Result<(u64, i64), RateLimitError>;
629
630 fn get(&self, key: &str) -> Result<Option<(u64, i64)>, RateLimitError>;
632
633 fn reset_key(&self, key: &str) -> Result<(), RateLimitError>;
635}
636
637pub struct InMemoryBackend {
642 entries: RwLock<HashMap<String, (u64, i64)>>, }
644
645impl InMemoryBackend {
646 pub fn new() -> Self {
647 Self {
648 entries: RwLock::new(HashMap::new()),
649 }
650 }
651}
652
653impl Default for InMemoryBackend {
654 fn default() -> Self {
655 Self::new()
656 }
657}
658
659impl DistributedBackend for InMemoryBackend {
660 fn incr_and_get(
661 &self,
662 key: &str,
663 window_secs: u64,
664 window_start: i64,
665 _max_requests: u64,
666 ) -> Result<(u64, i64), RateLimitError> {
667 let mut entries = self
668 .entries
669 .write()
670 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
671
672 let entry = entries
673 .entry(key.to_string())
674 .or_insert_with(|| (0, window_start));
675
676 if window_start - entry.1 >= window_secs as i64 {
678 *entry = (0, window_start);
680 }
681
682 entry.0 += 1;
683 let reset_at = entry.1 + window_secs as i64;
684 Ok((entry.0, reset_at))
685 }
686
687 fn get(&self, key: &str) -> Result<Option<(u64, i64)>, RateLimitError> {
688 let entries = self
689 .entries
690 .read()
691 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
692 Ok(entries.get(key).copied())
693 }
694
695 fn reset_key(&self, key: &str) -> Result<(), RateLimitError> {
696 let mut entries = self
697 .entries
698 .write()
699 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
700 entries.remove(key);
701 Ok(())
702 }
703}
704
705pub struct DistributedRateLimiter {
710 backend: Arc<dyn DistributedBackend>,
711 max_requests: u64,
712 window_secs: u64,
713}
714
715impl DistributedRateLimiter {
716 pub fn new(backend: Arc<dyn DistributedBackend>, max_requests: u64, window_secs: u64) -> Self {
722 Self {
723 backend,
724 max_requests,
725 window_secs,
726 }
727 }
728
729 pub fn in_memory(max_requests: u64, window_secs: u64) -> Self {
731 Self::new(Arc::new(InMemoryBackend::new()), max_requests, window_secs)
732 }
733}
734
735impl RateLimiter for DistributedRateLimiter {
736 fn acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
737 let window_start = now_secs();
738 let (count, reset_at) =
739 self.backend
740 .incr_and_get(key, self.window_secs, window_start, self.max_requests)?;
741
742 if count <= self.max_requests {
743 let remaining = self.max_requests - count;
744 Ok(RateLimitResult::allowed(remaining, reset_at * 1000))
745 } else {
746 Ok(RateLimitResult::rejected(0, reset_at * 1000))
747 }
748 }
749
750 fn try_acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
751 self.acquire(key)
752 }
753
754 fn reset(&self, key: &str) -> Result<(), RateLimitError> {
755 self.backend.reset_key(key)
756 }
757}
758
759#[derive(Debug, Clone)]
777pub struct RateLimitHeaders {
778 pub limit: u64,
780 pub remaining: u64,
782 pub reset: i64,
784 pub retry_after: Option<u64>,
786}
787
788impl RateLimitHeaders {
789 pub fn from_result(result: &RateLimitResult, limit: u64) -> Self {
794 let reset_secs = result.reset_at / 1000;
795 let now_secs_val = now_secs();
796 let retry_after = if !result.allowed {
797 let diff = reset_secs - now_secs_val;
798 if diff > 0 {
799 Some(diff as u64)
800 } else {
801 Some(1)
802 }
803 } else {
804 None
805 };
806
807 Self {
808 limit,
809 remaining: result.remaining,
810 reset: reset_secs,
811 retry_after,
812 }
813 }
814
815 pub fn to_headers(&self) -> Vec<(String, String)> {
817 let mut headers = vec![
818 ("X-RateLimit-Limit".to_string(), self.limit.to_string()),
819 (
820 "X-RateLimit-Remaining".to_string(),
821 self.remaining.to_string(),
822 ),
823 ("X-RateLimit-Reset".to_string(), self.reset.to_string()),
824 ];
825 if let Some(retry) = self.retry_after {
826 headers.push(("Retry-After".to_string(), retry.to_string()));
827 }
828 headers
829 }
830
831 pub fn to_json(&self) -> serde_json::Value {
833 let mut map = serde_json::json!({
834 "X-RateLimit-Limit": self.limit,
835 "X-RateLimit-Remaining": self.remaining,
836 "X-RateLimit-Reset": self.reset,
837 });
838 if let Some(retry) = self.retry_after {
839 map["Retry-After"] = serde_json::json!(retry);
840 }
841 map
842 }
843}
844
845#[derive(Debug, Clone)]
849pub enum RateLimitResponseStrategy {
850 TooManyRequests,
852 ServiceUnavailable,
854 Custom(u16),
856}
857
858impl RateLimitResponseStrategy {
859 pub fn status_code(&self) -> u16 {
861 match self {
862 RateLimitResponseStrategy::TooManyRequests => 429,
863 RateLimitResponseStrategy::ServiceUnavailable => 503,
864 RateLimitResponseStrategy::Custom(code) => *code,
865 }
866 }
867}
868
869#[derive(Debug, Clone)]
874pub struct RateLimitResponse {
875 pub status_code: u16,
877 pub headers: RateLimitHeaders,
879 pub body: serde_json::Value,
881}
882
883impl RateLimitResponse {
884 pub fn rejected(
890 result: &RateLimitResult,
891 limit: u64,
892 strategy: RateLimitResponseStrategy,
893 ) -> Self {
894 let headers = RateLimitHeaders::from_result(result, limit);
895 let status_code = strategy.status_code();
896 let body = serde_json::json!({
897 "error": "rate_limit_exceeded",
898 "message": "Rate limit exceeded. Please retry later.",
899 "retry_after": headers.retry_after.unwrap_or(1),
900 });
901
902 Self {
903 status_code,
904 headers,
905 body,
906 }
907 }
908
909 pub fn allowed(result: &RateLimitResult, limit: u64) -> Self {
911 let headers = RateLimitHeaders::from_result(result, limit);
912 Self {
913 status_code: 200,
914 headers,
915 body: serde_json::Value::Null,
916 }
917 }
918}
919
920pub struct MultiRateLimiter {
931 limiters: Vec<Arc<dyn RateLimiter>>,
932}
933
934impl MultiRateLimiter {
935 pub fn new(limiters: Vec<Arc<dyn RateLimiter>>) -> Self {
937 Self { limiters }
938 }
939
940 pub fn with_limiter(mut self, limiter: Arc<dyn RateLimiter>) -> Self {
942 self.limiters.push(limiter);
943 self
944 }
945
946 pub fn limiter_count(&self) -> usize {
947 self.limiters.len()
948 }
949
950 pub fn check_all(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
955 let mut best_result: Option<RateLimitResult> = None;
956 for limiter in &self.limiters {
957 let result = limiter.acquire(key)?;
958 match &best_result {
959 None => best_result = Some(result),
960 Some(current) => {
961 if !result.allowed {
963 if !current.allowed {
965 if result.remaining <= current.remaining {
967 best_result = Some(result);
968 }
969 } else {
970 best_result = Some(result);
972 }
973 } else if current.allowed && result.remaining < current.remaining {
974 best_result = Some(result);
976 }
977 }
978 }
979 }
980
981 best_result.ok_or_else(|| RateLimitError::Internal("No limiters configured".to_string()))
982 }
983}
984
985#[cfg(feature = "prod-rate-limit-tuning")]
990mod prod {
991 use super::DEFAULT_MAX_KEYS;
992 use serde::{Deserialize, Serialize};
993 use std::time::Duration;
994
995 #[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
997 pub enum RateLimitProdError {
998 #[error("rate limit capacity must be positive")]
999 CapacityNotPositive,
1000 #[error("rate limit rate must be positive")]
1001 RateNotPositive,
1002 #[error("rate limit window_size must be positive")]
1003 WindowSizeNotPositive,
1004 #[error("rate limit max_keys too small, minimum 100 recommended")]
1005 MaxKeysTooSmall,
1006 }
1007
1008 #[derive(Debug, Clone, Serialize, Deserialize)]
1010 pub struct RateLimitProdConfig {
1011 pub capacity: u64,
1012 pub rate: u64,
1013 pub window_size: Duration,
1014 pub max_keys: usize,
1015 }
1016
1017 impl Default for RateLimitProdConfig {
1018 fn default() -> Self {
1019 Self {
1020 capacity: 100,
1021 rate: 10,
1022 window_size: Duration::from_secs(1),
1023 max_keys: DEFAULT_MAX_KEYS,
1024 }
1025 }
1026 }
1027
1028 impl RateLimitProdConfig {
1029 pub fn new(capacity: u64, rate: u64, window_size: Duration, max_keys: usize) -> Self {
1030 Self {
1031 capacity,
1032 rate,
1033 window_size,
1034 max_keys,
1035 }
1036 }
1037
1038 pub fn validate(&self) -> Result<(), RateLimitProdError> {
1040 if self.capacity == 0 {
1041 return Err(RateLimitProdError::CapacityNotPositive);
1042 }
1043 if self.rate == 0 {
1044 return Err(RateLimitProdError::RateNotPositive);
1045 }
1046 if self.window_size.is_zero() {
1047 return Err(RateLimitProdError::WindowSizeNotPositive);
1048 }
1049 if self.max_keys < 100 {
1050 return Err(RateLimitProdError::MaxKeysTooSmall);
1051 }
1052 Ok(())
1053 }
1054 }
1055
1056 #[derive(Debug, Clone, Serialize, Deserialize)]
1058 pub struct RateLimitStats {
1059 pub capacity: u64,
1060 pub allowed_count: u64,
1061 pub rejected_count: u64,
1062 }
1063}
1064
1065#[cfg(feature = "prod-rate-limit-tuning")]
1066pub use prod::{RateLimitProdConfig, RateLimitProdError, RateLimitStats};
1067
1068#[cfg(test)]
1069mod tests {
1070 use super::*;
1071
1072 #[test]
1073 fn test_rate_limit_result_allowed() {
1074 let result = RateLimitResult::allowed(5, 1000);
1075 assert!(result.allowed);
1076 assert_eq!(result.remaining, 5);
1077 assert_eq!(result.reset_at, 1000);
1078 }
1079
1080 #[test]
1081 fn test_rate_limit_result_rejected() {
1082 let result = RateLimitResult::rejected(0, 2000);
1083 assert!(!result.allowed);
1084 assert_eq!(result.remaining, 0);
1085 }
1086
1087 #[test]
1088 fn test_sliding_window_limiter_new() {
1089 let limiter = SlidingWindowRateLimiter::new(10, Duration::from_secs(60));
1090 let result = limiter.acquire("test-key");
1091 assert!(result.is_ok());
1092 assert!(result.unwrap().allowed);
1093 }
1094
1095 #[test]
1096 fn test_sliding_window_limiter_full() {
1097 let limiter = SlidingWindowRateLimiter::new(2, Duration::from_secs(60));
1098
1099 let r1 = limiter.acquire("key1").unwrap();
1100 assert!(r1.allowed);
1101
1102 let r2 = limiter.acquire("key1").unwrap();
1103 assert!(r2.allowed);
1104
1105 let r3 = limiter.acquire("key1").unwrap();
1106 assert!(!r3.allowed);
1107 }
1108
1109 #[test]
1110 fn test_sliding_window_different_keys() {
1111 let limiter = SlidingWindowRateLimiter::new(1, Duration::from_secs(60));
1112
1113 let r1 = limiter.acquire("key-a").unwrap();
1114 assert!(r1.allowed);
1115
1116 let r2 = limiter.acquire("key-b").unwrap();
1117 assert!(r2.allowed);
1118 }
1119
1120 #[test]
1121 fn test_sliding_window_reset() {
1122 let limiter = SlidingWindowRateLimiter::new(1, Duration::from_secs(60));
1123
1124 limiter.acquire("reset-key").unwrap();
1125 limiter.acquire("reset-key").unwrap();
1126
1127 limiter.reset("reset-key").unwrap();
1128
1129 let result = limiter.acquire("reset-key").unwrap();
1130 assert!(result.allowed);
1131 }
1132
1133 #[test]
1134 fn test_token_bucket_limiter_new() {
1135 let limiter = TokenBucketRateLimiter::new(10, 1.0);
1136 let result = limiter.acquire("test-key");
1137 assert!(result.is_ok());
1138 assert!(result.unwrap().allowed);
1139 }
1140
1141 #[test]
1142 fn test_token_bucket_limiter_depletes() {
1143 let limiter = TokenBucketRateLimiter::new(2, 1.0);
1144
1145 let r1 = limiter.acquire("key1").unwrap();
1146 assert!(r1.allowed);
1147 assert_eq!(r1.remaining, 1);
1148
1149 let r2 = limiter.acquire("key1").unwrap();
1150 assert!(r2.allowed);
1151 assert_eq!(r2.remaining, 0);
1152
1153 let r3 = limiter.acquire("key1").unwrap();
1154 assert!(!r3.allowed);
1155 }
1156
1157 #[test]
1158 fn test_token_bucket_different_keys() {
1159 let limiter = TokenBucketRateLimiter::new(1, 1.0);
1160
1161 let r1 = limiter.acquire("key-a").unwrap();
1162 assert!(r1.allowed);
1163
1164 let r2 = limiter.acquire("key-b").unwrap();
1165 assert!(r2.allowed);
1166 }
1167
1168 #[test]
1169 fn test_token_bucket_reset() {
1170 let limiter = TokenBucketRateLimiter::new(1, 1.0);
1171
1172 limiter.acquire("reset-key").unwrap();
1173 limiter.acquire("reset-key").unwrap();
1174
1175 limiter.reset("reset-key").unwrap();
1176
1177 let result = limiter.acquire("reset-key").unwrap();
1178 assert!(result.allowed);
1179 }
1180
1181 #[test]
1182 fn test_limiter_try_acquire() {
1183 let limiter = SlidingWindowRateLimiter::new(1, Duration::from_secs(60));
1184
1185 let r1 = limiter.try_acquire("key").unwrap();
1186 assert!(r1.allowed);
1187
1188 let r2 = limiter.try_acquire("key").unwrap();
1189 assert!(!r2.allowed);
1190 }
1191
1192 #[test]
1195 fn test_token_bucket_zero_refill_rate_does_not_panic() {
1196 let limiter = TokenBucketRateLimiter::new(1, 0.0);
1199
1200 let r1 = limiter.acquire("zero-refill").unwrap();
1201 assert!(r1.allowed, "first acquire should be allowed");
1202
1203 let r2 = limiter.acquire("zero-refill").unwrap();
1205 assert!(!r2.allowed, "second acquire should be rejected");
1206 assert!(
1208 r2.reset_at > 0,
1209 "reset_at should be a valid timestamp, got: {}",
1210 r2.reset_at
1211 );
1212 }
1213
1214 #[test]
1215 fn test_token_bucket_negative_refill_rate_does_not_panic() {
1216 let limiter = TokenBucketRateLimiter::new(1, -1.0);
1218
1219 let r1 = limiter.acquire("neg-refill").unwrap();
1220 assert!(r1.allowed, "first acquire should be allowed");
1221
1222 let r2 = limiter.acquire("neg-refill").unwrap();
1223 assert!(!r2.allowed, "second acquire should be rejected");
1224 assert!(
1225 r2.reset_at > 0,
1226 "reset_at should be a valid timestamp, got: {}",
1227 r2.reset_at
1228 );
1229 }
1230
1231 #[test]
1234 fn test_fixed_window_limiter_allows_within_limit() {
1235 let limiter = FixedWindowRateLimiter::new(5, Duration::from_secs(60));
1236 for i in 0..5 {
1237 let r = limiter.acquire("key").unwrap();
1238 assert!(r.allowed, "request {} should be allowed", i);
1239 }
1240 }
1241
1242 #[test]
1243 fn test_fixed_window_limiter_rejects_over_limit() {
1244 let limiter = FixedWindowRateLimiter::new(2, Duration::from_secs(60));
1245 assert!(limiter.acquire("key").unwrap().allowed);
1246 assert!(limiter.acquire("key").unwrap().allowed);
1247 let r3 = limiter.acquire("key").unwrap();
1248 assert!(!r3.allowed);
1249 assert_eq!(r3.remaining, 0);
1250 }
1251
1252 #[test]
1253 fn test_fixed_window_limiter_different_keys() {
1254 let limiter = FixedWindowRateLimiter::new(1, Duration::from_secs(60));
1255 assert!(limiter.acquire("key-a").unwrap().allowed);
1256 assert!(limiter.acquire("key-b").unwrap().allowed);
1257 }
1258
1259 #[test]
1260 fn test_fixed_window_limiter_reset() {
1261 let limiter = FixedWindowRateLimiter::new(1, Duration::from_secs(60));
1262 limiter.acquire("key").unwrap();
1263 assert!(!limiter.acquire("key").unwrap().allowed);
1264 limiter.reset("key").unwrap();
1265 assert!(limiter.acquire("key").unwrap().allowed);
1266 }
1267
1268 #[test]
1269 fn test_fixed_window_limiter_remaining_decreases() {
1270 let limiter = FixedWindowRateLimiter::new(3, Duration::from_secs(60));
1271 let r1 = limiter.acquire("key").unwrap();
1272 assert_eq!(r1.remaining, 2);
1273 let r2 = limiter.acquire("key").unwrap();
1274 assert_eq!(r2.remaining, 1);
1275 let r3 = limiter.acquire("key").unwrap();
1276 assert_eq!(r3.remaining, 0);
1277 }
1278
1279 #[test]
1280 fn test_fixed_window_limiter_reset_at_positive() {
1281 let limiter = FixedWindowRateLimiter::new(1, Duration::from_secs(60));
1282 let r = limiter.acquire("key").unwrap();
1283 assert!(r.reset_at > 0);
1284 }
1285
1286 #[test]
1287 fn test_fixed_window_limiter_try_acquire() {
1288 let limiter = FixedWindowRateLimiter::new(1, Duration::from_secs(60));
1289 assert!(limiter.try_acquire("key").unwrap().allowed);
1290 assert!(!limiter.try_acquire("key").unwrap().allowed);
1291 }
1292
1293 #[test]
1296 fn test_in_memory_backend_new() {
1297 let backend = InMemoryBackend::new();
1298 let result = backend.get("key").unwrap();
1299 assert!(result.is_none());
1300 }
1301
1302 #[test]
1303 fn test_in_memory_backend_incr_and_get() {
1304 let backend = InMemoryBackend::new();
1305 let now = now_secs();
1306 let (count1, reset1) = backend.incr_and_get("key", 60, now, 10).unwrap();
1307 assert_eq!(count1, 1);
1308 assert_eq!(reset1, now + 60);
1309
1310 let (count2, _) = backend.incr_and_get("key", 60, now, 10).unwrap();
1311 assert_eq!(count2, 2);
1312 }
1313
1314 #[test]
1315 fn test_in_memory_backend_get() {
1316 let backend = InMemoryBackend::new();
1317 let now = now_secs();
1318 backend.incr_and_get("key", 60, now, 10).unwrap();
1319 let result = backend.get("key").unwrap();
1320 assert!(result.is_some());
1321 assert_eq!(result.unwrap().0, 1);
1322 }
1323
1324 #[test]
1325 fn test_in_memory_backend_reset_key() {
1326 let backend = InMemoryBackend::new();
1327 let now = now_secs();
1328 backend.incr_and_get("key", 60, now, 10).unwrap();
1329 assert!(backend.get("key").unwrap().is_some());
1330 backend.reset_key("key").unwrap();
1331 assert!(backend.get("key").unwrap().is_none());
1332 }
1333
1334 #[test]
1335 fn test_in_memory_backend_window_expiry() {
1336 let backend = InMemoryBackend::new();
1337 let now = now_secs();
1338 backend.incr_and_get("key", 60, now, 10).unwrap();
1340 backend.incr_and_get("key", 60, now, 10).unwrap();
1341 let (count, _) = backend.incr_and_get("key", 60, now + 61, 10).unwrap();
1343 assert_eq!(count, 1);
1344 }
1345
1346 #[test]
1347 fn test_distributed_rate_limiter_allows() {
1348 let limiter = DistributedRateLimiter::in_memory(5, 60);
1349 for i in 0..5 {
1350 let r = limiter.acquire("key").unwrap();
1351 assert!(r.allowed, "request {} should be allowed", i);
1352 }
1353 }
1354
1355 #[test]
1356 fn test_distributed_rate_limiter_rejects() {
1357 let limiter = DistributedRateLimiter::in_memory(2, 60);
1358 assert!(limiter.acquire("key").unwrap().allowed);
1359 assert!(limiter.acquire("key").unwrap().allowed);
1360 assert!(!limiter.acquire("key").unwrap().allowed);
1361 }
1362
1363 #[test]
1364 fn test_distributed_rate_limiter_shared_backend() {
1365 let backend = Arc::new(InMemoryBackend::new());
1367 let limiter1 = DistributedRateLimiter::new(backend.clone(), 2, 60);
1368 let limiter2 = DistributedRateLimiter::new(backend.clone(), 2, 60);
1369
1370 assert!(limiter1.acquire("key").unwrap().allowed);
1372 assert!(limiter2.acquire("key").unwrap().allowed);
1374 assert!(!limiter1.acquire("key").unwrap().allowed);
1376 }
1377
1378 #[test]
1379 fn test_distributed_rate_limiter_reset() {
1380 let limiter = DistributedRateLimiter::in_memory(1, 60);
1381 limiter.acquire("key").unwrap();
1382 assert!(!limiter.acquire("key").unwrap().allowed);
1383 limiter.reset("key").unwrap();
1384 assert!(limiter.acquire("key").unwrap().allowed);
1385 }
1386
1387 #[test]
1388 fn test_distributed_rate_limiter_different_keys() {
1389 let limiter = DistributedRateLimiter::in_memory(1, 60);
1390 assert!(limiter.acquire("key-a").unwrap().allowed);
1391 assert!(limiter.acquire("key-b").unwrap().allowed);
1392 }
1393
1394 #[test]
1395 fn test_distributed_rate_limiter_remaining() {
1396 let limiter = DistributedRateLimiter::in_memory(3, 60);
1397 let r1 = limiter.acquire("key").unwrap();
1398 assert_eq!(r1.remaining, 2);
1399 let r2 = limiter.acquire("key").unwrap();
1400 assert_eq!(r2.remaining, 1);
1401 }
1402
1403 #[test]
1406 fn test_rate_limit_headers_allowed() {
1407 let result = RateLimitResult::allowed(5, now_secs() * 1000 + 60000);
1408 let headers = RateLimitHeaders::from_result(&result, 10);
1409 assert_eq!(headers.limit, 10);
1410 assert_eq!(headers.remaining, 5);
1411 assert!(headers.retry_after.is_none());
1412 }
1413
1414 #[test]
1415 fn test_rate_limit_headers_rejected() {
1416 let result = RateLimitResult::rejected(0, now_secs() * 1000 + 60000);
1417 let headers = RateLimitHeaders::from_result(&result, 10);
1418 assert_eq!(headers.limit, 10);
1419 assert_eq!(headers.remaining, 0);
1420 assert!(headers.retry_after.is_some());
1421 assert!(headers.retry_after.unwrap() > 0);
1422 }
1423
1424 #[test]
1425 fn test_rate_limit_headers_to_headers_allowed() {
1426 let result = RateLimitResult::allowed(5, now_secs() * 1000 + 60000);
1427 let headers = RateLimitHeaders::from_result(&result, 10);
1428 let hdrs = headers.to_headers();
1429 assert_eq!(hdrs.len(), 3); assert!(hdrs
1431 .iter()
1432 .any(|(k, v)| k == "X-RateLimit-Limit" && v == "10"));
1433 assert!(hdrs
1434 .iter()
1435 .any(|(k, v)| k == "X-RateLimit-Remaining" && v == "5"));
1436 }
1437
1438 #[test]
1439 fn test_rate_limit_headers_to_headers_rejected() {
1440 let result = RateLimitResult::rejected(0, now_secs() * 1000 + 60000);
1441 let headers = RateLimitHeaders::from_result(&result, 10);
1442 let hdrs = headers.to_headers();
1443 assert_eq!(hdrs.len(), 4); assert!(hdrs.iter().any(|(k, _)| k == "Retry-After"));
1445 }
1446
1447 #[test]
1448 fn test_rate_limit_headers_to_json() {
1449 let result = RateLimitResult::allowed(5, now_secs() * 1000 + 60000);
1450 let headers = RateLimitHeaders::from_result(&result, 10);
1451 let json = headers.to_json();
1452 assert_eq!(json["X-RateLimit-Limit"], 10);
1453 assert_eq!(json["X-RateLimit-Remaining"], 5);
1454 assert!(json.get("Retry-After").is_none());
1455 }
1456
1457 #[test]
1458 fn test_rate_limit_headers_to_json_rejected() {
1459 let result = RateLimitResult::rejected(0, now_secs() * 1000 + 60000);
1460 let headers = RateLimitHeaders::from_result(&result, 10);
1461 let json = headers.to_json();
1462 assert!(json.get("Retry-After").is_some());
1463 }
1464
1465 #[test]
1466 fn test_rate_limit_response_strategy_status_codes() {
1467 assert_eq!(
1468 RateLimitResponseStrategy::TooManyRequests.status_code(),
1469 429
1470 );
1471 assert_eq!(
1472 RateLimitResponseStrategy::ServiceUnavailable.status_code(),
1473 503
1474 );
1475 assert_eq!(RateLimitResponseStrategy::Custom(502).status_code(), 502);
1476 }
1477
1478 #[test]
1479 fn test_rate_limit_response_rejected() {
1480 let result = RateLimitResult::rejected(0, now_secs() * 1000 + 60000);
1481 let response =
1482 RateLimitResponse::rejected(&result, 10, RateLimitResponseStrategy::TooManyRequests);
1483 assert_eq!(response.status_code, 429);
1484 assert_eq!(response.body["error"], "rate_limit_exceeded");
1485 assert!(response.body["retry_after"].as_u64().unwrap() > 0);
1486 assert!(response.headers.retry_after.is_some());
1487 }
1488
1489 #[test]
1490 fn test_rate_limit_response_allowed() {
1491 let result = RateLimitResult::allowed(5, now_secs() * 1000 + 60000);
1492 let response = RateLimitResponse::allowed(&result, 10);
1493 assert_eq!(response.status_code, 200);
1494 assert!(response.body.is_null());
1495 assert!(response.headers.retry_after.is_none());
1496 }
1497
1498 #[test]
1499 fn test_rate_limit_response_custom_strategy() {
1500 let result = RateLimitResult::rejected(0, now_secs() * 1000 + 60000);
1501 let response =
1502 RateLimitResponse::rejected(&result, 10, RateLimitResponseStrategy::Custom(503));
1503 assert_eq!(response.status_code, 503);
1504 }
1505
1506 #[test]
1509 fn test_multi_rate_limiter_all_allowed() {
1510 let l1 = Arc::new(SlidingWindowRateLimiter::new(10, Duration::from_secs(60)));
1511 let l2 = Arc::new(TokenBucketRateLimiter::new(10, 1.0));
1512 let multi = MultiRateLimiter::new(vec![l1, l2]);
1513
1514 let result = multi.check_all("key").unwrap();
1515 assert!(result.allowed);
1516 }
1517
1518 #[test]
1519 fn test_multi_rate_limiter_one_rejects() {
1520 let l1 = Arc::new(SlidingWindowRateLimiter::new(10, Duration::from_secs(60)));
1521 let l2 = Arc::new(TokenBucketRateLimiter::new(1, 1.0));
1522 let multi = MultiRateLimiter::new(vec![l1, l2]);
1523
1524 assert!(multi.check_all("key").unwrap().allowed);
1526 let result = multi.check_all("key").unwrap();
1528 assert!(!result.allowed);
1529 }
1530
1531 #[test]
1532 fn test_multi_rate_limiter_takes_strictest() {
1533 let l1 = Arc::new(SlidingWindowRateLimiter::new(5, Duration::from_secs(60)));
1534 let l2 = Arc::new(SlidingWindowRateLimiter::new(2, Duration::from_secs(60)));
1535 let multi = MultiRateLimiter::new(vec![l1, l2]);
1536
1537 multi.check_all("key").unwrap();
1539 multi.check_all("key").unwrap();
1540 assert!(!multi.check_all("key").unwrap().allowed);
1542 }
1543
1544 #[test]
1545 fn test_multi_rate_limiter_with_limiter() {
1546 let l1 = Arc::new(SlidingWindowRateLimiter::new(10, Duration::from_secs(60)));
1547 let multi = MultiRateLimiter::new(vec![]).with_limiter(l1);
1548 assert!(multi.check_all("key").unwrap().allowed);
1549 }
1550
1551 #[test]
1552 fn test_multi_rate_limiter_empty_errors() {
1553 let multi = MultiRateLimiter::new(vec![]);
1554 let result = multi.check_all("key");
1555 assert!(result.is_err());
1556 }
1557}
1558
1559#[cfg(all(test, feature = "prod-rate-limit-tuning"))]
1560mod prod_tests {
1561 use super::*;
1562
1563 #[test]
1564 fn test_rate_limit_prod_config_validate_ok() {
1565 let config = RateLimitProdConfig::new(100, 10, Duration::from_secs(1), 1000);
1566 assert!(config.validate().is_ok());
1567 }
1568
1569 #[test]
1570 fn test_rate_limit_prod_config_capacity_zero_rejected() {
1571 let config = RateLimitProdConfig::new(0, 10, Duration::from_secs(1), 1000);
1572 let err = config.validate().unwrap_err();
1573 assert!(err.to_string().contains("capacity must be positive"));
1574 }
1575
1576 #[test]
1577 fn test_rate_limit_prod_config_rate_zero_rejected() {
1578 let config = RateLimitProdConfig::new(100, 0, Duration::from_secs(1), 1000);
1579 assert!(config.validate().is_err());
1580 }
1581
1582 #[test]
1583 fn test_rate_limit_prod_config_max_keys_too_small() {
1584 let config = RateLimitProdConfig::new(100, 10, Duration::from_secs(1), 50);
1585 let err = config.validate().unwrap_err();
1586 assert!(err.to_string().contains("max_keys too small"));
1587 }
1588
1589 #[test]
1590 fn test_sliding_window_set_capacity() {
1591 let limiter = SlidingWindowRateLimiter::new(100, Duration::from_secs(60));
1592 assert_eq!(limiter.capacity(), 100);
1593 limiter.set_capacity(200);
1594 assert_eq!(limiter.capacity(), 200);
1595 }
1596
1597 #[test]
1598 fn test_sliding_window_set_rate() {
1599 let limiter = SlidingWindowRateLimiter::new(100, Duration::from_secs(10));
1600 limiter.set_rate(20);
1601 assert_eq!(limiter.capacity(), 200);
1603 }
1604
1605 #[test]
1606 fn test_sliding_window_stats_after_acquire() {
1607 let limiter = SlidingWindowRateLimiter::new(2, Duration::from_secs(60));
1608 limiter.acquire("k1").unwrap();
1609 limiter.acquire("k1").unwrap();
1610 limiter.acquire("k1").unwrap(); let stats = limiter.stats();
1612 assert_eq!(stats.capacity, 2);
1613 assert_eq!(stats.allowed_count, 2);
1614 assert_eq!(stats.rejected_count, 1);
1615 }
1616
1617 #[test]
1618 fn test_sliding_window_dynamic_capacity_takes_effect() {
1619 let limiter = SlidingWindowRateLimiter::new(2, Duration::from_secs(60));
1620 limiter.acquire("k").unwrap();
1621 limiter.acquire("k").unwrap();
1622 assert!(!limiter.acquire("k").unwrap().allowed);
1623 limiter.set_capacity(5);
1624 assert!(limiter.acquire("k").unwrap().allowed);
1625 }
1626}