Skip to main content

algocline_core/
budget.rs

1use std::sync::{Arc, Mutex};
2
3use crate::metrics::SessionStatus;
4
5// ─── Budget ──────────────────────────────────────────────────
6
7/// Session-level resource limits.
8///
9/// Extracted from `ctx.budget` at session start. When a limit is reached,
10/// `alc.llm()` raises a catchable Lua error (`"budget_exceeded: ..."`)
11/// **before** sending the request to the host. The check happens at
12/// call-site, not after the LLM response arrives.
13///
14/// Budget is shared across the entire session — if `alc.pipe()` chains
15/// multiple strategies, they all draw from the same budget.
16#[derive(Debug, Clone, Default)]
17pub struct Budget {
18    /// Maximum number of LLM calls allowed in this session.
19    /// Checked against `SessionStatus::llm_calls` (incremented in `on_paused`).
20    pub max_llm_calls: Option<u64>,
21    /// Maximum wall-clock time (ms) allowed for this session.
22    /// Measured from session start (`Instant::now()` at construction).
23    /// Note: this is wall-clock, not CPU time. Includes time spent
24    /// waiting for host LLM responses.
25    pub max_elapsed_ms: Option<u64>,
26    /// Maximum total tokens (prompt + response) allowed in this session.
27    /// Checked against accumulated `prompt_tokens + response_tokens`.
28    /// Token counts may be estimated (±30%) or host-provided depending
29    /// on `TokenSource`. Budget check uses whatever is available.
30    pub max_tokens: Option<u64>,
31}
32
33impl Budget {
34    /// Extract budget from ctx JSON. Returns None if no budget field present.
35    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    /// Check if the session is within budget given current counters.
51    /// Returns Err with a structured message if any limit is exceeded.
52    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    /// Remaining budget as JSON given current counters.
78    /// Returns `{ llm_calls: N|null, elapsed_ms: N|null, tokens: N|null }`.
79    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    /// Serialize budget limits to JSON (for stats output).
93    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/// Cheap, cloneable handle for budget checking from the Lua bridge.
109///
110/// Wraps the shared `SessionStatus` to expose only budget-related queries.
111/// Passed to `bridge::register_llm()` where it gates every `alc.llm()`
112/// and `alc.llm_batch()` call.
113///
114/// # Call site and threading
115///
116/// Both `check()` and `remaining()` are called exclusively from Lua
117/// closures registered in `bridge.rs`, which run on the Lua OS thread.
118/// They acquire `std::sync::Mutex<SessionStatus>` for a few microseconds
119/// (read-only field comparisons). See `SessionStatus` doc for full
120/// locking design.
121///
122/// # TOCTOU safety of Lua-side `alc.budget_check()`
123///
124/// The prelude's `alc.budget_check()` calls `remaining()` then the
125/// caller decides whether to call `alc.llm()` (which calls `check()`).
126/// Between `remaining()` release and `check()` acquire, `llm_calls`
127/// could theoretically change — but within a single session this is
128/// structurally impossible: the Lua thread is the only writer path
129/// (via observer callbacks), and observer callbacks only fire after
130/// the Lua thread yields control through the mpsc channel. Lua is
131/// single-threaded and does not yield between `budget_check()` and
132/// `alc.llm()`.
133///
134/// # Poison policy
135///
136/// `check()` propagates poison as `Err` — this surfaces as a Lua error,
137/// which is the correct behavior since a poisoned mutex indicates an
138/// unrecoverable state (OOM panic under lock). `remaining()` returns
139/// `Null` on poison — it is observational and non-fatal.
140#[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    /// Check if the session is within budget. Returns Err with a message if exceeded.
151    ///
152    /// Propagates mutex poison as `Err` (see BudgetHandle doc for rationale).
153    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    /// Remaining budget as JSON: `{ llm_calls: N|null, elapsed_ms: N|null }`.
162    /// Returns `serde_json::Value::Null` if no budget is set.
163    ///
164    /// Returns `Null` on mutex poison (observational, non-fatal).
165    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        // Simulate 2 LLM calls
239        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        // Now at 2 calls, budget is max_llm_calls=2
255        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        // Simulate an LLM call with a prompt that estimates to ≥10 tokens.
272        // "abcdefghijklmnopqrstuvwxyz0123456789abcd" = 40 ASCII chars → ceil(40/4) = 10 tokens
273        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        // "test" = 4 chars → ceil(4/4) = 1 token prompt, "r" → 1 token response
332        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        // 100 - (1 prompt + 1 response) = 98
348        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); // 5 - 1
392    }
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}