use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use std::sync::Arc;
const COST_EPSILON: f64 = 1e-9;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum BudgetKind {
CostUsd,
Tokens,
}
impl std::fmt::Display for BudgetKind {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
BudgetKind::CostUsd => f.write_str("cost budget (USD)"),
BudgetKind::Tokens => f.write_str("token budget"),
}
}
}
pub struct RouterBudget {
max_cost_usd: Option<f64>,
max_tokens: Option<u64>,
spent_bits: AtomicU64,
tokens: AtomicUsize,
trips: AtomicUsize,
}
impl std::fmt::Debug for RouterBudget {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("RouterBudget")
.field("max_cost_usd", &self.max_cost_usd)
.field("max_tokens", &self.max_tokens)
.field("spent_usd", &self.spent_usd())
.field("tokens", &self.tokens())
.field("trips", &self.trips())
.finish()
}
}
impl RouterBudget {
pub fn with_cost_limit(max_cost_usd: f64) -> Self {
Self::new(Some(max_cost_usd), None)
}
pub fn with_token_limit(max_tokens: u64) -> Self {
Self::new(None, Some(max_tokens))
}
pub fn with_cost_and_token_limits(max_cost_usd: f64, max_tokens: u64) -> Arc<Self> {
Arc::new(Self::new(Some(max_cost_usd), Some(max_tokens)))
}
fn new(max_cost_usd: Option<f64>, max_tokens: Option<u64>) -> Self {
Self {
max_cost_usd,
max_tokens,
spent_bits: AtomicU64::new(0.0f64.to_bits()),
tokens: AtomicUsize::new(0),
trips: AtomicUsize::new(0),
}
}
pub fn max_cost_usd(&self) -> Option<f64> {
self.max_cost_usd
}
pub fn max_tokens(&self) -> Option<u64> {
self.max_tokens
}
pub fn spent_usd(&self) -> f64 {
f64::from_bits(self.spent_bits.load(Ordering::Acquire))
}
pub fn tokens(&self) -> u64 {
self.tokens.load(Ordering::Acquire) as u64
}
pub fn trips(&self) -> usize {
self.trips.load(Ordering::Acquire)
}
pub fn is_tripped(&self) -> bool {
if let Some(limit) = self.max_cost_usd {
if self.spent_usd() > limit + COST_EPSILON {
return true;
}
}
if let Some(limit) = self.max_tokens {
if self.tokens() > limit {
return true;
}
}
false
}
pub fn precheck(
&self,
projected_cost_usd: f64,
projected_tokens: u64,
) -> Result<(), BudgetExceeded> {
if let Some(limit) = self.max_cost_usd {
if projected_cost_usd > 0.0
&& self.spent_usd() + projected_cost_usd > limit + COST_EPSILON
{
self.trips.fetch_add(1, Ordering::AcqRel);
return Err(BudgetExceeded {
kind: BudgetKind::CostUsd,
used: self.spent_usd(),
limit,
});
}
}
if let Some(limit) = self.max_tokens {
let used = self.tokens();
if used + projected_tokens > limit {
self.trips.fetch_add(1, Ordering::AcqRel);
return Err(BudgetExceeded {
kind: BudgetKind::Tokens,
used: used as f64,
limit: limit as f64,
});
}
}
Ok(())
}
pub fn record(&self, cost_usd: f64, total_tokens: u64) -> Option<BudgetExceeded> {
if cost_usd != 0.0 {
let mut cur = self.spent_bits.load(Ordering::Acquire);
loop {
let next = f64::from_bits(cur) + cost_usd;
match self.spent_bits.compare_exchange(
cur,
next.to_bits(),
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => break,
Err(actual) => cur = actual,
}
}
}
if total_tokens != 0 {
self.tokens
.fetch_add(total_tokens as usize, Ordering::AcqRel);
}
if let Some(limit) = self.max_cost_usd {
if self.spent_usd() > limit + COST_EPSILON {
self.trips.fetch_add(1, Ordering::AcqRel);
return Some(BudgetExceeded {
kind: BudgetKind::CostUsd,
used: self.spent_usd(),
limit,
});
}
}
if let Some(limit) = self.max_tokens {
if self.tokens() > limit {
self.trips.fetch_add(1, Ordering::AcqRel);
return Some(BudgetExceeded {
kind: BudgetKind::Tokens,
used: self.tokens() as f64,
limit: limit as f64,
});
}
}
None
}
pub fn reset(&self) {
self.spent_bits.store(0.0f64.to_bits(), Ordering::Release);
self.tokens.store(0, Ordering::Release);
self.trips.store(0, Ordering::Release);
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct BudgetExceeded {
pub kind: BudgetKind,
pub used: f64,
pub limit: f64,
}
impl std::fmt::Display for BudgetExceeded {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"{} exceeded: used {:.6}, limit {:.6}",
self.kind, self.used, self.limit
)
}
}
impl std::error::Error for BudgetExceeded {}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn allows_within_cost_cap_and_blocks_projected_overrun() {
let b = RouterBudget::with_cost_limit(1.0);
b.precheck(0.4, 10).unwrap();
b.record(0.4, 10);
assert_eq!(b.spent_usd(), 0.4);
let e = b.precheck(0.7, 0).unwrap_err();
assert_eq!(e.kind, BudgetKind::CostUsd);
assert_eq!(e.used, 0.4);
assert_eq!(e.limit, 1.0);
assert!(b.trips() >= 1);
}
#[test]
fn free_call_passes_cost_dimension_even_when_tripped() {
let b = RouterBudget::with_cost_limit(1.0);
b.record(2.0, 0);
assert!(b.is_tripped());
b.precheck(0.0, 0).unwrap();
assert!(b.precheck(0.01, 0).is_err());
}
#[test]
fn token_cap_counts_estimates_independent_of_price() {
let b = RouterBudget::with_token_limit(100);
b.precheck(0.0, 60).unwrap();
b.record(0.0, 60);
let e = b.precheck(0.0, 50).unwrap_err();
assert_eq!(e.kind, BudgetKind::Tokens);
assert_eq!(e.used, 60.0);
assert_eq!(e.limit, 100.0);
}
#[test]
fn record_latches_breaker_on_overshoot() {
let b = RouterBudget::with_cost_limit(1.0);
b.precheck(0.9, 0).unwrap();
let trip = b.record(1.5, 100).expect("should trip on record");
assert_eq!(trip.kind, BudgetKind::CostUsd);
assert!(b.is_tripped());
}
#[test]
fn concurrent_record_sums_without_losing_updates() {
let b = Arc::new(RouterBudget::with_cost_limit(f64::INFINITY));
let mut handles = Vec::new();
for _ in 0..8 {
let b = b.clone();
handles.push(std::thread::spawn(move || {
for _ in 0..1000 {
b.record(0.001, 1);
}
}));
}
for h in handles {
h.join().unwrap();
}
assert!((b.spent_usd() - 8.0).abs() < 1e-9);
assert_eq!(b.tokens(), 8000);
}
#[test]
fn reset_clears_totals_and_trip() {
let b = RouterBudget::with_cost_limit(1.0);
b.record(2.0, 50);
assert!(b.is_tripped());
b.reset();
assert!(!b.is_tripped());
assert_eq!(b.spent_usd(), 0.0);
assert_eq!(b.tokens(), 0);
assert_eq!(b.trips(), 0);
}
}