Skip to main content

sz_orm_limit/
lib.rs

1//! # SZ-ORM Limit — 限流器
2//!
3//! 提供令牌桶与滑动窗口限流,内置 OOM 防护(默认 max_keys=10000)。
4//!
5//! ## 主要类型
6//!
7//! - [`RateLimiter`] trait — 限流器接口
8//! - `TokenBucketLimiter` — 令牌桶实现
9//! - `SlidingWindowLimiter` — 滑动窗口实现
10//!
11//! v0.2.1 修复 Critical S-3:引入 `DEFAULT_MAX_KEYS` 防止无限 key 导致 OOM。
12
13pub mod composite;
14pub mod concurrency;
15pub mod config_builder;
16pub mod leaky_bucket;
17pub mod metrics;
18
19pub use composite::{
20    CompositeLimiter, CompositeStrategy, FallbackLimiter, LimitKeyBuilder, RateLimitRule, RuleSet,
21};
22pub use concurrency::{
23    ConcurrencyGuard, ConcurrencyLimiter, ConcurrencyStats, TimedConcurrencyLimiter,
24};
25pub use config_builder::{
26    ConfigError, LimitAlgorithm, RateLimitConfig, RateLimitConfigBuilder, TieredConfigBuilder,
27    TieredRateLimitConfig,
28};
29pub use leaky_bucket::{LeakyBucketLimiter, LeakyBucketStats, SlidingWindowLogLimiter};
30pub use metrics::{Alert, MetricsSnapshot, RateLimitMetrics, RateLimitMonitor};
31
32use std::collections::HashMap;
33use std::sync::atomic::{AtomicU64, Ordering};
34use std::sync::{Arc, RwLock};
35use std::time::{Duration, Instant};
36
37/// 默认最大 key 数量(v0.2.1 新增,修复 Critical S-3 OOM DoS)
38///
39/// 当 entries.len() 超过此值时,会强制淘汰一个 entry。
40/// 调用方可通过 `with_max_keys()` 调整。
41pub const DEFAULT_MAX_KEYS: usize = 10_000;
42
43pub trait RateLimiter: Send + Sync {
44    fn acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError>;
45    fn try_acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError>;
46    fn reset(&self, key: &str) -> Result<(), RateLimitError>;
47}
48
49#[derive(Debug, Clone)]
50pub struct RateLimitResult {
51    pub allowed: bool,
52    pub remaining: u64,
53    pub reset_at: i64,
54}
55
56impl RateLimitResult {
57    pub fn allowed(remaining: u64, reset_at: i64) -> Self {
58        Self {
59            allowed: true,
60            remaining,
61            reset_at,
62        }
63    }
64
65    pub fn rejected(remaining: u64, reset_at: i64) -> Self {
66        Self {
67            allowed: false,
68            remaining,
69            reset_at,
70        }
71    }
72
73    pub fn is_allowed(&self) -> bool {
74        self.allowed
75    }
76
77    pub fn is_rejected(&self) -> bool {
78        !self.allowed
79    }
80}
81
82pub struct SlidingWindowRateLimiter {
83    max_requests: Arc<AtomicU64>,
84    window_size: Duration,
85    entries: Arc<RwLock<HashMap<String, SlidingWindowEntry>>>,
86    /// 最大 key 数量(v0.2.1 新增,修复 Critical S-3 OOM DoS)
87    max_keys: usize,
88    /// v3.8.0: 已通过计数
89    allowed_count: AtomicU64,
90    /// v3.8.0: 已拒绝计数
91    rejected_count: AtomicU64,
92}
93
94#[derive(Clone)]
95struct SlidingWindowEntry {
96    requests: Vec<Instant>,
97}
98
99impl SlidingWindowRateLimiter {
100    pub fn new(max_requests: u64, window_size: Duration) -> Self {
101        Self {
102            max_requests: Arc::new(AtomicU64::new(max_requests)),
103            window_size,
104            entries: Arc::new(RwLock::new(HashMap::new())),
105            max_keys: DEFAULT_MAX_KEYS,
106            allowed_count: AtomicU64::new(0),
107            rejected_count: AtomicU64::new(0),
108        }
109    }
110
111    /// 配置最大 key 数量(v0.2.1 新增,修复 Critical S-3 OOM DoS)
112    ///
113    /// 当 entries.len() 超过 `max_keys` 时,会强制淘汰一个最旧的 entry。
114    /// 默认值为 `DEFAULT_MAX_KEYS`(10000)。
115    pub fn with_max_keys(mut self, max_keys: usize) -> Self {
116        self.max_keys = max_keys;
117        self
118    }
119
120    pub fn max_requests(&self) -> u64 {
121        self.max_requests.load(Ordering::Relaxed)
122    }
123
124    pub fn window_size(&self) -> Duration {
125        self.window_size
126    }
127
128    pub fn max_keys(&self) -> usize {
129        self.max_keys
130    }
131
132    pub fn key_count(&self) -> usize {
133        self.entries
134            .read()
135            .map(|entries| entries.len())
136            .unwrap_or(0)
137    }
138
139    fn cleanup_old_requests(&self, entry: &mut SlidingWindowEntry) {
140        let now = Instant::now();
141        entry
142            .requests
143            .retain(|&time| now.duration_since(time) < self.window_size);
144    }
145
146    /// 强制淘汰超出 `max_keys` 的最旧 entry(v0.2.1 新增,修复 Critical S-3)
147    ///
148    /// 策略:遍历所有 entry,找到 `requests[0]`(窗口内最早请求时间)最小的那个并删除。
149    /// 复杂度 O(n),但仅在 `entries.len() > max_keys` 时触发。
150    fn enforce_max_keys(&self, entries: &mut HashMap<String, SlidingWindowEntry>) {
151        while entries.len() > self.max_keys {
152            // 找到最旧的 entry(requests.first() 时间最早)
153            let now = Instant::now();
154            let oldest_key = entries
155                .iter()
156                .min_by_key(|(_, e)| e.requests.first().copied().unwrap_or(now))
157                .map(|(k, _)| k.clone());
158            match oldest_key {
159                Some(k) => {
160                    entries.remove(&k);
161                }
162                None => break,
163            }
164        }
165    }
166
167    /// v3.8.0: 运行时动态调整容量(max_requests)
168    #[cfg(feature = "prod-rate-limit-tuning")]
169    pub fn set_capacity(&self, capacity: u64) {
170        self.max_requests.store(capacity, Ordering::Relaxed);
171    }
172
173    /// v3.8.0: 运行时动态调整速率(requests per second)
174    #[cfg(feature = "prod-rate-limit-tuning")]
175    pub fn set_rate(&self, rate: u64) {
176        let window_secs = self.window_size.as_secs().max(1);
177        self.max_requests
178            .store(rate * window_secs, Ordering::Relaxed);
179    }
180
181    /// v3.8.0: 查询当前容量
182    #[cfg(feature = "prod-rate-limit-tuning")]
183    pub fn capacity(&self) -> u64 {
184        self.max_requests.load(Ordering::Relaxed)
185    }
186
187    /// v3.8.0: 查询统计信息
188    #[cfg(feature = "prod-rate-limit-tuning")]
189    pub fn stats(&self) -> RateLimitStats {
190        RateLimitStats {
191            capacity: self.max_requests.load(Ordering::Relaxed),
192            allowed_count: self.allowed_count.load(Ordering::Relaxed),
193            rejected_count: self.rejected_count.load(Ordering::Relaxed),
194        }
195    }
196}
197
198impl RateLimiter for SlidingWindowRateLimiter {
199    fn acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
200        let mut entries = self
201            .entries
202            .write()
203            .map_err(|e| RateLimitError::Internal(e.to_string()))?;
204
205        // v0.2.1 修复 Critical S-3:写入前强制淘汰超出 max_keys 的 entry
206        if entries.len() >= self.max_keys && !entries.contains_key(key) {
207            self.enforce_max_keys(&mut entries);
208        }
209
210        let entry = entries
211            .entry(key.to_string())
212            .or_insert_with(|| SlidingWindowEntry {
213                requests: Vec::new(),
214            });
215
216        self.cleanup_old_requests(entry);
217
218        let max_req = self.max_requests.load(Ordering::Relaxed);
219        if entry.requests.len() < max_req as usize {
220            entry.requests.push(Instant::now());
221            let remaining = max_req - entry.requests.len() as u64;
222            let reset_at = now_timestamp() + self.window_size.as_millis() as i64;
223            self.allowed_count.fetch_add(1, Ordering::Relaxed);
224            Ok(RateLimitResult::allowed(remaining, reset_at))
225        } else {
226            let oldest = entry
227                .requests
228                .first()
229                .map(|t| {
230                    let elapsed = t.elapsed().as_millis() as i64;
231                    let window_ms = self.window_size.as_millis() as i64;
232                    now_timestamp() + (window_ms - elapsed)
233                })
234                .unwrap_or(now_timestamp());
235
236            let remaining = 0;
237            self.rejected_count.fetch_add(1, Ordering::Relaxed);
238            Ok(RateLimitResult::rejected(remaining, oldest))
239        }
240    }
241
242    fn try_acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
243        self.acquire(key)
244    }
245
246    fn reset(&self, key: &str) -> Result<(), RateLimitError> {
247        let mut entries = self
248            .entries
249            .write()
250            .map_err(|e| RateLimitError::Internal(e.to_string()))?;
251        entries.remove(key);
252        Ok(())
253    }
254}
255
256pub struct TokenBucketRateLimiter {
257    capacity: f64,
258    refill_rate: f64,
259    entries: Arc<RwLock<HashMap<String, TokenBucketEntry>>>,
260    /// 最大 key 数量(v0.2.1 新增,修复 Critical S-3 OOM DoS)
261    max_keys: usize,
262}
263
264#[derive(Clone)]
265struct TokenBucketEntry {
266    tokens: f64,
267    last_refill: Instant,
268}
269
270impl TokenBucketRateLimiter {
271    pub fn new(capacity: u64, refill_per_second: f64) -> Self {
272        Self {
273            capacity: capacity as f64,
274            refill_rate: refill_per_second,
275            entries: Arc::new(RwLock::new(HashMap::new())),
276            max_keys: DEFAULT_MAX_KEYS,
277        }
278    }
279
280    /// 配置最大 key 数量(v0.2.1 新增,修复 Critical S-3 OOM DoS)
281    pub fn with_max_keys(mut self, max_keys: usize) -> Self {
282        self.max_keys = max_keys;
283        self
284    }
285
286    pub fn capacity(&self) -> u64 {
287        self.capacity as u64
288    }
289
290    pub fn refill_rate(&self) -> f64 {
291        self.refill_rate
292    }
293
294    pub fn max_keys(&self) -> usize {
295        self.max_keys
296    }
297
298    pub fn key_count(&self) -> usize {
299        self.entries
300            .read()
301            .map(|entries| entries.len())
302            .unwrap_or(0)
303    }
304
305    fn refill(&self, entry: &mut TokenBucketEntry) {
306        let now = Instant::now();
307        let elapsed = now.duration_since(entry.last_refill).as_secs_f64();
308        // 修复:refill_rate <= 0.0 时不补充令牌(也不消耗)
309        // 避免负数 refill_rate 导致 tokens 递减
310        let tokens_to_add = if self.refill_rate > 0.0 {
311            elapsed * self.refill_rate
312        } else {
313            0.0
314        };
315
316        entry.tokens = (entry.tokens + tokens_to_add).min(self.capacity);
317        entry.last_refill = now;
318    }
319
320    /// 强制淘汰超出 `max_keys` 的最旧 entry(v0.2.1 新增,修复 Critical S-3)
321    ///
322    /// 策略:找到 `last_refill` 最早的 entry(即最久未访问的)并删除。
323    /// 复杂度 O(n),但仅在 `entries.len() > max_keys` 时触发。
324    fn enforce_max_keys(&self, entries: &mut HashMap<String, TokenBucketEntry>) {
325        while entries.len() > self.max_keys {
326            let oldest_key = entries
327                .iter()
328                .min_by_key(|(_, e)| e.last_refill)
329                .map(|(k, _)| k.clone());
330            match oldest_key {
331                Some(k) => {
332                    entries.remove(&k);
333                }
334                None => break,
335            }
336        }
337    }
338}
339
340impl RateLimiter for TokenBucketRateLimiter {
341    fn acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
342        let mut entries = self
343            .entries
344            .write()
345            .map_err(|e| RateLimitError::Internal(e.to_string()))?;
346
347        // v0.2.1 修复 Critical S-3:写入前强制淘汰超出 max_keys 的 entry
348        if entries.len() >= self.max_keys && !entries.contains_key(key) {
349            self.enforce_max_keys(&mut entries);
350        }
351
352        let entry = entries
353            .entry(key.to_string())
354            .or_insert_with(|| TokenBucketEntry {
355                tokens: self.capacity,
356                last_refill: Instant::now(),
357            });
358
359        self.refill(entry);
360
361        if entry.tokens >= 1.0 {
362            entry.tokens -= 1.0;
363            let remaining = entry.tokens.floor() as u64;
364            // 修复:refill_rate <= 0.0 时令牌不补充,reset_at 设为远期时间
365            // 避免除零产生 inf,inf as i64 触发 panic
366            let reset_at = if self.refill_rate > 0.0 {
367                now_timestamp() + (1000.0 / self.refill_rate) as i64
368            } else {
369                // 令牌永不补充,reset_at 设为远期时间(约 292 年后的 i64::MAX)
370                i64::MAX
371            };
372            Ok(RateLimitResult::allowed(remaining, reset_at))
373        } else {
374            // 修复:refill_rate <= 0.0 时令牌不补充,永远等待
375            let reset_at = if self.refill_rate > 0.0 {
376                let wait_time = ((1.0 - entry.tokens) / self.refill_rate * 1000.0) as i64;
377                now_timestamp() + wait_time
378            } else {
379                // 令牌永不补充,永远等待
380                i64::MAX
381            };
382            Ok(RateLimitResult::rejected(0, reset_at))
383        }
384    }
385
386    fn try_acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
387        self.acquire(key)
388    }
389
390    fn reset(&self, key: &str) -> Result<(), RateLimitError> {
391        let mut entries = self
392            .entries
393            .write()
394            .map_err(|e| RateLimitError::Internal(e.to_string()))?;
395        entries.remove(key);
396        Ok(())
397    }
398}
399
400#[derive(Debug)]
401pub enum RateLimitError {
402    KeyNotFound(String),
403    Internal(String),
404}
405
406impl std::fmt::Display for RateLimitError {
407    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
408        match self {
409            RateLimitError::KeyNotFound(key) => write!(f, "Key not found: {}", key),
410            RateLimitError::Internal(msg) => write!(f, "Internal error: {}", msg),
411        }
412    }
413}
414
415impl std::error::Error for RateLimitError {}
416
417impl serde::Serialize for RateLimitError {
418    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
419    where
420        S: serde::Serializer,
421    {
422        serializer.serialize_str(&self.to_string())
423    }
424}
425
426fn now_timestamp() -> i64 {
427    use std::time::{SystemTime, UNIX_EPOCH};
428    SystemTime::now()
429        .duration_since(UNIX_EPOCH)
430        .unwrap_or_default()
431        .as_millis() as i64
432}
433
434fn now_secs() -> i64 {
435    use std::time::{SystemTime, UNIX_EPOCH};
436    SystemTime::now()
437        .duration_since(UNIX_EPOCH)
438        .unwrap_or_default()
439        .as_secs() as i64
440}
441
442// ============================================================================
443// 固定窗口算法(Fixed Window)
444//
445// 经典的固定窗口计数器限流:将时间划分为固定窗口(如每分钟),
446// 每个窗口内维护一个计数器,请求到来时计数器+1,超过阈值则拒绝。
447// 窗口结束时计数器重置。
448//
449// 优点:实现简单、内存占用低
450// 缺点:存在边界突刺问题(窗口切换瞬间可能通过 2 倍阈值的请求)
451// ============================================================================
452
453/// 固定窗口限流器
454///
455/// 将时间划分为固定大小的窗口,每个 key 在每个窗口内有独立的计数器。
456/// 当计数器超过 `max_requests` 时拒绝请求。
457///
458/// # 边界突刺
459///
460/// 固定窗口算法存在边界突刺问题:如果窗口结束前 1 秒通过了 max_requests
461/// 个请求,新窗口开始后第 1 秒又通过了 max_requests 个请求,则 2 秒内
462/// 通过了 2 * max_requests 个请求。如需更平滑的限流,请使用
463/// `SlidingWindowRateLimiter` 或 `TokenBucketRateLimiter`。
464pub struct FixedWindowRateLimiter {
465    max_requests: u64,
466    window_size: Duration,
467    entries: Arc<RwLock<HashMap<String, FixedWindowEntry>>>,
468    max_keys: usize,
469}
470
471#[derive(Clone)]
472struct FixedWindowEntry {
473    count: u64,
474    window_start: Instant,
475}
476
477impl FixedWindowRateLimiter {
478    /// 创建固定窗口限流器
479    ///
480    /// - `max_requests`:每个窗口内允许的最大请求数
481    /// - `window_size`:窗口大小(如 60 秒)
482    pub fn new(max_requests: u64, window_size: Duration) -> Self {
483        Self {
484            max_requests,
485            window_size,
486            entries: Arc::new(RwLock::new(HashMap::new())),
487            max_keys: DEFAULT_MAX_KEYS,
488        }
489    }
490
491    /// 配置最大 key 数量(OOM 防护)
492    pub fn with_max_keys(mut self, max_keys: usize) -> Self {
493        self.max_keys = max_keys;
494        self
495    }
496
497    pub fn max_requests(&self) -> u64 {
498        self.max_requests
499    }
500
501    pub fn window_size(&self) -> Duration {
502        self.window_size
503    }
504
505    pub fn max_keys(&self) -> usize {
506        self.max_keys
507    }
508
509    pub fn key_count(&self) -> usize {
510        self.entries
511            .read()
512            .map(|entries| entries.len())
513            .unwrap_or(0)
514    }
515
516    /// 强制淘汰超出 `max_keys` 的最旧 entry
517    fn enforce_max_keys(&self, entries: &mut HashMap<String, FixedWindowEntry>) {
518        while entries.len() > self.max_keys {
519            let oldest_key = entries
520                .iter()
521                .min_by_key(|(_, e)| e.window_start)
522                .map(|(k, _)| k.clone());
523            match oldest_key {
524                Some(k) => {
525                    entries.remove(&k);
526                }
527                None => break,
528            }
529        }
530    }
531}
532
533impl RateLimiter for FixedWindowRateLimiter {
534    fn acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
535        let mut entries = self
536            .entries
537            .write()
538            .map_err(|e| RateLimitError::Internal(e.to_string()))?;
539
540        if entries.len() >= self.max_keys && !entries.contains_key(key) {
541            self.enforce_max_keys(&mut entries);
542        }
543
544        let now = Instant::now();
545        let entry = entries
546            .entry(key.to_string())
547            .or_insert_with(|| FixedWindowEntry {
548                count: 0,
549                window_start: now,
550            });
551
552        // 检查窗口是否过期,过期则重置
553        if now.duration_since(entry.window_start) >= self.window_size {
554            entry.count = 0;
555            entry.window_start = now;
556        }
557
558        if entry.count < self.max_requests {
559            entry.count += 1;
560            let remaining = self.max_requests - entry.count;
561            let reset_at = now_timestamp() + self.window_size.as_millis() as i64;
562            Ok(RateLimitResult::allowed(remaining, reset_at))
563        } else {
564            // 计算窗口重置时间
565            let elapsed = now.duration_since(entry.window_start);
566            let remaining_window = self
567                .window_size
568                .checked_sub(elapsed)
569                .unwrap_or(Duration::ZERO);
570            let reset_at = now_timestamp() + remaining_window.as_millis() as i64;
571            Ok(RateLimitResult::rejected(0, reset_at))
572        }
573    }
574
575    fn try_acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
576        self.acquire(key)
577    }
578
579    fn reset(&self, key: &str) -> Result<(), RateLimitError> {
580        let mut entries = self
581            .entries
582            .write()
583            .map_err(|e| RateLimitError::Internal(e.to_string()))?;
584        entries.remove(key);
585        Ok(())
586    }
587}
588
589// ============================================================================
590// 分布式限流(Distributed Rate Limiting)
591//
592// 通过抽象后端存储实现分布式限流,支持多实例共享限流状态。
593// 提供内存后端(InMemoryBackend)模拟 Redis 行为,无需外部依赖。
594//
595// 接口设计参照 Redis 的 INCR + EXPIRE 原子操作模式:
596// 1. INCR key -> 获取当前计数
597// 2. 如果计数 == 1,设置 EXPIRE
598// 3. 根据计数判断是否允许
599// ============================================================================
600
601/// 分布式后端 trait
602///
603/// 抽象分布式存储后端(如 Redis),提供原子计数器操作。
604/// 内存实现 `InMemoryBackend` 可用于单机测试和开发。
605pub trait DistributedBackend: Send + Sync {
606    /// 原子递增并返回递增后的值
607    ///
608    /// 如果 key 不存在,创建并返回 1。
609    /// 如果 key 存在且未过期,递增并返回新值。
610    /// 如果 key 存在但已过期,重置为 1 并返回。
611    ///
612    /// - `key`:限流键
613    /// - `window_secs`:窗口大小(秒),仅在 key 新建或过期时设置 TTL
614    /// - `window_start`:当前窗口起始时间(Unix 秒)
615    /// - `max_requests`:窗口内最大请求数
616    ///
617    /// 返回 `(count, reset_at_secs)`:
618    /// - `count`:递增后的计数
619    /// - `reset_at_secs`:窗口重置时间(Unix 秒)
620    fn incr_and_get(
621        &self,
622        key: &str,
623        window_secs: u64,
624        window_start: i64,
625        max_requests: u64,
626    ) -> Result<(u64, i64), RateLimitError>;
627
628    /// 获取当前计数(不递增)
629    fn get(&self, key: &str) -> Result<Option<(u64, i64)>, RateLimitError>;
630
631    /// 重置 key(删除)
632    fn reset_key(&self, key: &str) -> Result<(), RateLimitError>;
633}
634
635/// 内存后端(模拟 Redis)
636///
637/// 使用 `RwLock<HashMap>` 存储计数器,模拟 Redis 的 INCR + EXPIRE 行为。
638/// 适用于单机场景和测试,不适用于真正的分布式环境。
639pub struct InMemoryBackend {
640    entries: RwLock<HashMap<String, (u64, i64)>>, // key -> (count, window_start_secs)
641}
642
643impl InMemoryBackend {
644    pub fn new() -> Self {
645        Self {
646            entries: RwLock::new(HashMap::new()),
647        }
648    }
649}
650
651impl Default for InMemoryBackend {
652    fn default() -> Self {
653        Self::new()
654    }
655}
656
657impl DistributedBackend for InMemoryBackend {
658    fn incr_and_get(
659        &self,
660        key: &str,
661        window_secs: u64,
662        window_start: i64,
663        _max_requests: u64,
664    ) -> Result<(u64, i64), RateLimitError> {
665        let mut entries = self
666            .entries
667            .write()
668            .map_err(|e| RateLimitError::Internal(e.to_string()))?;
669
670        let entry = entries
671            .entry(key.to_string())
672            .or_insert_with(|| (0, window_start));
673
674        // 检查窗口是否过期
675        if window_start - entry.1 >= window_secs as i64 {
676            // 窗口过期,重置
677            *entry = (0, window_start);
678        }
679
680        entry.0 += 1;
681        let reset_at = entry.1 + window_secs as i64;
682        Ok((entry.0, reset_at))
683    }
684
685    fn get(&self, key: &str) -> Result<Option<(u64, i64)>, RateLimitError> {
686        let entries = self
687            .entries
688            .read()
689            .map_err(|e| RateLimitError::Internal(e.to_string()))?;
690        Ok(entries.get(key).copied())
691    }
692
693    fn reset_key(&self, key: &str) -> Result<(), RateLimitError> {
694        let mut entries = self
695            .entries
696            .write()
697            .map_err(|e| RateLimitError::Internal(e.to_string()))?;
698        entries.remove(key);
699        Ok(())
700    }
701}
702
703/// 分布式限流器
704///
705/// 使用 `DistributedBackend` 实现跨实例共享的固定窗口限流。
706/// 适用于多实例部署场景。
707pub struct DistributedRateLimiter {
708    backend: Arc<dyn DistributedBackend>,
709    max_requests: u64,
710    window_secs: u64,
711}
712
713impl DistributedRateLimiter {
714    /// 创建分布式限流器
715    ///
716    /// - `backend`:分布式后端(如 `InMemoryBackend`)
717    /// - `max_requests`:每个窗口内允许的最大请求数
718    /// - `window_secs`:窗口大小(秒)
719    pub fn new(backend: Arc<dyn DistributedBackend>, max_requests: u64, window_secs: u64) -> Self {
720        Self {
721            backend,
722            max_requests,
723            window_secs,
724        }
725    }
726
727    /// 使用内存后端创建分布式限流器(便捷方法)
728    pub fn in_memory(max_requests: u64, window_secs: u64) -> Self {
729        Self::new(Arc::new(InMemoryBackend::new()), max_requests, window_secs)
730    }
731}
732
733impl RateLimiter for DistributedRateLimiter {
734    fn acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
735        let window_start = now_secs();
736        let (count, reset_at) =
737            self.backend
738                .incr_and_get(key, self.window_secs, window_start, self.max_requests)?;
739
740        if count <= self.max_requests {
741            let remaining = self.max_requests - count;
742            Ok(RateLimitResult::allowed(remaining, reset_at * 1000))
743        } else {
744            Ok(RateLimitResult::rejected(0, reset_at * 1000))
745        }
746    }
747
748    fn try_acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
749        self.acquire(key)
750    }
751
752    fn reset(&self, key: &str) -> Result<(), RateLimitError> {
753        self.backend.reset_key(key)
754    }
755}
756
757// ============================================================================
758// 限流响应策略(Rate Limit Response Strategy)
759//
760// 提供标准化的限流响应头和响应体生成,符合 IETF draft-ietf-httpapi-ratelimit-headers
761// 和常见 API 网关(如 Kong、AWS API Gateway)的实践。
762// ============================================================================
763
764/// 限流响应头
765///
766/// 生成标准的限流响应头,适用于 HTTP API 响应。
767///
768/// # 标准头
769///
770/// - `X-RateLimit-Limit`:窗口内最大请求数
771/// - `X-RateLimit-Remaining`:剩余请求数
772/// - `X-RateLimit-Reset`:窗口重置时间(Unix 秒)
773/// - `Retry-After`:拒绝时建议的重试等待时间(秒),仅在拒绝时包含
774#[derive(Debug, Clone)]
775pub struct RateLimitHeaders {
776    /// 窗口内最大请求数
777    pub limit: u64,
778    /// 剩余请求数
779    pub remaining: u64,
780    /// 窗口重置时间(Unix 秒)
781    pub reset: i64,
782    /// 重试等待时间(秒),仅在被拒绝时设置
783    pub retry_after: Option<u64>,
784}
785
786impl RateLimitHeaders {
787    /// 从限流结果创建响应头
788    ///
789    /// - `result`:限流结果
790    /// - `limit`:窗口内最大请求数
791    pub fn from_result(result: &RateLimitResult, limit: u64) -> Self {
792        let reset_secs = result.reset_at / 1000;
793        let now_secs_val = now_secs();
794        let retry_after = if !result.allowed {
795            let diff = reset_secs - now_secs_val;
796            if diff > 0 {
797                Some(diff as u64)
798            } else {
799                Some(1)
800            }
801        } else {
802            None
803        };
804
805        Self {
806            limit,
807            remaining: result.remaining,
808            reset: reset_secs,
809            retry_after,
810        }
811    }
812
813    /// 转换为 HTTP 头键值对
814    pub fn to_headers(&self) -> Vec<(String, String)> {
815        let mut headers = vec![
816            ("X-RateLimit-Limit".to_string(), self.limit.to_string()),
817            (
818                "X-RateLimit-Remaining".to_string(),
819                self.remaining.to_string(),
820            ),
821            ("X-RateLimit-Reset".to_string(), self.reset.to_string()),
822        ];
823        if let Some(retry) = self.retry_after {
824            headers.push(("Retry-After".to_string(), retry.to_string()));
825        }
826        headers
827    }
828
829    /// 转换为 JSON 对象
830    pub fn to_json(&self) -> serde_json::Value {
831        let mut map = serde_json::json!({
832            "X-RateLimit-Limit": self.limit,
833            "X-RateLimit-Remaining": self.remaining,
834            "X-RateLimit-Reset": self.reset,
835        });
836        if let Some(retry) = self.retry_after {
837            map["Retry-After"] = serde_json::json!(retry);
838        }
839        map
840    }
841}
842
843/// 限流响应策略
844///
845/// 定义被限流时的响应行为。
846#[derive(Debug, Clone)]
847pub enum RateLimitResponseStrategy {
848    /// 返回 429 Too Many Requests
849    TooManyRequests,
850    /// 返回 503 Service Unavailable
851    ServiceUnavailable,
852    /// 自定义状态码
853    Custom(u16),
854}
855
856impl RateLimitResponseStrategy {
857    /// 获取对应的 HTTP 状态码
858    pub fn status_code(&self) -> u16 {
859        match self {
860            RateLimitResponseStrategy::TooManyRequests => 429,
861            RateLimitResponseStrategy::ServiceUnavailable => 503,
862            RateLimitResponseStrategy::Custom(code) => *code,
863        }
864    }
865}
866
867/// 限流响应
868///
869/// 封装被限流时的完整响应信息,包括状态码、头和响应体。
870#[derive(Debug, Clone)]
871pub struct RateLimitResponse {
872    /// HTTP 状态码
873    pub status_code: u16,
874    /// 响应头
875    pub headers: RateLimitHeaders,
876    /// JSON 响应体
877    pub body: serde_json::Value,
878}
879
880impl RateLimitResponse {
881    /// 创建被限流的响应
882    ///
883    /// - `result`:限流结果(必须是 rejected)
884    /// - `limit`:窗口内最大请求数
885    /// - `strategy`:响应策略
886    pub fn rejected(
887        result: &RateLimitResult,
888        limit: u64,
889        strategy: RateLimitResponseStrategy,
890    ) -> Self {
891        let headers = RateLimitHeaders::from_result(result, limit);
892        let status_code = strategy.status_code();
893        let body = serde_json::json!({
894            "error": "rate_limit_exceeded",
895            "message": "Rate limit exceeded. Please retry later.",
896            "retry_after": headers.retry_after.unwrap_or(1),
897        });
898
899        Self {
900            status_code,
901            headers,
902            body,
903        }
904    }
905
906    /// 创建允许通过的响应(仅包含头,无响应体)
907    pub fn allowed(result: &RateLimitResult, limit: u64) -> Self {
908        let headers = RateLimitHeaders::from_result(result, limit);
909        Self {
910            status_code: 200,
911            headers,
912            body: serde_json::Value::Null,
913        }
914    }
915}
916
917// ============================================================================
918// 限流策略组合器(Rate Limit Policy)
919//
920// 允许将多个限流器组合,实现多维度限流(如同时限制 IP 和用户)。
921// ============================================================================
922
923/// 多维度限流策略
924///
925/// 对同一个请求应用多个限流器,只要任一限流器拒绝则拒绝。
926/// 适用于同时限制 IP 级别和用户级别的场景。
927pub struct MultiRateLimiter {
928    limiters: Vec<Arc<dyn RateLimiter>>,
929}
930
931impl MultiRateLimiter {
932    /// 创建多维度限流器
933    pub fn new(limiters: Vec<Arc<dyn RateLimiter>>) -> Self {
934        Self { limiters }
935    }
936
937    /// 添加限流器
938    pub fn with_limiter(mut self, limiter: Arc<dyn RateLimiter>) -> Self {
939        self.limiters.push(limiter);
940        self
941    }
942
943    pub fn limiter_count(&self) -> usize {
944        self.limiters.len()
945    }
946
947    /// 检查所有限流器,返回最严格的结果
948    ///
949    /// 如果任一限流器拒绝,则返回拒绝结果(取剩余最少的)。
950    /// 如果全部允许,则返回剩余最少的结果。
951    pub fn check_all(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
952        let mut best_result: Option<RateLimitResult> = None;
953        for limiter in &self.limiters {
954            let result = limiter.acquire(key)?;
955            match &best_result {
956                None => best_result = Some(result),
957                Some(current) => {
958                    // 取更严格的结果
959                    if !result.allowed {
960                        // 拒绝优先
961                        if !current.allowed {
962                            // 两个都拒绝,取剩余更少的
963                            if result.remaining <= current.remaining {
964                                best_result = Some(result);
965                            }
966                        } else {
967                            // 当前允许但新的拒绝 -> 用拒绝结果
968                            best_result = Some(result);
969                        }
970                    } else if current.allowed && result.remaining < current.remaining {
971                        // 两个都允许,取剩余更少的
972                        best_result = Some(result);
973                    }
974                }
975            }
976        }
977
978        best_result.ok_or_else(|| RateLimitError::Internal("No limiters configured".to_string()))
979    }
980}
981
982// ============================================================================
983// v3.8.0: 限流生产配置(prod-rate-limit-tuning feature)
984// ============================================================================
985
986#[cfg(feature = "prod-rate-limit-tuning")]
987mod prod {
988    use super::DEFAULT_MAX_KEYS;
989    use serde::{Deserialize, Serialize};
990    use std::time::Duration;
991
992    /// 限流生产配置错误
993    #[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
994    pub enum RateLimitProdError {
995        #[error("rate limit capacity must be positive")]
996        CapacityNotPositive,
997        #[error("rate limit rate must be positive")]
998        RateNotPositive,
999        #[error("rate limit window_size must be positive")]
1000        WindowSizeNotPositive,
1001        #[error("rate limit max_keys too small, minimum 100 recommended")]
1002        MaxKeysTooSmall,
1003    }
1004
1005    /// 限流生产配置
1006    #[derive(Debug, Clone, Serialize, Deserialize)]
1007    pub struct RateLimitProdConfig {
1008        pub capacity: u64,
1009        pub rate: u64,
1010        pub window_size: Duration,
1011        pub max_keys: usize,
1012    }
1013
1014    impl Default for RateLimitProdConfig {
1015        fn default() -> Self {
1016            Self {
1017                capacity: 100,
1018                rate: 10,
1019                window_size: Duration::from_secs(1),
1020                max_keys: DEFAULT_MAX_KEYS,
1021            }
1022        }
1023    }
1024
1025    impl RateLimitProdConfig {
1026        pub fn new(capacity: u64, rate: u64, window_size: Duration, max_keys: usize) -> Self {
1027            Self {
1028                capacity,
1029                rate,
1030                window_size,
1031                max_keys,
1032            }
1033        }
1034
1035        /// 校验阈值合理性
1036        pub fn validate(&self) -> Result<(), RateLimitProdError> {
1037            if self.capacity == 0 {
1038                return Err(RateLimitProdError::CapacityNotPositive);
1039            }
1040            if self.rate == 0 {
1041                return Err(RateLimitProdError::RateNotPositive);
1042            }
1043            if self.window_size.is_zero() {
1044                return Err(RateLimitProdError::WindowSizeNotPositive);
1045            }
1046            if self.max_keys < 100 {
1047                return Err(RateLimitProdError::MaxKeysTooSmall);
1048            }
1049            Ok(())
1050        }
1051    }
1052
1053    /// 限流统计信息
1054    #[derive(Debug, Clone, Serialize, Deserialize)]
1055    pub struct RateLimitStats {
1056        pub capacity: u64,
1057        pub allowed_count: u64,
1058        pub rejected_count: u64,
1059    }
1060}
1061
1062#[cfg(feature = "prod-rate-limit-tuning")]
1063pub use prod::{RateLimitProdConfig, RateLimitProdError, RateLimitStats};
1064
1065#[cfg(test)]
1066mod tests {
1067    use super::*;
1068
1069    #[test]
1070    fn test_rate_limit_result_allowed() {
1071        let result = RateLimitResult::allowed(5, 1000);
1072        assert!(result.allowed);
1073        assert_eq!(result.remaining, 5);
1074        assert_eq!(result.reset_at, 1000);
1075    }
1076
1077    #[test]
1078    fn test_rate_limit_result_rejected() {
1079        let result = RateLimitResult::rejected(0, 2000);
1080        assert!(!result.allowed);
1081        assert_eq!(result.remaining, 0);
1082    }
1083
1084    #[test]
1085    fn test_sliding_window_limiter_new() {
1086        let limiter = SlidingWindowRateLimiter::new(10, Duration::from_secs(60));
1087        let result = limiter.acquire("test-key");
1088        assert!(result.is_ok());
1089        assert!(result.unwrap().allowed);
1090    }
1091
1092    #[test]
1093    fn test_sliding_window_limiter_full() {
1094        let limiter = SlidingWindowRateLimiter::new(2, Duration::from_secs(60));
1095
1096        let r1 = limiter.acquire("key1").unwrap();
1097        assert!(r1.allowed);
1098
1099        let r2 = limiter.acquire("key1").unwrap();
1100        assert!(r2.allowed);
1101
1102        let r3 = limiter.acquire("key1").unwrap();
1103        assert!(!r3.allowed);
1104    }
1105
1106    #[test]
1107    fn test_sliding_window_different_keys() {
1108        let limiter = SlidingWindowRateLimiter::new(1, Duration::from_secs(60));
1109
1110        let r1 = limiter.acquire("key-a").unwrap();
1111        assert!(r1.allowed);
1112
1113        let r2 = limiter.acquire("key-b").unwrap();
1114        assert!(r2.allowed);
1115    }
1116
1117    #[test]
1118    fn test_sliding_window_reset() {
1119        let limiter = SlidingWindowRateLimiter::new(1, Duration::from_secs(60));
1120
1121        limiter.acquire("reset-key").unwrap();
1122        limiter.acquire("reset-key").unwrap();
1123
1124        limiter.reset("reset-key").unwrap();
1125
1126        let result = limiter.acquire("reset-key").unwrap();
1127        assert!(result.allowed);
1128    }
1129
1130    #[test]
1131    fn test_token_bucket_limiter_new() {
1132        let limiter = TokenBucketRateLimiter::new(10, 1.0);
1133        let result = limiter.acquire("test-key");
1134        assert!(result.is_ok());
1135        assert!(result.unwrap().allowed);
1136    }
1137
1138    #[test]
1139    fn test_token_bucket_limiter_depletes() {
1140        let limiter = TokenBucketRateLimiter::new(2, 1.0);
1141
1142        let r1 = limiter.acquire("key1").unwrap();
1143        assert!(r1.allowed);
1144        assert_eq!(r1.remaining, 1);
1145
1146        let r2 = limiter.acquire("key1").unwrap();
1147        assert!(r2.allowed);
1148        assert_eq!(r2.remaining, 0);
1149
1150        let r3 = limiter.acquire("key1").unwrap();
1151        assert!(!r3.allowed);
1152    }
1153
1154    #[test]
1155    fn test_token_bucket_different_keys() {
1156        let limiter = TokenBucketRateLimiter::new(1, 1.0);
1157
1158        let r1 = limiter.acquire("key-a").unwrap();
1159        assert!(r1.allowed);
1160
1161        let r2 = limiter.acquire("key-b").unwrap();
1162        assert!(r2.allowed);
1163    }
1164
1165    #[test]
1166    fn test_token_bucket_reset() {
1167        let limiter = TokenBucketRateLimiter::new(1, 1.0);
1168
1169        limiter.acquire("reset-key").unwrap();
1170        limiter.acquire("reset-key").unwrap();
1171
1172        limiter.reset("reset-key").unwrap();
1173
1174        let result = limiter.acquire("reset-key").unwrap();
1175        assert!(result.allowed);
1176    }
1177
1178    #[test]
1179    fn test_limiter_try_acquire() {
1180        let limiter = SlidingWindowRateLimiter::new(1, Duration::from_secs(60));
1181
1182        let r1 = limiter.try_acquire("key").unwrap();
1183        assert!(r1.allowed);
1184
1185        let r2 = limiter.try_acquire("key").unwrap();
1186        assert!(!r2.allowed);
1187    }
1188
1189    // ===== TDD RED:refill_rate=0 panic 修复测试 =====
1190
1191    #[test]
1192    fn test_token_bucket_zero_refill_rate_does_not_panic() {
1193        // refill_rate=0.0 表示令牌永不补充
1194        // capacity=1,第一次 acquire 应允许,第二次应拒绝且不 panic
1195        let limiter = TokenBucketRateLimiter::new(1, 0.0);
1196
1197        let r1 = limiter.acquire("zero-refill").unwrap();
1198        assert!(r1.allowed, "first acquire should be allowed");
1199
1200        // 第二次 acquire:令牌耗尽且不补充,应拒绝而非 panic
1201        let r2 = limiter.acquire("zero-refill").unwrap();
1202        assert!(!r2.allowed, "second acquire should be rejected");
1203        // reset_at 应为一个合理的远期时间(令牌不补充,永远等待)
1204        assert!(
1205            r2.reset_at > 0,
1206            "reset_at should be a valid timestamp, got: {}",
1207            r2.reset_at
1208        );
1209    }
1210
1211    #[test]
1212    fn test_token_bucket_negative_refill_rate_does_not_panic() {
1213        // refill_rate=-1.0 是错误配置,应被当作 0.0 处理而非 panic
1214        let limiter = TokenBucketRateLimiter::new(1, -1.0);
1215
1216        let r1 = limiter.acquire("neg-refill").unwrap();
1217        assert!(r1.allowed, "first acquire should be allowed");
1218
1219        let r2 = limiter.acquire("neg-refill").unwrap();
1220        assert!(!r2.allowed, "second acquire should be rejected");
1221        assert!(
1222            r2.reset_at > 0,
1223            "reset_at should be a valid timestamp, got: {}",
1224            r2.reset_at
1225        );
1226    }
1227
1228    // ===== 固定窗口算法测试 =====
1229
1230    #[test]
1231    fn test_fixed_window_limiter_allows_within_limit() {
1232        let limiter = FixedWindowRateLimiter::new(5, Duration::from_secs(60));
1233        for i in 0..5 {
1234            let r = limiter.acquire("key").unwrap();
1235            assert!(r.allowed, "request {} should be allowed", i);
1236        }
1237    }
1238
1239    #[test]
1240    fn test_fixed_window_limiter_rejects_over_limit() {
1241        let limiter = FixedWindowRateLimiter::new(2, Duration::from_secs(60));
1242        assert!(limiter.acquire("key").unwrap().allowed);
1243        assert!(limiter.acquire("key").unwrap().allowed);
1244        let r3 = limiter.acquire("key").unwrap();
1245        assert!(!r3.allowed);
1246        assert_eq!(r3.remaining, 0);
1247    }
1248
1249    #[test]
1250    fn test_fixed_window_limiter_different_keys() {
1251        let limiter = FixedWindowRateLimiter::new(1, Duration::from_secs(60));
1252        assert!(limiter.acquire("key-a").unwrap().allowed);
1253        assert!(limiter.acquire("key-b").unwrap().allowed);
1254    }
1255
1256    #[test]
1257    fn test_fixed_window_limiter_reset() {
1258        let limiter = FixedWindowRateLimiter::new(1, Duration::from_secs(60));
1259        limiter.acquire("key").unwrap();
1260        assert!(!limiter.acquire("key").unwrap().allowed);
1261        limiter.reset("key").unwrap();
1262        assert!(limiter.acquire("key").unwrap().allowed);
1263    }
1264
1265    #[test]
1266    fn test_fixed_window_limiter_remaining_decreases() {
1267        let limiter = FixedWindowRateLimiter::new(3, Duration::from_secs(60));
1268        let r1 = limiter.acquire("key").unwrap();
1269        assert_eq!(r1.remaining, 2);
1270        let r2 = limiter.acquire("key").unwrap();
1271        assert_eq!(r2.remaining, 1);
1272        let r3 = limiter.acquire("key").unwrap();
1273        assert_eq!(r3.remaining, 0);
1274    }
1275
1276    #[test]
1277    fn test_fixed_window_limiter_reset_at_positive() {
1278        let limiter = FixedWindowRateLimiter::new(1, Duration::from_secs(60));
1279        let r = limiter.acquire("key").unwrap();
1280        assert!(r.reset_at > 0);
1281    }
1282
1283    #[test]
1284    fn test_fixed_window_limiter_try_acquire() {
1285        let limiter = FixedWindowRateLimiter::new(1, Duration::from_secs(60));
1286        assert!(limiter.try_acquire("key").unwrap().allowed);
1287        assert!(!limiter.try_acquire("key").unwrap().allowed);
1288    }
1289
1290    // ===== 分布式限流测试 =====
1291
1292    #[test]
1293    fn test_in_memory_backend_new() {
1294        let backend = InMemoryBackend::new();
1295        let result = backend.get("key").unwrap();
1296        assert!(result.is_none());
1297    }
1298
1299    #[test]
1300    fn test_in_memory_backend_incr_and_get() {
1301        let backend = InMemoryBackend::new();
1302        let now = now_secs();
1303        let (count1, reset1) = backend.incr_and_get("key", 60, now, 10).unwrap();
1304        assert_eq!(count1, 1);
1305        assert_eq!(reset1, now + 60);
1306
1307        let (count2, _) = backend.incr_and_get("key", 60, now, 10).unwrap();
1308        assert_eq!(count2, 2);
1309    }
1310
1311    #[test]
1312    fn test_in_memory_backend_get() {
1313        let backend = InMemoryBackend::new();
1314        let now = now_secs();
1315        backend.incr_and_get("key", 60, now, 10).unwrap();
1316        let result = backend.get("key").unwrap();
1317        assert!(result.is_some());
1318        assert_eq!(result.unwrap().0, 1);
1319    }
1320
1321    #[test]
1322    fn test_in_memory_backend_reset_key() {
1323        let backend = InMemoryBackend::new();
1324        let now = now_secs();
1325        backend.incr_and_get("key", 60, now, 10).unwrap();
1326        assert!(backend.get("key").unwrap().is_some());
1327        backend.reset_key("key").unwrap();
1328        assert!(backend.get("key").unwrap().is_none());
1329    }
1330
1331    #[test]
1332    fn test_in_memory_backend_window_expiry() {
1333        let backend = InMemoryBackend::new();
1334        let now = now_secs();
1335        // 第一次请求,窗口开始
1336        backend.incr_and_get("key", 60, now, 10).unwrap();
1337        backend.incr_and_get("key", 60, now, 10).unwrap();
1338        // 窗口过期后(now + 61),应重置
1339        let (count, _) = backend.incr_and_get("key", 60, now + 61, 10).unwrap();
1340        assert_eq!(count, 1);
1341    }
1342
1343    #[test]
1344    fn test_distributed_rate_limiter_allows() {
1345        let limiter = DistributedRateLimiter::in_memory(5, 60);
1346        for i in 0..5 {
1347            let r = limiter.acquire("key").unwrap();
1348            assert!(r.allowed, "request {} should be allowed", i);
1349        }
1350    }
1351
1352    #[test]
1353    fn test_distributed_rate_limiter_rejects() {
1354        let limiter = DistributedRateLimiter::in_memory(2, 60);
1355        assert!(limiter.acquire("key").unwrap().allowed);
1356        assert!(limiter.acquire("key").unwrap().allowed);
1357        assert!(!limiter.acquire("key").unwrap().allowed);
1358    }
1359
1360    #[test]
1361    fn test_distributed_rate_limiter_shared_backend() {
1362        // 两个限流器共享同一个后端
1363        let backend = Arc::new(InMemoryBackend::new());
1364        let limiter1 = DistributedRateLimiter::new(backend.clone(), 2, 60);
1365        let limiter2 = DistributedRateLimiter::new(backend.clone(), 2, 60);
1366
1367        // limiter1 消耗 1 个
1368        assert!(limiter1.acquire("key").unwrap().allowed);
1369        // limiter2 消耗 1 个(共享计数)
1370        assert!(limiter2.acquire("key").unwrap().allowed);
1371        // 第 3 个应被拒绝(共享计数已达 2)
1372        assert!(!limiter1.acquire("key").unwrap().allowed);
1373    }
1374
1375    #[test]
1376    fn test_distributed_rate_limiter_reset() {
1377        let limiter = DistributedRateLimiter::in_memory(1, 60);
1378        limiter.acquire("key").unwrap();
1379        assert!(!limiter.acquire("key").unwrap().allowed);
1380        limiter.reset("key").unwrap();
1381        assert!(limiter.acquire("key").unwrap().allowed);
1382    }
1383
1384    #[test]
1385    fn test_distributed_rate_limiter_different_keys() {
1386        let limiter = DistributedRateLimiter::in_memory(1, 60);
1387        assert!(limiter.acquire("key-a").unwrap().allowed);
1388        assert!(limiter.acquire("key-b").unwrap().allowed);
1389    }
1390
1391    #[test]
1392    fn test_distributed_rate_limiter_remaining() {
1393        let limiter = DistributedRateLimiter::in_memory(3, 60);
1394        let r1 = limiter.acquire("key").unwrap();
1395        assert_eq!(r1.remaining, 2);
1396        let r2 = limiter.acquire("key").unwrap();
1397        assert_eq!(r2.remaining, 1);
1398    }
1399
1400    // ===== 限流响应策略测试 =====
1401
1402    #[test]
1403    fn test_rate_limit_headers_allowed() {
1404        let result = RateLimitResult::allowed(5, now_secs() * 1000 + 60000);
1405        let headers = RateLimitHeaders::from_result(&result, 10);
1406        assert_eq!(headers.limit, 10);
1407        assert_eq!(headers.remaining, 5);
1408        assert!(headers.retry_after.is_none());
1409    }
1410
1411    #[test]
1412    fn test_rate_limit_headers_rejected() {
1413        let result = RateLimitResult::rejected(0, now_secs() * 1000 + 60000);
1414        let headers = RateLimitHeaders::from_result(&result, 10);
1415        assert_eq!(headers.limit, 10);
1416        assert_eq!(headers.remaining, 0);
1417        assert!(headers.retry_after.is_some());
1418        assert!(headers.retry_after.unwrap() > 0);
1419    }
1420
1421    #[test]
1422    fn test_rate_limit_headers_to_headers_allowed() {
1423        let result = RateLimitResult::allowed(5, now_secs() * 1000 + 60000);
1424        let headers = RateLimitHeaders::from_result(&result, 10);
1425        let hdrs = headers.to_headers();
1426        assert_eq!(hdrs.len(), 3); // 不包含 Retry-After
1427        assert!(hdrs
1428            .iter()
1429            .any(|(k, v)| k == "X-RateLimit-Limit" && v == "10"));
1430        assert!(hdrs
1431            .iter()
1432            .any(|(k, v)| k == "X-RateLimit-Remaining" && v == "5"));
1433    }
1434
1435    #[test]
1436    fn test_rate_limit_headers_to_headers_rejected() {
1437        let result = RateLimitResult::rejected(0, now_secs() * 1000 + 60000);
1438        let headers = RateLimitHeaders::from_result(&result, 10);
1439        let hdrs = headers.to_headers();
1440        assert_eq!(hdrs.len(), 4); // 包含 Retry-After
1441        assert!(hdrs.iter().any(|(k, _)| k == "Retry-After"));
1442    }
1443
1444    #[test]
1445    fn test_rate_limit_headers_to_json() {
1446        let result = RateLimitResult::allowed(5, now_secs() * 1000 + 60000);
1447        let headers = RateLimitHeaders::from_result(&result, 10);
1448        let json = headers.to_json();
1449        assert_eq!(json["X-RateLimit-Limit"], 10);
1450        assert_eq!(json["X-RateLimit-Remaining"], 5);
1451        assert!(json.get("Retry-After").is_none());
1452    }
1453
1454    #[test]
1455    fn test_rate_limit_headers_to_json_rejected() {
1456        let result = RateLimitResult::rejected(0, now_secs() * 1000 + 60000);
1457        let headers = RateLimitHeaders::from_result(&result, 10);
1458        let json = headers.to_json();
1459        assert!(json.get("Retry-After").is_some());
1460    }
1461
1462    #[test]
1463    fn test_rate_limit_response_strategy_status_codes() {
1464        assert_eq!(
1465            RateLimitResponseStrategy::TooManyRequests.status_code(),
1466            429
1467        );
1468        assert_eq!(
1469            RateLimitResponseStrategy::ServiceUnavailable.status_code(),
1470            503
1471        );
1472        assert_eq!(RateLimitResponseStrategy::Custom(502).status_code(), 502);
1473    }
1474
1475    #[test]
1476    fn test_rate_limit_response_rejected() {
1477        let result = RateLimitResult::rejected(0, now_secs() * 1000 + 60000);
1478        let response =
1479            RateLimitResponse::rejected(&result, 10, RateLimitResponseStrategy::TooManyRequests);
1480        assert_eq!(response.status_code, 429);
1481        assert_eq!(response.body["error"], "rate_limit_exceeded");
1482        assert!(response.body["retry_after"].as_u64().unwrap() > 0);
1483        assert!(response.headers.retry_after.is_some());
1484    }
1485
1486    #[test]
1487    fn test_rate_limit_response_allowed() {
1488        let result = RateLimitResult::allowed(5, now_secs() * 1000 + 60000);
1489        let response = RateLimitResponse::allowed(&result, 10);
1490        assert_eq!(response.status_code, 200);
1491        assert!(response.body.is_null());
1492        assert!(response.headers.retry_after.is_none());
1493    }
1494
1495    #[test]
1496    fn test_rate_limit_response_custom_strategy() {
1497        let result = RateLimitResult::rejected(0, now_secs() * 1000 + 60000);
1498        let response =
1499            RateLimitResponse::rejected(&result, 10, RateLimitResponseStrategy::Custom(503));
1500        assert_eq!(response.status_code, 503);
1501    }
1502
1503    // ===== 多维度限流策略测试 =====
1504
1505    #[test]
1506    fn test_multi_rate_limiter_all_allowed() {
1507        let l1 = Arc::new(SlidingWindowRateLimiter::new(10, Duration::from_secs(60)));
1508        let l2 = Arc::new(TokenBucketRateLimiter::new(10, 1.0));
1509        let multi = MultiRateLimiter::new(vec![l1, l2]);
1510
1511        let result = multi.check_all("key").unwrap();
1512        assert!(result.allowed);
1513    }
1514
1515    #[test]
1516    fn test_multi_rate_limiter_one_rejects() {
1517        let l1 = Arc::new(SlidingWindowRateLimiter::new(10, Duration::from_secs(60)));
1518        let l2 = Arc::new(TokenBucketRateLimiter::new(1, 1.0));
1519        let multi = MultiRateLimiter::new(vec![l1, l2]);
1520
1521        // 第一次通过
1522        assert!(multi.check_all("key").unwrap().allowed);
1523        // 第二次:l2 拒绝
1524        let result = multi.check_all("key").unwrap();
1525        assert!(!result.allowed);
1526    }
1527
1528    #[test]
1529    fn test_multi_rate_limiter_takes_strictest() {
1530        let l1 = Arc::new(SlidingWindowRateLimiter::new(5, Duration::from_secs(60)));
1531        let l2 = Arc::new(SlidingWindowRateLimiter::new(2, Duration::from_secs(60)));
1532        let multi = MultiRateLimiter::new(vec![l1, l2]);
1533
1534        // 两次请求后 l2 达到上限
1535        multi.check_all("key").unwrap();
1536        multi.check_all("key").unwrap();
1537        // 第三次:l2 拒绝
1538        assert!(!multi.check_all("key").unwrap().allowed);
1539    }
1540
1541    #[test]
1542    fn test_multi_rate_limiter_with_limiter() {
1543        let l1 = Arc::new(SlidingWindowRateLimiter::new(10, Duration::from_secs(60)));
1544        let multi = MultiRateLimiter::new(vec![]).with_limiter(l1);
1545        assert!(multi.check_all("key").unwrap().allowed);
1546    }
1547
1548    #[test]
1549    fn test_multi_rate_limiter_empty_errors() {
1550        let multi = MultiRateLimiter::new(vec![]);
1551        let result = multi.check_all("key");
1552        assert!(result.is_err());
1553    }
1554}
1555
1556#[cfg(all(test, feature = "prod-rate-limit-tuning"))]
1557mod prod_tests {
1558    use super::*;
1559
1560    #[test]
1561    fn test_rate_limit_prod_config_validate_ok() {
1562        let config = RateLimitProdConfig::new(100, 10, Duration::from_secs(1), 1000);
1563        assert!(config.validate().is_ok());
1564    }
1565
1566    #[test]
1567    fn test_rate_limit_prod_config_capacity_zero_rejected() {
1568        let config = RateLimitProdConfig::new(0, 10, Duration::from_secs(1), 1000);
1569        let err = config.validate().unwrap_err();
1570        assert!(err.to_string().contains("capacity must be positive"));
1571    }
1572
1573    #[test]
1574    fn test_rate_limit_prod_config_rate_zero_rejected() {
1575        let config = RateLimitProdConfig::new(100, 0, Duration::from_secs(1), 1000);
1576        assert!(config.validate().is_err());
1577    }
1578
1579    #[test]
1580    fn test_rate_limit_prod_config_max_keys_too_small() {
1581        let config = RateLimitProdConfig::new(100, 10, Duration::from_secs(1), 50);
1582        let err = config.validate().unwrap_err();
1583        assert!(err.to_string().contains("max_keys too small"));
1584    }
1585
1586    #[test]
1587    fn test_sliding_window_set_capacity() {
1588        let limiter = SlidingWindowRateLimiter::new(100, Duration::from_secs(60));
1589        assert_eq!(limiter.capacity(), 100);
1590        limiter.set_capacity(200);
1591        assert_eq!(limiter.capacity(), 200);
1592    }
1593
1594    #[test]
1595    fn test_sliding_window_set_rate() {
1596        let limiter = SlidingWindowRateLimiter::new(100, Duration::from_secs(10));
1597        limiter.set_rate(20);
1598        // rate=20, window=10s → capacity=200
1599        assert_eq!(limiter.capacity(), 200);
1600    }
1601
1602    #[test]
1603    fn test_sliding_window_stats_after_acquire() {
1604        let limiter = SlidingWindowRateLimiter::new(2, Duration::from_secs(60));
1605        limiter.acquire("k1").unwrap();
1606        limiter.acquire("k1").unwrap();
1607        limiter.acquire("k1").unwrap(); // rejected
1608        let stats = limiter.stats();
1609        assert_eq!(stats.capacity, 2);
1610        assert_eq!(stats.allowed_count, 2);
1611        assert_eq!(stats.rejected_count, 1);
1612    }
1613
1614    #[test]
1615    fn test_sliding_window_dynamic_capacity_takes_effect() {
1616        let limiter = SlidingWindowRateLimiter::new(2, Duration::from_secs(60));
1617        limiter.acquire("k").unwrap();
1618        limiter.acquire("k").unwrap();
1619        assert!(!limiter.acquire("k").unwrap().allowed);
1620        limiter.set_capacity(5);
1621        assert!(limiter.acquire("k").unwrap().allowed);
1622    }
1623}