1use dashmap::DashMap;
7use prometheus::{register_int_counter_vec, register_int_gauge_vec, IntCounterVec, IntGaugeVec};
8use std::sync::atomic::{AtomicU64, AtomicU8, Ordering};
9use std::sync::LazyLock;
10use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};
11use tracing::{debug, info, trace, warn};
12
13use zentinel_common::budget::{
14 BudgetAlert, BudgetCheckResult, BudgetPeriod, TenantBudgetStatus, TokenBudgetConfig,
15};
16
17struct BudgetTrackerMetrics {
19 tenants: IntGaugeVec,
21 evictions: IntCounterVec,
23}
24
25static TRACKER_METRICS: LazyLock<Option<BudgetTrackerMetrics>> = LazyLock::new(|| {
26 let tenants = register_int_gauge_vec!(
27 "zentinel_inference_budget_tenants",
28 "Distinct tenants currently tracked by token budget trackers",
29 &["route"]
30 )
31 .ok()?;
32 let evictions = register_int_counter_vec!(
33 "zentinel_inference_budget_tenant_evictions_total",
34 "Tenant budget states evicted to enforce the max-tenants bound",
35 &["route"]
36 )
37 .ok()?;
38 Some(BudgetTrackerMetrics { tenants, evictions })
39});
40
41struct TenantBudgetState {
43 period_start: Instant,
45 period_start_unix: u64,
47 tokens_used: AtomicU64,
49 alerts_fired: AtomicU8,
52}
53
54impl TenantBudgetState {
55 fn new() -> Self {
56 let now_unix = SystemTime::now()
57 .duration_since(UNIX_EPOCH)
58 .unwrap_or_default()
59 .as_secs();
60
61 Self {
62 period_start: Instant::now(),
63 period_start_unix: now_unix,
64 tokens_used: AtomicU64::new(0),
65 alerts_fired: AtomicU8::new(0),
66 }
67 }
68
69 fn tokens_used(&self) -> u64 {
70 self.tokens_used.load(Ordering::Acquire)
71 }
72
73 fn add_tokens(&self, tokens: u64) {
74 self.tokens_used.fetch_add(tokens, Ordering::AcqRel);
75 }
76
77 fn elapsed(&self) -> Duration {
78 self.period_start.elapsed()
79 }
80
81 fn reset(&mut self) {
82 let now_unix = SystemTime::now()
83 .duration_since(UNIX_EPOCH)
84 .unwrap_or_default()
85 .as_secs();
86
87 self.period_start = Instant::now();
88 self.period_start_unix = now_unix;
89 self.tokens_used.store(0, Ordering::Release);
90 self.alerts_fired.store(0, Ordering::Release);
91 }
92
93 fn has_fired_alert(&self, threshold_index: u8) -> bool {
94 let mask = 1u8 << threshold_index;
95 (self.alerts_fired.load(Ordering::Acquire) & mask) != 0
96 }
97
98 fn mark_alert_fired(&self, threshold_index: u8) {
99 let mask = 1u8 << threshold_index;
100 self.alerts_fired.fetch_or(mask, Ordering::AcqRel);
101 }
102}
103
104pub struct TokenBudgetTracker {
113 config: TokenBudgetConfig,
115 tenants: DashMap<String, TenantBudgetState>,
117 route_id: String,
119}
120
121impl TokenBudgetTracker {
122 pub fn new(config: TokenBudgetConfig, route_id: impl Into<String>) -> Self {
124 let route_id = route_id.into();
125
126 info!(
127 route_id = %route_id,
128 period = ?config.period,
129 limit = config.limit,
130 enforce = config.enforce,
131 rollover = config.rollover,
132 "Created token budget tracker"
133 );
134
135 Self {
136 config,
137 tenants: DashMap::new(),
138 route_id,
139 }
140 }
141
142 pub fn check(&self, tenant: &str, estimated_tokens: u64) -> BudgetCheckResult {
146 let state = self.get_or_create_tenant(tenant);
147 let period_secs = self.config.period.as_secs();
148
149 let elapsed = state.elapsed();
151 if elapsed.as_secs() >= period_secs {
152 drop(state);
153 self.reset_period(tenant);
154 return self.check(tenant, estimated_tokens);
155 }
156
157 let current_used = state.tokens_used();
158 let would_use = current_used + estimated_tokens;
159
160 if would_use <= self.config.limit {
162 let remaining = self.config.limit.saturating_sub(would_use);
163 trace!(
164 route_id = %self.route_id,
165 tenant = tenant,
166 current_used = current_used,
167 estimated_tokens = estimated_tokens,
168 remaining = remaining,
169 "Budget check: allowed"
170 );
171 return BudgetCheckResult::Allowed { remaining };
172 }
173
174 if let Some(burst) = self.config.burst_allowance {
176 let burst_limit = self.config.limit + (self.config.limit as f64 * burst) as u64;
177 if would_use <= burst_limit {
178 let over_by = would_use - self.config.limit;
179 let remaining = (self.config.limit as i64) - (would_use as i64);
180 trace!(
181 route_id = %self.route_id,
182 tenant = tenant,
183 over_by = over_by,
184 "Budget check: soft limit (burst)"
185 );
186 return BudgetCheckResult::Soft { remaining, over_by };
187 }
188 }
189
190 if self.config.enforce {
192 let retry_after = period_secs.saturating_sub(elapsed.as_secs());
193 debug!(
194 route_id = %self.route_id,
195 tenant = tenant,
196 current_used = current_used,
197 limit = self.config.limit,
198 retry_after_secs = retry_after,
199 "Budget exhausted"
200 );
201 BudgetCheckResult::Exhausted {
202 retry_after_secs: retry_after,
203 }
204 } else {
205 let over_by = would_use - self.config.limit;
207 let remaining = (self.config.limit as i64) - (would_use as i64);
208 debug!(
209 route_id = %self.route_id,
210 tenant = tenant,
211 over_by = over_by,
212 "Budget exceeded (not enforced)"
213 );
214 BudgetCheckResult::Soft { remaining, over_by }
215 }
216 }
217
218 pub fn record(&self, tenant: &str, actual_tokens: u64) -> Vec<BudgetAlert> {
222 let state = self.get_or_create_tenant(tenant);
223 let period_secs = self.config.period.as_secs();
224
225 let elapsed = state.elapsed();
227 if elapsed.as_secs() >= period_secs {
228 drop(state);
229 self.reset_period(tenant);
230 return self.record(tenant, actual_tokens);
231 }
232
233 state.add_tokens(actual_tokens);
235 let new_total = state.tokens_used();
236
237 trace!(
238 route_id = %self.route_id,
239 tenant = tenant,
240 tokens = actual_tokens,
241 total = new_total,
242 limit = self.config.limit,
243 "Recorded token usage"
244 );
245
246 let mut alerts = Vec::new();
248 let usage_pct = new_total as f64 / self.config.limit as f64;
249
250 for (idx, &threshold) in self.config.alert_thresholds.iter().enumerate() {
251 if usage_pct >= threshold && !state.has_fired_alert(idx as u8) {
252 state.mark_alert_fired(idx as u8);
253
254 let alert = BudgetAlert {
255 tenant: tenant.to_string(),
256 threshold,
257 tokens_used: new_total,
258 tokens_limit: self.config.limit,
259 period_start: state.period_start_unix,
260 };
261
262 info!(
263 route_id = %self.route_id,
264 tenant = tenant,
265 threshold_pct = threshold * 100.0,
266 tokens_used = new_total,
267 tokens_limit = self.config.limit,
268 "Budget alert threshold crossed"
269 );
270
271 alerts.push(alert);
272 }
273 }
274
275 alerts
276 }
277
278 pub fn status(&self, tenant: &str) -> TenantBudgetStatus {
280 let state = self.get_or_create_tenant(tenant);
281 let period_secs = self.config.period.as_secs();
282 let elapsed = state.elapsed();
283
284 let tokens_used = state.tokens_used();
285 let tokens_remaining = self.config.limit.saturating_sub(tokens_used);
286 let usage_percent = (tokens_used as f64 / self.config.limit as f64) * 100.0;
287 let period_end = state.period_start_unix + period_secs;
288
289 TenantBudgetStatus {
290 tokens_used,
291 tokens_limit: self.config.limit,
292 tokens_remaining,
293 usage_percent,
294 period_start: state.period_start_unix,
295 period_end,
296 exhausted: tokens_used >= self.config.limit && self.config.enforce,
297 }
298 }
299
300 pub fn reset_period(&self, tenant: &str) {
302 if let Some(mut state) = self.tenants.get_mut(tenant) {
303 let old_tokens = state.tokens_used();
304
305 if self.config.rollover && old_tokens < self.config.limit {
307 let unused = self.config.limit - old_tokens;
308 state.reset();
309 let rollover = unused.min(self.config.limit);
311 state.add_tokens(rollover);
312 info!(
313 route_id = %self.route_id,
314 tenant = tenant,
315 rollover_tokens = rollover,
316 "Period reset with rollover"
317 );
318 } else {
319 state.reset();
320 debug!(
321 route_id = %self.route_id,
322 tenant = tenant,
323 previous_tokens = old_tokens,
324 "Period reset"
325 );
326 }
327 }
328 }
329
330 pub fn tenant_count(&self) -> usize {
332 self.tenants.len()
333 }
334
335 pub fn period_secs(&self) -> u64 {
337 self.config.period.as_secs()
338 }
339
340 pub fn limit(&self) -> u64 {
342 self.config.limit
343 }
344
345 pub fn is_enforced(&self) -> bool {
347 self.config.enforce
348 }
349
350 fn get_or_create_tenant(
351 &self,
352 tenant: &str,
353 ) -> dashmap::mapref::one::Ref<'_, String, TenantBudgetState> {
354 if !self.tenants.contains_key(tenant) && self.tenants.len() >= self.config.max_tenants {
355 self.evict_tenants();
356 }
357
358 self.tenants
359 .entry(tenant.to_string())
360 .or_insert_with(TenantBudgetState::new);
361
362 if let Some(metrics) = TRACKER_METRICS.as_ref() {
363 metrics
364 .tenants
365 .with_label_values(&[&self.route_id])
366 .set(self.tenants.len() as i64);
367 }
368
369 self.tenants.get(tenant).expect("Just inserted")
370 }
371
372 fn evict_tenants(&self) {
380 let period_secs = self.config.period.as_secs();
381 let keep_secs = if self.config.rollover {
382 period_secs.saturating_mul(2)
383 } else {
384 period_secs
385 };
386
387 let before = self.tenants.len();
388 self.tenants
389 .retain(|_, state| state.elapsed().as_secs() < keep_secs);
390
391 if self.tenants.len() >= self.config.max_tenants {
392 let target = (self.config.max_tenants * 9).div_ceil(10);
393 let mut entries: Vec<(String, u64)> = self
394 .tenants
395 .iter()
396 .map(|e| (e.key().clone(), e.value().period_start_unix))
397 .collect();
398 entries.sort_by_key(|(_, start)| *start);
399 let excess = self.tenants.len().saturating_sub(target);
400 for (tenant, _) in entries.iter().take(excess) {
401 self.tenants.remove(tenant);
402 }
403 }
404
405 let evicted = before.saturating_sub(self.tenants.len());
406 if evicted > 0 {
407 warn!(
408 route_id = %self.route_id,
409 evicted = evicted,
410 remaining = self.tenants.len(),
411 max_tenants = self.config.max_tenants,
412 "Tenant budget map at capacity, evicted tenant state"
413 );
414 if let Some(metrics) = TRACKER_METRICS.as_ref() {
415 metrics
416 .evictions
417 .with_label_values(&[&self.route_id])
418 .inc_by(evicted as u64);
419 }
420 }
421 }
422}
423
424#[cfg(test)]
429mod tests {
430 use super::*;
431
432 fn test_config() -> TokenBudgetConfig {
433 TokenBudgetConfig {
434 period: BudgetPeriod::Custom { seconds: 60 },
435 limit: 1000,
436 alert_thresholds: vec![0.50, 0.80, 0.95],
437 enforce: true,
438 rollover: false,
439 burst_allowance: None,
440 max_tenants: 10_000,
441 }
442 }
443
444 #[test]
445 fn test_check_allowed() {
446 let tracker = TokenBudgetTracker::new(test_config(), "test-route");
447
448 let result = tracker.check("tenant-1", 100);
449 assert!(result.is_allowed());
450
451 if let BudgetCheckResult::Allowed { remaining } = result {
452 assert_eq!(remaining, 900);
453 } else {
454 panic!("Expected Allowed result");
455 }
456 }
457
458 #[test]
459 fn test_check_exhausted() {
460 let tracker = TokenBudgetTracker::new(test_config(), "test-route");
461
462 tracker.record("tenant-1", 1000);
464
465 let result = tracker.check("tenant-1", 100);
467 assert!(!result.is_allowed());
468
469 if let BudgetCheckResult::Exhausted { retry_after_secs } = result {
470 assert!(retry_after_secs > 0);
471 } else {
472 panic!("Expected Exhausted result");
473 }
474 }
475
476 #[test]
477 fn test_record_alerts() {
478 let tracker = TokenBudgetTracker::new(test_config(), "test-route");
479
480 let alerts = tracker.record("tenant-1", 500);
482 assert_eq!(alerts.len(), 1);
483 assert!((alerts[0].threshold - 0.50).abs() < 0.001);
484
485 let alerts = tracker.record("tenant-1", 300);
487 assert_eq!(alerts.len(), 1);
488 assert!((alerts[0].threshold - 0.80).abs() < 0.001);
489
490 let alerts = tracker.record("tenant-1", 200);
492 assert_eq!(alerts.len(), 1);
493 assert!((alerts[0].threshold - 0.95).abs() < 0.001);
494
495 let alerts = tracker.record("tenant-1", 100);
497 assert!(alerts.is_empty());
498 }
499
500 #[test]
501 fn test_status() {
502 let tracker = TokenBudgetTracker::new(test_config(), "test-route");
503
504 tracker.record("tenant-1", 400);
505
506 let status = tracker.status("tenant-1");
507 assert_eq!(status.tokens_used, 400);
508 assert_eq!(status.tokens_limit, 1000);
509 assert_eq!(status.tokens_remaining, 600);
510 assert!((status.usage_percent - 40.0).abs() < 0.001);
511 assert!(!status.exhausted);
512 }
513
514 #[test]
515 fn test_burst_allowance() {
516 let mut config = test_config();
517 config.burst_allowance = Some(0.10); let tracker = TokenBudgetTracker::new(config, "test-route");
520
521 tracker.record("tenant-1", 950);
523
524 let result = tracker.check("tenant-1", 100);
525 assert!(result.is_allowed());
526
527 if let BudgetCheckResult::Soft { remaining, over_by } = result {
528 assert_eq!(over_by, 50);
529 assert_eq!(remaining, -50);
530 } else {
531 panic!("Expected Soft result");
532 }
533 }
534
535 #[test]
536 fn test_no_enforcement() {
537 let mut config = test_config();
538 config.enforce = false;
539
540 let tracker = TokenBudgetTracker::new(config, "test-route");
541
542 tracker.record("tenant-1", 1000);
544
545 let result = tracker.check("tenant-1", 100);
547 assert!(result.is_allowed());
548 }
549
550 #[test]
551 fn test_period_reset() {
552 let tracker = TokenBudgetTracker::new(test_config(), "test-route");
553
554 tracker.record("tenant-1", 500);
555 assert_eq!(tracker.status("tenant-1").tokens_used, 500);
556
557 tracker.reset_period("tenant-1");
558 assert_eq!(tracker.status("tenant-1").tokens_used, 0);
559 }
560
561 #[test]
562 fn test_rollover() {
563 let mut config = test_config();
564 config.rollover = true;
565
566 let tracker = TokenBudgetTracker::new(config, "test-route");
567
568 tracker.record("tenant-1", 300);
570
571 tracker.reset_period("tenant-1");
573
574 let status = tracker.status("tenant-1");
576 assert_eq!(status.tokens_used, 700);
577 }
578
579 #[test]
580 fn test_multiple_tenants() {
581 let tracker = TokenBudgetTracker::new(test_config(), "test-route");
582
583 tracker.record("tenant-1", 500);
584 tracker.record("tenant-2", 200);
585
586 assert_eq!(tracker.status("tenant-1").tokens_used, 500);
587 assert_eq!(tracker.status("tenant-2").tokens_used, 200);
588 assert_eq!(tracker.tenant_count(), 2);
589 }
590
591 #[test]
592 fn tenant_map_never_exceeds_max_tenants() {
593 let mut config = test_config();
594 config.max_tenants = 10;
595 let tracker = TokenBudgetTracker::new(config, "test-route");
596
597 for i in 0..100 {
598 tracker.record(&format!("tenant-{i}"), 1);
599 }
600
601 assert!(
602 tracker.tenant_count() <= 10,
603 "tenant map grew past max_tenants: {}",
604 tracker.tenant_count()
605 );
606 }
607
608 #[test]
609 fn tracker_still_enforces_budget_after_eviction() {
610 let mut config = test_config();
611 config.max_tenants = 5;
612 let tracker = TokenBudgetTracker::new(config, "test-route");
613
614 for i in 0..50 {
616 tracker.record(&format!("tenant-{i}"), 1);
617 }
618
619 tracker.record("fresh", 1000);
621 let result = tracker.check("fresh", 100);
622 assert!(
623 !result.is_allowed(),
624 "exhausted tenant must still be blocked after evictions"
625 );
626 }
627}