use std::sync::Arc;
#[cfg(feature = "loom-check")]
use loom::sync::atomic::{AtomicUsize, Ordering};
#[cfg(not(feature = "loom-check"))]
use std::sync::atomic::{AtomicUsize, Ordering};
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum LimiterError {
ConcurrencyLimit { max: usize },
}
impl std::fmt::Display for LimiterError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::ConcurrencyLimit { max } => {
write!(f, "max concurrency reached (limit: {})", max)
}
}
}
}
impl std::error::Error for LimiterError {}
pub struct AgentExecutionLimiter {
max_concurrency: usize,
current: AtomicUsize,
}
impl std::fmt::Debug for AgentExecutionLimiter {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("AgentExecutionLimiter")
.field("max_concurrency", &self.max_concurrency)
.field("current", &self.current.load(Ordering::Relaxed))
.finish()
}
}
impl AgentExecutionLimiter {
pub fn new(max_concurrency: usize) -> Self {
Self {
max_concurrency,
current: AtomicUsize::new(0),
}
}
pub fn unlimited() -> Self {
Self::new(usize::MAX)
}
pub fn max_concurrency(&self) -> usize {
self.max_concurrency
}
pub fn try_acquire(self: &Arc<Self>) -> Result<ExecutionSlot, LimiterError> {
let mut cur = self.current.load(Ordering::Acquire);
loop {
if cur >= self.max_concurrency {
return Err(LimiterError::ConcurrencyLimit {
max: self.max_concurrency,
});
}
match self.current.compare_exchange_weak(
cur,
cur + 1,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => {
return Ok(ExecutionSlot {
limiter: Arc::clone(self),
});
}
Err(actual) => cur = actual,
}
}
}
pub fn current(&self) -> usize {
self.current.load(Ordering::Acquire)
}
}
impl Default for AgentExecutionLimiter {
fn default() -> Self {
Self::unlimited()
}
}
pub struct ExecutionSlot {
limiter: Arc<AgentExecutionLimiter>,
}
impl std::fmt::Debug for ExecutionSlot {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ExecutionSlot")
.field("limiter", &self.limiter)
.finish()
}
}
impl ExecutionSlot {
pub fn limiter(&self) -> &Arc<AgentExecutionLimiter> {
&self.limiter
}
}
impl Drop for ExecutionSlot {
fn drop(&mut self) {
self.limiter.current.fetch_sub(1, Ordering::AcqRel);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn acquire_release_round_trip() {
let limiter = Arc::new(AgentExecutionLimiter::new(2));
assert_eq!(limiter.current(), 0);
let a = limiter.try_acquire().expect("first slot");
assert_eq!(limiter.current(), 1);
let b = limiter.try_acquire().expect("second slot");
assert_eq!(limiter.current(), 2);
assert_eq!(
limiter.try_acquire().unwrap_err(),
LimiterError::ConcurrencyLimit { max: 2 }
);
assert_eq!(limiter.current(), 2);
drop(a);
assert_eq!(limiter.current(), 1);
let c = limiter.try_acquire().expect("slot after release");
assert_eq!(limiter.current(), 2);
drop(c);
drop(b);
assert_eq!(limiter.current(), 0);
}
#[test]
fn unlimited_never_rejects_but_tracks() {
let limiter = Arc::new(AgentExecutionLimiter::unlimited());
let a = limiter.try_acquire().expect("always ok");
let b = limiter.try_acquire().expect("always ok");
assert_eq!(limiter.current(), 2);
drop(a);
drop(b);
assert_eq!(limiter.current(), 0);
}
#[test]
fn zero_cap_rejects_immediately() {
let limiter = Arc::new(AgentExecutionLimiter::new(0));
assert_eq!(
limiter.try_acquire().unwrap_err(),
LimiterError::ConcurrencyLimit { max: 0 }
);
assert_eq!(limiter.current(), 0);
}
#[test]
fn concurrent_acquire_release_conservation() {
const THREADS: usize = 8;
const ITERS: usize = 500;
let limiter = Arc::new(AgentExecutionLimiter::new(THREADS)); let max_seen = Arc::new(AtomicUsize::new(0));
let handles: Vec<_> = (0..THREADS)
.map(|_| {
let limiter = Arc::clone(&limiter);
let max_seen = Arc::clone(&max_seen);
std::thread::spawn(move || {
for _ in 0..ITERS {
let slot = limiter.try_acquire().expect("cap == threads, must succeed");
let cur = limiter.current();
max_seen.fetch_max(cur, Ordering::AcqRel);
drop(slot);
}
})
})
.collect();
for h in handles {
h.join().unwrap();
}
assert_eq!(limiter.current(), 0, "counter conserved to exactly 0");
assert!(max_seen.load(Ordering::Acquire) <= THREADS);
}
#[test]
#[cfg(feature = "loom-check")]
fn loom_cap_one_exclusivity_across_serialized_rounds() {
loom::model(|| {
let limiter = Arc::new(AgentExecutionLimiter::new(1));
let wins = Arc::new(loom::sync::atomic::AtomicUsize::new(0));
let mut handles = Vec::new();
for _ in 0..2 {
let l = Arc::clone(&limiter);
let w = Arc::clone(&wins);
handles.push(loom::thread::spawn(move || {
if let Ok(slot) = l.try_acquire() {
w.fetch_add(1, Ordering::Relaxed);
assert_eq!(l.current(), 1, "the holder is the only booking");
drop(slot);
if let Ok(_s) = l.try_acquire() {
w.fetch_add(1, Ordering::Relaxed);
}
}
}));
}
for h in handles {
h.join().unwrap();
}
let total = wins.load(Ordering::Relaxed);
assert!((2..=4).contains(&total), "bookings bounded: {total}");
assert_eq!(limiter.current(), 0, "every slot released");
});
}
}