1use serde::{Deserialize, Serialize};
4
5#[derive(Debug, Clone, Serialize, Deserialize)]
7pub struct Budget {
8 pub max_input_tokens: Option<u64>,
10
11 pub max_output_tokens: Option<u64>,
13
14 pub max_total_tokens: Option<u64>,
16
17 pub max_tool_calls: Option<u32>,
19
20 pub max_wall_time_ms: Option<u64>,
22
23 pub max_cost_cents: Option<u64>,
25}
26
27impl Default for Budget {
28 fn default() -> Self {
29 Self {
30 max_input_tokens: Some(100_000),
31 max_output_tokens: Some(50_000),
32 max_total_tokens: Some(150_000),
33 max_tool_calls: Some(50),
34 max_wall_time_ms: Some(5 * 60 * 1000), max_cost_cents: Some(500), }
37 }
38}
39
40impl Budget {
41 pub fn has_cost_headroom(&self, usage: &BudgetUsage, additional_cents: u64) -> bool {
48 match self.max_cost_cents {
49 Some(cap) => usage.cost_cents.saturating_add(additional_cents) <= cap,
50 None => true,
51 }
52 }
53}
54
55#[derive(Debug, Clone, Default, Serialize, Deserialize)]
57pub struct BudgetUsage {
58 pub input_tokens: u64,
59 pub output_tokens: u64,
60 pub tool_calls: u32,
61 pub wall_time_ms: u64,
62 pub cost_cents: u64,
63}
64
65impl BudgetUsage {
66 pub fn total_tokens(&self) -> u64 {
67 self.input_tokens + self.output_tokens
68 }
69
70 pub fn check_against(&self, budget: &Budget) -> Option<BudgetExceeded> {
72 if let Some(limit) = budget.max_input_tokens {
73 if self.input_tokens > limit {
74 return Some(BudgetExceeded::InputTokens {
75 used: self.input_tokens,
76 limit,
77 });
78 }
79 }
80
81 if let Some(limit) = budget.max_output_tokens {
82 if self.output_tokens > limit {
83 return Some(BudgetExceeded::OutputTokens {
84 used: self.output_tokens,
85 limit,
86 });
87 }
88 }
89
90 if let Some(limit) = budget.max_total_tokens {
91 if self.total_tokens() > limit {
92 return Some(BudgetExceeded::TotalTokens {
93 used: self.total_tokens(),
94 limit,
95 });
96 }
97 }
98
99 if let Some(limit) = budget.max_tool_calls {
100 if self.tool_calls > limit {
101 return Some(BudgetExceeded::ToolCalls {
102 used: self.tool_calls,
103 limit,
104 });
105 }
106 }
107
108 if let Some(limit) = budget.max_wall_time_ms {
109 if self.wall_time_ms > limit {
110 return Some(BudgetExceeded::WallTime {
111 used_ms: self.wall_time_ms,
112 limit_ms: limit,
113 });
114 }
115 }
116
117 if let Some(limit) = budget.max_cost_cents {
118 if self.cost_cents > limit {
119 return Some(BudgetExceeded::Cost {
120 used_cents: self.cost_cents,
121 limit_cents: limit,
122 });
123 }
124 }
125
126 None
127 }
128}
129
130#[derive(Debug, Clone, Serialize, Deserialize)]
132#[serde(tag = "type", rename_all = "snake_case")]
133pub enum BudgetExceeded {
134 InputTokens { used: u64, limit: u64 },
135 OutputTokens { used: u64, limit: u64 },
136 TotalTokens { used: u64, limit: u64 },
137 ToolCalls { used: u32, limit: u32 },
138 WallTime { used_ms: u64, limit_ms: u64 },
139 Cost { used_cents: u64, limit_cents: u64 },
140}
141
142impl std::fmt::Display for BudgetExceeded {
143 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
144 match self {
145 BudgetExceeded::InputTokens { used, limit } => {
146 write!(f, "input tokens exceeded: {used}/{limit}")
147 }
148 BudgetExceeded::OutputTokens { used, limit } => {
149 write!(f, "output tokens exceeded: {used}/{limit}")
150 }
151 BudgetExceeded::TotalTokens { used, limit } => {
152 write!(f, "total tokens exceeded: {used}/{limit}")
153 }
154 BudgetExceeded::ToolCalls { used, limit } => {
155 write!(f, "tool calls exceeded: {used}/{limit}")
156 }
157 BudgetExceeded::WallTime { used_ms, limit_ms } => {
158 write!(f, "wall time exceeded: {used_ms}ms/{limit_ms}ms")
159 }
160 BudgetExceeded::Cost {
161 used_cents,
162 limit_cents,
163 } => {
164 write!(
165 f,
166 "cost exceeded: ${:.2}/${:.2}",
167 *used_cents as f64 / 100.0,
168 *limit_cents as f64 / 100.0
169 )
170 }
171 }
172 }
173}
174
175#[cfg(test)]
176mod tests {
177 use super::*;
178
179 #[test]
180 fn cost_headroom_respects_cap() {
181 let budget = Budget {
182 max_cost_cents: Some(100),
183 ..Budget::default()
184 };
185 let usage = BudgetUsage {
186 cost_cents: 80,
187 ..BudgetUsage::default()
188 };
189 assert!(budget.has_cost_headroom(&usage, 20));
191 assert!(!budget.has_cost_headroom(&usage, 21));
193 }
194
195 #[test]
196 fn cost_headroom_true_when_no_cap() {
197 let budget = Budget {
198 max_cost_cents: None,
199 ..Budget::default()
200 };
201 let usage = BudgetUsage {
202 cost_cents: 10_000,
203 ..BudgetUsage::default()
204 };
205 assert!(budget.has_cost_headroom(&usage, u64::MAX));
206 }
207
208 #[test]
209 fn cost_headroom_saturates_on_overflow() {
210 let budget = Budget {
211 max_cost_cents: Some(u64::MAX),
212 ..Budget::default()
213 };
214 let usage = BudgetUsage {
215 cost_cents: u64::MAX,
216 ..BudgetUsage::default()
217 };
218 assert!(budget.has_cost_headroom(&usage, 5));
220 }
221}