Skip to main content

sz_orm_limit/
concurrency.rs

1//! 并发限流器(Concurrency Limiter)
2//!
3//! 限制同时进行的请求数量,而非时间窗口内的请求总数。
4//! 适用于保护下游资源(如数据库连接池、外部 API)不被压垮。
5//!
6//! 与令牌桶/滑动窗口的区别:
7//! - 令牌桶/滑动窗口:限制 QPS(每秒请求数)
8//! - 并发限流器:限制并发数(同时进行的请求数)
9
10use std::collections::HashMap;
11use std::sync::atomic::{AtomicI64, AtomicU64, Ordering};
12use std::sync::{Arc, RwLock};
13use std::time::Duration;
14
15use crate::{now_timestamp, RateLimitError, RateLimitResult};
16
17/// 并发限流器
18///
19/// 限制每个 key 同时进行的请求数量。请求完成后必须调用 `release` 释放。
20/// 适用于保护下游资源不被压垮。
21///
22/// # 示例
23///
24/// ```rust
25/// use sz_orm_limit::concurrency::ConcurrencyLimiter;
26/// use std::time::Duration;
27///
28/// let limiter = ConcurrencyLimiter::new(10);
29/// let r = limiter.acquire("user-1").unwrap();
30/// assert!(r.allowed);
31/// limiter.release("user-1");
32/// ```
33pub struct ConcurrencyLimiter {
34    max_concurrent: u64,
35    counters: Arc<RwLock<HashMap<String, AtomicI64>>>,
36    /// 全局并发上限(跨所有 key)
37    global_limit: AtomicI64,
38    /// 当前全局并发数
39    global_current: AtomicI64,
40    /// 最大 key 数量(OOM 防护)
41    max_keys: usize,
42    /// 总获取次数
43    total_acquires: AtomicU64,
44    /// 总释放次数
45    total_releases: AtomicU64,
46    /// 总拒绝次数
47    total_rejections: AtomicU64,
48}
49
50impl ConcurrencyLimiter {
51    /// 创建并发限流器
52    ///
53    /// - `max_concurrent`:每个 key 允许的最大并发数
54    pub fn new(max_concurrent: u64) -> Self {
55        Self {
56            max_concurrent,
57            counters: Arc::new(RwLock::new(HashMap::new())),
58            global_limit: AtomicI64::new(max_concurrent as i64 * 100),
59            global_current: AtomicI64::new(0),
60            max_keys: crate::DEFAULT_MAX_KEYS,
61            total_acquires: AtomicU64::new(0),
62            total_releases: AtomicU64::new(0),
63            total_rejections: AtomicU64::new(0),
64        }
65    }
66
67    /// 配置全局并发上限
68    pub fn with_global_limit(mut self, limit: u64) -> Self {
69        self.global_limit = AtomicI64::new(limit as i64);
70        self
71    }
72
73    /// 配置最大 key 数量
74    pub fn with_max_keys(mut self, max_keys: usize) -> Self {
75        self.max_keys = max_keys;
76        self
77    }
78
79    /// 获取 key 的当前并发数
80    pub fn current_concurrent(&self, key: &str) -> i64 {
81        let counters = self.counters.read().map_err(|e| e.to_string());
82        match counters {
83            Ok(map) => map.get(key).map(|a| a.load(Ordering::Relaxed)).unwrap_or(0),
84            Err(_) => 0,
85        }
86    }
87
88    /// 获取全局当前并发数
89    pub fn global_current(&self) -> i64 {
90        self.global_current.load(Ordering::Relaxed)
91    }
92
93    /// 获取最大并发数
94    pub fn max_concurrent(&self) -> u64 {
95        self.max_concurrent
96    }
97
98    /// 获取全局并发上限
99    pub fn global_limit(&self) -> i64 {
100        self.global_limit.load(Ordering::Relaxed)
101    }
102
103    /// 获取当前 key 数量
104    pub fn key_count(&self) -> usize {
105        self.counters.read().map(|m| m.len()).unwrap_or(0)
106    }
107
108    /// 尝试获取一个并发槽位
109    ///
110    /// 如果当前 key 的并发数未超过 `max_concurrent` 且全局并发数未超过
111    /// `global_limit`,则并发数 +1 并返回 allowed。
112    /// 否则返回 rejected。
113    pub fn acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
114        self.total_acquires.fetch_add(1, Ordering::Relaxed);
115        let mut counters = self
116            .counters
117            .write()
118            .map_err(|e| RateLimitError::Internal(e.to_string()))?;
119
120        // OOM 防护:超出 max_keys 时淘汰一个
121        if counters.len() >= self.max_keys && !counters.contains_key(key) {
122            let oldest = counters.keys().next().cloned();
123            if let Some(k) = oldest {
124                counters.remove(&k);
125            }
126        }
127
128        let counter = counters
129            .entry(key.to_string())
130            .or_insert_with(|| AtomicI64::new(0));
131
132        let current = counter.load(Ordering::Relaxed);
133        let global = self.global_current.load(Ordering::Relaxed);
134        let global_limit = self.global_limit.load(Ordering::Relaxed);
135
136        if current >= self.max_concurrent as i64 {
137            self.total_rejections.fetch_add(1, Ordering::Relaxed);
138            return Ok(RateLimitResult::rejected(0, now_timestamp() + 1000));
139        }
140        if global >= global_limit {
141            self.total_rejections.fetch_add(1, Ordering::Relaxed);
142            return Ok(RateLimitResult::rejected(0, now_timestamp() + 1000));
143        }
144
145        counter.fetch_add(1, Ordering::Relaxed);
146        self.global_current.fetch_add(1, Ordering::Relaxed);
147        let remaining = self.max_concurrent - counter.load(Ordering::Relaxed) as u64;
148        Ok(RateLimitResult::allowed(remaining, now_timestamp() + 60000))
149    }
150
151    /// 释放一个并发槽位
152    ///
153    /// 请求完成后必须调用此方法释放槽位,否则并发数会持续增长。
154    pub fn release(&self, key: &str) -> Result<(), RateLimitError> {
155        self.total_releases.fetch_add(1, Ordering::Relaxed);
156        let counters = self
157            .counters
158            .read()
159            .map_err(|e| RateLimitError::Internal(e.to_string()))?;
160        if let Some(counter) = counters.get(key) {
161            let current = counter.load(Ordering::Relaxed);
162            if current > 0 {
163                counter.fetch_sub(1, Ordering::Relaxed);
164                self.global_current.fetch_sub(1, Ordering::Relaxed);
165            }
166        }
167        Ok(())
168    }
169
170    /// 重置 key 的并发计数
171    pub fn reset(&self, key: &str) -> Result<(), RateLimitError> {
172        let counters = self
173            .counters
174            .read()
175            .map_err(|e| RateLimitError::Internal(e.to_string()))?;
176        if let Some(counter) = counters.get(key) {
177            let current = counter.load(Ordering::Relaxed);
178            if current > 0 {
179                self.global_current.fetch_sub(current, Ordering::Relaxed);
180                counter.store(0, Ordering::Relaxed);
181            }
182        }
183        Ok(())
184    }
185
186    /// 获取统计信息
187    pub fn stats(&self) -> ConcurrencyStats {
188        ConcurrencyStats {
189            max_concurrent: self.max_concurrent,
190            global_limit: self.global_limit.load(Ordering::Relaxed),
191            global_current: self.global_current.load(Ordering::Relaxed),
192            key_count: self.key_count(),
193            total_acquires: self.total_acquires.load(Ordering::Relaxed),
194            total_releases: self.total_releases.load(Ordering::Relaxed),
195            total_rejections: self.total_rejections.load(Ordering::Relaxed),
196        }
197    }
198}
199
200/// 并发限流统计信息
201#[derive(Debug, Clone, serde::Serialize)]
202pub struct ConcurrencyStats {
203    pub max_concurrent: u64,
204    pub global_limit: i64,
205    pub global_current: i64,
206    pub key_count: usize,
207    pub total_acquires: u64,
208    pub total_releases: u64,
209    pub total_rejections: u64,
210}
211
212/// 带超时的并发槽位守卫
213///
214/// 获取后自动占用槽位,drop 时自动释放。
215/// 适用于 `with` 模式确保资源释放。
216pub struct ConcurrencyGuard<'a> {
217    limiter: &'a ConcurrencyLimiter,
218    key: String,
219    acquired: bool,
220}
221
222impl<'a> ConcurrencyGuard<'a> {
223    /// 尝试获取并发槽位
224    pub fn acquire(limiter: &'a ConcurrencyLimiter, key: &str) -> Result<Self, RateLimitError> {
225        let result = limiter.acquire(key)?;
226        Ok(Self {
227            limiter,
228            key: key.to_string(),
229            acquired: result.allowed,
230        })
231    }
232
233    /// 是否成功获取
234    pub fn is_acquired(&self) -> bool {
235        self.acquired
236    }
237
238    /// 手动释放(drop 时也会自动释放)
239    pub fn release(mut self) {
240        if self.acquired {
241            let _ = self.limiter.release(&self.key);
242            self.acquired = false;
243        }
244    }
245}
246
247impl<'a> Drop for ConcurrencyGuard<'a> {
248    fn drop(&mut self) {
249        if self.acquired {
250            let _ = self.limiter.release(&self.key);
251        }
252    }
253}
254
255/// 带等待超时的并发限流器
256///
257/// 在 `acquire` 时如果并发已满,会等待最多 `wait_timeout` 时间,
258/// 期间不断重试。超时后返回 rejected。
259pub struct TimedConcurrencyLimiter {
260    inner: ConcurrencyLimiter,
261    wait_timeout: Duration,
262    retry_interval: Duration,
263}
264
265impl TimedConcurrencyLimiter {
266    /// 创建带超时的并发限流器
267    ///
268    /// - `max_concurrent`:每个 key 的最大并发数
269    /// - `wait_timeout`:获取槽位的最长等待时间
270    pub fn new(max_concurrent: u64, wait_timeout: Duration) -> Self {
271        Self {
272            inner: ConcurrencyLimiter::new(max_concurrent),
273            wait_timeout,
274            retry_interval: Duration::from_millis(10),
275        }
276    }
277
278    /// 配置重试间隔
279    pub fn with_retry_interval(mut self, interval: Duration) -> Self {
280        self.retry_interval = interval;
281        self
282    }
283
284    /// 配置全局并发上限
285    pub fn with_global_limit(self, limit: u64) -> Self {
286        Self {
287            inner: self.inner.with_global_limit(limit),
288            ..self
289        }
290    }
291
292    /// 尝试获取槽位,最多等待 `wait_timeout`
293    ///
294    /// 注意:此方法会忙等重试,不适用于高并发场景。
295    /// 生产环境建议使用异步版本或消息队列。
296    pub fn acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
297        let start = std::time::Instant::now();
298        loop {
299            let result = self.inner.acquire(key)?;
300            if result.allowed {
301                return Ok(result);
302            }
303            if start.elapsed() >= self.wait_timeout {
304                return Ok(result);
305            }
306            std::thread::sleep(self.retry_interval);
307        }
308    }
309
310    /// 释放槽位
311    pub fn release(&self, key: &str) -> Result<(), RateLimitError> {
312        self.inner.release(key)
313    }
314
315    /// 获取内部限流器
316    pub fn inner(&self) -> &ConcurrencyLimiter {
317        &self.inner
318    }
319}
320
321#[cfg(test)]
322mod tests {
323    use super::*;
324
325    #[test]
326    fn test_concurrency_limiter_basic_acquire() {
327        let limiter = ConcurrencyLimiter::new(5);
328        let r = limiter.acquire("k").unwrap();
329        assert!(r.allowed);
330        assert_eq!(r.remaining, 4);
331    }
332
333    #[test]
334    fn test_concurrency_limiter_max_concurrent() {
335        let limiter = ConcurrencyLimiter::new(2);
336        assert!(limiter.acquire("k").unwrap().allowed);
337        assert!(limiter.acquire("k").unwrap().allowed);
338        let r3 = limiter.acquire("k").unwrap();
339        assert!(!r3.allowed);
340    }
341
342    #[test]
343    fn test_concurrency_limiter_release() {
344        let limiter = ConcurrencyLimiter::new(1);
345        assert!(limiter.acquire("k").unwrap().allowed);
346        assert!(!limiter.acquire("k").unwrap().allowed);
347        limiter.release("k").unwrap();
348        assert!(limiter.acquire("k").unwrap().allowed);
349    }
350
351    #[test]
352    fn test_concurrency_limiter_different_keys() {
353        let limiter = ConcurrencyLimiter::new(1);
354        assert!(limiter.acquire("a").unwrap().allowed);
355        assert!(limiter.acquire("b").unwrap().allowed);
356    }
357
358    #[test]
359    fn test_concurrency_limiter_current_concurrent() {
360        let limiter = ConcurrencyLimiter::new(5);
361        assert_eq!(limiter.current_concurrent("k"), 0);
362        limiter.acquire("k").unwrap();
363        assert_eq!(limiter.current_concurrent("k"), 1);
364        limiter.acquire("k").unwrap();
365        assert_eq!(limiter.current_concurrent("k"), 2);
366    }
367
368    #[test]
369    fn test_concurrency_limiter_global_current() {
370        let limiter = ConcurrencyLimiter::new(5);
371        assert_eq!(limiter.global_current(), 0);
372        limiter.acquire("a").unwrap();
373        limiter.acquire("b").unwrap();
374        assert_eq!(limiter.global_current(), 2);
375        limiter.release("a").unwrap();
376        assert_eq!(limiter.global_current(), 1);
377    }
378
379    #[test]
380    fn test_concurrency_limiter_global_limit() {
381        let limiter = ConcurrencyLimiter::new(10).with_global_limit(2);
382        assert!(limiter.acquire("a").unwrap().allowed);
383        assert!(limiter.acquire("b").unwrap().allowed);
384        let r3 = limiter.acquire("c").unwrap();
385        assert!(!r3.allowed);
386    }
387
388    #[test]
389    fn test_concurrency_limiter_reset() {
390        let limiter = ConcurrencyLimiter::new(5);
391        limiter.acquire("k").unwrap();
392        limiter.acquire("k").unwrap();
393        assert_eq!(limiter.current_concurrent("k"), 2);
394        limiter.reset("k").unwrap();
395        assert_eq!(limiter.current_concurrent("k"), 0);
396    }
397
398    #[test]
399    fn test_concurrency_limiter_stats() {
400        let limiter = ConcurrencyLimiter::new(3);
401        limiter.acquire("k").unwrap();
402        limiter.acquire("k").unwrap();
403        limiter.release("k").unwrap();
404        limiter.acquire("k").unwrap();
405        let stats = limiter.stats();
406        assert_eq!(stats.max_concurrent, 3);
407        assert_eq!(stats.total_acquires, 3);
408        assert_eq!(stats.total_releases, 1);
409        assert_eq!(stats.global_current, 2);
410    }
411
412    #[test]
413    fn test_concurrency_limiter_key_count() {
414        let limiter = ConcurrencyLimiter::new(5);
415        assert_eq!(limiter.key_count(), 0);
416        limiter.acquire("a").unwrap();
417        limiter.acquire("b").unwrap();
418        assert_eq!(limiter.key_count(), 2);
419    }
420
421    #[test]
422    fn test_concurrency_guard_auto_release() {
423        let limiter = ConcurrencyLimiter::new(1);
424        {
425            let guard = ConcurrencyGuard::acquire(&limiter, "k").unwrap();
426            assert!(guard.is_acquired());
427            assert_eq!(limiter.current_concurrent("k"), 1);
428        }
429        assert_eq!(limiter.current_concurrent("k"), 0);
430    }
431
432    #[test]
433    fn test_concurrency_guard_manual_release() {
434        let limiter = ConcurrencyLimiter::new(1);
435        let guard = ConcurrencyGuard::acquire(&limiter, "k").unwrap();
436        assert_eq!(limiter.current_concurrent("k"), 1);
437        guard.release();
438        assert_eq!(limiter.current_concurrent("k"), 0);
439    }
440
441    #[test]
442    fn test_concurrency_guard_rejected() {
443        let limiter = ConcurrencyLimiter::new(1);
444        let g1 = ConcurrencyGuard::acquire(&limiter, "k").unwrap();
445        let g2 = ConcurrencyGuard::acquire(&limiter, "k").unwrap();
446        assert!(g1.is_acquired());
447        assert!(!g2.is_acquired());
448    }
449
450    #[test]
451    fn test_timed_concurrency_limiter_immediate() {
452        let limiter = TimedConcurrencyLimiter::new(2, Duration::from_millis(100));
453        let r = limiter.acquire("k").unwrap();
454        assert!(r.allowed);
455        limiter.release("k").unwrap();
456    }
457
458    #[test]
459    fn test_timed_concurrency_limiter_timeout() {
460        let limiter = TimedConcurrencyLimiter::new(1, Duration::from_millis(50))
461            .with_retry_interval(Duration::from_millis(5));
462        let r1 = limiter.acquire("k").unwrap();
463        assert!(r1.allowed);
464        let r2 = limiter.acquire("k").unwrap();
465        assert!(!r2.allowed);
466    }
467
468    #[test]
469    fn test_concurrency_limiter_release_below_zero_guard() {
470        let limiter = ConcurrencyLimiter::new(5);
471        limiter.release("k").unwrap();
472        assert_eq!(limiter.current_concurrent("k"), 0);
473    }
474
475    #[test]
476    fn test_concurrency_limiter_double_release() {
477        let limiter = ConcurrencyLimiter::new(5);
478        limiter.acquire("k").unwrap();
479        limiter.release("k").unwrap();
480        limiter.release("k").unwrap();
481        assert_eq!(limiter.current_concurrent("k"), 0);
482        assert_eq!(limiter.global_current(), 0);
483    }
484}