Skip to main content

sz_orm_pool/
rate_limiter.rs

1//! # 限流器(Rate Limiter)
2//!
3//! P1-4 修复:将限流器抽象从扩展层(sz-orm-limit)提升到核心层,
4//! 消除 core → limit 的反向依赖。
5//!
6//! ## 设计
7//!
8//! - [`RateLimiter`] trait:限流器抽象接口,核心层定义
9//! - [`RateLimitResult`]:限流判定结果
10//! - [`RateLimitError`]:限流器内部错误
11//!
12//! 扩展包 sz-orm-limit 提供多种具体实现(令牌桶/滑动窗口/固定窗口/分布式),
13//! 但必须实现核心 trait,保证核心层不依赖扩展层。
14//!
15//! 连接池 `Pool` 仅依赖核心 trait,调用方通过 `set_rate_limiter`
16//! 注入任意实现(sz-orm-limit 提供的或自定义的)。
17
18/// 限流判定结果
19///
20/// P1-4:从 sz-orm-limit 迁移到核心层,作为限流器抽象的一部分。
21#[derive(Debug, Clone)]
22pub struct RateLimitResult {
23    /// 是否放行
24    pub allowed: bool,
25    /// 当前窗口剩余配额
26    pub remaining: u64,
27    /// 窗口重置时间戳(毫秒)
28    pub reset_at: i64,
29}
30
31impl RateLimitResult {
32    /// 构造“放行”结果
33    ///
34    /// - `remaining`:剩余配额
35    /// - `reset_at`:窗口重置时间戳(毫秒)
36    pub fn allowed(remaining: u64, reset_at: i64) -> Self {
37        Self {
38            allowed: true,
39            remaining,
40            reset_at,
41        }
42    }
43
44    /// 构造“拒绝”结果
45    ///
46    /// - `remaining`:剩余配额(通常为 0)
47    /// - `reset_at`:窗口重置时间戳(毫秒)
48    pub fn rejected(remaining: u64, reset_at: i64) -> Self {
49        Self {
50            allowed: false,
51            remaining,
52            reset_at,
53        }
54    }
55}
56
57/// 限流器错误
58///
59/// P1-4:从 sz-orm-limit 迁移到核心层。
60#[derive(Debug)]
61pub enum RateLimitError {
62    /// 指定 key 未找到
63    KeyNotFound(String),
64    /// 限流器内部错误(如锁中毒、后端不可用)
65    Internal(String),
66}
67
68impl std::fmt::Display for RateLimitError {
69    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
70        match self {
71            RateLimitError::KeyNotFound(key) => write!(f, "Key not found: {}", key),
72            RateLimitError::Internal(msg) => write!(f, "Internal error: {}", msg),
73        }
74    }
75}
76
77impl std::error::Error for RateLimitError {}
78
79impl serde::Serialize for RateLimitError {
80    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
81    where
82        S: serde::Serializer,
83    {
84        serializer.serialize_str(&self.to_string())
85    }
86}
87
88/// 限流器抽象 trait
89///
90/// P1-4:核心层定义的限流器接口,消除对 sz-orm-limit 的反向依赖。
91///
92/// 实现方需保证线程安全(`Send + Sync`),以便在连接池中作为
93/// `Arc<RwLock<Option<Arc<dyn RateLimiter>>>>` 使用。
94///
95/// 常见实现(位于 sz-orm-limit 扩展包):
96/// - `SlidingWindowRateLimiter`:滑动窗口
97/// - `TokenBucketRateLimiter`:令牌桶
98/// - `FixedWindowRateLimiter`:固定窗口
99/// - `DistributedRateLimiter`:分布式限流(基于 Redis 等共享存储)
100pub trait RateLimiter: Send + Sync {
101    /// 获取一个令牌(阻塞语义:可能等待)
102    ///
103    /// 返回 [`RateLimitResult`] 表示放行或拒绝,或 [`RateLimitError`] 表示内部错误。
104    fn acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError>;
105
106    /// 尝试获取一个令牌(非阻塞语义:立即返回)
107    ///
108    /// 默认实现委托给 `acquire`;具体实现可覆盖以提供真正的非阻塞路径。
109    fn try_acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
110        self.acquire(key)
111    }
112
113    /// 重置指定 key 的限流状态
114    fn reset(&self, key: &str) -> Result<(), RateLimitError>;
115}
116
117#[cfg(test)]
118mod tests {
119    use super::*;
120    use std::sync::atomic::{AtomicU64, Ordering};
121    use std::sync::Arc;
122
123    #[test]
124    fn test_rate_limit_result_allowed() {
125        let result = RateLimitResult::allowed(5, 1000);
126        assert!(result.allowed);
127        assert_eq!(result.remaining, 5);
128        assert_eq!(result.reset_at, 1000);
129    }
130
131    #[test]
132    fn test_rate_limit_result_rejected() {
133        let result = RateLimitResult::rejected(0, 2000);
134        assert!(!result.allowed);
135        assert_eq!(result.remaining, 0);
136    }
137
138    #[test]
139    fn test_rate_limit_error_display() {
140        let e = RateLimitError::KeyNotFound("user-1".to_string());
141        assert!(format!("{}", e).contains("user-1"));
142        let e = RateLimitError::Internal("lock poisoned".to_string());
143        assert!(format!("{}", e).contains("lock poisoned"));
144    }
145
146    /// 简单计数器限流器(仅用于测试 trait 抽象可用性)
147    struct CounterLimiter {
148        max: u64,
149        count: AtomicU64,
150    }
151
152    impl CounterLimiter {
153        fn new(max: u64) -> Self {
154            Self {
155                max,
156                count: AtomicU64::new(0),
157            }
158        }
159    }
160
161    impl RateLimiter for CounterLimiter {
162        fn acquire(&self, _key: &str) -> Result<RateLimitResult, RateLimitError> {
163            let prev = self.count.fetch_add(1, Ordering::SeqCst);
164            if prev < self.max {
165                Ok(RateLimitResult::allowed(self.max - prev - 1, 0))
166            } else {
167                Ok(RateLimitResult::rejected(0, 0))
168            }
169        }
170
171        fn reset(&self, _key: &str) -> Result<(), RateLimitError> {
172            self.count.store(0, Ordering::SeqCst);
173            Ok(())
174        }
175    }
176
177    #[test]
178    fn test_counter_limiter_via_trait() {
179        let limiter: Arc<dyn RateLimiter> = Arc::new(CounterLimiter::new(2));
180        let r1 = limiter.acquire("k").unwrap();
181        assert!(r1.allowed);
182        let r2 = limiter.acquire("k").unwrap();
183        assert!(r2.allowed);
184        let r3 = limiter.acquire("k").unwrap();
185        assert!(!r3.allowed);
186    }
187
188    #[test]
189    fn test_counter_limiter_reset() {
190        let limiter: Arc<dyn RateLimiter> = Arc::new(CounterLimiter::new(1));
191        assert!(limiter.acquire("k").unwrap().allowed);
192        assert!(!limiter.acquire("k").unwrap().allowed);
193        limiter.reset("k").unwrap();
194        assert!(limiter.acquire("k").unwrap().allowed);
195    }
196
197    #[test]
198    fn test_default_try_acquire_delegates_to_acquire() {
199        let limiter: Arc<dyn RateLimiter> = Arc::new(CounterLimiter::new(1));
200        // 默认 try_acquire 应委托给 acquire
201        let r1 = limiter.try_acquire("k").unwrap();
202        assert!(r1.allowed);
203        let r2 = limiter.try_acquire("k").unwrap();
204        assert!(!r2.allowed);
205    }
206
207    #[test]
208    fn test_rate_limiter_is_send_sync() {
209        fn assert_send_sync<T: Send + Sync>() {}
210        assert_send_sync::<Arc<dyn RateLimiter>>();
211    }
212}