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 pub fn cost_remaining_cents(&self, usage: &BudgetUsage) -> Option<u64> {
61 self.max_cost_cents
62 .map(|cap| cap.saturating_sub(usage.cost_cents))
63 }
64}
65
66#[derive(Debug, Clone, Default, Serialize, Deserialize)]
68pub struct BudgetUsage {
69 pub input_tokens: u64,
70 pub output_tokens: u64,
71 pub tool_calls: u32,
72 pub wall_time_ms: u64,
73 pub cost_cents: u64,
74}
75
76impl BudgetUsage {
77 pub fn total_tokens(&self) -> u64 {
78 self.input_tokens + self.output_tokens
79 }
80
81 pub fn check_against(&self, budget: &Budget) -> Option<BudgetExceeded> {
83 if let Some(limit) = budget.max_input_tokens {
84 if self.input_tokens > limit {
85 return Some(BudgetExceeded::InputTokens {
86 used: self.input_tokens,
87 limit,
88 });
89 }
90 }
91
92 if let Some(limit) = budget.max_output_tokens {
93 if self.output_tokens > limit {
94 return Some(BudgetExceeded::OutputTokens {
95 used: self.output_tokens,
96 limit,
97 });
98 }
99 }
100
101 if let Some(limit) = budget.max_total_tokens {
102 if self.total_tokens() > limit {
103 return Some(BudgetExceeded::TotalTokens {
104 used: self.total_tokens(),
105 limit,
106 });
107 }
108 }
109
110 if let Some(limit) = budget.max_tool_calls {
111 if self.tool_calls > limit {
112 return Some(BudgetExceeded::ToolCalls {
113 used: self.tool_calls,
114 limit,
115 });
116 }
117 }
118
119 if let Some(limit) = budget.max_wall_time_ms {
120 if self.wall_time_ms > limit {
121 return Some(BudgetExceeded::WallTime {
122 used_ms: self.wall_time_ms,
123 limit_ms: limit,
124 });
125 }
126 }
127
128 if let Some(limit) = budget.max_cost_cents {
129 if self.cost_cents > limit {
130 return Some(BudgetExceeded::Cost {
131 used_cents: self.cost_cents,
132 limit_cents: limit,
133 });
134 }
135 }
136
137 None
138 }
139}
140
141#[derive(Debug, Clone, Serialize, Deserialize)]
143#[serde(tag = "type", rename_all = "snake_case")]
144pub enum BudgetExceeded {
145 InputTokens { used: u64, limit: u64 },
146 OutputTokens { used: u64, limit: u64 },
147 TotalTokens { used: u64, limit: u64 },
148 ToolCalls { used: u32, limit: u32 },
149 WallTime { used_ms: u64, limit_ms: u64 },
150 Cost { used_cents: u64, limit_cents: u64 },
151}
152
153impl std::fmt::Display for BudgetExceeded {
154 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
155 match self {
156 BudgetExceeded::InputTokens { used, limit } => {
157 write!(f, "input tokens exceeded: {used}/{limit}")
158 }
159 BudgetExceeded::OutputTokens { used, limit } => {
160 write!(f, "output tokens exceeded: {used}/{limit}")
161 }
162 BudgetExceeded::TotalTokens { used, limit } => {
163 write!(f, "total tokens exceeded: {used}/{limit}")
164 }
165 BudgetExceeded::ToolCalls { used, limit } => {
166 write!(f, "tool calls exceeded: {used}/{limit}")
167 }
168 BudgetExceeded::WallTime { used_ms, limit_ms } => {
169 write!(f, "wall time exceeded: {used_ms}ms/{limit_ms}ms")
170 }
171 BudgetExceeded::Cost {
172 used_cents,
173 limit_cents,
174 } => {
175 write!(
176 f,
177 "cost exceeded: ${:.2}/${:.2}",
178 *used_cents as f64 / 100.0,
179 *limit_cents as f64 / 100.0
180 )
181 }
182 }
183 }
184}
185
186#[cfg(test)]
187mod tests {
188 use super::*;
189
190 #[test]
191 fn cost_headroom_respects_cap() {
192 let budget = Budget {
193 max_cost_cents: Some(100),
194 ..Budget::default()
195 };
196 let usage = BudgetUsage {
197 cost_cents: 80,
198 ..BudgetUsage::default()
199 };
200 assert!(budget.has_cost_headroom(&usage, 20));
202 assert!(!budget.has_cost_headroom(&usage, 21));
204 }
205
206 #[test]
207 fn cost_headroom_true_when_no_cap() {
208 let budget = Budget {
209 max_cost_cents: None,
210 ..Budget::default()
211 };
212 let usage = BudgetUsage {
213 cost_cents: 10_000,
214 ..BudgetUsage::default()
215 };
216 assert!(budget.has_cost_headroom(&usage, u64::MAX));
217 }
218
219 #[test]
220 fn cost_headroom_saturates_on_overflow() {
221 let budget = Budget {
222 max_cost_cents: Some(u64::MAX),
223 ..Budget::default()
224 };
225 let usage = BudgetUsage {
226 cost_cents: u64::MAX,
227 ..BudgetUsage::default()
228 };
229 assert!(budget.has_cost_headroom(&usage, 5));
231 }
232
233 #[test]
234 fn cost_remaining_reports_headroom_and_saturates() {
235 let budget = Budget {
236 max_cost_cents: Some(500),
237 ..Budget::default()
238 };
239 let usage = BudgetUsage {
241 cost_cents: 200,
242 ..BudgetUsage::default()
243 };
244 assert_eq!(budget.cost_remaining_cents(&usage), Some(300));
245 let spent = BudgetUsage {
247 cost_cents: 600,
248 ..BudgetUsage::default()
249 };
250 assert_eq!(budget.cost_remaining_cents(&spent), Some(0));
251 let uncapped = Budget {
253 max_cost_cents: None,
254 ..Budget::default()
255 };
256 assert_eq!(uncapped.cost_remaining_cents(&usage), None);
257 }
258}