1use std::sync::{Arc, Mutex};
2
3use crate::metrics::SessionStatus;
4
5#[derive(Debug, Clone, Default)]
17pub struct Budget {
18 pub max_llm_calls: Option<u64>,
21 pub max_elapsed_ms: Option<u64>,
26 pub max_tokens: Option<u64>,
31}
32
33impl Budget {
34 pub fn from_ctx(ctx: &serde_json::Value) -> Option<Self> {
36 let obj = ctx.as_object()?.get("budget")?.as_object()?;
37 let max_llm_calls = obj.get("max_llm_calls").and_then(|v| v.as_u64());
38 let max_elapsed_ms = obj.get("max_elapsed_ms").and_then(|v| v.as_u64());
39 let max_tokens = obj.get("max_tokens").and_then(|v| v.as_u64());
40 if max_llm_calls.is_none() && max_elapsed_ms.is_none() && max_tokens.is_none() {
41 return None;
42 }
43 Some(Self {
44 max_llm_calls,
45 max_elapsed_ms,
46 max_tokens,
47 })
48 }
49
50 pub fn check(&self, llm_calls: u64, elapsed_ms: u64, total_tokens: u64) -> Result<(), String> {
53 if let Some(max) = self.max_llm_calls {
54 if llm_calls >= max {
55 return Err(format!(
56 "budget_exceeded: max_llm_calls ({max}) reached ({llm_calls} used)"
57 ));
58 }
59 }
60 if let Some(max_ms) = self.max_elapsed_ms {
61 if elapsed_ms >= max_ms {
62 return Err(format!(
63 "budget_exceeded: max_elapsed_ms ({max_ms}ms) reached ({elapsed_ms}ms elapsed)"
64 ));
65 }
66 }
67 if let Some(max) = self.max_tokens {
68 if total_tokens >= max {
69 return Err(format!(
70 "budget_exceeded: max_tokens ({max}) reached ({total_tokens} used)"
71 ));
72 }
73 }
74 Ok(())
75 }
76
77 pub fn remaining_json(
80 &self,
81 llm_calls: u64,
82 elapsed_ms: u64,
83 total_tokens: u64,
84 ) -> serde_json::Value {
85 serde_json::json!({
86 "llm_calls": self.max_llm_calls.map(|max| max.saturating_sub(llm_calls)),
87 "elapsed_ms": self.max_elapsed_ms.map(|max| max.saturating_sub(elapsed_ms)),
88 "tokens": self.max_tokens.map(|max| max.saturating_sub(total_tokens)),
89 })
90 }
91
92 pub fn to_json(&self) -> serde_json::Value {
94 let mut map = serde_json::Map::new();
95 if let Some(max) = self.max_llm_calls {
96 map.insert("max_llm_calls".into(), max.into());
97 }
98 if let Some(max) = self.max_elapsed_ms {
99 map.insert("max_elapsed_ms".into(), max.into());
100 }
101 if let Some(max) = self.max_tokens {
102 map.insert("max_tokens".into(), max.into());
103 }
104 serde_json::Value::Object(map)
105 }
106}
107
108#[derive(Clone)]
141pub struct BudgetHandle {
142 auto: Arc<Mutex<SessionStatus>>,
143}
144
145impl BudgetHandle {
146 pub(crate) fn new(auto: Arc<Mutex<SessionStatus>>) -> Self {
147 Self { auto }
148 }
149
150 pub fn check(&self) -> Result<(), String> {
154 let m = self
155 .auto
156 .lock()
157 .map_err(|_| "budget check: mutex poisoned".to_string())?;
158 m.check_budget()
159 }
160
161 pub fn remaining(&self) -> serde_json::Value {
166 let m = match self.auto.lock() {
167 Ok(m) => m,
168 Err(_) => return serde_json::Value::Null,
169 };
170 m.budget_remaining()
171 }
172}
173
174#[cfg(test)]
175mod tests {
176 use super::*;
177 use crate::{ExecutionMetrics, ExecutionObserver, LlmQuery, QueryId};
178
179 #[test]
180 fn budget_from_ctx_none_when_missing() {
181 let ctx = serde_json::json!({"task": "test"});
182 assert!(Budget::from_ctx(&ctx).is_none());
183 }
184
185 #[test]
186 fn budget_from_ctx_none_when_empty() {
187 let ctx = serde_json::json!({"budget": {}});
188 assert!(Budget::from_ctx(&ctx).is_none());
189 }
190
191 #[test]
192 fn budget_from_ctx_extracts_llm_calls() {
193 let ctx = serde_json::json!({"budget": {"max_llm_calls": 10}});
194 let budget = Budget::from_ctx(&ctx).expect("should parse");
195 assert_eq!(budget.max_llm_calls, Some(10));
196 assert_eq!(budget.max_elapsed_ms, None);
197 }
198
199 #[test]
200 fn budget_from_ctx_extracts_elapsed_ms() {
201 let ctx = serde_json::json!({"budget": {"max_elapsed_ms": 5000}});
202 let budget = Budget::from_ctx(&ctx).expect("should parse");
203 assert_eq!(budget.max_llm_calls, None);
204 assert_eq!(budget.max_elapsed_ms, Some(5000));
205 }
206
207 #[test]
208 fn budget_from_ctx_extracts_both() {
209 let ctx = serde_json::json!({"budget": {"max_llm_calls": 5, "max_elapsed_ms": 30000}});
210 let budget = Budget::from_ctx(&ctx).expect("should parse");
211 assert_eq!(budget.max_llm_calls, Some(5));
212 assert_eq!(budget.max_elapsed_ms, Some(30000));
213 }
214
215 #[test]
216 fn budget_check_passes_when_within_limits() {
217 let metrics = ExecutionMetrics::new();
218 metrics.set_budget(Budget {
219 max_llm_calls: Some(5),
220 max_elapsed_ms: None,
221 max_tokens: None,
222 });
223 let handle = metrics.budget_handle();
224 assert!(handle.check().is_ok());
225 }
226
227 #[test]
228 fn budget_check_fails_when_llm_calls_exceeded() {
229 let metrics = ExecutionMetrics::new();
230 metrics.set_budget(Budget {
231 max_llm_calls: Some(2),
232 max_elapsed_ms: None,
233 max_tokens: None,
234 });
235 let observer = metrics.create_observer();
236 let handle = metrics.budget_handle();
237
238 let q = vec![LlmQuery {
240 id: QueryId::single(),
241 prompt: "p".into(),
242 system: None,
243 max_tokens: 10,
244 grounded: false,
245 underspecified: false,
246 cache_breakpoint: None,
247 role: None,
248 }];
249 observer.on_paused(&q);
250 observer.on_response_fed(&QueryId::single(), "r", None);
251 observer.on_resumed();
252 observer.on_paused(&q);
253
254 let result = handle.check();
256 assert!(result.is_err());
257 assert!(result.unwrap_err().contains("budget_exceeded"));
258 }
259
260 #[test]
261 fn budget_check_fails_when_tokens_exceeded() {
262 let metrics = ExecutionMetrics::new();
263 metrics.set_budget(Budget {
264 max_llm_calls: None,
265 max_elapsed_ms: None,
266 max_tokens: Some(10),
267 });
268 let observer = metrics.create_observer();
269 let handle = metrics.budget_handle();
270
271 let q = vec![LlmQuery {
274 id: QueryId::single(),
275 prompt: "abcdefghijklmnopqrstuvwxyz0123456789abcd".into(),
276 system: None,
277 max_tokens: 100,
278 grounded: false,
279 underspecified: false,
280 cache_breakpoint: None,
281 role: None,
282 }];
283 observer.on_paused(&q);
284 observer.on_response_fed(&QueryId::single(), "r", None);
285 observer.on_resumed();
286
287 let result = handle.check();
288 assert!(result.is_err());
289 assert!(result.unwrap_err().contains("max_tokens"));
290 }
291
292 #[test]
293 fn budget_check_passes_when_tokens_within_limit() {
294 let metrics = ExecutionMetrics::new();
295 metrics.set_budget(Budget {
296 max_llm_calls: None,
297 max_elapsed_ms: None,
298 max_tokens: Some(1000),
299 });
300 let observer = metrics.create_observer();
301 let handle = metrics.budget_handle();
302
303 let q = vec![LlmQuery {
304 id: QueryId::single(),
305 prompt: "short".into(),
306 system: None,
307 max_tokens: 100,
308 grounded: false,
309 underspecified: false,
310 cache_breakpoint: None,
311 role: None,
312 }];
313 observer.on_paused(&q);
314 observer.on_response_fed(&QueryId::single(), "reply", None);
315 observer.on_resumed();
316
317 assert!(handle.check().is_ok());
318 }
319
320 #[test]
321 fn budget_remaining_tracks_tokens() {
322 let metrics = ExecutionMetrics::new();
323 metrics.set_budget(Budget {
324 max_llm_calls: None,
325 max_elapsed_ms: None,
326 max_tokens: Some(100),
327 });
328 let observer = metrics.create_observer();
329 let handle = metrics.budget_handle();
330
331 let q = vec![LlmQuery {
333 id: QueryId::single(),
334 prompt: "test".into(),
335 system: None,
336 max_tokens: 10,
337 grounded: false,
338 underspecified: false,
339 cache_breakpoint: None,
340 role: None,
341 }];
342 observer.on_paused(&q);
343 observer.on_response_fed(&QueryId::single(), "r", None);
344 observer.on_resumed();
345
346 let remaining = handle.remaining();
347 assert_eq!(remaining["tokens"], 98);
349 }
350
351 #[test]
352 fn budget_from_ctx_extracts_max_tokens() {
353 let ctx = serde_json::json!({"budget": {"max_tokens": 5000}});
354 let budget = Budget::from_ctx(&ctx).expect("should parse");
355 assert_eq!(budget.max_llm_calls, None);
356 assert_eq!(budget.max_elapsed_ms, None);
357 assert_eq!(budget.max_tokens, Some(5000));
358 }
359
360 #[test]
361 fn budget_remaining_null_when_no_budget() {
362 let metrics = ExecutionMetrics::new();
363 let handle = metrics.budget_handle();
364 assert!(handle.remaining().is_null());
365 }
366
367 #[test]
368 fn budget_remaining_tracks_llm_calls() {
369 let metrics = ExecutionMetrics::new();
370 metrics.set_budget(Budget {
371 max_llm_calls: Some(5),
372 max_elapsed_ms: None,
373 max_tokens: None,
374 });
375 let observer = metrics.create_observer();
376 let handle = metrics.budget_handle();
377
378 let q = vec![LlmQuery {
379 id: QueryId::single(),
380 prompt: "p".into(),
381 system: None,
382 max_tokens: 10,
383 grounded: false,
384 underspecified: false,
385 cache_breakpoint: None,
386 role: None,
387 }];
388 observer.on_paused(&q);
389
390 let remaining = handle.remaining();
391 assert_eq!(remaining["llm_calls"], 4); }
393
394 #[test]
395 fn budget_in_stats_json() {
396 let metrics = ExecutionMetrics::new();
397 metrics.set_budget(Budget {
398 max_llm_calls: Some(10),
399 max_elapsed_ms: Some(60000),
400 max_tokens: None,
401 });
402 let observer = metrics.create_observer();
403 observer.on_completed(&serde_json::json!(null));
404
405 let json = metrics.to_json();
406 let budget = &json["auto"]["budget"];
407 assert_eq!(budget["max_llm_calls"], 10);
408 assert_eq!(budget["max_elapsed_ms"], 60000);
409 }
410
411 #[test]
412 fn no_budget_in_stats_json_when_not_set() {
413 let metrics = ExecutionMetrics::new();
414 let observer = metrics.create_observer();
415 observer.on_completed(&serde_json::json!(null));
416
417 let json = metrics.to_json();
418 assert!(json["auto"].get("budget").is_none());
419 }
420}