Skip to main content

sz_orm_limit/
leaky_bucket.rs

1//! 漏桶限流器(Leaky Bucket Limiter)
2//!
3//! 漏桶算法以恒定速率"漏出"请求,多余请求被拒绝或排队。
4//! 与令牌桶对偶:令牌桶限制突发量,漏桶平滑输出速率。
5//!
6//! 适用于需要平滑输出速率的场景(如消息队列生产者)。
7
8use std::collections::HashMap;
9use std::sync::atomic::{AtomicU64, Ordering};
10use std::sync::{Arc, RwLock};
11use std::time::{Duration, Instant};
12
13use crate::{now_timestamp, RateLimitError, RateLimitResult};
14
15/// 漏桶限流器
16///
17/// 桶容量为 `capacity`,以 `leak_rate`(请求/秒)的速率漏出。
18/// 请求到来时加入桶,如果桶满则拒绝。
19pub struct LeakyBucketLimiter {
20    capacity: u64,
21    leak_rate: f64,
22    buckets: Arc<RwLock<HashMap<String, LeakyBucketEntry>>>,
23    max_keys: usize,
24    total_allowed: AtomicU64,
25    total_rejected: AtomicU64,
26}
27
28#[derive(Clone)]
29struct LeakyBucketEntry {
30    water: f64,
31    last_leak: Instant,
32}
33
34impl LeakyBucketLimiter {
35    /// 创建漏桶限流器
36    ///
37    /// - `capacity`:桶容量(最大排队请求数)
38    /// - `leak_rate`:漏出速率(请求/秒)
39    pub fn new(capacity: u64, leak_rate: f64) -> Self {
40        Self {
41            capacity,
42            leak_rate,
43            buckets: Arc::new(RwLock::new(HashMap::new())),
44            max_keys: crate::DEFAULT_MAX_KEYS,
45            total_allowed: AtomicU64::new(0),
46            total_rejected: AtomicU64::new(0),
47        }
48    }
49
50    /// 配置最大 key 数量
51    pub fn with_max_keys(mut self, max_keys: usize) -> Self {
52        self.max_keys = max_keys;
53        self
54    }
55
56    /// 桶容量
57    pub fn capacity(&self) -> u64 {
58        self.capacity
59    }
60
61    /// 漏出速率
62    pub fn leak_rate(&self) -> f64 {
63        self.leak_rate
64    }
65
66    /// 当前 key 数量
67    pub fn key_count(&self) -> usize {
68        self.buckets.read().map(|m| m.len()).unwrap_or(0)
69    }
70
71    /// 获取 key 的当前水位
72    pub fn water_level(&self, key: &str) -> f64 {
73        let buckets = self.buckets.read().map_err(|e| e.to_string());
74        match buckets {
75            Ok(map) => map.get(key).map(|e| e.water).unwrap_or(0.0),
76            Err(_) => 0.0,
77        }
78    }
79
80    fn leak(&self, entry: &mut LeakyBucketEntry) {
81        let now = Instant::now();
82        let elapsed = now.duration_since(entry.last_leak).as_secs_f64();
83        let leaked = if self.leak_rate > 0.0 {
84            elapsed * self.leak_rate
85        } else {
86            0.0
87        };
88        entry.water = (entry.water - leaked).max(0.0);
89        entry.last_leak = now;
90    }
91
92    fn enforce_max_keys(&self, buckets: &mut HashMap<String, LeakyBucketEntry>) {
93        while buckets.len() > self.max_keys {
94            let oldest = buckets
95                .iter()
96                .min_by_key(|(_, e)| e.last_leak)
97                .map(|(k, _)| k.clone());
98            match oldest {
99                Some(k) => {
100                    buckets.remove(&k);
101                }
102                None => break,
103            }
104        }
105    }
106
107    /// 尝试加入一个请求
108    pub fn acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
109        let mut buckets = self
110            .buckets
111            .write()
112            .map_err(|e| RateLimitError::Internal(e.to_string()))?;
113
114        if buckets.len() >= self.max_keys && !buckets.contains_key(key) {
115            self.enforce_max_keys(&mut buckets);
116        }
117
118        let entry = buckets
119            .entry(key.to_string())
120            .or_insert_with(|| LeakyBucketEntry {
121                water: 0.0,
122                last_leak: Instant::now(),
123            });
124
125        self.leak(entry);
126
127        if entry.water + 1.0 <= self.capacity as f64 {
128            entry.water += 1.0;
129            let remaining = (self.capacity as f64 - entry.water).floor() as u64;
130            self.total_allowed.fetch_add(1, Ordering::Relaxed);
131            Ok(RateLimitResult::allowed(remaining, now_timestamp() + 1000))
132        } else {
133            self.total_rejected.fetch_add(1, Ordering::Relaxed);
134            Ok(RateLimitResult::rejected(0, now_timestamp() + 1000))
135        }
136    }
137
138    /// 重置 key
139    pub fn reset(&self, key: &str) -> Result<(), RateLimitError> {
140        let mut buckets = self
141            .buckets
142            .write()
143            .map_err(|e| RateLimitError::Internal(e.to_string()))?;
144        buckets.remove(key);
145        Ok(())
146    }
147
148    /// 统计信息
149    pub fn stats(&self) -> LeakyBucketStats {
150        LeakyBucketStats {
151            capacity: self.capacity,
152            leak_rate: self.leak_rate,
153            key_count: self.key_count(),
154            total_allowed: self.total_allowed.load(Ordering::Relaxed),
155            total_rejected: self.total_rejected.load(Ordering::Relaxed),
156        }
157    }
158}
159
160/// 漏桶统计信息
161#[derive(Debug, Clone, serde::Serialize)]
162pub struct LeakyBucketStats {
163    pub capacity: u64,
164    pub leak_rate: f64,
165    pub key_count: usize,
166    pub total_allowed: u64,
167    pub total_rejected: u64,
168}
169
170/// 滑动窗口日志算法(Sliding Window Log)
171///
172/// 与滑动窗口计数器不同,日志算法记录每个请求的精确时间戳,
173/// 提供更精确的限流(无边界突刺),但内存占用更高。
174pub struct SlidingWindowLogLimiter {
175    max_requests: u64,
176    window_size: Duration,
177    entries: Arc<RwLock<HashMap<String, Vec<Instant>>>>,
178    max_keys: usize,
179}
180
181impl SlidingWindowLogLimiter {
182    /// 创建滑动窗口日志限流器
183    pub fn new(max_requests: u64, window_size: Duration) -> Self {
184        Self {
185            max_requests,
186            window_size,
187            entries: Arc::new(RwLock::new(HashMap::new())),
188            max_keys: crate::DEFAULT_MAX_KEYS,
189        }
190    }
191
192    /// 配置最大 key 数量
193    pub fn with_max_keys(mut self, max_keys: usize) -> Self {
194        self.max_keys = max_keys;
195        self
196    }
197
198    /// 最大请求数
199    pub fn max_requests(&self) -> u64 {
200        self.max_requests
201    }
202
203    /// 窗口大小
204    pub fn window_size(&self) -> Duration {
205        self.window_size
206    }
207
208    /// 当前 key 数量
209    pub fn key_count(&self) -> usize {
210        self.entries.read().map(|m| m.len()).unwrap_or(0)
211    }
212
213    /// 当前窗口内请求数
214    pub fn current_count(&self, key: &str) -> usize {
215        let entries = self.entries.read().map_err(|e| e.to_string());
216        match entries {
217            Ok(map) => {
218                if let Some(log) = map.get(key) {
219                    let now = Instant::now();
220                    log.iter()
221                        .filter(|&&t| now.duration_since(t) < self.window_size)
222                        .count()
223                } else {
224                    0
225                }
226            }
227            Err(_) => 0,
228        }
229    }
230
231    fn enforce_max_keys(&self, entries: &mut HashMap<String, Vec<Instant>>) {
232        while entries.len() > self.max_keys {
233            let now = Instant::now();
234            let oldest = entries
235                .iter()
236                .min_by_key(|(_, log)| log.first().copied().unwrap_or(now))
237                .map(|(k, _)| k.clone());
238            match oldest {
239                Some(k) => {
240                    entries.remove(&k);
241                }
242                None => break,
243            }
244        }
245    }
246
247    /// 尝试获取一个请求
248    pub fn acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
249        let mut entries = self
250            .entries
251            .write()
252            .map_err(|e| RateLimitError::Internal(e.to_string()))?;
253
254        if entries.len() >= self.max_keys && !entries.contains_key(key) {
255            self.enforce_max_keys(&mut entries);
256        }
257
258        let log = entries.entry(key.to_string()).or_insert_with(Vec::new);
259
260        let now = Instant::now();
261        log.retain(|&t| now.duration_since(t) < self.window_size);
262
263        if log.len() < self.max_requests as usize {
264            log.push(now);
265            let remaining = self.max_requests - log.len() as u64;
266            Ok(RateLimitResult::allowed(remaining, now_timestamp() + 1000))
267        } else {
268            Ok(RateLimitResult::rejected(0, now_timestamp() + 1000))
269        }
270    }
271
272    /// 重置 key
273    pub fn reset(&self, key: &str) -> Result<(), RateLimitError> {
274        let mut entries = self
275            .entries
276            .write()
277            .map_err(|e| RateLimitError::Internal(e.to_string()))?;
278        entries.remove(key);
279        Ok(())
280    }
281}
282
283#[cfg(test)]
284mod tests {
285    use super::*;
286
287    #[test]
288    fn test_leaky_bucket_basic() {
289        let limiter = LeakyBucketLimiter::new(5, 1.0);
290        let r = limiter.acquire("k").unwrap();
291        assert!(r.allowed);
292    }
293
294    #[test]
295    fn test_leaky_bucket_full() {
296        let limiter = LeakyBucketLimiter::new(2, 0.0);
297        assert!(limiter.acquire("k").unwrap().allowed);
298        assert!(limiter.acquire("k").unwrap().allowed);
299        assert!(!limiter.acquire("k").unwrap().allowed);
300    }
301
302    #[test]
303    fn test_leaky_bucket_different_keys() {
304        let limiter = LeakyBucketLimiter::new(1, 0.0);
305        assert!(limiter.acquire("a").unwrap().allowed);
306        assert!(limiter.acquire("b").unwrap().allowed);
307    }
308
309    #[test]
310    fn test_leaky_bucket_water_level() {
311        let limiter = LeakyBucketLimiter::new(5, 0.0);
312        assert_eq!(limiter.water_level("k"), 0.0);
313        limiter.acquire("k").unwrap();
314        assert!(limiter.water_level("k") > 0.0);
315    }
316
317    #[test]
318    fn test_leaky_bucket_reset() {
319        let limiter = LeakyBucketLimiter::new(1, 0.0);
320        limiter.acquire("k").unwrap();
321        assert!(!limiter.acquire("k").unwrap().allowed);
322        limiter.reset("k").unwrap();
323        assert!(limiter.acquire("k").unwrap().allowed);
324    }
325
326    #[test]
327    fn test_leaky_bucket_stats() {
328        let limiter = LeakyBucketLimiter::new(3, 1.0);
329        limiter.acquire("k").unwrap();
330        limiter.acquire("k").unwrap();
331        let stats = limiter.stats();
332        assert_eq!(stats.capacity, 3);
333        assert_eq!(stats.total_allowed, 2);
334        assert_eq!(stats.total_rejected, 0);
335    }
336
337    #[test]
338    fn test_leaky_bucket_key_count() {
339        let limiter = LeakyBucketLimiter::new(5, 1.0);
340        assert_eq!(limiter.key_count(), 0);
341        limiter.acquire("a").unwrap();
342        limiter.acquire("b").unwrap();
343        assert_eq!(limiter.key_count(), 2);
344    }
345
346    #[test]
347    fn test_sliding_window_log_basic() {
348        let limiter = SlidingWindowLogLimiter::new(5, Duration::from_secs(60));
349        let r = limiter.acquire("k").unwrap();
350        assert!(r.allowed);
351    }
352
353    #[test]
354    fn test_sliding_window_log_full() {
355        let limiter = SlidingWindowLogLimiter::new(2, Duration::from_secs(60));
356        assert!(limiter.acquire("k").unwrap().allowed);
357        assert!(limiter.acquire("k").unwrap().allowed);
358        assert!(!limiter.acquire("k").unwrap().allowed);
359    }
360
361    #[test]
362    fn test_sliding_window_log_different_keys() {
363        let limiter = SlidingWindowLogLimiter::new(1, Duration::from_secs(60));
364        assert!(limiter.acquire("a").unwrap().allowed);
365        assert!(limiter.acquire("b").unwrap().allowed);
366    }
367
368    #[test]
369    fn test_sliding_window_log_current_count() {
370        let limiter = SlidingWindowLogLimiter::new(5, Duration::from_secs(60));
371        assert_eq!(limiter.current_count("k"), 0);
372        limiter.acquire("k").unwrap();
373        limiter.acquire("k").unwrap();
374        assert_eq!(limiter.current_count("k"), 2);
375    }
376
377    #[test]
378    fn test_sliding_window_log_reset() {
379        let limiter = SlidingWindowLogLimiter::new(1, Duration::from_secs(60));
380        limiter.acquire("k").unwrap();
381        assert!(!limiter.acquire("k").unwrap().allowed);
382        limiter.reset("k").unwrap();
383        assert!(limiter.acquire("k").unwrap().allowed);
384    }
385
386    #[test]
387    fn test_sliding_window_log_key_count() {
388        let limiter = SlidingWindowLogLimiter::new(5, Duration::from_secs(60));
389        assert_eq!(limiter.key_count(), 0);
390        limiter.acquire("a").unwrap();
391        assert_eq!(limiter.key_count(), 1);
392    }
393
394    #[test]
395    fn test_sliding_window_log_remaining() {
396        let limiter = SlidingWindowLogLimiter::new(3, Duration::from_secs(60));
397        let r1 = limiter.acquire("k").unwrap();
398        assert_eq!(r1.remaining, 2);
399        let r2 = limiter.acquire("k").unwrap();
400        assert_eq!(r2.remaining, 1);
401    }
402}