Skip to main content

turnframe_tasks/
budget.rs

1//! What a turn's model calls may spend, and which bound stopped them.
2//!
3//! A call is reserved before it is sent, so parallel tasks cannot overshoot a
4//! bound; tokens are counted from what providers report, after each call. A
5//! bound that is reached stops new calls and is reported, never a partial effect:
6//! what the caller does with an unfinished task is its own decision.
7
8use 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/// The limits of one phase of a turn. Every bound is optional except parallelism.
17#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
18#[serde(default, deny_unknown_fields)]
19#[non_exhaustive]
20pub struct Budget {
21    /// Model calls, each reserved before it is sent.
22    pub max_model_calls: Option<u32>,
23    /// Dependent calls in a row; a repair, a vote round and an escalation count one.
24    pub max_chain_depth: Option<u8>,
25    /// Calls in flight at once.
26    pub max_parallel: usize,
27    /// Prompt tokens as providers report them; no call starts once reached.
28    pub max_prompt_tokens: Option<u64>,
29    /// Wall clock of the whole phase, in seconds.
30    pub max_wall_clock_secs: Option<u64>,
31    /// Deadline of one call, in seconds.
32    pub per_call_timeout_secs: u64,
33}
34
35impl Budget {
36    /// Understanding a message: 32 calls, depth 8, 6 in flight, 100k tokens, 30 s.
37    #[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    /// Writing the reply: 12 calls, depth 4, 4 in flight, 40k tokens, 20 s.
50    #[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    /// No bound but parallelism and the per-call deadline: an explicit choice.
63    #[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/// The bound that stopped a call.
83#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
84#[serde(rename_all = "snake_case")]
85#[non_exhaustive]
86pub enum BudgetBound {
87    /// `max_model_calls`.
88    ModelCalls,
89    /// `max_chain_depth`.
90    ChainDepth,
91    /// `max_prompt_tokens`.
92    PromptTokens,
93    /// `max_wall_clock_secs`.
94    WallClock,
95}
96
97impl BudgetBound {
98    /// Stable label, for records and metrics.
99    #[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/// What one phase has spent so far. Shared by every task the phase runs.
111#[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    /// A tracker starting now.
124    #[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    /// The limits in force.
138    #[must_use]
139    pub const fn budget(&self) -> &Budget {
140        &self.budget
141    }
142
143    /// Reserves one call at `depth`.
144    ///
145    /// # Errors
146    ///
147    /// The [`BudgetBound`] that forbids it; the bound is also remembered for the report.
148    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    /// Waits for a slot among the calls in flight.
194    pub async fn permit(&self) -> Option<SemaphorePermit<'_>> {
195        self.permits.acquire().await.ok()
196    }
197
198    /// Counts the prompt tokens a provider reported.
199    pub fn record_tokens(&self, tokens: u64) {
200        self.tokens.fetch_add(tokens, Ordering::SeqCst);
201    }
202
203    /// Time left on the wall clock, when it is bounded.
204    #[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    /// The deadline of the next call: the per-call timeout, cut to the time left.
212    #[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    /// The first bound reached, if any.
223    #[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    /// What was spent, for the replay record.
232    #[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}