1use std::sync::Mutex;
9use std::sync::atomic::{AtomicU8, AtomicU32, AtomicU64, Ordering};
10use std::time::{Duration, Instant};
11
12use serde::{Deserialize, Serialize};
13use tokio::sync::{Semaphore, SemaphorePermit};
14use turnframe_core::replay::BudgetReport;
15
16#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
18#[serde(default, deny_unknown_fields)]
19#[non_exhaustive]
20pub struct Budget {
21 pub max_model_calls: Option<u32>,
23 pub max_chain_depth: Option<u8>,
25 pub max_parallel: usize,
27 pub max_prompt_tokens: Option<u64>,
29 pub max_wall_clock_secs: Option<u64>,
31 pub per_call_timeout_secs: u64,
33}
34
35impl Budget {
36 #[must_use]
38 pub const fn understanding() -> Self {
39 Self {
40 max_model_calls: Some(32),
41 max_chain_depth: Some(8),
42 max_parallel: 6,
43 max_prompt_tokens: Some(100_000),
44 max_wall_clock_secs: Some(30),
45 per_call_timeout_secs: 20,
46 }
47 }
48
49 #[must_use]
51 pub const fn narration() -> Self {
52 Self {
53 max_model_calls: Some(12),
54 max_chain_depth: Some(4),
55 max_parallel: 4,
56 max_prompt_tokens: Some(40_000),
57 max_wall_clock_secs: Some(20),
58 per_call_timeout_secs: 20,
59 }
60 }
61
62 #[must_use]
64 pub const fn unbounded() -> Self {
65 Self {
66 max_model_calls: None,
67 max_chain_depth: None,
68 max_parallel: 6,
69 max_prompt_tokens: None,
70 max_wall_clock_secs: None,
71 per_call_timeout_secs: 60,
72 }
73 }
74}
75
76impl Default for Budget {
77 fn default() -> Self {
78 Self::understanding()
79 }
80}
81
82#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
84#[serde(rename_all = "snake_case")]
85#[non_exhaustive]
86pub enum BudgetBound {
87 ModelCalls,
89 ChainDepth,
91 PromptTokens,
93 WallClock,
95}
96
97impl BudgetBound {
98 #[must_use]
100 pub const fn as_str(self) -> &'static str {
101 match self {
102 Self::ModelCalls => "model_calls",
103 Self::ChainDepth => "chain_depth",
104 Self::PromptTokens => "prompt_tokens",
105 Self::WallClock => "wall_clock",
106 }
107 }
108}
109
110#[derive(Debug)]
112pub struct BudgetTracker {
113 budget: Budget,
114 started: Instant,
115 calls: AtomicU32,
116 tokens: AtomicU64,
117 max_depth: AtomicU8,
118 exhausted: Mutex<Option<BudgetBound>>,
119 permits: Semaphore,
120}
121
122impl BudgetTracker {
123 #[must_use]
125 pub fn new(budget: Budget) -> Self {
126 Self {
127 budget,
128 started: Instant::now(),
129 calls: AtomicU32::new(0),
130 tokens: AtomicU64::new(0),
131 max_depth: AtomicU8::new(0),
132 exhausted: Mutex::new(None),
133 permits: Semaphore::new(budget.max_parallel.max(1)),
134 }
135 }
136
137 #[must_use]
139 pub const fn budget(&self) -> &Budget {
140 &self.budget
141 }
142
143 pub fn reserve(&self, depth: u8) -> Result<(), BudgetBound> {
149 let refused = self.first_bound(depth);
150 if let Some(bound) = refused {
151 self.exhaust(bound);
152 return Err(bound);
153 }
154 let limit = self.budget.max_model_calls.unwrap_or(u32::MAX);
155 let reserved = self
156 .calls
157 .fetch_update(Ordering::SeqCst, Ordering::SeqCst, |calls| {
158 (calls < limit).then_some(calls + 1)
159 });
160 if reserved.is_err() {
161 self.exhaust(BudgetBound::ModelCalls);
162 return Err(BudgetBound::ModelCalls);
163 }
164 self.max_depth.fetch_max(depth, Ordering::SeqCst);
165 Ok(())
166 }
167
168 fn first_bound(&self, depth: u8) -> Option<BudgetBound> {
169 if self.remaining_wall_clock() == Some(Duration::ZERO) {
170 return Some(BudgetBound::WallClock);
171 }
172 if self.budget.max_chain_depth.is_some_and(|max| depth > max) {
173 return Some(BudgetBound::ChainDepth);
174 }
175 if self
176 .budget
177 .max_prompt_tokens
178 .is_some_and(|max| self.tokens.load(Ordering::SeqCst) >= max)
179 {
180 return Some(BudgetBound::PromptTokens);
181 }
182 None
183 }
184
185 fn exhaust(&self, bound: BudgetBound) {
186 let mut exhausted = self
187 .exhausted
188 .lock()
189 .unwrap_or_else(std::sync::PoisonError::into_inner);
190 exhausted.get_or_insert(bound);
191 }
192
193 pub async fn permit(&self) -> Option<SemaphorePermit<'_>> {
195 self.permits.acquire().await.ok()
196 }
197
198 pub fn record_tokens(&self, tokens: u64) {
200 self.tokens.fetch_add(tokens, Ordering::SeqCst);
201 }
202
203 #[must_use]
205 pub fn remaining_wall_clock(&self) -> Option<Duration> {
206 self.budget
207 .max_wall_clock_secs
208 .map(|secs| Duration::from_secs(secs).saturating_sub(self.started.elapsed()))
209 }
210
211 #[must_use]
213 pub fn call_timeout(&self, requested: Option<Duration>) -> Duration {
214 let per_call = Duration::from_secs(self.budget.per_call_timeout_secs);
215 let wanted = requested.map_or(per_call, |requested| requested.min(per_call));
216 match self.remaining_wall_clock() {
217 Some(left) => wanted.min(left),
218 None => wanted,
219 }
220 }
221
222 #[must_use]
224 pub fn exhausted(&self) -> Option<BudgetBound> {
225 *self
226 .exhausted
227 .lock()
228 .unwrap_or_else(std::sync::PoisonError::into_inner)
229 }
230
231 #[must_use]
233 pub fn report(&self) -> BudgetReport {
234 let mut report = BudgetReport::default();
235 report.model_calls = self.calls.load(Ordering::SeqCst);
236 report.prompt_tokens = self.tokens.load(Ordering::SeqCst);
237 report.max_depth = self.max_depth.load(Ordering::SeqCst);
238 report.exhausted = self.exhausted().map(|bound| bound.as_str().to_owned());
239 report
240 }
241}
242
243#[cfg(test)]
244mod tests {
245 use super::*;
246
247 #[test]
248 fn calls_are_refused_once_the_bound_is_reached_and_the_bound_is_reported() {
249 let tracker = BudgetTracker::new(Budget {
250 max_model_calls: Some(2),
251 ..Budget::understanding()
252 });
253 assert!(tracker.reserve(1).is_ok());
254 assert!(tracker.reserve(2).is_ok());
255 assert_eq!(tracker.reserve(1), Err(BudgetBound::ModelCalls));
256 let report = tracker.report();
257 assert_eq!(report.model_calls, 2);
258 assert_eq!(report.max_depth, 2);
259 assert_eq!(report.exhausted.as_deref(), Some("model_calls"));
260 }
261
262 #[test]
263 fn a_chain_too_deep_is_refused_before_a_call_is_counted() {
264 let tracker = BudgetTracker::new(Budget::understanding());
265 assert_eq!(tracker.reserve(9), Err(BudgetBound::ChainDepth));
266 assert_eq!(tracker.report().model_calls, 0);
267 }
268
269 #[test]
270 fn reported_tokens_stop_the_next_call() {
271 let tracker = BudgetTracker::new(Budget {
272 max_prompt_tokens: Some(100),
273 ..Budget::understanding()
274 });
275 tracker.record_tokens(100);
276 assert_eq!(tracker.reserve(1), Err(BudgetBound::PromptTokens));
277 }
278
279 #[test]
280 fn unbounded_is_a_choice_that_still_keeps_a_deadline() {
281 let tracker = BudgetTracker::new(Budget::unbounded());
282 for _ in 0..100 {
283 assert!(tracker.reserve(200).is_ok());
284 }
285 assert_eq!(tracker.call_timeout(None), Duration::from_secs(60));
286 }
287}