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 }];
248 observer.on_paused(&q);
249 observer.on_response_fed(&QueryId::single(), "r", None);
250 observer.on_resumed();
251 observer.on_paused(&q);
252
253 let result = handle.check();
255 assert!(result.is_err());
256 assert!(result.unwrap_err().contains("budget_exceeded"));
257 }
258
259 #[test]
260 fn budget_check_fails_when_tokens_exceeded() {
261 let metrics = ExecutionMetrics::new();
262 metrics.set_budget(Budget {
263 max_llm_calls: None,
264 max_elapsed_ms: None,
265 max_tokens: Some(10),
266 });
267 let observer = metrics.create_observer();
268 let handle = metrics.budget_handle();
269
270 let q = vec![LlmQuery {
273 id: QueryId::single(),
274 prompt: "abcdefghijklmnopqrstuvwxyz0123456789abcd".into(),
275 system: None,
276 max_tokens: 100,
277 grounded: false,
278 underspecified: false,
279 cache_breakpoint: None,
280 }];
281 observer.on_paused(&q);
282 observer.on_response_fed(&QueryId::single(), "r", None);
283 observer.on_resumed();
284
285 let result = handle.check();
286 assert!(result.is_err());
287 assert!(result.unwrap_err().contains("max_tokens"));
288 }
289
290 #[test]
291 fn budget_check_passes_when_tokens_within_limit() {
292 let metrics = ExecutionMetrics::new();
293 metrics.set_budget(Budget {
294 max_llm_calls: None,
295 max_elapsed_ms: None,
296 max_tokens: Some(1000),
297 });
298 let observer = metrics.create_observer();
299 let handle = metrics.budget_handle();
300
301 let q = vec![LlmQuery {
302 id: QueryId::single(),
303 prompt: "short".into(),
304 system: None,
305 max_tokens: 100,
306 grounded: false,
307 underspecified: false,
308 cache_breakpoint: None,
309 }];
310 observer.on_paused(&q);
311 observer.on_response_fed(&QueryId::single(), "reply", None);
312 observer.on_resumed();
313
314 assert!(handle.check().is_ok());
315 }
316
317 #[test]
318 fn budget_remaining_tracks_tokens() {
319 let metrics = ExecutionMetrics::new();
320 metrics.set_budget(Budget {
321 max_llm_calls: None,
322 max_elapsed_ms: None,
323 max_tokens: Some(100),
324 });
325 let observer = metrics.create_observer();
326 let handle = metrics.budget_handle();
327
328 let q = vec![LlmQuery {
330 id: QueryId::single(),
331 prompt: "test".into(),
332 system: None,
333 max_tokens: 10,
334 grounded: false,
335 underspecified: false,
336 cache_breakpoint: None,
337 }];
338 observer.on_paused(&q);
339 observer.on_response_fed(&QueryId::single(), "r", None);
340 observer.on_resumed();
341
342 let remaining = handle.remaining();
343 assert_eq!(remaining["tokens"], 98);
345 }
346
347 #[test]
348 fn budget_from_ctx_extracts_max_tokens() {
349 let ctx = serde_json::json!({"budget": {"max_tokens": 5000}});
350 let budget = Budget::from_ctx(&ctx).expect("should parse");
351 assert_eq!(budget.max_llm_calls, None);
352 assert_eq!(budget.max_elapsed_ms, None);
353 assert_eq!(budget.max_tokens, Some(5000));
354 }
355
356 #[test]
357 fn budget_remaining_null_when_no_budget() {
358 let metrics = ExecutionMetrics::new();
359 let handle = metrics.budget_handle();
360 assert!(handle.remaining().is_null());
361 }
362
363 #[test]
364 fn budget_remaining_tracks_llm_calls() {
365 let metrics = ExecutionMetrics::new();
366 metrics.set_budget(Budget {
367 max_llm_calls: Some(5),
368 max_elapsed_ms: None,
369 max_tokens: None,
370 });
371 let observer = metrics.create_observer();
372 let handle = metrics.budget_handle();
373
374 let q = vec![LlmQuery {
375 id: QueryId::single(),
376 prompt: "p".into(),
377 system: None,
378 max_tokens: 10,
379 grounded: false,
380 underspecified: false,
381 cache_breakpoint: None,
382 }];
383 observer.on_paused(&q);
384
385 let remaining = handle.remaining();
386 assert_eq!(remaining["llm_calls"], 4); }
388
389 #[test]
390 fn budget_in_stats_json() {
391 let metrics = ExecutionMetrics::new();
392 metrics.set_budget(Budget {
393 max_llm_calls: Some(10),
394 max_elapsed_ms: Some(60000),
395 max_tokens: None,
396 });
397 let observer = metrics.create_observer();
398 observer.on_completed(&serde_json::json!(null));
399
400 let json = metrics.to_json();
401 let budget = &json["auto"]["budget"];
402 assert_eq!(budget["max_llm_calls"], 10);
403 assert_eq!(budget["max_elapsed_ms"], 60000);
404 }
405
406 #[test]
407 fn no_budget_in_stats_json_when_not_set() {
408 let metrics = ExecutionMetrics::new();
409 let observer = metrics.create_observer();
410 observer.on_completed(&serde_json::json!(null));
411
412 let json = metrics.to_json();
413 assert!(json["auto"].get("budget").is_none());
414 }
415}