Skip to main content

llm_browser_testkit/
budgets.rs

1//! Budget tracking and enforcement.
2
3use crate::costs::UsageSnapshot;
4use crate::scenario::{BudgetDef, BudgetEnforcement, BudgetsConfig};
5
6/// Result of a budget check before or after a call.
7#[derive(Debug, Clone, PartialEq, Eq)]
8pub enum BudgetStatus {
9    /// Budget is within limits.
10    Ok,
11    /// Budget exceeded, enforcement is soft — log but continue.
12    SoftExceeded {
13        /// Which budget was exceeded.
14        budget: String,
15        /// Human-readable message.
16        message: String,
17    },
18    /// Budget exceeded, enforcement is hard — abort execution.
19    HardExceeded {
20        /// Which budget was exceeded.
21        budget: String,
22        /// Human-readable message.
23        message: String,
24    },
25}
26
27/// Tracks overall and per-test budgets across a scenario run.
28#[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    /// Creates a new budget tracker from the scenario's budget config.
55    #[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    /// Checks global budgets against the global usage snapshot.
67    #[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    /// Checks per-test budgets against a test's usage snapshot.
76    ///
77    /// `override_budget` allows the test to specify tighter (or looser)
78    /// limits.
79    #[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    /// Estimates the cost of a pending LLM call to check if it would exceed
97    /// the budget.
98    ///
99    /// Uses `max_tokens` from the request as a worst-case estimate for
100    /// output tokens, plus estimated input tokens.
101    #[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    /// Checks the budget for a flat-cost call (MCP, agent).
127    #[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    /// Checks both per-test and global budgets. Returns the most severe
193    /// violation.
194    #[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}