use std::collections::HashMap;
use std::sync::atomic::{AtomicI64, AtomicU64, Ordering};
use std::sync::{Arc, RwLock};
use std::time::Duration;
use crate::{now_timestamp, RateLimitError, RateLimitResult};
pub struct ConcurrencyLimiter {
max_concurrent: u64,
counters: Arc<RwLock<HashMap<String, AtomicI64>>>,
global_limit: AtomicI64,
global_current: AtomicI64,
max_keys: usize,
total_acquires: AtomicU64,
total_releases: AtomicU64,
total_rejections: AtomicU64,
}
impl ConcurrencyLimiter {
pub fn new(max_concurrent: u64) -> Self {
Self {
max_concurrent,
counters: Arc::new(RwLock::new(HashMap::new())),
global_limit: AtomicI64::new(max_concurrent as i64 * 100),
global_current: AtomicI64::new(0),
max_keys: crate::DEFAULT_MAX_KEYS,
total_acquires: AtomicU64::new(0),
total_releases: AtomicU64::new(0),
total_rejections: AtomicU64::new(0),
}
}
pub fn with_global_limit(mut self, limit: u64) -> Self {
self.global_limit = AtomicI64::new(limit as i64);
self
}
pub fn with_max_keys(mut self, max_keys: usize) -> Self {
self.max_keys = max_keys;
self
}
pub fn current_concurrent(&self, key: &str) -> i64 {
let counters = self.counters.read().map_err(|e| e.to_string());
match counters {
Ok(map) => map.get(key).map(|a| a.load(Ordering::Relaxed)).unwrap_or(0),
Err(_) => 0,
}
}
pub fn global_current(&self) -> i64 {
self.global_current.load(Ordering::Relaxed)
}
pub fn max_concurrent(&self) -> u64 {
self.max_concurrent
}
pub fn global_limit(&self) -> i64 {
self.global_limit.load(Ordering::Relaxed)
}
pub fn key_count(&self) -> usize {
self.counters.read().map(|m| m.len()).unwrap_or(0)
}
pub fn acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
self.total_acquires.fetch_add(1, Ordering::Relaxed);
let mut counters = self
.counters
.write()
.map_err(|e| RateLimitError::Internal(e.to_string()))?;
if counters.len() >= self.max_keys && !counters.contains_key(key) {
let oldest = counters.keys().next().cloned();
if let Some(k) = oldest {
counters.remove(&k);
}
}
let counter = counters
.entry(key.to_string())
.or_insert_with(|| AtomicI64::new(0));
let current = counter.load(Ordering::Relaxed);
let global = self.global_current.load(Ordering::Relaxed);
let global_limit = self.global_limit.load(Ordering::Relaxed);
if current >= self.max_concurrent as i64 {
self.total_rejections.fetch_add(1, Ordering::Relaxed);
return Ok(RateLimitResult::rejected(0, now_timestamp() + 1000));
}
if global >= global_limit {
self.total_rejections.fetch_add(1, Ordering::Relaxed);
return Ok(RateLimitResult::rejected(0, now_timestamp() + 1000));
}
counter.fetch_add(1, Ordering::Relaxed);
self.global_current.fetch_add(1, Ordering::Relaxed);
let remaining = self.max_concurrent - counter.load(Ordering::Relaxed) as u64;
Ok(RateLimitResult::allowed(remaining, now_timestamp() + 60000))
}
pub fn release(&self, key: &str) -> Result<(), RateLimitError> {
self.total_releases.fetch_add(1, Ordering::Relaxed);
let counters = self
.counters
.read()
.map_err(|e| RateLimitError::Internal(e.to_string()))?;
if let Some(counter) = counters.get(key) {
let current = counter.load(Ordering::Relaxed);
if current > 0 {
counter.fetch_sub(1, Ordering::Relaxed);
self.global_current.fetch_sub(1, Ordering::Relaxed);
}
}
Ok(())
}
pub fn reset(&self, key: &str) -> Result<(), RateLimitError> {
let counters = self
.counters
.read()
.map_err(|e| RateLimitError::Internal(e.to_string()))?;
if let Some(counter) = counters.get(key) {
let current = counter.load(Ordering::Relaxed);
if current > 0 {
self.global_current.fetch_sub(current, Ordering::Relaxed);
counter.store(0, Ordering::Relaxed);
}
}
Ok(())
}
pub fn stats(&self) -> ConcurrencyStats {
ConcurrencyStats {
max_concurrent: self.max_concurrent,
global_limit: self.global_limit.load(Ordering::Relaxed),
global_current: self.global_current.load(Ordering::Relaxed),
key_count: self.key_count(),
total_acquires: self.total_acquires.load(Ordering::Relaxed),
total_releases: self.total_releases.load(Ordering::Relaxed),
total_rejections: self.total_rejections.load(Ordering::Relaxed),
}
}
}
#[derive(Debug, Clone, serde::Serialize)]
pub struct ConcurrencyStats {
pub max_concurrent: u64,
pub global_limit: i64,
pub global_current: i64,
pub key_count: usize,
pub total_acquires: u64,
pub total_releases: u64,
pub total_rejections: u64,
}
pub struct ConcurrencyGuard<'a> {
limiter: &'a ConcurrencyLimiter,
key: String,
acquired: bool,
}
impl<'a> ConcurrencyGuard<'a> {
pub fn acquire(limiter: &'a ConcurrencyLimiter, key: &str) -> Result<Self, RateLimitError> {
let result = limiter.acquire(key)?;
Ok(Self {
limiter,
key: key.to_string(),
acquired: result.allowed,
})
}
pub fn is_acquired(&self) -> bool {
self.acquired
}
pub fn release(mut self) {
if self.acquired {
let _ = self.limiter.release(&self.key);
self.acquired = false;
}
}
}
impl<'a> Drop for ConcurrencyGuard<'a> {
fn drop(&mut self) {
if self.acquired {
let _ = self.limiter.release(&self.key);
}
}
}
pub struct TimedConcurrencyLimiter {
inner: ConcurrencyLimiter,
wait_timeout: Duration,
retry_interval: Duration,
}
impl TimedConcurrencyLimiter {
pub fn new(max_concurrent: u64, wait_timeout: Duration) -> Self {
Self {
inner: ConcurrencyLimiter::new(max_concurrent),
wait_timeout,
retry_interval: Duration::from_millis(10),
}
}
pub fn with_retry_interval(mut self, interval: Duration) -> Self {
self.retry_interval = interval;
self
}
pub fn with_global_limit(self, limit: u64) -> Self {
Self {
inner: self.inner.with_global_limit(limit),
..self
}
}
pub fn acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
let start = std::time::Instant::now();
loop {
let result = self.inner.acquire(key)?;
if result.allowed {
return Ok(result);
}
if start.elapsed() >= self.wait_timeout {
return Ok(result);
}
std::thread::sleep(self.retry_interval);
}
}
pub fn release(&self, key: &str) -> Result<(), RateLimitError> {
self.inner.release(key)
}
pub fn inner(&self) -> &ConcurrencyLimiter {
&self.inner
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_concurrency_limiter_basic_acquire() {
let limiter = ConcurrencyLimiter::new(5);
let r = limiter.acquire("k").unwrap();
assert!(r.allowed);
assert_eq!(r.remaining, 4);
}
#[test]
fn test_concurrency_limiter_max_concurrent() {
let limiter = ConcurrencyLimiter::new(2);
assert!(limiter.acquire("k").unwrap().allowed);
assert!(limiter.acquire("k").unwrap().allowed);
let r3 = limiter.acquire("k").unwrap();
assert!(!r3.allowed);
}
#[test]
fn test_concurrency_limiter_release() {
let limiter = ConcurrencyLimiter::new(1);
assert!(limiter.acquire("k").unwrap().allowed);
assert!(!limiter.acquire("k").unwrap().allowed);
limiter.release("k").unwrap();
assert!(limiter.acquire("k").unwrap().allowed);
}
#[test]
fn test_concurrency_limiter_different_keys() {
let limiter = ConcurrencyLimiter::new(1);
assert!(limiter.acquire("a").unwrap().allowed);
assert!(limiter.acquire("b").unwrap().allowed);
}
#[test]
fn test_concurrency_limiter_current_concurrent() {
let limiter = ConcurrencyLimiter::new(5);
assert_eq!(limiter.current_concurrent("k"), 0);
limiter.acquire("k").unwrap();
assert_eq!(limiter.current_concurrent("k"), 1);
limiter.acquire("k").unwrap();
assert_eq!(limiter.current_concurrent("k"), 2);
}
#[test]
fn test_concurrency_limiter_global_current() {
let limiter = ConcurrencyLimiter::new(5);
assert_eq!(limiter.global_current(), 0);
limiter.acquire("a").unwrap();
limiter.acquire("b").unwrap();
assert_eq!(limiter.global_current(), 2);
limiter.release("a").unwrap();
assert_eq!(limiter.global_current(), 1);
}
#[test]
fn test_concurrency_limiter_global_limit() {
let limiter = ConcurrencyLimiter::new(10).with_global_limit(2);
assert!(limiter.acquire("a").unwrap().allowed);
assert!(limiter.acquire("b").unwrap().allowed);
let r3 = limiter.acquire("c").unwrap();
assert!(!r3.allowed);
}
#[test]
fn test_concurrency_limiter_reset() {
let limiter = ConcurrencyLimiter::new(5);
limiter.acquire("k").unwrap();
limiter.acquire("k").unwrap();
assert_eq!(limiter.current_concurrent("k"), 2);
limiter.reset("k").unwrap();
assert_eq!(limiter.current_concurrent("k"), 0);
}
#[test]
fn test_concurrency_limiter_stats() {
let limiter = ConcurrencyLimiter::new(3);
limiter.acquire("k").unwrap();
limiter.acquire("k").unwrap();
limiter.release("k").unwrap();
limiter.acquire("k").unwrap();
let stats = limiter.stats();
assert_eq!(stats.max_concurrent, 3);
assert_eq!(stats.total_acquires, 3);
assert_eq!(stats.total_releases, 1);
assert_eq!(stats.global_current, 2);
}
#[test]
fn test_concurrency_limiter_key_count() {
let limiter = ConcurrencyLimiter::new(5);
assert_eq!(limiter.key_count(), 0);
limiter.acquire("a").unwrap();
limiter.acquire("b").unwrap();
assert_eq!(limiter.key_count(), 2);
}
#[test]
fn test_concurrency_guard_auto_release() {
let limiter = ConcurrencyLimiter::new(1);
{
let guard = ConcurrencyGuard::acquire(&limiter, "k").unwrap();
assert!(guard.is_acquired());
assert_eq!(limiter.current_concurrent("k"), 1);
}
assert_eq!(limiter.current_concurrent("k"), 0);
}
#[test]
fn test_concurrency_guard_manual_release() {
let limiter = ConcurrencyLimiter::new(1);
let guard = ConcurrencyGuard::acquire(&limiter, "k").unwrap();
assert_eq!(limiter.current_concurrent("k"), 1);
guard.release();
assert_eq!(limiter.current_concurrent("k"), 0);
}
#[test]
fn test_concurrency_guard_rejected() {
let limiter = ConcurrencyLimiter::new(1);
let g1 = ConcurrencyGuard::acquire(&limiter, "k").unwrap();
let g2 = ConcurrencyGuard::acquire(&limiter, "k").unwrap();
assert!(g1.is_acquired());
assert!(!g2.is_acquired());
}
#[test]
fn test_timed_concurrency_limiter_immediate() {
let limiter = TimedConcurrencyLimiter::new(2, Duration::from_millis(100));
let r = limiter.acquire("k").unwrap();
assert!(r.allowed);
limiter.release("k").unwrap();
}
#[test]
fn test_timed_concurrency_limiter_timeout() {
let limiter = TimedConcurrencyLimiter::new(1, Duration::from_millis(50))
.with_retry_interval(Duration::from_millis(5));
let r1 = limiter.acquire("k").unwrap();
assert!(r1.allowed);
let r2 = limiter.acquire("k").unwrap();
assert!(!r2.allowed);
}
#[test]
fn test_concurrency_limiter_release_below_zero_guard() {
let limiter = ConcurrencyLimiter::new(5);
limiter.release("k").unwrap();
assert_eq!(limiter.current_concurrent("k"), 0);
}
#[test]
fn test_concurrency_limiter_double_release() {
let limiter = ConcurrencyLimiter::new(5);
limiter.acquire("k").unwrap();
limiter.release("k").unwrap();
limiter.release("k").unwrap();
assert_eq!(limiter.current_concurrent("k"), 0);
assert_eq!(limiter.global_current(), 0);
}
}