use std::sync::Arc;
#[cfg(feature = "loom-check")]
use loom::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
#[cfg(not(feature = "loom-check"))]
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum BudgetError {
MaxSpawnCountReached { max: usize },
TokenBudgetExhausted { used: u64, max: u64 },
}
impl std::fmt::Display for BudgetError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::MaxSpawnCountReached { max } => {
write!(f, "max spawn count reached (limit: {})", max)
}
Self::TokenBudgetExhausted { used, max } => {
write!(f, "child token budget exhausted ({} / {})", used, max)
}
}
}
}
impl std::error::Error for BudgetError {}
pub fn usage_total(u: &agent_base::UsageInfo) -> u64 {
match u.total_tokens {
Some(t) => t as u64,
None => u.prompt_tokens.unwrap_or(0) as u64 + u.completion_tokens.unwrap_or(0) as u64,
}
}
pub struct RolloutBudget {
child_max_tokens: u64, used_tokens: AtomicU64,
max_spawns: usize, spawn_count: AtomicUsize,
}
impl std::fmt::Debug for RolloutBudget {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("RolloutBudget")
.field("child_max_tokens", &self.child_max_tokens)
.field("used_tokens", &self.used_tokens.load(Ordering::Relaxed))
.field("max_spawns", &self.max_spawns)
.field("spawn_count", &self.spawn_count.load(Ordering::Relaxed))
.finish()
}
}
impl RolloutBudget {
pub fn new(child_max_tokens: Option<u64>, max_spawns: Option<usize>) -> Self {
Self {
child_max_tokens: child_max_tokens.unwrap_or(u64::MAX),
used_tokens: AtomicU64::new(0),
max_spawns: max_spawns.unwrap_or(usize::MAX),
spawn_count: AtomicUsize::new(0),
}
}
pub fn try_reserve_spawn(self: &Arc<Self>) -> Result<SpawnTicket, BudgetError> {
let prev = self.spawn_count.fetch_add(1, Ordering::AcqRel);
if prev >= self.max_spawns {
self.spawn_count.fetch_sub(1, Ordering::AcqRel);
return Err(BudgetError::MaxSpawnCountReached {
max: self.max_spawns,
});
}
let used = self.used_tokens.load(Ordering::Acquire);
if used >= self.child_max_tokens {
self.spawn_count.fetch_sub(1, Ordering::AcqRel);
return Err(BudgetError::TokenBudgetExhausted {
used,
max: self.child_max_tokens,
});
}
Ok(SpawnTicket {
budget: Some(Arc::clone(self)),
})
}
pub fn record_usage(&self, tokens: u64) {
self.used_tokens.fetch_add(tokens, Ordering::Relaxed);
}
pub fn used_tokens(&self) -> u64 {
self.used_tokens.load(Ordering::Acquire)
}
pub fn spawn_count(&self) -> usize {
self.spawn_count.load(Ordering::Acquire)
}
pub fn child_max_tokens(&self) -> u64 {
self.child_max_tokens
}
pub fn max_spawns(&self) -> usize {
self.max_spawns
}
}
#[must_use = "an uncommitted ticket rolls back the spawn count on drop"]
pub struct SpawnTicket {
budget: Option<Arc<RolloutBudget>>,
}
impl std::fmt::Debug for SpawnTicket {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("SpawnTicket")
.field("committed", &self.budget.is_none())
.finish()
}
}
impl SpawnTicket {
pub fn commit(mut self) {
self.budget = None;
}
}
impl Drop for SpawnTicket {
fn drop(&mut self) {
if let Some(budget) = &self.budget {
budget.spawn_count.fetch_sub(1, Ordering::AcqRel);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn budget(max_spawns: Option<usize>, child_max_tokens: Option<u64>) -> Arc<RolloutBudget> {
Arc::new(RolloutBudget::new(child_max_tokens, max_spawns))
}
#[test]
fn spawn_cap_rejects_and_restores_exact() {
let b = budget(Some(2), None);
let t1 = b.try_reserve_spawn().expect("1st");
let t2 = b.try_reserve_spawn().expect("2nd");
assert_eq!(b.spawn_count(), 2);
assert_eq!(
b.try_reserve_spawn().unwrap_err(),
BudgetError::MaxSpawnCountReached { max: 2 }
);
assert_eq!(b.spawn_count(), 2);
t1.commit();
assert_eq!(b.spawn_count(), 2);
drop(t2);
assert_eq!(b.spawn_count(), 1);
let t3 = b.try_reserve_spawn().expect("room after rollback");
assert_eq!(b.spawn_count(), 2);
t3.commit();
}
#[tokio::test(flavor = "multi_thread")]
async fn committed_spawns_never_released_by_close() {
let b = budget(Some(1), None);
let t = b.try_reserve_spawn().expect("first");
t.commit();
assert_eq!(
b.try_reserve_spawn().unwrap_err(),
BudgetError::MaxSpawnCountReached { max: 1 }
);
}
#[test]
fn token_exhaustion_rejects_and_restores_count() {
let b = budget(None, Some(100));
let t = b.try_reserve_spawn().expect("under budget");
t.commit();
b.record_usage(100);
assert_eq!(
b.try_reserve_spawn().unwrap_err(),
BudgetError::TokenBudgetExhausted {
used: 100,
max: 100
}
);
assert_eq!(b.spawn_count(), 1);
}
#[test]
fn unlimited_by_default() {
let b = budget(None, None);
let mut ts = Vec::new();
for _ in 0..1000 {
ts.push(b.try_reserve_spawn().expect("unlimited"));
}
assert_eq!(b.spawn_count(), 1000);
for t in ts {
t.commit();
}
assert_eq!(b.spawn_count(), 1000);
}
#[test]
fn usage_total_prefers_provider_total() {
let u = agent_base::UsageInfo {
prompt_tokens: Some(10),
completion_tokens: Some(5),
total_tokens: Some(20),
reasoning_tokens: None,
};
assert_eq!(usage_total(&u), 20);
let u2 = agent_base::UsageInfo {
total_tokens: None,
..u
};
assert_eq!(usage_total(&u2), 15);
let none = agent_base::UsageInfo {
prompt_tokens: None,
completion_tokens: None,
total_tokens: None,
reasoning_tokens: None,
};
assert_eq!(usage_total(&none), 0);
}
#[tokio::test(flavor = "multi_thread")]
async fn cas_conservation_thousand_mixed_commit_rollback() {
const N: usize = 1000;
let b = budget(None, None);
let mut set = tokio::task::JoinSet::new();
for i in 0..N {
let b = Arc::clone(&b);
set.spawn(async move {
let t = b.try_reserve_spawn().expect("unlimited");
tokio::task::yield_now().await;
if i % 3 != 0 {
t.commit();
}
});
}
while set.join_next().await.is_some() {}
let committed = (0..N).filter(|i| i % 3 != 0).count();
assert_eq!(
b.spawn_count(),
committed,
"CAS conservation: commits minus rollbacks exactly"
);
}
#[test]
#[cfg(feature = "loom-check")]
fn loom_reserve_commit_vs_rollback_race() {
loom::model(|| {
let b = Arc::new(RolloutBudget::new(None, Some(1)));
let h1 = {
let b = Arc::clone(&b);
loom::thread::spawn(move || match b.try_reserve_spawn() {
Ok(t) => {
t.commit();
true
}
Err(_) => false,
})
};
let h2 = {
let b = Arc::clone(&b);
loom::thread::spawn(move || match b.try_reserve_spawn() {
Ok(t) => {
drop(t); true
}
Err(_) => false,
})
};
let committed_by_t1 = h1.join().unwrap();
let got_and_returned_t2 = h2.join().unwrap();
assert!(
usize::from(committed_by_t1) <= 1,
"cumulative commits ≤ cap"
);
assert_eq!(b.spawn_count(), usize::from(committed_by_t1));
let _ = got_and_returned_t2;
});
}
#[test]
#[cfg(feature = "loom-check")]
fn loom_token_gate_is_a_throttle_not_a_quota() {
loom::model(|| {
let b = Arc::new(RolloutBudget::new(Some(1), None));
let u = {
let b = Arc::clone(&b);
loom::thread::spawn(move || b.record_usage(1))
};
let r = {
let b = Arc::clone(&b);
loom::thread::spawn(move || match b.try_reserve_spawn() {
Ok(t) => {
t.commit();
true
}
Err(_) => false,
})
};
u.join().unwrap();
let admitted = r.join().unwrap();
assert_eq!(b.used_tokens(), 1);
assert_eq!(
b.spawn_count(),
usize::from(admitted),
"a rejected reservation rolled back exactly"
);
});
}
}