use crate::error::AuthError;
use std::collections::{BTreeMap, HashMap};
use std::future::Future;
use std::pin::Pin;
use std::sync::{Arc, Mutex};
use std::task::{Context, Poll, Waker};
use std::time::{Duration, Instant};
pub struct ConcurrencyLimiter {
inner: Mutex<SemState>,
max: usize,
}
struct SemState {
available: usize,
next_id: u64,
waiters: BTreeMap<u64, Waker>,
}
impl SemState {
fn wake_front(&self) {
if let Some((_, waker)) = self.waiters.iter().next() {
waker.wake_by_ref();
}
}
}
impl ConcurrencyLimiter {
pub fn new(max: usize) -> Self {
Self {
inner: Mutex::new(SemState {
available: max,
next_id: 0,
waiters: BTreeMap::new(),
}),
max,
}
}
pub fn acquire(self: &Arc<Self>) -> Acquire {
Acquire {
limiter: Arc::clone(self),
id: None,
}
}
pub fn current(&self) -> usize {
let st = self.inner.lock().expect("concurrency limiter poisoned");
self.max - st.available
}
fn release(&self) {
let mut st = self.inner.lock().expect("concurrency limiter poisoned");
st.available += 1;
st.wake_front();
}
}
pub struct Acquire {
limiter: Arc<ConcurrencyLimiter>,
id: Option<u64>,
}
impl Future for Acquire {
type Output = ConcurrencyPermit;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let this = self.get_mut();
let mut st = this.limiter.inner.lock().expect("concurrency limiter poisoned");
let is_front = match this.id {
Some(id) => st.waiters.keys().next() == Some(&id),
None => st.waiters.is_empty(),
};
if st.available > 0 && is_front {
st.available -= 1;
if let Some(id) = this.id.take() {
st.waiters.remove(&id);
}
if st.available > 0 {
st.wake_front();
}
return Poll::Ready(ConcurrencyPermit {
limiter: Arc::clone(&this.limiter),
});
}
let id = match this.id {
Some(id) => id,
None => {
let id = st.next_id;
st.next_id = st.next_id.wrapping_add(1);
this.id = Some(id);
id
}
};
st.waiters.insert(id, cx.waker().clone());
Poll::Pending
}
}
impl Drop for Acquire {
fn drop(&mut self) {
if let Some(id) = self.id.take() {
let mut st = self.limiter.inner.lock().expect("concurrency limiter poisoned");
let was_front = st.waiters.keys().next() == Some(&id);
st.waiters.remove(&id);
if was_front && st.available > 0 {
st.wake_front();
}
}
}
}
pub struct ConcurrencyPermit {
limiter: Arc<ConcurrencyLimiter>,
}
impl Drop for ConcurrencyPermit {
fn drop(&mut self) {
self.limiter.release();
}
}
pub struct RateLimiter {
max_requests: usize,
window: Duration,
state: Mutex<RateState>,
}
struct RateState {
entries: HashMap<String, RateEntry>,
last_prune: Instant,
}
#[derive(Clone, Copy)]
struct RateEntry {
count: u64,
window_start: Instant,
}
impl RateLimiter {
pub fn new(max_requests: usize, window: Duration) -> Self {
Self {
max_requests,
window,
state: Mutex::new(RateState {
entries: HashMap::new(),
last_prune: Instant::now(),
}),
}
}
pub fn check(&self, ip: &str) -> Result<(), AuthError> {
let now = Instant::now();
let window = self.window;
let mut state = self.state.lock().expect("rate limiter state poisoned");
if now.duration_since(state.last_prune) >= window {
state
.entries
.retain(|_, entry| now.duration_since(entry.window_start) < window);
state.last_prune = now;
}
if let Some(entry) = state.entries.get_mut(ip) {
if now.duration_since(entry.window_start) >= window {
entry.window_start = now;
entry.count = 1;
return Ok(());
}
entry.count += 1;
return if entry.count > self.max_requests as u64 {
Err(AuthError::RateLimited)
} else {
Ok(())
};
}
state.entries.insert(
ip.to_string(),
RateEntry {
count: 1,
window_start: now,
},
);
Ok(())
}
pub fn prune(&self) {
let now = Instant::now();
let window = self.window;
let mut state = self.state.lock().expect("rate limiter state poisoned");
state
.entries
.retain(|_, entry| now.duration_since(entry.window_start) < window);
state.last_prune = now;
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_concurrency_limiter_releases_on_drop() {
let limiter = Arc::new(ConcurrencyLimiter::new(2));
assert_eq!(limiter.current(), 0);
let p1 = limiter.acquire().await;
let _p2 = limiter.acquire().await;
assert_eq!(limiter.current(), 2);
drop(p1);
assert_eq!(limiter.current(), 1);
}
#[tokio::test]
async fn test_concurrency_limiter_parks_and_wakes_when_full() {
let limiter = Arc::new(ConcurrencyLimiter::new(2));
let p1 = limiter.acquire().await;
let _p2 = limiter.acquire().await;
let mut fut = Box::pin(limiter.acquire());
let waker = Waker::from(Arc::new(NoopWake));
let mut cx = Context::from_waker(&waker);
assert!(fut.as_mut().poll(&mut cx).is_pending());
assert_eq!(limiter.current(), 2);
drop(p1);
let p3 = match fut.as_mut().poll(&mut cx) {
Poll::Ready(permit) => permit,
Poll::Pending => panic!("acquire should resolve once a permit is freed"),
};
assert_eq!(limiter.current(), 2); drop(p3);
assert_eq!(limiter.current(), 1); }
#[tokio::test]
async fn test_concurrency_limiter_fifo_no_barging() {
let limiter = Arc::new(ConcurrencyLimiter::new(1));
let p1 = limiter.acquire().await;
let waker = Waker::from(Arc::new(NoopWake));
let mut cx = Context::from_waker(&waker);
let mut a = Box::pin(limiter.acquire());
assert!(a.as_mut().poll(&mut cx).is_pending());
drop(p1);
let mut b = Box::pin(limiter.acquire());
assert!(
b.as_mut().poll(&mut cx).is_pending(),
"b must not take the permit ahead of the queued front waiter a"
);
assert!(
a.as_mut().poll(&mut cx).is_ready(),
"front waiter a should receive the freed permit"
);
}
struct NoopWake;
impl std::task::Wake for NoopWake {
fn wake(self: Arc<Self>) {}
}
#[test]
fn test_rate_limiter_allows_until_limit() {
let limiter = RateLimiter::new(2, Duration::from_secs(60));
assert!(limiter.check("1.2.3.4").is_ok());
assert!(limiter.check("1.2.3.4").is_ok());
assert!(matches!(limiter.check("1.2.3.4"), Err(AuthError::RateLimited)));
assert!(limiter.check("5.6.7.8").is_ok());
}
#[test]
fn test_rate_limiter_prune_bounds_map() {
let limiter = RateLimiter::new(5, Duration::from_millis(1));
for i in 0..100 {
let _ = limiter.check(&format!("10.0.0.{i}"));
}
std::thread::sleep(Duration::from_millis(5));
limiter.prune();
let state = limiter.state.lock().unwrap();
assert_eq!(state.entries.len(), 0);
}
}