1use 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
17pub struct ConcurrencyLimiter {
34 max_concurrent: u64,
35 counters: Arc<RwLock<HashMap<String, AtomicI64>>>,
36 global_limit: AtomicI64,
38 global_current: AtomicI64,
40 max_keys: usize,
42 total_acquires: AtomicU64,
44 total_releases: AtomicU64,
46 total_rejections: AtomicU64,
48}
49
50impl ConcurrencyLimiter {
51 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 pub fn with_global_limit(mut self, limit: u64) -> Self {
69 self.global_limit = AtomicI64::new(limit as i64);
70 self
71 }
72
73 pub fn with_max_keys(mut self, max_keys: usize) -> Self {
75 self.max_keys = max_keys;
76 self
77 }
78
79 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 pub fn global_current(&self) -> i64 {
90 self.global_current.load(Ordering::Relaxed)
91 }
92
93 pub fn max_concurrent(&self) -> u64 {
95 self.max_concurrent
96 }
97
98 pub fn global_limit(&self) -> i64 {
100 self.global_limit.load(Ordering::Relaxed)
101 }
102
103 pub fn key_count(&self) -> usize {
105 self.counters.read().map(|m| m.len()).unwrap_or(0)
106 }
107
108 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 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 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 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 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#[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
212pub struct ConcurrencyGuard<'a> {
217 limiter: &'a ConcurrencyLimiter,
218 key: String,
219 acquired: bool,
220}
221
222impl<'a> ConcurrencyGuard<'a> {
223 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 pub fn is_acquired(&self) -> bool {
235 self.acquired
236 }
237
238 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
255pub struct TimedConcurrencyLimiter {
260 inner: ConcurrencyLimiter,
261 wait_timeout: Duration,
262 retry_interval: Duration,
263}
264
265impl TimedConcurrencyLimiter {
266 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 pub fn with_retry_interval(mut self, interval: Duration) -> Self {
280 self.retry_interval = interval;
281 self
282 }
283
284 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 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 pub fn release(&self, key: &str) -> Result<(), RateLimitError> {
312 self.inner.release(key)
313 }
314
315 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}