Skip to main content

sz_orm_limit/
lib.rs

1//! # SZ-ORM Limit — Rate Limiter
2//!
3//! Provides token bucket and sliding window rate limiting, with built-in OOM protection
4//! (default max_keys=10000).
5//!
6//! ## Main types
7//!
8//! - [`RateLimiter`] trait — rate limiter interface
9//! - `TokenBucketLimiter` — token bucket implementation
10//! - `SlidingWindowLimiter` — sliding window implementation
11//!
12//! v0.2.1 fixes Critical S-3: introduces `DEFAULT_MAX_KEYS` to prevent unbounded keys from causing OOM.
13
14pub 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
38/// Default maximum number of keys (added in v0.2.1, fixes Critical S-3 OOM DoS)
39///
40/// When entries.len() exceeds this value, an entry is forcibly evicted.
41/// Callers can adjust this via `with_max_keys()`.
42pub 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    /// Maximum number of keys (added in v0.2.1, fixes Critical S-3 OOM DoS)
88    max_keys: usize,
89    /// v3.8.0: allowed count
90    allowed_count: AtomicU64,
91    /// v3.8.0: rejected count
92    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    /// Configures the maximum number of keys (added in v0.2.1, fixes Critical S-3 OOM DoS)
113    ///
114    /// When entries.len() exceeds `max_keys`, the oldest entry is forcibly evicted.
115    /// The default value is `DEFAULT_MAX_KEYS` (10000).
116    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    /// Forcibly evicts the oldest entries that exceed `max_keys` (added in v0.2.1, fixes Critical S-3)
148    ///
149    /// Strategy: iterates over all entries, finds the one with the smallest `requests[0]`
150    /// (earliest request time within the window), and deletes it.
151    /// Complexity is O(n), but it is only triggered when `entries.len() > max_keys`.
152    fn enforce_max_keys(&self, entries: &mut HashMap<String, SlidingWindowEntry>) {
153        while entries.len() > self.max_keys {
154            // 找到最旧的 entry(requests.first() 时间最早)
155            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    /// v3.8.0: dynamically adjust capacity (max_requests) at runtime
170    #[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    /// v3.8.0: dynamically adjust rate (requests per second) at runtime
176    #[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    /// v3.8.0: query current capacity
184    #[cfg(feature = "prod-rate-limit-tuning")]
185    pub fn capacity(&self) -> u64 {
186        self.max_requests.load(Ordering::Relaxed)
187    }
188
189    /// v3.8.0: query statistics
190    #[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        // v0.2.1 修复 Critical S-3:写入前强制淘汰超出 max_keys 的 entry
208        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    /// Maximum number of keys (added in v0.2.1, fixes Critical S-3 OOM DoS)
263    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    /// Configures the maximum number of keys (added in v0.2.1, fixes Critical S-3 OOM DoS)
283    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        // 修复:refill_rate <= 0.0 时不补充令牌(也不消耗)
311        // 避免负数 refill_rate 导致 tokens 递减
312        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    /// Forcibly evicts the oldest entries that exceed `max_keys` (added in v0.2.1, fixes Critical S-3)
323    ///
324    /// Strategy: finds the entry with the earliest `last_refill` (i.e., the least recently accessed) and deletes it.
325    /// Complexity is O(n), but it is only triggered when `entries.len() > max_keys`.
326    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        // v0.2.1 修复 Critical S-3:写入前强制淘汰超出 max_keys 的 entry
350        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            // 修复:refill_rate <= 0.0 时令牌不补充,reset_at 设为远期时间
367            // 避免除零产生 inf,inf as i64 触发 panic
368            let reset_at = if self.refill_rate > 0.0 {
369                now_timestamp() + (1000.0 / self.refill_rate) as i64
370            } else {
371                // 令牌永不补充,reset_at 设为远期时间(约 292 年后的 i64::MAX)
372                i64::MAX
373            };
374            Ok(RateLimitResult::allowed(remaining, reset_at))
375        } else {
376            // 修复:refill_rate <= 0.0 时令牌不补充,永远等待
377            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                // 令牌永不补充,永远等待
382                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
444// ============================================================================
445// 固定窗口算法(Fixed Window)
446//
447// 经典的固定窗口计数器限流:将时间划分为固定窗口(如每分钟),
448// 每个窗口内维护一个计数器,请求到来时计数器+1,超过阈值则拒绝。
449// 窗口结束时计数器重置。
450//
451// 优点:实现简单、内存占用低
452// 缺点:存在边界突刺问题(窗口切换瞬间可能通过 2 倍阈值的请求)
453// ============================================================================
454
455/// Fixed window rate limiter
456///
457/// Divides time into fixed-size windows; each key has an independent counter within each window.
458/// Requests are rejected when the counter exceeds `max_requests`.
459///
460/// # Boundary burst
461///
462/// The fixed window algorithm has the boundary burst problem: if max_requests requests pass
463/// in the last 1 second before the window ends, and another max_requests requests pass in the
464/// first 1 second of the new window, then 2 * max_requests requests pass within 2 seconds.
465/// For smoother rate limiting, use `SlidingWindowRateLimiter` or `TokenBucketRateLimiter`.
466pub 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    /// Creates a fixed window rate limiter
481    ///
482    /// - `max_requests`: maximum number of requests allowed within each window
483    /// - `window_size`: window size (e.g., 60 seconds)
484    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    /// Configures the maximum number of keys (OOM protection)
494    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    /// Forcibly evicts the oldest entries that exceed `max_keys`
519    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        // 检查窗口是否过期,过期则重置
555        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            // 计算窗口重置时间
567            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
591// ============================================================================
592// 分布式限流(Distributed Rate Limiting)
593//
594// 通过抽象后端存储实现分布式限流,支持多实例共享限流状态。
595// 提供内存后端(InMemoryBackend)模拟 Redis 行为,无需外部依赖。
596//
597// 接口设计参照 Redis 的 INCR + EXPIRE 原子操作模式:
598// 1. INCR key -> 获取当前计数
599// 2. 如果计数 == 1,设置 EXPIRE
600// 3. 根据计数判断是否允许
601// ============================================================================
602
603/// Distributed backend trait
604///
605/// Abstracts a distributed storage backend (such as Redis) and provides atomic counter operations.
606/// The in-memory implementation `InMemoryBackend` can be used for single-machine testing and development.
607pub trait DistributedBackend: Send + Sync {
608    /// Atomically increments and returns the value after increment
609    ///
610    /// If the key does not exist, creates it and returns 1.
611    /// If the key exists and has not expired, increments and returns the new value.
612    /// If the key exists but has expired, resets to 1 and returns.
613    ///
614    /// - `key`: rate limit key
615    /// - `window_secs`: window size (seconds); TTL is set only when the key is newly created or expired
616    /// - `window_start`: current window start time (Unix seconds)
617    /// - `max_requests`: maximum number of requests within the window
618    ///
619    /// Returns `(count, reset_at_secs)`:
620    /// - `count`: the count after increment
621    /// - `reset_at_secs`: window reset time (Unix seconds)
622    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    /// Returns the current count (without incrementing)
631    fn get(&self, key: &str) -> Result<Option<(u64, i64)>, RateLimitError>;
632
633    /// Resets the key (deletes it)
634    fn reset_key(&self, key: &str) -> Result<(), RateLimitError>;
635}
636
637/// In-memory backend (simulates Redis)
638///
639/// Uses `RwLock<HashMap>` to store counters, simulating Redis's INCR + EXPIRE behavior.
640/// Suitable for single-machine scenarios and testing; not suitable for real distributed environments.
641pub struct InMemoryBackend {
642    entries: RwLock<HashMap<String, (u64, i64)>>, // key -> (count, window_start_secs)
643}
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        // 检查窗口是否过期
677        if window_start - entry.1 >= window_secs as i64 {
678            // 窗口过期,重置
679            *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
705/// Distributed rate limiter
706///
707/// Uses `DistributedBackend` to implement cross-instance shared fixed window rate limiting.
708/// Suitable for multi-instance deployment scenarios.
709pub struct DistributedRateLimiter {
710    backend: Arc<dyn DistributedBackend>,
711    max_requests: u64,
712    window_secs: u64,
713}
714
715impl DistributedRateLimiter {
716    /// Creates a distributed rate limiter
717    ///
718    /// - `backend`: distributed backend (e.g., `InMemoryBackend`)
719    /// - `max_requests`: maximum number of requests allowed within each window
720    /// - `window_secs`: window size (seconds)
721    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    /// Creates a distributed rate limiter using an in-memory backend (convenience method)
730    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// ============================================================================
760// 限流响应策略(Rate Limit Response Strategy)
761//
762// 提供标准化的限流响应头和响应体生成,符合 IETF draft-ietf-httpapi-ratelimit-headers
763// 和常见 API 网关(如 Kong、AWS API Gateway)的实践。
764// ============================================================================
765
766/// Rate limit response headers
767///
768/// Generates standard rate limit response headers, suitable for HTTP API responses.
769///
770/// # Standard headers
771///
772/// - `X-RateLimit-Limit`: maximum number of requests within the window
773/// - `X-RateLimit-Remaining`: remaining number of requests
774/// - `X-RateLimit-Reset`: window reset time (Unix seconds)
775/// - `Retry-After`: suggested retry wait time (seconds) when rejected; included only on rejection
776#[derive(Debug, Clone)]
777pub struct RateLimitHeaders {
778    /// Maximum number of requests within the window
779    pub limit: u64,
780    /// Remaining number of requests
781    pub remaining: u64,
782    /// Window reset time (Unix seconds)
783    pub reset: i64,
784    /// Retry wait time (seconds); set only when rejected
785    pub retry_after: Option<u64>,
786}
787
788impl RateLimitHeaders {
789    /// Creates response headers from a rate limit result
790    ///
791    /// - `result`: rate limit result
792    /// - `limit`: maximum number of requests within the window
793    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    /// Converts to HTTP header key-value pairs
816    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    /// Converts to a JSON object
832    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/// Rate limit response strategy
846///
847/// Defines the response behavior when rate limited.
848#[derive(Debug, Clone)]
849pub enum RateLimitResponseStrategy {
850    /// Returns 429 Too Many Requests
851    TooManyRequests,
852    /// Returns 503 Service Unavailable
853    ServiceUnavailable,
854    /// Custom status code
855    Custom(u16),
856}
857
858impl RateLimitResponseStrategy {
859    /// Returns the corresponding HTTP status code
860    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/// Rate limit response
870///
871/// Encapsulates the complete response information when rate limited, including status code,
872/// headers, and response body.
873#[derive(Debug, Clone)]
874pub struct RateLimitResponse {
875    /// HTTP status code
876    pub status_code: u16,
877    /// Response headers
878    pub headers: RateLimitHeaders,
879    /// JSON response body
880    pub body: serde_json::Value,
881}
882
883impl RateLimitResponse {
884    /// Creates a rate-limited response
885    ///
886    /// - `result`: rate limit result (must be rejected)
887    /// - `limit`: maximum number of requests within the window
888    /// - `strategy`: response strategy
889    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    /// Creates an allowed response (contains only headers, no body)
910    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
920// ============================================================================
921// 限流策略组合器(Rate Limit Policy)
922//
923// 允许将多个限流器组合,实现多维度限流(如同时限制 IP 和用户)。
924// ============================================================================
925
926/// Multi-dimensional rate limiting strategy
927///
928/// Applies multiple limiters to the same request; rejects if any limiter rejects.
929/// Suitable for scenarios that simultaneously restrict at both the IP level and the user level.
930pub struct MultiRateLimiter {
931    limiters: Vec<Arc<dyn RateLimiter>>,
932}
933
934impl MultiRateLimiter {
935    /// Creates a multi-dimensional rate limiter
936    pub fn new(limiters: Vec<Arc<dyn RateLimiter>>) -> Self {
937        Self { limiters }
938    }
939
940    /// Adds a rate limiter
941    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    /// Checks all limiters and returns the strictest result
951    ///
952    /// If any limiter rejects, returns a rejected result (with the fewest remaining).
953    /// If all allow, returns the result with the fewest remaining.
954    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                    // 取更严格的结果
962                    if !result.allowed {
963                        // 拒绝优先
964                        if !current.allowed {
965                            // 两个都拒绝,取剩余更少的
966                            if result.remaining <= current.remaining {
967                                best_result = Some(result);
968                            }
969                        } else {
970                            // 当前允许但新的拒绝 -> 用拒绝结果
971                            best_result = Some(result);
972                        }
973                    } else if current.allowed && result.remaining < current.remaining {
974                        // 两个都允许,取剩余更少的
975                        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// ============================================================================
986// v3.8.0: 限流生产配置(prod-rate-limit-tuning feature)
987// ============================================================================
988
989#[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    /// Rate limit production configuration error
996    #[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    /// Rate limit production configuration
1009    #[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        /// Validates the reasonableness of the thresholds
1039        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    /// Rate limit statistics
1057    #[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    // ===== TDD RED:refill_rate=0 panic 修复测试 =====
1193
1194    #[test]
1195    fn test_token_bucket_zero_refill_rate_does_not_panic() {
1196        // refill_rate=0.0 表示令牌永不补充
1197        // capacity=1,第一次 acquire 应允许,第二次应拒绝且不 panic
1198        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        // 第二次 acquire:令牌耗尽且不补充,应拒绝而非 panic
1204        let r2 = limiter.acquire("zero-refill").unwrap();
1205        assert!(!r2.allowed, "second acquire should be rejected");
1206        // reset_at 应为一个合理的远期时间(令牌不补充,永远等待)
1207        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        // refill_rate=-1.0 是错误配置,应被当作 0.0 处理而非 panic
1217        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    // ===== 固定窗口算法测试 =====
1232
1233    #[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    // ===== 分布式限流测试 =====
1294
1295    #[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        // 第一次请求,窗口开始
1339        backend.incr_and_get("key", 60, now, 10).unwrap();
1340        backend.incr_and_get("key", 60, now, 10).unwrap();
1341        // 窗口过期后(now + 61),应重置
1342        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        // 两个限流器共享同一个后端
1366        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        // limiter1 消耗 1 个
1371        assert!(limiter1.acquire("key").unwrap().allowed);
1372        // limiter2 消耗 1 个(共享计数)
1373        assert!(limiter2.acquire("key").unwrap().allowed);
1374        // 第 3 个应被拒绝(共享计数已达 2)
1375        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    // ===== 限流响应策略测试 =====
1404
1405    #[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); // 不包含 Retry-After
1430        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); // 包含 Retry-After
1444        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    // ===== 多维度限流策略测试 =====
1507
1508    #[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        // 第一次通过
1525        assert!(multi.check_all("key").unwrap().allowed);
1526        // 第二次:l2 拒绝
1527        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        // 两次请求后 l2 达到上限
1538        multi.check_all("key").unwrap();
1539        multi.check_all("key").unwrap();
1540        // 第三次:l2 拒绝
1541        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        // rate=20, window=10s → capacity=200
1602        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(); // rejected
1611        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}