use std::sync::Arc;
use std::sync::Mutex;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::Duration;
use tokio::sync::{OwnedSemaphorePermit, Semaphore};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Outcome {
Success,
Overload,
}
#[derive(Debug)]
pub struct AdaptiveLimiter {
sem: Arc<Semaphore>,
limit: AtomicUsize,
debt: AtomicUsize,
in_flight: AtomicUsize,
resize: Mutex<()>,
min_limit: usize,
max_limit: usize,
increase_by: usize,
decrease_factor: f64,
}
const GROW_UTILISATION_PERCENT: usize = 80;
impl AdaptiveLimiter {
#[must_use]
pub fn new(
initial: usize,
min_limit: usize,
max_limit: usize,
increase_by: usize,
decrease_factor: f64,
) -> Arc<Self> {
let min = min_limit.max(1);
let max = max_limit.max(min);
let initial = initial.clamp(min, max);
let factor = decrease_factor.clamp(0.5, 0.999);
Arc::new(Self {
sem: Arc::new(Semaphore::new(initial)),
limit: AtomicUsize::new(initial),
debt: AtomicUsize::new(0),
in_flight: AtomicUsize::new(0),
resize: Mutex::new(()),
min_limit: min,
max_limit: max,
increase_by: increase_by.max(1),
decrease_factor: factor,
})
}
#[must_use]
pub fn limit(&self) -> usize {
self.limit.load(Ordering::Acquire)
}
#[must_use]
pub fn in_flight(&self) -> usize {
self.in_flight.load(Ordering::Acquire)
}
pub async fn acquire(self: &Arc<Self>) -> Permit {
let permit = Arc::clone(&self.sem)
.acquire_owned()
.await
.expect("adaptive limiter semaphore is never closed");
self.in_flight.fetch_add(1, Ordering::AcqRel);
Permit {
inner: Some(permit),
limiter: Arc::clone(self),
}
}
pub async fn acquire_timeout(self: &Arc<Self>, timeout: Duration) -> Option<Permit> {
tokio::time::timeout(timeout, self.acquire()).await.ok()
}
#[must_use]
pub fn try_acquire(self: &Arc<Self>) -> Option<Permit> {
let permit = Arc::clone(&self.sem).try_acquire_owned().ok()?;
self.in_flight.fetch_add(1, Ordering::AcqRel);
Some(Permit {
inner: Some(permit),
limiter: Arc::clone(self),
})
}
pub fn record(&self, outcome: Outcome) {
let _guard = self
.resize
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let current = self.limit.load(Ordering::Acquire);
let new = match outcome {
Outcome::Overload => {
#[allow(clippy::cast_possible_truncation, clippy::cast_sign_loss)]
let shrunk = (current as f64 * self.decrease_factor) as usize;
shrunk.max(self.min_limit)
}
Outcome::Success => {
let in_flight = self.in_flight.load(Ordering::Acquire);
if in_flight * 100 >= current * GROW_UTILISATION_PERCENT {
(current + self.increase_by).min(self.max_limit)
} else {
current
}
}
};
if new != current {
self.apply_limit(current, new);
}
}
fn apply_limit(&self, current: usize, new: usize) {
if new > current {
let mut grow = new - current;
let debt = self.debt.load(Ordering::Acquire);
let cancel = grow.min(debt);
if cancel > 0 {
self.debt.fetch_sub(cancel, Ordering::AcqRel);
grow -= cancel;
}
if grow > 0 {
self.sem.add_permits(grow);
}
} else {
let shrink = current - new;
let forgot = self.sem.forget_permits(shrink);
if shrink > forgot {
self.debt.fetch_add(shrink - forgot, Ordering::AcqRel);
}
}
self.limit.store(new, Ordering::Release);
}
}
#[derive(Debug)]
pub struct Permit {
inner: Option<OwnedSemaphorePermit>,
limiter: Arc<AdaptiveLimiter>,
}
impl Drop for Permit {
fn drop(&mut self) {
self.limiter.in_flight.fetch_sub(1, Ordering::AcqRel);
let permit = self.inner.take().expect("permit held until drop");
let took_debt = self.debt_take().is_some();
if took_debt {
permit.forget();
} else {
drop(permit);
}
}
}
impl Permit {
fn debt_take(&self) -> Option<()> {
self.limiter
.debt
.fetch_update(Ordering::AcqRel, Ordering::Acquire, |d| {
if d > 0 { Some(d - 1) } else { None }
})
.ok()
.map(|_| ())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn gates_to_the_limit() {
let lim = AdaptiveLimiter::new(2, 1, 10, 1, 0.5);
let _p1 = lim.try_acquire().expect("1st");
let _p2 = lim.try_acquire().expect("2nd");
assert!(lim.try_acquire().is_none(), "3rd blocked at limit 2");
assert_eq!(lim.in_flight(), 2);
}
#[tokio::test]
async fn min_limit_floors_at_one_no_deadlock() {
let lim = AdaptiveLimiter::new(8, 1, 8, 1, 0.5);
for _ in 0..50 {
lim.record(Outcome::Overload);
assert!(lim.limit() >= 1, "limit floored at 1");
}
assert_eq!(lim.limit(), 1);
assert!(lim.try_acquire().is_some());
}
#[tokio::test]
async fn grows_only_when_well_utilised() {
let lim = AdaptiveLimiter::new(4, 1, 100, 1, 0.5);
lim.record(Outcome::Success);
assert_eq!(lim.limit(), 4, "idle success does not grow the limit");
let permits: Vec<_> = (0..4).map(|_| lim.try_acquire().unwrap()).collect();
lim.record(Outcome::Success);
assert_eq!(lim.limit(), 5, "well-utilised success grows the limit");
drop(permits);
}
#[tokio::test]
async fn shrink_never_over_admits() {
let lim = AdaptiveLimiter::new(10, 1, 10, 1, 0.5);
let permits: Vec<_> = (0..10).map(|_| lim.try_acquire().unwrap()).collect();
assert!(lim.try_acquire().is_none());
lim.record(Outcome::Overload); assert_eq!(lim.limit(), 5);
drop(permits);
let mut reacquired = Vec::new();
while let Some(p) = lim.try_acquire() {
reacquired.push(p);
}
assert_eq!(
reacquired.len(),
5,
"total capacity must equal the new limit, not over-admit"
);
}
#[tokio::test]
async fn grow_wakes_a_waiter_without_spin() {
let lim = AdaptiveLimiter::new(1, 1, 10, 1, 0.5);
let held = lim.try_acquire().unwrap();
let lim2 = Arc::clone(&lim);
let waiter = tokio::spawn(async move { lim2.acquire().await });
tokio::task::yield_now().await;
lim.apply_limit(1, 2);
let _woken = tokio::time::timeout(Duration::from_secs(1), waiter)
.await
.expect("waiter must wake promptly, not spin")
.unwrap();
drop(held);
}
#[tokio::test]
async fn acquire_timeout_sheds_when_saturated() {
let lim = AdaptiveLimiter::new(1, 1, 10, 1, 0.5);
let _held = lim.try_acquire().unwrap();
let shed = lim.acquire_timeout(Duration::from_millis(20)).await;
assert!(shed.is_none(), "saturated acquire sheds after the timeout");
}
}