1use crate::costs::UsageSnapshot;
4use crate::scenario::{BudgetDef, BudgetEnforcement, BudgetsConfig};
5
6#[derive(Debug, Clone, PartialEq, Eq)]
8pub enum BudgetStatus {
9 Ok,
11 SoftExceeded {
13 budget: String,
15 message: String,
17 },
18 HardExceeded {
20 budget: String,
22 message: String,
24 },
25}
26
27#[derive(Debug, Clone)]
29pub struct BudgetTracker {
30 global: Option<ResolvedBudget>,
31 per_test_default: Option<ResolvedBudget>,
32}
33
34#[derive(Debug, Clone)]
35pub(crate) struct ResolvedBudget {
36 max_cost: Option<f64>,
37 max_tokens: Option<u64>,
38 max_calls: Option<u64>,
39 enforcement: BudgetEnforcement,
40}
41
42impl ResolvedBudget {
43 fn from_def(def: &BudgetDef) -> Self {
44 Self {
45 max_cost: def.max_cost,
46 max_tokens: def.max_tokens,
47 max_calls: def.max_calls,
48 enforcement: def.enforcement.clone().unwrap_or(BudgetEnforcement::Hard),
49 }
50 }
51}
52
53impl BudgetTracker {
54 #[must_use]
56 pub fn from_config(budgets: &BudgetsConfig) -> Self {
57 Self {
58 global: budgets.global.as_ref().map(ResolvedBudget::from_def),
59 per_test_default: budgets
60 .per_test_default
61 .as_ref()
62 .map(ResolvedBudget::from_def),
63 }
64 }
65
66 #[must_use]
68 pub fn check_global(&self, usage: &UsageSnapshot) -> BudgetStatus {
69 let Some(global) = &self.global else {
70 return BudgetStatus::Ok;
71 };
72 Self::check_budget("global", global, usage)
73 }
74
75 #[must_use]
80 pub fn check_per_test(
81 &self,
82 test_name: &str,
83 usage: &UsageSnapshot,
84 override_budget: Option<&BudgetDef>,
85 ) -> BudgetStatus {
86 let budget = match override_budget {
87 Some(def) => ResolvedBudget::from_def(def),
88 None => match &self.per_test_default {
89 Some(def) => def.clone(),
90 None => return BudgetStatus::Ok,
91 },
92 };
93 Self::check_budget(test_name, &budget, usage)
94 }
95
96 #[must_use]
102 #[allow(dead_code)]
103 #[allow(clippy::cast_precision_loss, clippy::suboptimal_flops)]
104 pub(crate) fn check_pre_flight_llm(
105 budget: &ResolvedBudget,
106 usage: &UsageSnapshot,
107 estimated_input_tokens: u64,
108 max_tokens: u64,
109 input_price_per_1m: f64,
110 output_price_per_1m: f64,
111 ) -> BudgetStatus {
112 let estimated_cost = (estimated_input_tokens as f64 / 1_000_000.0) * input_price_per_1m
113 + (max_tokens as f64 / 1_000_000.0) * output_price_per_1m;
114 let estimated_total_tokens = usage.total_tokens + estimated_input_tokens + max_tokens;
115
116 Self::check_limits(
117 budget,
118 "pre-flight",
119 usage,
120 estimated_cost,
121 estimated_total_tokens,
122 1,
123 )
124 }
125
126 #[must_use]
128 #[allow(dead_code)]
129 pub(crate) fn check_pre_flight_flat(
130 budget: &ResolvedBudget,
131 usage: &UsageSnapshot,
132 per_call_price: f64,
133 ) -> BudgetStatus {
134 Self::check_limits(budget, "pre-flight", usage, per_call_price, 0, 1)
135 }
136
137 fn check_budget(name: &str, budget: &ResolvedBudget, usage: &UsageSnapshot) -> BudgetStatus {
138 Self::check_limits(budget, name, usage, 0.0, 0, 0)
139 }
140
141 #[allow(clippy::cast_precision_loss)]
142 fn check_limits(
143 budget: &ResolvedBudget,
144 name: &str,
145 usage: &UsageSnapshot,
146 additional_cost: f64,
147 additional_tokens: u64,
148 additional_calls: u64,
149 ) -> BudgetStatus {
150 let projected_cost = usage.total_cost + additional_cost;
151 let projected_tokens = usage.total_tokens + additional_tokens;
152 let projected_calls = usage.total_calls + additional_calls;
153
154 let exceeded = |limit_name: &str, current: f64, limit: f64| -> Option<BudgetStatus> {
155 if current > limit {
156 let msg =
157 format!("{limit_name} budget exceeded for '{name}': {current:.6} > {limit:.6}");
158 Some(match budget.enforcement {
159 BudgetEnforcement::Hard => BudgetStatus::HardExceeded {
160 budget: name.to_owned(),
161 message: msg,
162 },
163 BudgetEnforcement::Soft => BudgetStatus::SoftExceeded {
164 budget: name.to_owned(),
165 message: msg,
166 },
167 })
168 } else {
169 None
170 }
171 };
172
173 if let Some(max) = budget.max_cost {
174 if let Some(status) = exceeded("Cost", projected_cost, max) {
175 return status;
176 }
177 }
178 if let Some(max) = budget.max_tokens {
179 if let Some(status) = exceeded("Token", projected_tokens as f64, max as f64) {
180 return status;
181 }
182 }
183 if let Some(max) = budget.max_calls {
184 if let Some(status) = exceeded("Call", projected_calls as f64, max as f64) {
185 return status;
186 }
187 }
188
189 BudgetStatus::Ok
190 }
191
192 #[must_use]
195 pub fn check_all(
196 &self,
197 test_name: &str,
198 test_usage: &UsageSnapshot,
199 global_usage: &UsageSnapshot,
200 test_budget_override: Option<&BudgetDef>,
201 ) -> BudgetStatus {
202 let per_test = self.check_per_test(test_name, test_usage, test_budget_override);
203 if matches!(per_test, BudgetStatus::HardExceeded { .. }) {
204 return per_test;
205 }
206 let global = self.check_global(global_usage);
207 if matches!(global, BudgetStatus::HardExceeded { .. }) {
208 return global;
209 }
210 if per_test != BudgetStatus::Ok {
211 return per_test;
212 }
213 global
214 }
215}
216
217#[cfg(test)]
218mod tests {
219 use crate::budgets::{BudgetStatus, BudgetTracker};
220 use crate::costs::UsageSnapshot;
221 use crate::scenario::{BudgetDef, BudgetEnforcement, BudgetsConfig};
222
223 #[test]
224 fn test_no_budgets_always_ok() {
225 let config = BudgetsConfig::default();
226 let tracker = BudgetTracker::from_config(&config);
227 let usage = UsageSnapshot::default();
228 assert_eq!(tracker.check_global(&usage), BudgetStatus::Ok);
229 assert_eq!(
230 tracker.check_per_test("test", &usage, None),
231 BudgetStatus::Ok
232 );
233 }
234
235 #[test]
236 fn test_global_cost_hard_limit() {
237 let config = BudgetsConfig {
238 global: Some(BudgetDef {
239 max_cost: Some(5.0),
240 max_tokens: None,
241 max_calls: None,
242 enforcement: Some(BudgetEnforcement::Hard),
243 }),
244 per_test_default: None,
245 };
246 let tracker = BudgetTracker::from_config(&config);
247 let under = UsageSnapshot {
248 total_cost: 3.0,
249 ..UsageSnapshot::default()
250 };
251 assert_eq!(tracker.check_global(&under), BudgetStatus::Ok);
252 let over = UsageSnapshot {
253 total_cost: 6.0,
254 ..UsageSnapshot::default()
255 };
256 assert!(matches!(
257 tracker.check_global(&over),
258 BudgetStatus::HardExceeded { .. }
259 ));
260 }
261
262 #[test]
263 fn test_global_cost_soft_limit() {
264 let config = BudgetsConfig {
265 global: Some(BudgetDef {
266 max_cost: Some(5.0),
267 max_tokens: None,
268 max_calls: None,
269 enforcement: Some(BudgetEnforcement::Soft),
270 }),
271 per_test_default: None,
272 };
273 let tracker = BudgetTracker::from_config(&config);
274 let over = UsageSnapshot {
275 total_cost: 6.0,
276 ..UsageSnapshot::default()
277 };
278 assert!(matches!(
279 tracker.check_global(&over),
280 BudgetStatus::SoftExceeded { .. }
281 ));
282 }
283
284 #[test]
285 fn test_per_test_token_limit() {
286 let config = BudgetsConfig {
287 global: None,
288 per_test_default: Some(BudgetDef {
289 max_cost: None,
290 max_tokens: Some(10000),
291 max_calls: None,
292 enforcement: Some(BudgetEnforcement::Hard),
293 }),
294 };
295 let tracker = BudgetTracker::from_config(&config);
296 let over = UsageSnapshot {
297 total_tokens: 15000,
298 ..UsageSnapshot::default()
299 };
300 assert!(matches!(
301 tracker.check_per_test("test", &over, None),
302 BudgetStatus::HardExceeded { .. }
303 ));
304 }
305
306 #[test]
307 fn test_check_all_global_priority() {
308 let config = BudgetsConfig {
309 global: Some(BudgetDef {
310 max_cost: Some(5.0),
311 max_tokens: None,
312 max_calls: None,
313 enforcement: Some(BudgetEnforcement::Hard),
314 }),
315 per_test_default: Some(BudgetDef {
316 max_cost: Some(10.0),
317 max_tokens: None,
318 max_calls: None,
319 enforcement: Some(BudgetEnforcement::Hard),
320 }),
321 };
322 let tracker = BudgetTracker::from_config(&config);
323 let global = UsageSnapshot {
324 total_cost: 6.0,
325 ..UsageSnapshot::default()
326 };
327 let test = UsageSnapshot::default();
328 assert!(matches!(
329 tracker.check_all("test", &test, &global, None),
330 BudgetStatus::HardExceeded { .. }
331 ));
332 }
333
334 #[test]
335 fn test_per_test_call_limit_hard() {
336 let config = BudgetsConfig {
337 global: None,
338 per_test_default: Some(BudgetDef {
339 max_cost: None,
340 max_tokens: None,
341 max_calls: Some(5),
342 enforcement: Some(BudgetEnforcement::Hard),
343 }),
344 };
345 let tracker = BudgetTracker::from_config(&config);
346 let ok_usage = UsageSnapshot {
347 total_calls: 3,
348 ..UsageSnapshot::default()
349 };
350 assert_eq!(
351 tracker.check_per_test("test", &ok_usage, None),
352 BudgetStatus::Ok
353 );
354 let exceeded = UsageSnapshot {
355 total_calls: 10,
356 ..UsageSnapshot::default()
357 };
358 assert!(matches!(
359 tracker.check_per_test("test", &exceeded, None),
360 BudgetStatus::HardExceeded { .. }
361 ));
362 }
363
364 #[test]
365 fn test_all_budget_types_at_once() {
366 let config = BudgetsConfig {
367 global: None,
368 per_test_default: Some(BudgetDef {
369 max_cost: Some(1.0),
370 max_tokens: Some(1000),
371 max_calls: Some(10),
372 enforcement: Some(BudgetEnforcement::Hard),
373 }),
374 };
375 let tracker = BudgetTracker::from_config(&config);
376 let fine = UsageSnapshot {
377 total_cost: 0.5,
378 total_tokens: 500,
379 total_calls: 5,
380 ..UsageSnapshot::default()
381 };
382 assert_eq!(
383 tracker.check_per_test("test", &fine, None),
384 BudgetStatus::Ok
385 );
386 }
387}