1use std::collections::BTreeMap;
2use std::sync::{Arc, Mutex, MutexGuard};
3use std::time::Duration;
4
5use serde::{Deserialize, Serialize};
6use thiserror::Error;
7
8#[derive(Clone, Copy, Debug, Default, Deserialize, Eq, PartialEq, Serialize)]
10pub struct Budget {
11 pub tokens: Option<u64>,
13 pub cost_microusd: Option<u64>,
15 pub duration: Option<Duration>,
17 pub turns: Option<u64>,
19 pub tool_calls: Option<u64>,
21 pub delegations: Option<u64>,
23}
24
25impl Budget {
26 #[must_use]
28 pub fn tighten(self, other: Self) -> Self {
29 Self {
30 tokens: minimum(self.tokens, other.tokens),
31 cost_microusd: minimum(self.cost_microusd, other.cost_microusd),
32 duration: minimum(self.duration, other.duration),
33 turns: minimum(self.turns, other.turns),
34 tool_calls: minimum(self.tool_calls, other.tool_calls),
35 delegations: minimum(self.delegations, other.delegations),
36 }
37 }
38}
39
40fn minimum<T: Ord>(left: Option<T>, right: Option<T>) -> Option<T> {
41 match (left, right) {
42 (Some(left), Some(right)) => Some(left.min(right)),
43 (Some(value), None) | (None, Some(value)) => Some(value),
44 (None, None) => None,
45 }
46}
47
48#[derive(Clone, Copy, Debug, Default, Deserialize, Eq, PartialEq, Serialize)]
50pub struct Usage {
51 pub tokens: u64,
53 pub cost_microusd: u64,
55 pub duration_micros: u64,
57 pub turns: u64,
59 pub tool_calls: u64,
61 pub delegations: u64,
63}
64
65impl Usage {
66 fn checked_add(self, delta: Self) -> Option<Self> {
67 Some(Self {
68 tokens: self.tokens.checked_add(delta.tokens)?,
69 cost_microusd: self.cost_microusd.checked_add(delta.cost_microusd)?,
70 duration_micros: self.duration_micros.checked_add(delta.duration_micros)?,
71 turns: self.turns.checked_add(delta.turns)?,
72 tool_calls: self.tool_calls.checked_add(delta.tool_calls)?,
73 delegations: self.delegations.checked_add(delta.delegations)?,
74 })
75 }
76
77 fn checked_sub(self, delta: Self) -> Option<Self> {
78 Some(Self {
79 tokens: self.tokens.checked_sub(delta.tokens)?,
80 cost_microusd: self.cost_microusd.checked_sub(delta.cost_microusd)?,
81 duration_micros: self.duration_micros.checked_sub(delta.duration_micros)?,
82 turns: self.turns.checked_sub(delta.turns)?,
83 tool_calls: self.tool_calls.checked_sub(delta.tool_calls)?,
84 delegations: self.delegations.checked_sub(delta.delegations)?,
85 })
86 }
87}
88
89#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)]
91#[non_exhaustive]
92pub enum BudgetResource {
93 Tokens,
95 Cost,
97 Duration,
99 Turns,
101 ToolCalls,
103 Delegations,
105 CounterOverflow,
107}
108
109#[derive(Clone, Debug, Error, Eq, PartialEq)]
111#[error("{resource:?} budget exceeded: limit={limit}, attempted={attempted}")]
112pub struct BudgetExceeded {
113 pub resource: BudgetResource,
115 pub limit: u128,
117 pub attempted: u128,
119}
120
121#[derive(Clone, Debug)]
123pub struct BudgetTracker {
124 inner: Arc<BudgetState>,
125 reservation: Option<Arc<ReservationLease>>,
126}
127
128#[derive(Debug)]
129struct BudgetState {
130 limit: Budget,
131 ledger: Mutex<BudgetLedger>,
132}
133
134#[derive(Debug, Default)]
135struct BudgetLedger {
136 usage: Usage,
137 reserved: Usage,
138 reservations: BTreeMap<u64, Usage>,
139 next_reservation_id: u64,
140}
141
142#[derive(Debug)]
143struct ReservationLease {
144 state: Arc<BudgetState>,
145 id: u64,
146}
147
148impl Drop for ReservationLease {
149 fn drop(&mut self) {
150 let mut ledger = lock_ledger(&self.state);
151 if let Some(remaining) = ledger.reservations.remove(&self.id) {
152 ledger.reserved = ledger
153 .reserved
154 .checked_sub(remaining)
155 .expect("reservation aggregate contains every live reservation");
156 }
157 }
158}
159
160#[derive(Clone, Debug)]
166pub struct BudgetReservation {
167 tracker: BudgetTracker,
168 reserved: Usage,
169}
170
171impl BudgetReservation {
172 pub fn tracker(&self) -> BudgetTracker {
174 self.tracker.clone()
175 }
176
177 pub const fn reserved(&self) -> Usage {
179 self.reserved
180 }
181
182 pub fn remaining(&self) -> Usage {
184 self.tracker.reservation_remaining().unwrap_or_default()
185 }
186
187 pub fn forfeit_remaining(&self) -> Result<Usage, BudgetExceeded> {
196 self.tracker.forfeit_reservation()
197 }
198
199 pub(crate) fn belongs_to(&self, tracker: &BudgetTracker) -> bool {
200 Arc::ptr_eq(&self.tracker.inner, &tracker.inner)
201 }
202}
203
204impl BudgetTracker {
205 pub fn new(limit: Budget) -> Self {
207 Self {
208 inner: Arc::new(BudgetState {
209 limit,
210 ledger: Mutex::new(BudgetLedger::default()),
211 }),
212 reservation: None,
213 }
214 }
215
216 pub fn restore(limit: Budget, usage: Usage) -> Result<Self, BudgetExceeded> {
222 validate(limit, usage)?;
223 Ok(Self {
224 inner: Arc::new(BudgetState {
225 limit,
226 ledger: Mutex::new(BudgetLedger {
227 usage,
228 ..BudgetLedger::default()
229 }),
230 }),
231 reservation: None,
232 })
233 }
234
235 pub fn limit(&self) -> Budget {
237 self.inner.limit
238 }
239
240 pub fn usage(&self) -> Usage {
242 self.lock_ledger().usage
243 }
244
245 pub fn try_consume(&self, delta: Usage) -> Result<Usage, BudgetExceeded> {
252 let mut ledger = self.lock_ledger();
253 if let Some(lease) = &self.reservation {
254 let remaining = ledger
255 .reservations
256 .get(&lease.id)
257 .copied()
258 .unwrap_or_default();
259 validate_reservation(remaining, delta)?;
260 let attempted = ledger.usage.checked_add(delta).ok_or(BudgetExceeded {
261 resource: BudgetResource::CounterOverflow,
262 limit: u128::from(u64::MAX),
263 attempted: u128::from(u64::MAX) + 1,
264 })?;
265 let remaining = remaining.checked_sub(delta).ok_or_else(counter_overflow)?;
266 ledger.reservations.insert(lease.id, remaining);
267 ledger.reserved = ledger
268 .reserved
269 .checked_sub(delta)
270 .ok_or_else(counter_overflow)?;
271 ledger.usage = attempted;
272 return Ok(attempted);
273 }
274
275 let committed_and_reserved =
276 ledger
277 .usage
278 .checked_add(ledger.reserved)
279 .ok_or(BudgetExceeded {
280 resource: BudgetResource::CounterOverflow,
281 limit: u128::from(u64::MAX),
282 attempted: u128::from(u64::MAX) + 1,
283 })?;
284 let attempted = committed_and_reserved
285 .checked_add(delta)
286 .ok_or(BudgetExceeded {
287 resource: BudgetResource::CounterOverflow,
288 limit: u128::from(u64::MAX),
289 attempted: u128::from(u64::MAX) + 1,
290 })?;
291
292 validate(self.inner.limit, attempted)?;
293 ledger.usage = ledger
294 .usage
295 .checked_add(delta)
296 .ok_or_else(counter_overflow)?;
297 Ok(ledger.usage)
298 }
299
300 pub fn try_reserve(&self, amount: Usage) -> Result<BudgetReservation, BudgetExceeded> {
307 self.try_reserve_batch([amount])?
308 .pop()
309 .ok_or_else(counter_overflow)
310 }
311
312 pub fn try_reserve_batch(
321 &self,
322 amounts: impl IntoIterator<Item = Usage>,
323 ) -> Result<Vec<BudgetReservation>, BudgetExceeded> {
324 let amounts = amounts.into_iter().collect::<Vec<_>>();
325 let mut requested = Usage::default();
326 for amount in &amounts {
327 requested = requested.checked_add(*amount).ok_or(BudgetExceeded {
328 resource: BudgetResource::CounterOverflow,
329 limit: u128::from(u64::MAX),
330 attempted: u128::from(u64::MAX) + 1,
331 })?;
332 }
333 let count = u64::try_from(amounts.len()).map_err(|_| counter_overflow())?;
334
335 let mut ledger = self.lock_ledger();
336 let start = ledger.next_reservation_id;
337 let next_reservation_id = start.checked_add(count).ok_or_else(counter_overflow)?;
338 if let Some(lease) = &self.reservation {
339 let remaining = ledger
340 .reservations
341 .get(&lease.id)
342 .copied()
343 .unwrap_or_default();
344 validate_reservation(remaining, requested)?;
345 let remaining = remaining
346 .checked_sub(requested)
347 .ok_or_else(counter_overflow)?;
348 ledger.reservations.insert(lease.id, remaining);
349 } else {
350 let attempted = ledger
351 .usage
352 .checked_add(ledger.reserved)
353 .and_then(|current| current.checked_add(requested))
354 .ok_or(BudgetExceeded {
355 resource: BudgetResource::CounterOverflow,
356 limit: u128::from(u64::MAX),
357 attempted: u128::from(u64::MAX) + 1,
358 })?;
359 validate(self.inner.limit, attempted)?;
360 ledger.reserved = ledger
361 .reserved
362 .checked_add(requested)
363 .ok_or_else(counter_overflow)?;
364 }
365
366 ledger.next_reservation_id = next_reservation_id;
367 let mut reservations = Vec::with_capacity(amounts.len());
368 for (id, amount) in (start..next_reservation_id).zip(amounts) {
369 ledger.reservations.insert(id, amount);
370 reservations.push(BudgetReservation {
371 tracker: Self {
372 inner: self.inner.clone(),
373 reservation: Some(Arc::new(ReservationLease {
374 state: self.inner.clone(),
375 id,
376 })),
377 },
378 reserved: amount,
379 });
380 }
381 Ok(reservations)
382 }
383
384 fn reservation_remaining(&self) -> Option<Usage> {
385 let lease = self.reservation.as_ref()?;
386 self.lock_ledger().reservations.get(&lease.id).copied()
387 }
388
389 fn forfeit_reservation(&self) -> Result<Usage, BudgetExceeded> {
390 let Some(lease) = &self.reservation else {
391 return Err(counter_overflow());
392 };
393 let mut ledger = self.lock_ledger();
394 let remaining = ledger
395 .reservations
396 .get(&lease.id)
397 .copied()
398 .unwrap_or_default();
399 let attempted = ledger
400 .usage
401 .checked_add(remaining)
402 .ok_or_else(counter_overflow)?;
403 ledger.reserved = ledger
404 .reserved
405 .checked_sub(remaining)
406 .ok_or_else(counter_overflow)?;
407 ledger.reservations.insert(lease.id, Usage::default());
408 ledger.usage = attempted;
409 Ok(attempted)
410 }
411
412 fn lock_ledger(&self) -> MutexGuard<'_, BudgetLedger> {
413 lock_ledger(&self.inner)
414 }
415}
416
417fn lock_ledger(state: &BudgetState) -> MutexGuard<'_, BudgetLedger> {
418 state
419 .ledger
420 .lock()
421 .unwrap_or_else(std::sync::PoisonError::into_inner)
422}
423
424fn validate_reservation(remaining: Usage, delta: Usage) -> Result<(), BudgetExceeded> {
425 check(BudgetResource::Tokens, Some(remaining.tokens), delta.tokens)?;
426 check(
427 BudgetResource::Cost,
428 Some(remaining.cost_microusd),
429 delta.cost_microusd,
430 )?;
431 check(
432 BudgetResource::Duration,
433 Some(remaining.duration_micros),
434 delta.duration_micros,
435 )?;
436 check(BudgetResource::Turns, Some(remaining.turns), delta.turns)?;
437 check(
438 BudgetResource::ToolCalls,
439 Some(remaining.tool_calls),
440 delta.tool_calls,
441 )?;
442 check(
443 BudgetResource::Delegations,
444 Some(remaining.delegations),
445 delta.delegations,
446 )
447}
448
449fn counter_overflow() -> BudgetExceeded {
450 BudgetExceeded {
451 resource: BudgetResource::CounterOverflow,
452 limit: u128::from(u64::MAX),
453 attempted: u128::from(u64::MAX) + 1,
454 }
455}
456
457fn validate(limit: Budget, attempted: Usage) -> Result<(), BudgetExceeded> {
458 check(BudgetResource::Tokens, limit.tokens, attempted.tokens)?;
459 check(
460 BudgetResource::Cost,
461 limit.cost_microusd,
462 attempted.cost_microusd,
463 )?;
464 check(
465 BudgetResource::Duration,
466 limit.duration.map(|duration| duration.as_micros()),
467 u128::from(attempted.duration_micros),
468 )?;
469 check(BudgetResource::Turns, limit.turns, attempted.turns)?;
470 check(
471 BudgetResource::ToolCalls,
472 limit.tool_calls,
473 attempted.tool_calls,
474 )?;
475 check(
476 BudgetResource::Delegations,
477 limit.delegations,
478 attempted.delegations,
479 )
480}
481
482fn check<T>(resource: BudgetResource, limit: Option<T>, attempted: T) -> Result<(), BudgetExceeded>
483where
484 T: Copy + Into<u128> + Ord,
485{
486 if let Some(limit) = limit
487 && attempted > limit
488 {
489 return Err(BudgetExceeded {
490 resource,
491 limit: limit.into(),
492 attempted: attempted.into(),
493 });
494 }
495 Ok(())
496}
497
498#[cfg(test)]
499mod tests {
500 use super::{Budget, BudgetResource, BudgetTracker, Usage};
501
502 #[test]
503 fn rejected_consumption_is_atomic() {
504 let tracker = BudgetTracker::new(Budget {
505 tokens: Some(10),
506 tool_calls: Some(1),
507 ..Budget::default()
508 });
509
510 tracker
511 .try_consume(Usage {
512 tokens: 6,
513 ..Usage::default()
514 })
515 .unwrap();
516
517 let error = tracker
518 .try_consume(Usage {
519 tokens: 5,
520 tool_calls: 1,
521 ..Usage::default()
522 })
523 .unwrap_err();
524
525 assert_eq!(error.resource, BudgetResource::Tokens);
526 assert_eq!(
527 tracker.usage(),
528 Usage {
529 tokens: 6,
530 ..Usage::default()
531 }
532 );
533 }
534
535 #[test]
536 fn tightening_keeps_stricter_limits() {
537 let first = Budget {
538 tokens: Some(100),
539 turns: None,
540 ..Budget::default()
541 };
542 let second = Budget {
543 tokens: Some(50),
544 turns: Some(3),
545 ..Budget::default()
546 };
547
548 let tightened = first.tighten(second);
549
550 assert_eq!(tightened.tokens, Some(50));
551 assert_eq!(tightened.turns, Some(3));
552 }
553
554 #[test]
555 fn restored_usage_is_validated_against_limits() {
556 let tracker = BudgetTracker::restore(
557 Budget {
558 turns: Some(3),
559 ..Budget::default()
560 },
561 Usage {
562 turns: 2,
563 ..Usage::default()
564 },
565 )
566 .unwrap();
567
568 assert_eq!(tracker.usage().turns, 2);
569
570 let error = BudgetTracker::restore(
571 Budget {
572 turns: Some(1),
573 ..Budget::default()
574 },
575 Usage {
576 turns: 2,
577 ..Usage::default()
578 },
579 )
580 .unwrap_err();
581 assert_eq!(error.resource, BudgetResource::Turns);
582 }
583
584 #[test]
585 fn reservations_isolate_parallel_budget_shares() {
586 let tracker = BudgetTracker::new(Budget {
587 tokens: Some(10),
588 ..Budget::default()
589 });
590 let reservations = tracker
591 .try_reserve_batch([
592 Usage {
593 tokens: 4,
594 ..Usage::default()
595 },
596 Usage {
597 tokens: 6,
598 ..Usage::default()
599 },
600 ])
601 .unwrap();
602
603 let unreserved = tracker
604 .try_consume(Usage {
605 tokens: 1,
606 ..Usage::default()
607 })
608 .unwrap_err();
609 assert_eq!(unreserved.resource, BudgetResource::Tokens);
610
611 reservations[0]
612 .tracker()
613 .try_consume(Usage {
614 tokens: 4,
615 ..Usage::default()
616 })
617 .unwrap();
618 let branch_error = reservations[1]
619 .tracker()
620 .try_consume(Usage {
621 tokens: 7,
622 ..Usage::default()
623 })
624 .unwrap_err();
625
626 assert_eq!(branch_error.limit, 6);
627 assert_eq!(tracker.usage().tokens, 4);
628 }
629
630 #[test]
631 fn unused_reservation_is_released_on_last_scoped_drop() {
632 let tracker = BudgetTracker::new(Budget {
633 turns: Some(2),
634 ..Budget::default()
635 });
636 let reservation = tracker
637 .try_reserve(Usage {
638 turns: 2,
639 ..Usage::default()
640 })
641 .unwrap();
642 let scoped = reservation.tracker();
643 drop(reservation);
644
645 assert!(
646 tracker
647 .try_consume(Usage {
648 turns: 1,
649 ..Usage::default()
650 })
651 .is_err()
652 );
653 drop(scoped);
654
655 tracker
656 .try_consume(Usage {
657 turns: 2,
658 ..Usage::default()
659 })
660 .unwrap();
661 }
662
663 #[test]
664 fn forfeiting_a_reservation_commits_its_entire_remaining_share() {
665 let tracker = BudgetTracker::new(Budget {
666 turns: Some(3),
667 ..Budget::default()
668 });
669 let reservation = tracker
670 .try_reserve(Usage {
671 turns: 2,
672 ..Usage::default()
673 })
674 .unwrap();
675 reservation
676 .tracker()
677 .try_consume(Usage {
678 turns: 1,
679 ..Usage::default()
680 })
681 .unwrap();
682
683 let usage = reservation.forfeit_remaining().unwrap();
684
685 assert_eq!(usage.turns, 2);
686 assert_eq!(reservation.remaining().turns, 0);
687 tracker
688 .try_consume(Usage {
689 turns: 1,
690 ..Usage::default()
691 })
692 .unwrap();
693 assert_eq!(tracker.usage().turns, 3);
694 }
695
696 #[test]
697 fn failed_batch_reservation_changes_nothing() {
698 let tracker = BudgetTracker::new(Budget {
699 tool_calls: Some(2),
700 ..Budget::default()
701 });
702
703 let error = tracker
704 .try_reserve_batch([
705 Usage {
706 tool_calls: 1,
707 ..Usage::default()
708 },
709 Usage {
710 tool_calls: 2,
711 ..Usage::default()
712 },
713 ])
714 .unwrap_err();
715
716 assert_eq!(error.resource, BudgetResource::ToolCalls);
717 tracker
718 .try_consume(Usage {
719 tool_calls: 2,
720 ..Usage::default()
721 })
722 .unwrap();
723 }
724
725 #[test]
726 fn nested_reservation_cannot_exceed_parent_share() {
727 let tracker = BudgetTracker::new(Budget {
728 delegations: Some(5),
729 ..Budget::default()
730 });
731 let parent = tracker
732 .try_reserve(Usage {
733 delegations: 3,
734 ..Usage::default()
735 })
736 .unwrap();
737 let child = parent
738 .tracker()
739 .try_reserve(Usage {
740 delegations: 2,
741 ..Usage::default()
742 })
743 .unwrap();
744
745 let error = parent
746 .tracker()
747 .try_consume(Usage {
748 delegations: 2,
749 ..Usage::default()
750 })
751 .unwrap_err();
752 assert_eq!(error.limit, 1);
753
754 child
755 .tracker()
756 .try_consume(Usage {
757 delegations: 2,
758 ..Usage::default()
759 })
760 .unwrap();
761 assert_eq!(tracker.usage().delegations, 2);
762 }
763}