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