use crate::costs::UsageSnapshot;
use crate::scenario::{BudgetDef, BudgetEnforcement, BudgetsConfig};
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum BudgetStatus {
Ok,
SoftExceeded {
budget: String,
message: String,
},
HardExceeded {
budget: String,
message: String,
},
}
#[derive(Debug, Clone)]
pub struct BudgetTracker {
global: Option<ResolvedBudget>,
per_test_default: Option<ResolvedBudget>,
}
#[derive(Debug, Clone)]
pub(crate) struct ResolvedBudget {
max_cost: Option<f64>,
max_tokens: Option<u64>,
max_calls: Option<u64>,
enforcement: BudgetEnforcement,
}
impl ResolvedBudget {
fn from_def(def: &BudgetDef) -> Self {
Self {
max_cost: def.max_cost,
max_tokens: def.max_tokens,
max_calls: def.max_calls,
enforcement: def.enforcement.clone().unwrap_or(BudgetEnforcement::Hard),
}
}
}
impl BudgetTracker {
#[must_use]
pub fn from_config(budgets: &BudgetsConfig) -> Self {
Self {
global: budgets.global.as_ref().map(ResolvedBudget::from_def),
per_test_default: budgets
.per_test_default
.as_ref()
.map(ResolvedBudget::from_def),
}
}
#[must_use]
pub fn check_global(&self, usage: &UsageSnapshot) -> BudgetStatus {
let Some(global) = &self.global else {
return BudgetStatus::Ok;
};
Self::check_budget("global", global, usage)
}
#[must_use]
pub fn check_per_test(
&self,
test_name: &str,
usage: &UsageSnapshot,
override_budget: Option<&BudgetDef>,
) -> BudgetStatus {
let budget = match override_budget {
Some(def) => ResolvedBudget::from_def(def),
None => match &self.per_test_default {
Some(def) => def.clone(),
None => return BudgetStatus::Ok,
},
};
Self::check_budget(test_name, &budget, usage)
}
#[must_use]
#[allow(dead_code)]
#[allow(clippy::cast_precision_loss, clippy::suboptimal_flops)]
pub(crate) fn check_pre_flight_llm(
budget: &ResolvedBudget,
usage: &UsageSnapshot,
estimated_input_tokens: u64,
max_tokens: u64,
input_price_per_1m: f64,
output_price_per_1m: f64,
) -> BudgetStatus {
let estimated_cost = (estimated_input_tokens as f64 / 1_000_000.0) * input_price_per_1m
+ (max_tokens as f64 / 1_000_000.0) * output_price_per_1m;
let estimated_total_tokens = usage.total_tokens + estimated_input_tokens + max_tokens;
Self::check_limits(
budget,
"pre-flight",
usage,
estimated_cost,
estimated_total_tokens,
1,
)
}
#[must_use]
#[allow(dead_code)]
pub(crate) fn check_pre_flight_flat(
budget: &ResolvedBudget,
usage: &UsageSnapshot,
per_call_price: f64,
) -> BudgetStatus {
Self::check_limits(budget, "pre-flight", usage, per_call_price, 0, 1)
}
fn check_budget(name: &str, budget: &ResolvedBudget, usage: &UsageSnapshot) -> BudgetStatus {
Self::check_limits(budget, name, usage, 0.0, 0, 0)
}
#[allow(clippy::cast_precision_loss)]
fn check_limits(
budget: &ResolvedBudget,
name: &str,
usage: &UsageSnapshot,
additional_cost: f64,
additional_tokens: u64,
additional_calls: u64,
) -> BudgetStatus {
let projected_cost = usage.total_cost + additional_cost;
let projected_tokens = usage.total_tokens + additional_tokens;
let projected_calls = usage.total_calls + additional_calls;
let exceeded = |limit_name: &str, current: f64, limit: f64| -> Option<BudgetStatus> {
if current > limit {
let msg =
format!("{limit_name} budget exceeded for '{name}': {current:.6} > {limit:.6}");
Some(match budget.enforcement {
BudgetEnforcement::Hard => BudgetStatus::HardExceeded {
budget: name.to_owned(),
message: msg,
},
BudgetEnforcement::Soft => BudgetStatus::SoftExceeded {
budget: name.to_owned(),
message: msg,
},
})
} else {
None
}
};
if let Some(max) = budget.max_cost {
if let Some(status) = exceeded("Cost", projected_cost, max) {
return status;
}
}
if let Some(max) = budget.max_tokens {
if let Some(status) = exceeded("Token", projected_tokens as f64, max as f64) {
return status;
}
}
if let Some(max) = budget.max_calls {
if let Some(status) = exceeded("Call", projected_calls as f64, max as f64) {
return status;
}
}
BudgetStatus::Ok
}
#[must_use]
pub fn check_all(
&self,
test_name: &str,
test_usage: &UsageSnapshot,
global_usage: &UsageSnapshot,
test_budget_override: Option<&BudgetDef>,
) -> BudgetStatus {
let per_test = self.check_per_test(test_name, test_usage, test_budget_override);
if matches!(per_test, BudgetStatus::HardExceeded { .. }) {
return per_test;
}
let global = self.check_global(global_usage);
if matches!(global, BudgetStatus::HardExceeded { .. }) {
return global;
}
if per_test != BudgetStatus::Ok {
return per_test;
}
global
}
}
#[cfg(test)]
mod tests {
use crate::budgets::{BudgetStatus, BudgetTracker};
use crate::costs::UsageSnapshot;
use crate::scenario::{BudgetDef, BudgetEnforcement, BudgetsConfig};
#[test]
fn test_no_budgets_always_ok() {
let config = BudgetsConfig::default();
let tracker = BudgetTracker::from_config(&config);
let usage = UsageSnapshot::default();
assert_eq!(tracker.check_global(&usage), BudgetStatus::Ok);
assert_eq!(
tracker.check_per_test("test", &usage, None),
BudgetStatus::Ok
);
}
#[test]
fn test_global_cost_hard_limit() {
let config = BudgetsConfig {
global: Some(BudgetDef {
max_cost: Some(5.0),
max_tokens: None,
max_calls: None,
enforcement: Some(BudgetEnforcement::Hard),
}),
per_test_default: None,
};
let tracker = BudgetTracker::from_config(&config);
let under = UsageSnapshot {
total_cost: 3.0,
..UsageSnapshot::default()
};
assert_eq!(tracker.check_global(&under), BudgetStatus::Ok);
let over = UsageSnapshot {
total_cost: 6.0,
..UsageSnapshot::default()
};
assert!(matches!(
tracker.check_global(&over),
BudgetStatus::HardExceeded { .. }
));
}
#[test]
fn test_global_cost_soft_limit() {
let config = BudgetsConfig {
global: Some(BudgetDef {
max_cost: Some(5.0),
max_tokens: None,
max_calls: None,
enforcement: Some(BudgetEnforcement::Soft),
}),
per_test_default: None,
};
let tracker = BudgetTracker::from_config(&config);
let over = UsageSnapshot {
total_cost: 6.0,
..UsageSnapshot::default()
};
assert!(matches!(
tracker.check_global(&over),
BudgetStatus::SoftExceeded { .. }
));
}
#[test]
fn test_per_test_token_limit() {
let config = BudgetsConfig {
global: None,
per_test_default: Some(BudgetDef {
max_cost: None,
max_tokens: Some(10000),
max_calls: None,
enforcement: Some(BudgetEnforcement::Hard),
}),
};
let tracker = BudgetTracker::from_config(&config);
let over = UsageSnapshot {
total_tokens: 15000,
..UsageSnapshot::default()
};
assert!(matches!(
tracker.check_per_test("test", &over, None),
BudgetStatus::HardExceeded { .. }
));
}
#[test]
fn test_check_all_global_priority() {
let config = BudgetsConfig {
global: Some(BudgetDef {
max_cost: Some(5.0),
max_tokens: None,
max_calls: None,
enforcement: Some(BudgetEnforcement::Hard),
}),
per_test_default: Some(BudgetDef {
max_cost: Some(10.0),
max_tokens: None,
max_calls: None,
enforcement: Some(BudgetEnforcement::Hard),
}),
};
let tracker = BudgetTracker::from_config(&config);
let global = UsageSnapshot {
total_cost: 6.0,
..UsageSnapshot::default()
};
let test = UsageSnapshot::default();
assert!(matches!(
tracker.check_all("test", &test, &global, None),
BudgetStatus::HardExceeded { .. }
));
}
#[test]
fn test_per_test_call_limit_hard() {
let config = BudgetsConfig {
global: None,
per_test_default: Some(BudgetDef {
max_cost: None,
max_tokens: None,
max_calls: Some(5),
enforcement: Some(BudgetEnforcement::Hard),
}),
};
let tracker = BudgetTracker::from_config(&config);
let ok_usage = UsageSnapshot {
total_calls: 3,
..UsageSnapshot::default()
};
assert_eq!(
tracker.check_per_test("test", &ok_usage, None),
BudgetStatus::Ok
);
let exceeded = UsageSnapshot {
total_calls: 10,
..UsageSnapshot::default()
};
assert!(matches!(
tracker.check_per_test("test", &exceeded, None),
BudgetStatus::HardExceeded { .. }
));
}
#[test]
fn test_all_budget_types_at_once() {
let config = BudgetsConfig {
global: None,
per_test_default: Some(BudgetDef {
max_cost: Some(1.0),
max_tokens: Some(1000),
max_calls: Some(10),
enforcement: Some(BudgetEnforcement::Hard),
}),
};
let tracker = BudgetTracker::from_config(&config);
let fine = UsageSnapshot {
total_cost: 0.5,
total_tokens: 500,
total_calls: 5,
..UsageSnapshot::default()
};
assert_eq!(
tracker.check_per_test("test", &fine, None),
BudgetStatus::Ok
);
}
}