tokio_prompt_orchestrator/
token_budget.rs1use std::sync::atomic::{AtomicU64, Ordering};
47use std::sync::{Arc, Mutex};
48use std::time::{Duration, Instant};
49use thiserror::Error;
50use tracing::{debug, warn};
51
52#[derive(Debug, Clone)]
54pub struct TokenBudgetConfig {
55 pub max_tokens_per_request: u64,
61
62 pub max_tokens_per_period: u64,
67
68 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#[derive(Debug, Error)]
84pub enum TokenBudgetError {
85 #[error("request exceeds per-request token limit: estimated {estimated} > max {limit}")]
87 PerRequestLimitExceeded { estimated: u64, limit: u64 },
88
89 #[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
100struct PeriodState {
103 tokens_used: u64,
104 window_start: Instant,
105}
106
107#[derive(Clone)]
112pub struct TokenBudgetGuard {
113 config: TokenBudgetConfig,
114 total_tokens_estimated: Arc<AtomicU64>,
118 requests_allowed: Arc<AtomicU64>,
120 requests_rejected: Arc<AtomicU64>,
122 period: Arc<Mutex<PeriodState>>,
123}
124
125impl TokenBudgetGuard {
126 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 pub fn check(&self, text: &str) -> Result<u64, TokenBudgetError> {
150 let estimated = estimate_tokens(text);
151
152 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 if self.config.max_tokens_per_period > 0 {
170 let mut period = self.period.lock().unwrap_or_else(|e| e.into_inner());
171
172 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 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 period.tokens_used = period.tokens_used.saturating_sub(estimated - actual);
221 } else {
222 period.tokens_used = period
224 .tokens_used
225 .saturating_add(actual - estimated)
226 .min(self.config.max_tokens_per_period);
227 }
228 }
229
230 pub fn total_tokens_estimated(&self) -> u64 {
232 self.total_tokens_estimated.load(Ordering::Relaxed)
233 }
234
235 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 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 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 pub fn requests_allowed(&self) -> u64 {
260 self.requests_allowed.load(Ordering::Relaxed)
261 }
262
263 pub fn requests_rejected(&self) -> u64 {
265 self.requests_rejected.load(Ordering::Relaxed)
266 }
267
268 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
281pub 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 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 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 let _ = guard.check("Hi"); 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(); let used_before = guard.period_tokens_used();
360 guard.release(estimated, 1); 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"); let _ = guard.check("Hello world this is way too long for the limit"); 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}