Skip to main content

tokio_prompt_orchestrator/
token_budget.rs

1//! Token Budget Middleware
2//!
3//! Pre-estimates token counts before requests reach the LLM, enforcing spend
4//! limits at the orchestration layer rather than discovering them via a billing
5//! surprise.
6//!
7//! ## Why estimate before sending?
8//!
9//! LLM APIs charge per token.  Without pre-flight counting:
10//! - A single runaway prompt (e.g. 200 KB user-uploaded doc) consumes the
11//!   entire daily budget in one call.
12//! - Cost anomalies are invisible until the invoice arrives.
13//! - There is no way to shed low-priority requests before they hit the API.
14//!
15//! `TokenBudgetGuard` adds a zero-network-round-trip gate that estimates token
16//! count, checks it against configurable per-request and period limits, and
17//! rejects requests that would overflow the budget before they reach the wire.
18//!
19//! ## Token estimation
20//!
21//! The estimator uses the `⌈len / 4⌉` heuristic (one token ≈ 4 UTF-8 bytes
22//! in English prose).  This intentionally over-estimates by ~10–15 % on code
23//! and ~5 % on English, providing a conservative safety margin without
24//! requiring a tokenizer dependency.  For accurate accounting of actual tokens
25//! used, pair this with [`crate::metrics`] which records real token counts
26//! from provider responses.
27//!
28//! ## Example
29//!
30//! ```rust
31//! use tokio_prompt_orchestrator::token_budget::{TokenBudgetGuard, TokenBudgetConfig};
32//!
33//! let guard = TokenBudgetGuard::new(TokenBudgetConfig {
34//!     max_tokens_per_request: 4_096,
35//!     max_tokens_per_period: 1_000_000,
36//!     period: std::time::Duration::from_secs(3600), // 1-hour rolling window
37//! });
38//!
39//! let prompt = "Summarise this document in three bullet points.";
40//! match guard.check(prompt) {
41//!     Ok(estimated) => println!("Allowed — estimated {estimated} tokens"),
42//!     Err(e) => eprintln!("Budget exceeded: {e}"),
43//! }
44//! ```
45
46use std::sync::atomic::{AtomicU64, Ordering};
47use std::sync::{Arc, Mutex};
48use std::time::{Duration, Instant};
49use thiserror::Error;
50use tracing::{debug, warn};
51
52/// Configuration for [`TokenBudgetGuard`].
53#[derive(Debug, Clone)]
54pub struct TokenBudgetConfig {
55    /// Maximum estimated tokens allowed in a single request.
56    ///
57    /// Requests whose estimated prompt token count exceeds this limit are
58    /// rejected before they reach the LLM.  A value of `0` disables the
59    /// per-request cap.
60    pub max_tokens_per_request: u64,
61
62    /// Maximum total tokens consumed within a rolling `period` window.
63    ///
64    /// The period budget resets automatically when `period` elapses.  A value
65    /// of `0` disables the period cap.
66    pub max_tokens_per_period: u64,
67
68    /// Rolling window duration for the period budget.
69    pub period: Duration,
70}
71
72impl Default for TokenBudgetConfig {
73    fn default() -> Self {
74        Self {
75            max_tokens_per_request: 8_192,
76            max_tokens_per_period: 1_000_000,
77            period: Duration::from_secs(3600),
78        }
79    }
80}
81
82/// Token budget rejection reasons.
83#[derive(Debug, Error)]
84pub enum TokenBudgetError {
85    /// The single request exceeds `max_tokens_per_request`.
86    #[error("request exceeds per-request token limit: estimated {estimated} > max {limit}")]
87    PerRequestLimitExceeded { estimated: u64, limit: u64 },
88
89    /// The period budget would be exceeded by this request.
90    #[error(
91        "request would exceed period token budget: used {used} + estimated {estimated} > max {limit}"
92    )]
93    PeriodLimitExceeded {
94        used: u64,
95        estimated: u64,
96        limit: u64,
97    },
98}
99
100/// Mutable period-tracking state (protected by a `Mutex` for atomic
101/// window-reset + counter-increment without a separate compare-and-swap loop).
102struct PeriodState {
103    tokens_used: u64,
104    window_start: Instant,
105}
106
107/// Guards LLM requests against token over-spend.
108///
109/// `TokenBudgetGuard` is `Clone + Send + Sync`.  All clones share the same
110/// budget counters via an `Arc`.
111#[derive(Clone)]
112pub struct TokenBudgetGuard {
113    config: TokenBudgetConfig,
114    /// Total tokens consumed across all requests since the process started
115    /// (not reset on window rollover — use `period_tokens_used()` for
116    /// the rolling window value).
117    total_tokens_estimated: Arc<AtomicU64>,
118    /// Number of requests that passed the budget gate.
119    requests_allowed: Arc<AtomicU64>,
120    /// Number of requests rejected by the budget gate.
121    requests_rejected: Arc<AtomicU64>,
122    period: Arc<Mutex<PeriodState>>,
123}
124
125impl TokenBudgetGuard {
126    /// Create a new guard with the given configuration.
127    pub fn new(config: TokenBudgetConfig) -> Self {
128        Self {
129            config,
130            total_tokens_estimated: Arc::new(AtomicU64::new(0)),
131            requests_allowed: Arc::new(AtomicU64::new(0)),
132            requests_rejected: Arc::new(AtomicU64::new(0)),
133            period: Arc::new(Mutex::new(PeriodState {
134                tokens_used: 0,
135                window_start: Instant::now(),
136            })),
137        }
138    }
139
140    /// Check whether `text` fits within the budget limits.
141    ///
142    /// If the check passes, the estimated token count is **reserved** against
143    /// the period budget.  Call [`release`](Self::release) with the actual
144    /// token count from the provider response to correct the reservation.
145    ///
146    /// # Returns
147    /// - `Ok(estimated_tokens)` — request is within budget; proceed to LLM.
148    /// - `Err(TokenBudgetError::*)` — budget exceeded; reject the request.
149    pub fn check(&self, text: &str) -> Result<u64, TokenBudgetError> {
150        let estimated = estimate_tokens(text);
151
152        // Per-request cap.
153        if self.config.max_tokens_per_request > 0
154            && estimated > self.config.max_tokens_per_request
155        {
156            warn!(
157                estimated,
158                limit = self.config.max_tokens_per_request,
159                "token_budget: per-request limit exceeded"
160            );
161            self.requests_rejected.fetch_add(1, Ordering::Relaxed);
162            return Err(TokenBudgetError::PerRequestLimitExceeded {
163                estimated,
164                limit: self.config.max_tokens_per_request,
165            });
166        }
167
168        // Period cap — needs mutex for atomic read-modify-write.
169        if self.config.max_tokens_per_period > 0 {
170            let mut period = self.period.lock().unwrap_or_else(|e| e.into_inner());
171
172            // Roll the window if the period has elapsed.
173            if period.window_start.elapsed() >= self.config.period {
174                debug!(
175                    previous_used = period.tokens_used,
176                    "token_budget: rolling period window"
177                );
178                period.tokens_used = 0;
179                period.window_start = Instant::now();
180            }
181
182            if period.tokens_used + estimated > self.config.max_tokens_per_period {
183                warn!(
184                    used = period.tokens_used,
185                    estimated,
186                    limit = self.config.max_tokens_per_period,
187                    "token_budget: period limit exceeded"
188                );
189                self.requests_rejected.fetch_add(1, Ordering::Relaxed);
190                return Err(TokenBudgetError::PeriodLimitExceeded {
191                    used: period.tokens_used,
192                    estimated,
193                    limit: self.config.max_tokens_per_period,
194                });
195            }
196
197            period.tokens_used += estimated;
198        }
199
200        self.total_tokens_estimated
201            .fetch_add(estimated, Ordering::Relaxed);
202        self.requests_allowed.fetch_add(1, Ordering::Relaxed);
203        debug!(estimated, "token_budget: request allowed");
204        Ok(estimated)
205    }
206
207    /// Adjust the period budget by the difference between the estimated and
208    /// actual token counts returned by the provider.
209    ///
210    /// Call this after receiving the LLM response with the real token count.
211    /// If `actual < estimated`, the surplus is credited back.  If
212    /// `actual > estimated` (rare for input tokens), the overage is debited.
213    pub fn release(&self, estimated: u64, actual: u64) {
214        if actual == estimated {
215            return;
216        }
217        let mut period = self.period.lock().unwrap_or_else(|e| e.into_inner());
218        if actual < estimated {
219            // Credit back the over-reservation.
220            period.tokens_used = period.tokens_used.saturating_sub(estimated - actual);
221        } else {
222            // Charge the extra tokens consumed.
223            period.tokens_used = period
224                .tokens_used
225                .saturating_add(actual - estimated)
226                .min(self.config.max_tokens_per_period);
227        }
228    }
229
230    /// Total tokens estimated across all requests since creation (lifetime).
231    pub fn total_tokens_estimated(&self) -> u64 {
232        self.total_tokens_estimated.load(Ordering::Relaxed)
233    }
234
235    /// Tokens consumed in the current rolling period window.
236    pub fn period_tokens_used(&self) -> u64 {
237        let period = self.period.lock().unwrap_or_else(|e| e.into_inner());
238        period.tokens_used
239    }
240
241    /// Remaining token budget in the current period.
242    pub fn period_tokens_remaining(&self) -> u64 {
243        let period = self.period.lock().unwrap_or_else(|e| e.into_inner());
244        self.config
245            .max_tokens_per_period
246            .saturating_sub(period.tokens_used)
247    }
248
249    /// Fraction of the period budget consumed (0.0 = empty, 1.0 = exhausted).
250    pub fn period_utilization(&self) -> f64 {
251        if self.config.max_tokens_per_period == 0 {
252            return 0.0;
253        }
254        let used = self.period_tokens_used() as f64;
255        (used / self.config.max_tokens_per_period as f64).min(1.0)
256    }
257
258    /// Number of requests allowed through the budget gate.
259    pub fn requests_allowed(&self) -> u64 {
260        self.requests_allowed.load(Ordering::Relaxed)
261    }
262
263    /// Number of requests rejected by the budget gate.
264    pub fn requests_rejected(&self) -> u64 {
265        self.requests_rejected.load(Ordering::Relaxed)
266    }
267
268    /// Budget rejection rate: `rejected / (allowed + rejected)`.
269    pub fn rejection_rate(&self) -> f64 {
270        let allowed = self.requests_allowed.load(Ordering::Relaxed) as f64;
271        let rejected = self.requests_rejected.load(Ordering::Relaxed) as f64;
272        let total = allowed + rejected;
273        if total == 0.0 {
274            0.0
275        } else {
276            rejected / total
277        }
278    }
279}
280
281/// Estimate the number of tokens in `text`.
282///
283/// Uses the `⌈len / 4⌉` heuristic: one token ≈ 4 UTF-8 bytes in English
284/// prose.  Over-estimates by ~10–15 % on code; ~5 % on English.  Never
285/// returns 0 for non-empty input (minimum 1 token).
286pub fn estimate_tokens(text: &str) -> u64 {
287    let bytes = text.len() as u64;
288    if bytes == 0 {
289        return 0;
290    }
291    bytes.div_ceil(4).max(1)
292}
293
294#[cfg(test)]
295mod tests {
296    use super::*;
297
298    #[test]
299    fn estimate_tokens_empty() {
300        assert_eq!(estimate_tokens(""), 0);
301    }
302
303    #[test]
304    fn estimate_tokens_short() {
305        // "Hello" = 5 bytes → ceil(5/4) = 2
306        assert_eq!(estimate_tokens("Hello"), 2);
307    }
308
309    #[test]
310    fn estimate_tokens_typical_prompt() {
311        let prompt = "Summarise this article in three bullet points.";
312        let est = estimate_tokens(prompt);
313        assert!(est > 0 && est < 20, "estimate was {est}");
314    }
315
316    #[test]
317    fn per_request_limit_rejected() {
318        let guard = TokenBudgetGuard::new(TokenBudgetConfig {
319            max_tokens_per_request: 5,
320            max_tokens_per_period: 1_000_000,
321            period: Duration::from_secs(3600),
322        });
323        // "Hello world this is a long prompt" is well over 5 tokens.
324        let result = guard.check("Hello world this is a long prompt that clearly exceeds five tokens");
325        assert!(matches!(result, Err(TokenBudgetError::PerRequestLimitExceeded { .. })));
326    }
327
328    #[test]
329    fn short_request_passes() {
330        let guard = TokenBudgetGuard::new(TokenBudgetConfig {
331            max_tokens_per_request: 1000,
332            max_tokens_per_period: 1_000_000,
333            period: Duration::from_secs(3600),
334        });
335        assert!(guard.check("Hi").is_ok());
336    }
337
338    #[test]
339    fn period_budget_exhausted() {
340        let guard = TokenBudgetGuard::new(TokenBudgetConfig {
341            max_tokens_per_request: 1000,
342            max_tokens_per_period: 10,
343            period: Duration::from_secs(3600),
344        });
345        // First request consumes up to 10 tokens.
346        let _ = guard.check("Hi"); // 1 token
347        let result = guard.check("Hello world this is a very long prompt that will overflow the budget");
348        assert!(matches!(result, Err(TokenBudgetError::PeriodLimitExceeded { .. })));
349    }
350
351    #[test]
352    fn release_credits_back_surplus() {
353        let guard = TokenBudgetGuard::new(TokenBudgetConfig {
354            max_tokens_per_request: 1000,
355            max_tokens_per_period: 100,
356            period: Duration::from_secs(3600),
357        });
358        let estimated = guard.check("Hello world").unwrap(); // small estimate
359        let used_before = guard.period_tokens_used();
360        guard.release(estimated, 1); // actual was 1 token less than estimated
361        assert!(guard.period_tokens_used() < used_before);
362    }
363
364    #[test]
365    fn rejection_rate_tracks_correctly() {
366        let guard = TokenBudgetGuard::new(TokenBudgetConfig {
367            max_tokens_per_request: 3,
368            max_tokens_per_period: 1_000_000,
369            period: Duration::from_secs(3600),
370        });
371        let _ = guard.check("Hi");  // allowed (≤ 3 tokens)
372        let _ = guard.check("Hello world this is way too long for the limit");  // rejected
373        assert_eq!(guard.requests_allowed(), 1);
374        assert_eq!(guard.requests_rejected(), 1);
375        assert!((guard.rejection_rate() - 0.5).abs() < 1e-10);
376    }
377}