1use std::collections::HashMap;
14use std::sync::atomic::{AtomicU64, Ordering};
15use std::sync::{Arc, RwLock};
16use std::time::{Duration, Instant};
17
18pub const DEFAULT_MAX_KEYS: usize = 10_000;
23
24pub trait RateLimiter: Send + Sync {
25 fn acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError>;
26 fn try_acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError>;
27 fn reset(&self, key: &str) -> Result<(), RateLimitError>;
28}
29
30#[derive(Debug, Clone)]
31pub struct RateLimitResult {
32 pub allowed: bool,
33 pub remaining: u64,
34 pub reset_at: i64,
35}
36
37impl RateLimitResult {
38 pub fn allowed(remaining: u64, reset_at: i64) -> Self {
39 Self {
40 allowed: true,
41 remaining,
42 reset_at,
43 }
44 }
45
46 pub fn rejected(remaining: u64, reset_at: i64) -> Self {
47 Self {
48 allowed: false,
49 remaining,
50 reset_at,
51 }
52 }
53}
54
55pub struct SlidingWindowRateLimiter {
56 max_requests: Arc<AtomicU64>,
57 window_size: Duration,
58 entries: Arc<RwLock<HashMap<String, SlidingWindowEntry>>>,
59 max_keys: usize,
61 allowed_count: AtomicU64,
63 rejected_count: AtomicU64,
65}
66
67#[derive(Clone)]
68struct SlidingWindowEntry {
69 requests: Vec<Instant>,
70}
71
72impl SlidingWindowRateLimiter {
73 pub fn new(max_requests: u64, window_size: Duration) -> Self {
74 Self {
75 max_requests: Arc::new(AtomicU64::new(max_requests)),
76 window_size,
77 entries: Arc::new(RwLock::new(HashMap::new())),
78 max_keys: DEFAULT_MAX_KEYS,
79 allowed_count: AtomicU64::new(0),
80 rejected_count: AtomicU64::new(0),
81 }
82 }
83
84 pub fn with_max_keys(mut self, max_keys: usize) -> Self {
89 self.max_keys = max_keys;
90 self
91 }
92
93 fn cleanup_old_requests(&self, entry: &mut SlidingWindowEntry) {
94 let now = Instant::now();
95 entry
96 .requests
97 .retain(|&time| now.duration_since(time) < self.window_size);
98 }
99
100 fn enforce_max_keys(&self, entries: &mut HashMap<String, SlidingWindowEntry>) {
105 while entries.len() > self.max_keys {
106 let now = Instant::now();
108 let oldest_key = entries
109 .iter()
110 .min_by_key(|(_, e)| e.requests.first().copied().unwrap_or(now))
111 .map(|(k, _)| k.clone());
112 match oldest_key {
113 Some(k) => {
114 entries.remove(&k);
115 }
116 None => break,
117 }
118 }
119 }
120
121 #[cfg(feature = "prod-rate-limit-tuning")]
123 pub fn set_capacity(&self, capacity: u64) {
124 self.max_requests.store(capacity, Ordering::Relaxed);
125 }
126
127 #[cfg(feature = "prod-rate-limit-tuning")]
129 pub fn set_rate(&self, rate: u64) {
130 let window_secs = self.window_size.as_secs().max(1);
131 self.max_requests
132 .store(rate * window_secs, Ordering::Relaxed);
133 }
134
135 #[cfg(feature = "prod-rate-limit-tuning")]
137 pub fn capacity(&self) -> u64 {
138 self.max_requests.load(Ordering::Relaxed)
139 }
140
141 #[cfg(feature = "prod-rate-limit-tuning")]
143 pub fn stats(&self) -> RateLimitStats {
144 RateLimitStats {
145 capacity: self.max_requests.load(Ordering::Relaxed),
146 allowed_count: self.allowed_count.load(Ordering::Relaxed),
147 rejected_count: self.rejected_count.load(Ordering::Relaxed),
148 }
149 }
150}
151
152impl RateLimiter for SlidingWindowRateLimiter {
153 fn acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
154 let mut entries = self
155 .entries
156 .write()
157 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
158
159 if entries.len() >= self.max_keys && !entries.contains_key(key) {
161 self.enforce_max_keys(&mut entries);
162 }
163
164 let entry = entries
165 .entry(key.to_string())
166 .or_insert_with(|| SlidingWindowEntry {
167 requests: Vec::new(),
168 });
169
170 self.cleanup_old_requests(entry);
171
172 let max_req = self.max_requests.load(Ordering::Relaxed);
173 if entry.requests.len() < max_req as usize {
174 entry.requests.push(Instant::now());
175 let remaining = max_req - entry.requests.len() as u64;
176 let reset_at = now_timestamp() + self.window_size.as_millis() as i64;
177 self.allowed_count.fetch_add(1, Ordering::Relaxed);
178 Ok(RateLimitResult::allowed(remaining, reset_at))
179 } else {
180 let oldest = entry
181 .requests
182 .first()
183 .map(|t| {
184 let elapsed = t.elapsed().as_millis() as i64;
185 let window_ms = self.window_size.as_millis() as i64;
186 now_timestamp() + (window_ms - elapsed)
187 })
188 .unwrap_or(now_timestamp());
189
190 let remaining = 0;
191 self.rejected_count.fetch_add(1, Ordering::Relaxed);
192 Ok(RateLimitResult::rejected(remaining, oldest))
193 }
194 }
195
196 fn try_acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
197 self.acquire(key)
198 }
199
200 fn reset(&self, key: &str) -> Result<(), RateLimitError> {
201 let mut entries = self
202 .entries
203 .write()
204 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
205 entries.remove(key);
206 Ok(())
207 }
208}
209
210pub struct TokenBucketRateLimiter {
211 capacity: f64,
212 refill_rate: f64,
213 entries: Arc<RwLock<HashMap<String, TokenBucketEntry>>>,
214 max_keys: usize,
216}
217
218#[derive(Clone)]
219struct TokenBucketEntry {
220 tokens: f64,
221 last_refill: Instant,
222}
223
224impl TokenBucketRateLimiter {
225 pub fn new(capacity: u64, refill_per_second: f64) -> Self {
226 Self {
227 capacity: capacity as f64,
228 refill_rate: refill_per_second,
229 entries: Arc::new(RwLock::new(HashMap::new())),
230 max_keys: DEFAULT_MAX_KEYS,
231 }
232 }
233
234 pub fn with_max_keys(mut self, max_keys: usize) -> Self {
236 self.max_keys = max_keys;
237 self
238 }
239
240 fn refill(&self, entry: &mut TokenBucketEntry) {
241 let now = Instant::now();
242 let elapsed = now.duration_since(entry.last_refill).as_secs_f64();
243 let tokens_to_add = if self.refill_rate > 0.0 {
246 elapsed * self.refill_rate
247 } else {
248 0.0
249 };
250
251 entry.tokens = (entry.tokens + tokens_to_add).min(self.capacity);
252 entry.last_refill = now;
253 }
254
255 fn enforce_max_keys(&self, entries: &mut HashMap<String, TokenBucketEntry>) {
260 while entries.len() > self.max_keys {
261 let oldest_key = entries
262 .iter()
263 .min_by_key(|(_, e)| e.last_refill)
264 .map(|(k, _)| k.clone());
265 match oldest_key {
266 Some(k) => {
267 entries.remove(&k);
268 }
269 None => break,
270 }
271 }
272 }
273}
274
275impl RateLimiter for TokenBucketRateLimiter {
276 fn acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
277 let mut entries = self
278 .entries
279 .write()
280 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
281
282 if entries.len() >= self.max_keys && !entries.contains_key(key) {
284 self.enforce_max_keys(&mut entries);
285 }
286
287 let entry = entries
288 .entry(key.to_string())
289 .or_insert_with(|| TokenBucketEntry {
290 tokens: self.capacity,
291 last_refill: Instant::now(),
292 });
293
294 self.refill(entry);
295
296 if entry.tokens >= 1.0 {
297 entry.tokens -= 1.0;
298 let remaining = entry.tokens.floor() as u64;
299 let reset_at = if self.refill_rate > 0.0 {
302 now_timestamp() + (1000.0 / self.refill_rate) as i64
303 } else {
304 i64::MAX
306 };
307 Ok(RateLimitResult::allowed(remaining, reset_at))
308 } else {
309 let reset_at = if self.refill_rate > 0.0 {
311 let wait_time = ((1.0 - entry.tokens) / self.refill_rate * 1000.0) as i64;
312 now_timestamp() + wait_time
313 } else {
314 i64::MAX
316 };
317 Ok(RateLimitResult::rejected(0, reset_at))
318 }
319 }
320
321 fn try_acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
322 self.acquire(key)
323 }
324
325 fn reset(&self, key: &str) -> Result<(), RateLimitError> {
326 let mut entries = self
327 .entries
328 .write()
329 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
330 entries.remove(key);
331 Ok(())
332 }
333}
334
335#[derive(Debug)]
336pub enum RateLimitError {
337 KeyNotFound(String),
338 Internal(String),
339}
340
341impl std::fmt::Display for RateLimitError {
342 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
343 match self {
344 RateLimitError::KeyNotFound(key) => write!(f, "Key not found: {}", key),
345 RateLimitError::Internal(msg) => write!(f, "Internal error: {}", msg),
346 }
347 }
348}
349
350impl std::error::Error for RateLimitError {}
351
352impl serde::Serialize for RateLimitError {
353 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
354 where
355 S: serde::Serializer,
356 {
357 serializer.serialize_str(&self.to_string())
358 }
359}
360
361fn now_timestamp() -> i64 {
362 use std::time::{SystemTime, UNIX_EPOCH};
363 SystemTime::now()
364 .duration_since(UNIX_EPOCH)
365 .unwrap_or_default()
366 .as_millis() as i64
367}
368
369fn now_secs() -> i64 {
370 use std::time::{SystemTime, UNIX_EPOCH};
371 SystemTime::now()
372 .duration_since(UNIX_EPOCH)
373 .unwrap_or_default()
374 .as_secs() as i64
375}
376
377pub struct FixedWindowRateLimiter {
400 max_requests: u64,
401 window_size: Duration,
402 entries: Arc<RwLock<HashMap<String, FixedWindowEntry>>>,
403 max_keys: usize,
404}
405
406#[derive(Clone)]
407struct FixedWindowEntry {
408 count: u64,
409 window_start: Instant,
410}
411
412impl FixedWindowRateLimiter {
413 pub fn new(max_requests: u64, window_size: Duration) -> Self {
418 Self {
419 max_requests,
420 window_size,
421 entries: Arc::new(RwLock::new(HashMap::new())),
422 max_keys: DEFAULT_MAX_KEYS,
423 }
424 }
425
426 pub fn with_max_keys(mut self, max_keys: usize) -> Self {
428 self.max_keys = max_keys;
429 self
430 }
431
432 fn enforce_max_keys(&self, entries: &mut HashMap<String, FixedWindowEntry>) {
434 while entries.len() > self.max_keys {
435 let oldest_key = entries
436 .iter()
437 .min_by_key(|(_, e)| e.window_start)
438 .map(|(k, _)| k.clone());
439 match oldest_key {
440 Some(k) => {
441 entries.remove(&k);
442 }
443 None => break,
444 }
445 }
446 }
447}
448
449impl RateLimiter for FixedWindowRateLimiter {
450 fn acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
451 let mut entries = self
452 .entries
453 .write()
454 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
455
456 if entries.len() >= self.max_keys && !entries.contains_key(key) {
457 self.enforce_max_keys(&mut entries);
458 }
459
460 let now = Instant::now();
461 let entry = entries
462 .entry(key.to_string())
463 .or_insert_with(|| FixedWindowEntry {
464 count: 0,
465 window_start: now,
466 });
467
468 if now.duration_since(entry.window_start) >= self.window_size {
470 entry.count = 0;
471 entry.window_start = now;
472 }
473
474 if entry.count < self.max_requests {
475 entry.count += 1;
476 let remaining = self.max_requests - entry.count;
477 let reset_at = now_timestamp() + self.window_size.as_millis() as i64;
478 Ok(RateLimitResult::allowed(remaining, reset_at))
479 } else {
480 let elapsed = now.duration_since(entry.window_start);
482 let remaining_window = self
483 .window_size
484 .checked_sub(elapsed)
485 .unwrap_or(Duration::ZERO);
486 let reset_at = now_timestamp() + remaining_window.as_millis() as i64;
487 Ok(RateLimitResult::rejected(0, reset_at))
488 }
489 }
490
491 fn try_acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
492 self.acquire(key)
493 }
494
495 fn reset(&self, key: &str) -> Result<(), RateLimitError> {
496 let mut entries = self
497 .entries
498 .write()
499 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
500 entries.remove(key);
501 Ok(())
502 }
503}
504
505pub trait DistributedBackend: Send + Sync {
522 fn incr_and_get(
537 &self,
538 key: &str,
539 window_secs: u64,
540 window_start: i64,
541 max_requests: u64,
542 ) -> Result<(u64, i64), RateLimitError>;
543
544 fn get(&self, key: &str) -> Result<Option<(u64, i64)>, RateLimitError>;
546
547 fn reset_key(&self, key: &str) -> Result<(), RateLimitError>;
549}
550
551pub struct InMemoryBackend {
556 entries: RwLock<HashMap<String, (u64, i64)>>, }
558
559impl InMemoryBackend {
560 pub fn new() -> Self {
561 Self {
562 entries: RwLock::new(HashMap::new()),
563 }
564 }
565}
566
567impl Default for InMemoryBackend {
568 fn default() -> Self {
569 Self::new()
570 }
571}
572
573impl DistributedBackend for InMemoryBackend {
574 fn incr_and_get(
575 &self,
576 key: &str,
577 window_secs: u64,
578 window_start: i64,
579 _max_requests: u64,
580 ) -> Result<(u64, i64), RateLimitError> {
581 let mut entries = self
582 .entries
583 .write()
584 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
585
586 let entry = entries
587 .entry(key.to_string())
588 .or_insert_with(|| (0, window_start));
589
590 if window_start - entry.1 >= window_secs as i64 {
592 *entry = (0, window_start);
594 }
595
596 entry.0 += 1;
597 let reset_at = entry.1 + window_secs as i64;
598 Ok((entry.0, reset_at))
599 }
600
601 fn get(&self, key: &str) -> Result<Option<(u64, i64)>, RateLimitError> {
602 let entries = self
603 .entries
604 .read()
605 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
606 Ok(entries.get(key).copied())
607 }
608
609 fn reset_key(&self, key: &str) -> Result<(), RateLimitError> {
610 let mut entries = self
611 .entries
612 .write()
613 .map_err(|e| RateLimitError::Internal(e.to_string()))?;
614 entries.remove(key);
615 Ok(())
616 }
617}
618
619pub struct DistributedRateLimiter {
624 backend: Arc<dyn DistributedBackend>,
625 max_requests: u64,
626 window_secs: u64,
627}
628
629impl DistributedRateLimiter {
630 pub fn new(backend: Arc<dyn DistributedBackend>, max_requests: u64, window_secs: u64) -> Self {
636 Self {
637 backend,
638 max_requests,
639 window_secs,
640 }
641 }
642
643 pub fn in_memory(max_requests: u64, window_secs: u64) -> Self {
645 Self::new(Arc::new(InMemoryBackend::new()), max_requests, window_secs)
646 }
647}
648
649impl RateLimiter for DistributedRateLimiter {
650 fn acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
651 let window_start = now_secs();
652 let (count, reset_at) =
653 self.backend
654 .incr_and_get(key, self.window_secs, window_start, self.max_requests)?;
655
656 if count <= self.max_requests {
657 let remaining = self.max_requests - count;
658 Ok(RateLimitResult::allowed(remaining, reset_at * 1000))
659 } else {
660 Ok(RateLimitResult::rejected(0, reset_at * 1000))
661 }
662 }
663
664 fn try_acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
665 self.acquire(key)
666 }
667
668 fn reset(&self, key: &str) -> Result<(), RateLimitError> {
669 self.backend.reset_key(key)
670 }
671}
672
673#[derive(Debug, Clone)]
691pub struct RateLimitHeaders {
692 pub limit: u64,
694 pub remaining: u64,
696 pub reset: i64,
698 pub retry_after: Option<u64>,
700}
701
702impl RateLimitHeaders {
703 pub fn from_result(result: &RateLimitResult, limit: u64) -> Self {
708 let reset_secs = result.reset_at / 1000;
709 let now_secs_val = now_secs();
710 let retry_after = if !result.allowed {
711 let diff = reset_secs - now_secs_val;
712 if diff > 0 {
713 Some(diff as u64)
714 } else {
715 Some(1)
716 }
717 } else {
718 None
719 };
720
721 Self {
722 limit,
723 remaining: result.remaining,
724 reset: reset_secs,
725 retry_after,
726 }
727 }
728
729 pub fn to_headers(&self) -> Vec<(String, String)> {
731 let mut headers = vec![
732 ("X-RateLimit-Limit".to_string(), self.limit.to_string()),
733 (
734 "X-RateLimit-Remaining".to_string(),
735 self.remaining.to_string(),
736 ),
737 ("X-RateLimit-Reset".to_string(), self.reset.to_string()),
738 ];
739 if let Some(retry) = self.retry_after {
740 headers.push(("Retry-After".to_string(), retry.to_string()));
741 }
742 headers
743 }
744
745 pub fn to_json(&self) -> serde_json::Value {
747 let mut map = serde_json::json!({
748 "X-RateLimit-Limit": self.limit,
749 "X-RateLimit-Remaining": self.remaining,
750 "X-RateLimit-Reset": self.reset,
751 });
752 if let Some(retry) = self.retry_after {
753 map["Retry-After"] = serde_json::json!(retry);
754 }
755 map
756 }
757}
758
759#[derive(Debug, Clone)]
763pub enum RateLimitResponseStrategy {
764 TooManyRequests,
766 ServiceUnavailable,
768 Custom(u16),
770}
771
772impl RateLimitResponseStrategy {
773 pub fn status_code(&self) -> u16 {
775 match self {
776 RateLimitResponseStrategy::TooManyRequests => 429,
777 RateLimitResponseStrategy::ServiceUnavailable => 503,
778 RateLimitResponseStrategy::Custom(code) => *code,
779 }
780 }
781}
782
783#[derive(Debug, Clone)]
787pub struct RateLimitResponse {
788 pub status_code: u16,
790 pub headers: RateLimitHeaders,
792 pub body: serde_json::Value,
794}
795
796impl RateLimitResponse {
797 pub fn rejected(
803 result: &RateLimitResult,
804 limit: u64,
805 strategy: RateLimitResponseStrategy,
806 ) -> Self {
807 let headers = RateLimitHeaders::from_result(result, limit);
808 let status_code = strategy.status_code();
809 let body = serde_json::json!({
810 "error": "rate_limit_exceeded",
811 "message": "Rate limit exceeded. Please retry later.",
812 "retry_after": headers.retry_after.unwrap_or(1),
813 });
814
815 Self {
816 status_code,
817 headers,
818 body,
819 }
820 }
821
822 pub fn allowed(result: &RateLimitResult, limit: u64) -> Self {
824 let headers = RateLimitHeaders::from_result(result, limit);
825 Self {
826 status_code: 200,
827 headers,
828 body: serde_json::Value::Null,
829 }
830 }
831}
832
833pub struct MultiRateLimiter {
844 limiters: Vec<Arc<dyn RateLimiter>>,
845}
846
847impl MultiRateLimiter {
848 pub fn new(limiters: Vec<Arc<dyn RateLimiter>>) -> Self {
850 Self { limiters }
851 }
852
853 pub fn with_limiter(mut self, limiter: Arc<dyn RateLimiter>) -> Self {
855 self.limiters.push(limiter);
856 self
857 }
858
859 pub fn check_all(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
864 let mut best_result: Option<RateLimitResult> = None;
865 for limiter in &self.limiters {
866 let result = limiter.acquire(key)?;
867 match &best_result {
868 None => best_result = Some(result),
869 Some(current) => {
870 if !result.allowed {
872 if !current.allowed {
874 if result.remaining <= current.remaining {
876 best_result = Some(result);
877 }
878 } else {
879 best_result = Some(result);
881 }
882 } else if current.allowed && result.remaining < current.remaining {
883 best_result = Some(result);
885 }
886 }
887 }
888 }
889
890 best_result.ok_or_else(|| RateLimitError::Internal("No limiters configured".to_string()))
891 }
892}
893
894#[cfg(feature = "prod-rate-limit-tuning")]
899mod prod {
900 use super::DEFAULT_MAX_KEYS;
901 use serde::{Deserialize, Serialize};
902 use std::time::Duration;
903
904 #[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
906 pub enum RateLimitProdError {
907 #[error("rate limit capacity must be positive")]
908 CapacityNotPositive,
909 #[error("rate limit rate must be positive")]
910 RateNotPositive,
911 #[error("rate limit window_size must be positive")]
912 WindowSizeNotPositive,
913 #[error("rate limit max_keys too small, minimum 100 recommended")]
914 MaxKeysTooSmall,
915 }
916
917 #[derive(Debug, Clone, Serialize, Deserialize)]
919 pub struct RateLimitProdConfig {
920 pub capacity: u64,
921 pub rate: u64,
922 pub window_size: Duration,
923 pub max_keys: usize,
924 }
925
926 impl Default for RateLimitProdConfig {
927 fn default() -> Self {
928 Self {
929 capacity: 100,
930 rate: 10,
931 window_size: Duration::from_secs(1),
932 max_keys: DEFAULT_MAX_KEYS,
933 }
934 }
935 }
936
937 impl RateLimitProdConfig {
938 pub fn new(capacity: u64, rate: u64, window_size: Duration, max_keys: usize) -> Self {
939 Self {
940 capacity,
941 rate,
942 window_size,
943 max_keys,
944 }
945 }
946
947 pub fn validate(&self) -> Result<(), RateLimitProdError> {
949 if self.capacity == 0 {
950 return Err(RateLimitProdError::CapacityNotPositive);
951 }
952 if self.rate == 0 {
953 return Err(RateLimitProdError::RateNotPositive);
954 }
955 if self.window_size.is_zero() {
956 return Err(RateLimitProdError::WindowSizeNotPositive);
957 }
958 if self.max_keys < 100 {
959 return Err(RateLimitProdError::MaxKeysTooSmall);
960 }
961 Ok(())
962 }
963 }
964
965 #[derive(Debug, Clone, Serialize, Deserialize)]
967 pub struct RateLimitStats {
968 pub capacity: u64,
969 pub allowed_count: u64,
970 pub rejected_count: u64,
971 }
972}
973
974#[cfg(feature = "prod-rate-limit-tuning")]
975pub use prod::{RateLimitProdConfig, RateLimitProdError, RateLimitStats};
976
977#[cfg(test)]
978mod tests {
979 use super::*;
980
981 #[test]
982 fn test_rate_limit_result_allowed() {
983 let result = RateLimitResult::allowed(5, 1000);
984 assert!(result.allowed);
985 assert_eq!(result.remaining, 5);
986 assert_eq!(result.reset_at, 1000);
987 }
988
989 #[test]
990 fn test_rate_limit_result_rejected() {
991 let result = RateLimitResult::rejected(0, 2000);
992 assert!(!result.allowed);
993 assert_eq!(result.remaining, 0);
994 }
995
996 #[test]
997 fn test_sliding_window_limiter_new() {
998 let limiter = SlidingWindowRateLimiter::new(10, Duration::from_secs(60));
999 let result = limiter.acquire("test-key");
1000 assert!(result.is_ok());
1001 assert!(result.unwrap().allowed);
1002 }
1003
1004 #[test]
1005 fn test_sliding_window_limiter_full() {
1006 let limiter = SlidingWindowRateLimiter::new(2, Duration::from_secs(60));
1007
1008 let r1 = limiter.acquire("key1").unwrap();
1009 assert!(r1.allowed);
1010
1011 let r2 = limiter.acquire("key1").unwrap();
1012 assert!(r2.allowed);
1013
1014 let r3 = limiter.acquire("key1").unwrap();
1015 assert!(!r3.allowed);
1016 }
1017
1018 #[test]
1019 fn test_sliding_window_different_keys() {
1020 let limiter = SlidingWindowRateLimiter::new(1, Duration::from_secs(60));
1021
1022 let r1 = limiter.acquire("key-a").unwrap();
1023 assert!(r1.allowed);
1024
1025 let r2 = limiter.acquire("key-b").unwrap();
1026 assert!(r2.allowed);
1027 }
1028
1029 #[test]
1030 fn test_sliding_window_reset() {
1031 let limiter = SlidingWindowRateLimiter::new(1, Duration::from_secs(60));
1032
1033 limiter.acquire("reset-key").unwrap();
1034 limiter.acquire("reset-key").unwrap();
1035
1036 limiter.reset("reset-key").unwrap();
1037
1038 let result = limiter.acquire("reset-key").unwrap();
1039 assert!(result.allowed);
1040 }
1041
1042 #[test]
1043 fn test_token_bucket_limiter_new() {
1044 let limiter = TokenBucketRateLimiter::new(10, 1.0);
1045 let result = limiter.acquire("test-key");
1046 assert!(result.is_ok());
1047 assert!(result.unwrap().allowed);
1048 }
1049
1050 #[test]
1051 fn test_token_bucket_limiter_depletes() {
1052 let limiter = TokenBucketRateLimiter::new(2, 1.0);
1053
1054 let r1 = limiter.acquire("key1").unwrap();
1055 assert!(r1.allowed);
1056 assert_eq!(r1.remaining, 1);
1057
1058 let r2 = limiter.acquire("key1").unwrap();
1059 assert!(r2.allowed);
1060 assert_eq!(r2.remaining, 0);
1061
1062 let r3 = limiter.acquire("key1").unwrap();
1063 assert!(!r3.allowed);
1064 }
1065
1066 #[test]
1067 fn test_token_bucket_different_keys() {
1068 let limiter = TokenBucketRateLimiter::new(1, 1.0);
1069
1070 let r1 = limiter.acquire("key-a").unwrap();
1071 assert!(r1.allowed);
1072
1073 let r2 = limiter.acquire("key-b").unwrap();
1074 assert!(r2.allowed);
1075 }
1076
1077 #[test]
1078 fn test_token_bucket_reset() {
1079 let limiter = TokenBucketRateLimiter::new(1, 1.0);
1080
1081 limiter.acquire("reset-key").unwrap();
1082 limiter.acquire("reset-key").unwrap();
1083
1084 limiter.reset("reset-key").unwrap();
1085
1086 let result = limiter.acquire("reset-key").unwrap();
1087 assert!(result.allowed);
1088 }
1089
1090 #[test]
1091 fn test_limiter_try_acquire() {
1092 let limiter = SlidingWindowRateLimiter::new(1, Duration::from_secs(60));
1093
1094 let r1 = limiter.try_acquire("key").unwrap();
1095 assert!(r1.allowed);
1096
1097 let r2 = limiter.try_acquire("key").unwrap();
1098 assert!(!r2.allowed);
1099 }
1100
1101 #[test]
1104 fn test_token_bucket_zero_refill_rate_does_not_panic() {
1105 let limiter = TokenBucketRateLimiter::new(1, 0.0);
1108
1109 let r1 = limiter.acquire("zero-refill").unwrap();
1110 assert!(r1.allowed, "first acquire should be allowed");
1111
1112 let r2 = limiter.acquire("zero-refill").unwrap();
1114 assert!(!r2.allowed, "second acquire should be rejected");
1115 assert!(
1117 r2.reset_at > 0,
1118 "reset_at should be a valid timestamp, got: {}",
1119 r2.reset_at
1120 );
1121 }
1122
1123 #[test]
1124 fn test_token_bucket_negative_refill_rate_does_not_panic() {
1125 let limiter = TokenBucketRateLimiter::new(1, -1.0);
1127
1128 let r1 = limiter.acquire("neg-refill").unwrap();
1129 assert!(r1.allowed, "first acquire should be allowed");
1130
1131 let r2 = limiter.acquire("neg-refill").unwrap();
1132 assert!(!r2.allowed, "second acquire should be rejected");
1133 assert!(
1134 r2.reset_at > 0,
1135 "reset_at should be a valid timestamp, got: {}",
1136 r2.reset_at
1137 );
1138 }
1139
1140 #[test]
1143 fn test_fixed_window_limiter_allows_within_limit() {
1144 let limiter = FixedWindowRateLimiter::new(5, Duration::from_secs(60));
1145 for i in 0..5 {
1146 let r = limiter.acquire("key").unwrap();
1147 assert!(r.allowed, "request {} should be allowed", i);
1148 }
1149 }
1150
1151 #[test]
1152 fn test_fixed_window_limiter_rejects_over_limit() {
1153 let limiter = FixedWindowRateLimiter::new(2, Duration::from_secs(60));
1154 assert!(limiter.acquire("key").unwrap().allowed);
1155 assert!(limiter.acquire("key").unwrap().allowed);
1156 let r3 = limiter.acquire("key").unwrap();
1157 assert!(!r3.allowed);
1158 assert_eq!(r3.remaining, 0);
1159 }
1160
1161 #[test]
1162 fn test_fixed_window_limiter_different_keys() {
1163 let limiter = FixedWindowRateLimiter::new(1, Duration::from_secs(60));
1164 assert!(limiter.acquire("key-a").unwrap().allowed);
1165 assert!(limiter.acquire("key-b").unwrap().allowed);
1166 }
1167
1168 #[test]
1169 fn test_fixed_window_limiter_reset() {
1170 let limiter = FixedWindowRateLimiter::new(1, Duration::from_secs(60));
1171 limiter.acquire("key").unwrap();
1172 assert!(!limiter.acquire("key").unwrap().allowed);
1173 limiter.reset("key").unwrap();
1174 assert!(limiter.acquire("key").unwrap().allowed);
1175 }
1176
1177 #[test]
1178 fn test_fixed_window_limiter_remaining_decreases() {
1179 let limiter = FixedWindowRateLimiter::new(3, Duration::from_secs(60));
1180 let r1 = limiter.acquire("key").unwrap();
1181 assert_eq!(r1.remaining, 2);
1182 let r2 = limiter.acquire("key").unwrap();
1183 assert_eq!(r2.remaining, 1);
1184 let r3 = limiter.acquire("key").unwrap();
1185 assert_eq!(r3.remaining, 0);
1186 }
1187
1188 #[test]
1189 fn test_fixed_window_limiter_reset_at_positive() {
1190 let limiter = FixedWindowRateLimiter::new(1, Duration::from_secs(60));
1191 let r = limiter.acquire("key").unwrap();
1192 assert!(r.reset_at > 0);
1193 }
1194
1195 #[test]
1196 fn test_fixed_window_limiter_try_acquire() {
1197 let limiter = FixedWindowRateLimiter::new(1, Duration::from_secs(60));
1198 assert!(limiter.try_acquire("key").unwrap().allowed);
1199 assert!(!limiter.try_acquire("key").unwrap().allowed);
1200 }
1201
1202 #[test]
1205 fn test_in_memory_backend_new() {
1206 let backend = InMemoryBackend::new();
1207 let result = backend.get("key").unwrap();
1208 assert!(result.is_none());
1209 }
1210
1211 #[test]
1212 fn test_in_memory_backend_incr_and_get() {
1213 let backend = InMemoryBackend::new();
1214 let now = now_secs();
1215 let (count1, reset1) = backend.incr_and_get("key", 60, now, 10).unwrap();
1216 assert_eq!(count1, 1);
1217 assert_eq!(reset1, now + 60);
1218
1219 let (count2, _) = backend.incr_and_get("key", 60, now, 10).unwrap();
1220 assert_eq!(count2, 2);
1221 }
1222
1223 #[test]
1224 fn test_in_memory_backend_get() {
1225 let backend = InMemoryBackend::new();
1226 let now = now_secs();
1227 backend.incr_and_get("key", 60, now, 10).unwrap();
1228 let result = backend.get("key").unwrap();
1229 assert!(result.is_some());
1230 assert_eq!(result.unwrap().0, 1);
1231 }
1232
1233 #[test]
1234 fn test_in_memory_backend_reset_key() {
1235 let backend = InMemoryBackend::new();
1236 let now = now_secs();
1237 backend.incr_and_get("key", 60, now, 10).unwrap();
1238 assert!(backend.get("key").unwrap().is_some());
1239 backend.reset_key("key").unwrap();
1240 assert!(backend.get("key").unwrap().is_none());
1241 }
1242
1243 #[test]
1244 fn test_in_memory_backend_window_expiry() {
1245 let backend = InMemoryBackend::new();
1246 let now = now_secs();
1247 backend.incr_and_get("key", 60, now, 10).unwrap();
1249 backend.incr_and_get("key", 60, now, 10).unwrap();
1250 let (count, _) = backend.incr_and_get("key", 60, now + 61, 10).unwrap();
1252 assert_eq!(count, 1);
1253 }
1254
1255 #[test]
1256 fn test_distributed_rate_limiter_allows() {
1257 let limiter = DistributedRateLimiter::in_memory(5, 60);
1258 for i in 0..5 {
1259 let r = limiter.acquire("key").unwrap();
1260 assert!(r.allowed, "request {} should be allowed", i);
1261 }
1262 }
1263
1264 #[test]
1265 fn test_distributed_rate_limiter_rejects() {
1266 let limiter = DistributedRateLimiter::in_memory(2, 60);
1267 assert!(limiter.acquire("key").unwrap().allowed);
1268 assert!(limiter.acquire("key").unwrap().allowed);
1269 assert!(!limiter.acquire("key").unwrap().allowed);
1270 }
1271
1272 #[test]
1273 fn test_distributed_rate_limiter_shared_backend() {
1274 let backend = Arc::new(InMemoryBackend::new());
1276 let limiter1 = DistributedRateLimiter::new(backend.clone(), 2, 60);
1277 let limiter2 = DistributedRateLimiter::new(backend.clone(), 2, 60);
1278
1279 assert!(limiter1.acquire("key").unwrap().allowed);
1281 assert!(limiter2.acquire("key").unwrap().allowed);
1283 assert!(!limiter1.acquire("key").unwrap().allowed);
1285 }
1286
1287 #[test]
1288 fn test_distributed_rate_limiter_reset() {
1289 let limiter = DistributedRateLimiter::in_memory(1, 60);
1290 limiter.acquire("key").unwrap();
1291 assert!(!limiter.acquire("key").unwrap().allowed);
1292 limiter.reset("key").unwrap();
1293 assert!(limiter.acquire("key").unwrap().allowed);
1294 }
1295
1296 #[test]
1297 fn test_distributed_rate_limiter_different_keys() {
1298 let limiter = DistributedRateLimiter::in_memory(1, 60);
1299 assert!(limiter.acquire("key-a").unwrap().allowed);
1300 assert!(limiter.acquire("key-b").unwrap().allowed);
1301 }
1302
1303 #[test]
1304 fn test_distributed_rate_limiter_remaining() {
1305 let limiter = DistributedRateLimiter::in_memory(3, 60);
1306 let r1 = limiter.acquire("key").unwrap();
1307 assert_eq!(r1.remaining, 2);
1308 let r2 = limiter.acquire("key").unwrap();
1309 assert_eq!(r2.remaining, 1);
1310 }
1311
1312 #[test]
1315 fn test_rate_limit_headers_allowed() {
1316 let result = RateLimitResult::allowed(5, now_secs() * 1000 + 60000);
1317 let headers = RateLimitHeaders::from_result(&result, 10);
1318 assert_eq!(headers.limit, 10);
1319 assert_eq!(headers.remaining, 5);
1320 assert!(headers.retry_after.is_none());
1321 }
1322
1323 #[test]
1324 fn test_rate_limit_headers_rejected() {
1325 let result = RateLimitResult::rejected(0, now_secs() * 1000 + 60000);
1326 let headers = RateLimitHeaders::from_result(&result, 10);
1327 assert_eq!(headers.limit, 10);
1328 assert_eq!(headers.remaining, 0);
1329 assert!(headers.retry_after.is_some());
1330 assert!(headers.retry_after.unwrap() > 0);
1331 }
1332
1333 #[test]
1334 fn test_rate_limit_headers_to_headers_allowed() {
1335 let result = RateLimitResult::allowed(5, now_secs() * 1000 + 60000);
1336 let headers = RateLimitHeaders::from_result(&result, 10);
1337 let hdrs = headers.to_headers();
1338 assert_eq!(hdrs.len(), 3); assert!(hdrs
1340 .iter()
1341 .any(|(k, v)| k == "X-RateLimit-Limit" && v == "10"));
1342 assert!(hdrs
1343 .iter()
1344 .any(|(k, v)| k == "X-RateLimit-Remaining" && v == "5"));
1345 }
1346
1347 #[test]
1348 fn test_rate_limit_headers_to_headers_rejected() {
1349 let result = RateLimitResult::rejected(0, now_secs() * 1000 + 60000);
1350 let headers = RateLimitHeaders::from_result(&result, 10);
1351 let hdrs = headers.to_headers();
1352 assert_eq!(hdrs.len(), 4); assert!(hdrs.iter().any(|(k, _)| k == "Retry-After"));
1354 }
1355
1356 #[test]
1357 fn test_rate_limit_headers_to_json() {
1358 let result = RateLimitResult::allowed(5, now_secs() * 1000 + 60000);
1359 let headers = RateLimitHeaders::from_result(&result, 10);
1360 let json = headers.to_json();
1361 assert_eq!(json["X-RateLimit-Limit"], 10);
1362 assert_eq!(json["X-RateLimit-Remaining"], 5);
1363 assert!(json.get("Retry-After").is_none());
1364 }
1365
1366 #[test]
1367 fn test_rate_limit_headers_to_json_rejected() {
1368 let result = RateLimitResult::rejected(0, now_secs() * 1000 + 60000);
1369 let headers = RateLimitHeaders::from_result(&result, 10);
1370 let json = headers.to_json();
1371 assert!(json.get("Retry-After").is_some());
1372 }
1373
1374 #[test]
1375 fn test_rate_limit_response_strategy_status_codes() {
1376 assert_eq!(
1377 RateLimitResponseStrategy::TooManyRequests.status_code(),
1378 429
1379 );
1380 assert_eq!(
1381 RateLimitResponseStrategy::ServiceUnavailable.status_code(),
1382 503
1383 );
1384 assert_eq!(RateLimitResponseStrategy::Custom(502).status_code(), 502);
1385 }
1386
1387 #[test]
1388 fn test_rate_limit_response_rejected() {
1389 let result = RateLimitResult::rejected(0, now_secs() * 1000 + 60000);
1390 let response =
1391 RateLimitResponse::rejected(&result, 10, RateLimitResponseStrategy::TooManyRequests);
1392 assert_eq!(response.status_code, 429);
1393 assert_eq!(response.body["error"], "rate_limit_exceeded");
1394 assert!(response.body["retry_after"].as_u64().unwrap() > 0);
1395 assert!(response.headers.retry_after.is_some());
1396 }
1397
1398 #[test]
1399 fn test_rate_limit_response_allowed() {
1400 let result = RateLimitResult::allowed(5, now_secs() * 1000 + 60000);
1401 let response = RateLimitResponse::allowed(&result, 10);
1402 assert_eq!(response.status_code, 200);
1403 assert!(response.body.is_null());
1404 assert!(response.headers.retry_after.is_none());
1405 }
1406
1407 #[test]
1408 fn test_rate_limit_response_custom_strategy() {
1409 let result = RateLimitResult::rejected(0, now_secs() * 1000 + 60000);
1410 let response =
1411 RateLimitResponse::rejected(&result, 10, RateLimitResponseStrategy::Custom(503));
1412 assert_eq!(response.status_code, 503);
1413 }
1414
1415 #[test]
1418 fn test_multi_rate_limiter_all_allowed() {
1419 let l1 = Arc::new(SlidingWindowRateLimiter::new(10, Duration::from_secs(60)));
1420 let l2 = Arc::new(TokenBucketRateLimiter::new(10, 1.0));
1421 let multi = MultiRateLimiter::new(vec![l1, l2]);
1422
1423 let result = multi.check_all("key").unwrap();
1424 assert!(result.allowed);
1425 }
1426
1427 #[test]
1428 fn test_multi_rate_limiter_one_rejects() {
1429 let l1 = Arc::new(SlidingWindowRateLimiter::new(10, Duration::from_secs(60)));
1430 let l2 = Arc::new(TokenBucketRateLimiter::new(1, 1.0));
1431 let multi = MultiRateLimiter::new(vec![l1, l2]);
1432
1433 assert!(multi.check_all("key").unwrap().allowed);
1435 let result = multi.check_all("key").unwrap();
1437 assert!(!result.allowed);
1438 }
1439
1440 #[test]
1441 fn test_multi_rate_limiter_takes_strictest() {
1442 let l1 = Arc::new(SlidingWindowRateLimiter::new(5, Duration::from_secs(60)));
1443 let l2 = Arc::new(SlidingWindowRateLimiter::new(2, Duration::from_secs(60)));
1444 let multi = MultiRateLimiter::new(vec![l1, l2]);
1445
1446 multi.check_all("key").unwrap();
1448 multi.check_all("key").unwrap();
1449 assert!(!multi.check_all("key").unwrap().allowed);
1451 }
1452
1453 #[test]
1454 fn test_multi_rate_limiter_with_limiter() {
1455 let l1 = Arc::new(SlidingWindowRateLimiter::new(10, Duration::from_secs(60)));
1456 let multi = MultiRateLimiter::new(vec![]).with_limiter(l1);
1457 assert!(multi.check_all("key").unwrap().allowed);
1458 }
1459
1460 #[test]
1461 fn test_multi_rate_limiter_empty_errors() {
1462 let multi = MultiRateLimiter::new(vec![]);
1463 let result = multi.check_all("key");
1464 assert!(result.is_err());
1465 }
1466}
1467
1468#[cfg(all(test, feature = "prod-rate-limit-tuning"))]
1469mod prod_tests {
1470 use super::*;
1471
1472 #[test]
1473 fn test_rate_limit_prod_config_validate_ok() {
1474 let config = RateLimitProdConfig::new(100, 10, Duration::from_secs(1), 1000);
1475 assert!(config.validate().is_ok());
1476 }
1477
1478 #[test]
1479 fn test_rate_limit_prod_config_capacity_zero_rejected() {
1480 let config = RateLimitProdConfig::new(0, 10, Duration::from_secs(1), 1000);
1481 let err = config.validate().unwrap_err();
1482 assert!(err.to_string().contains("capacity must be positive"));
1483 }
1484
1485 #[test]
1486 fn test_rate_limit_prod_config_rate_zero_rejected() {
1487 let config = RateLimitProdConfig::new(100, 0, Duration::from_secs(1), 1000);
1488 assert!(config.validate().is_err());
1489 }
1490
1491 #[test]
1492 fn test_rate_limit_prod_config_max_keys_too_small() {
1493 let config = RateLimitProdConfig::new(100, 10, Duration::from_secs(1), 50);
1494 let err = config.validate().unwrap_err();
1495 assert!(err.to_string().contains("max_keys too small"));
1496 }
1497
1498 #[test]
1499 fn test_sliding_window_set_capacity() {
1500 let limiter = SlidingWindowRateLimiter::new(100, Duration::from_secs(60));
1501 assert_eq!(limiter.capacity(), 100);
1502 limiter.set_capacity(200);
1503 assert_eq!(limiter.capacity(), 200);
1504 }
1505
1506 #[test]
1507 fn test_sliding_window_set_rate() {
1508 let limiter = SlidingWindowRateLimiter::new(100, Duration::from_secs(10));
1509 limiter.set_rate(20);
1510 assert_eq!(limiter.capacity(), 200);
1512 }
1513
1514 #[test]
1515 fn test_sliding_window_stats_after_acquire() {
1516 let limiter = SlidingWindowRateLimiter::new(2, Duration::from_secs(60));
1517 limiter.acquire("k1").unwrap();
1518 limiter.acquire("k1").unwrap();
1519 limiter.acquire("k1").unwrap(); let stats = limiter.stats();
1521 assert_eq!(stats.capacity, 2);
1522 assert_eq!(stats.allowed_count, 2);
1523 assert_eq!(stats.rejected_count, 1);
1524 }
1525
1526 #[test]
1527 fn test_sliding_window_dynamic_capacity_takes_effect() {
1528 let limiter = SlidingWindowRateLimiter::new(2, Duration::from_secs(60));
1529 limiter.acquire("k").unwrap();
1530 limiter.acquire("k").unwrap();
1531 assert!(!limiter.acquire("k").unwrap().allowed);
1532 limiter.set_capacity(5);
1533 assert!(limiter.acquire("k").unwrap().allowed);
1534 }
1535}