use std::collections::HashMap;
use std::net::IpAddr;
use std::sync::{Arc, Mutex};
#[derive(Debug)]
pub struct Limiter {
global_cap: usize,
per_ip_cap: usize,
inner: Mutex<LimiterInner>,
}
#[derive(Debug, Default)]
struct LimiterInner {
global: usize,
per_ip: HashMap<IpAddr, usize>,
}
pub struct Permit {
limiter: Arc<Limiter>,
ip: IpAddr,
}
impl std::fmt::Debug for Permit {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Permit").field("ip", &self.ip).finish()
}
}
impl Limiter {
pub fn new(global_cap: usize, per_ip_cap: usize) -> Arc<Self> {
Arc::new(Self {
global_cap,
per_ip_cap,
inner: Mutex::new(LimiterInner::default()),
})
}
pub fn try_acquire(self: &Arc<Self>, ip: IpAddr) -> Option<Permit> {
let mut inner = self.inner.lock().expect("limiter mutex poisoned");
if inner.global >= self.global_cap {
return None;
}
let per_ip = inner.per_ip.get(&ip).copied().unwrap_or(0);
if per_ip >= self.per_ip_cap {
return None;
}
inner.global += 1;
*inner.per_ip.entry(ip).or_insert(0) += 1;
drop(inner);
Some(Permit {
limiter: Arc::clone(self),
ip,
})
}
pub fn snapshot(&self) -> (usize, usize) {
let inner = self.inner.lock().expect("limiter mutex poisoned");
(inner.global, inner.per_ip.len())
}
}
impl Drop for Permit {
fn drop(&mut self) {
let mut inner = self.limiter.inner.lock().expect("limiter mutex poisoned");
inner.global = inner.global.saturating_sub(1);
if let Some(c) = inner.per_ip.get_mut(&self.ip) {
*c = c.saturating_sub(1);
if *c == 0 {
inner.per_ip.remove(&self.ip);
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn ip(n: u8) -> IpAddr {
IpAddr::from([127, 0, 0, n])
}
#[test]
fn releases_on_drop_decrements_counters() {
let l = Limiter::new(10, 10);
{
let _p = l.try_acquire(ip(1)).expect("acquire");
assert_eq!(l.snapshot().0, 1);
}
assert_eq!(l.snapshot().0, 0);
}
#[test]
fn global_cap_rejects_excess() {
let l = Limiter::new(2, 10);
let _a = l.try_acquire(ip(1)).expect("first");
let _b = l.try_acquire(ip(2)).expect("second");
assert!(l.try_acquire(ip(3)).is_none(), "third must be rejected");
}
#[test]
fn per_ip_cap_rejects_excess_from_same_ip() {
let l = Limiter::new(100, 2);
let _a = l.try_acquire(ip(1)).expect("first");
let _b = l.try_acquire(ip(1)).expect("second");
assert!(
l.try_acquire(ip(1)).is_none(),
"third from same ip rejected"
);
let _c = l.try_acquire(ip(2)).expect("different ip ok");
}
#[test]
fn rejection_does_not_increment_counters() {
let l = Limiter::new(1, 10);
let _a = l.try_acquire(ip(1)).expect("first");
assert_eq!(l.snapshot().0, 1);
let rejected = l.try_acquire(ip(2));
assert!(rejected.is_none());
assert_eq!(l.snapshot().0, 1, "rejection must not increment");
}
}