Skip to main content

zentinel_proxy/inference/
budget.rs

1//! Token budget tracker for per-tenant cumulative usage tracking.
2//!
3//! Unlike rate limiting (tokens per minute), budgets track cumulative usage
4//! over longer periods (hourly, daily, monthly) with optional enforcement.
5
6use 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
17/// Prometheus metrics for tenant budget maps.
18struct BudgetTrackerMetrics {
19    /// Distinct tenants currently tracked, per route
20    tenants: IntGaugeVec,
21    /// Tenants evicted to enforce the max-tenants bound, per route
22    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
41/// Per-tenant budget state tracking
42struct TenantBudgetState {
43    /// Period start time
44    period_start: Instant,
45    /// Period start time as Unix timestamp (for reporting)
46    period_start_unix: u64,
47    /// Tokens used in current period
48    tokens_used: AtomicU64,
49    /// Bitmask of alert thresholds that have been triggered
50    /// Bit 0 = first threshold, Bit 1 = second, etc.
51    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
104/// Token budget tracker for per-tenant usage tracking.
105///
106/// Tracks cumulative token usage over configurable periods (hourly, daily, monthly)
107/// with support for:
108/// - Configurable alert thresholds
109/// - Hard or soft enforcement
110/// - Optional burst allowance
111/// - Period rollover
112pub struct TokenBudgetTracker {
113    /// Budget configuration
114    config: TokenBudgetConfig,
115    /// Per-tenant budget state
116    tenants: DashMap<String, TenantBudgetState>,
117    /// Route ID for logging
118    route_id: String,
119}
120
121impl TokenBudgetTracker {
122    /// Create a new token budget tracker with the given configuration.
123    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    /// Check if a request with the given token count is allowed.
143    ///
144    /// This does NOT consume tokens - call `record()` after the request completes.
145    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        // Check if period has expired
150        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        // Check against limit
161        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        // Check burst allowance
175        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        // Budget exhausted
191        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            // Not enforcing, just log and allow
206            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    /// Record actual token usage after a request completes.
219    ///
220    /// Returns any budget alerts that should be fired.
221    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        // Check if period has expired
226        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        // Add tokens
234        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        // Check for alert thresholds
247        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    /// Get the current budget status for a tenant.
279    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    /// Reset the budget period for a tenant.
301    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            // Handle rollover
306            if self.config.rollover && old_tokens < self.config.limit {
307                let unused = self.config.limit - old_tokens;
308                state.reset();
309                // Add back unused tokens (capped at limit)
310                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    /// Get the number of tracked tenants.
331    pub fn tenant_count(&self) -> usize {
332        self.tenants.len()
333    }
334
335    /// Get the period duration in seconds.
336    pub fn period_secs(&self) -> u64 {
337        self.config.period.as_secs()
338    }
339
340    /// Get the budget limit.
341    pub fn limit(&self) -> u64 {
342        self.config.limit
343    }
344
345    /// Check if enforcement is enabled.
346    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    /// Evict tenant state so a new tenant can be admitted without unbounded growth.
373    ///
374    /// Tenants whose period has expired are dropped first — their state is
375    /// equivalent to a fresh one (with rollover enabled, one extra period is
376    /// kept so unused tokens can still roll over on next access). If the map
377    /// is still at capacity, the tenants with the oldest periods are evicted
378    /// down to 90% of the bound.
379    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// ============================================================================
425// Tests
426// ============================================================================
427
428#[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        // Use up the budget
463        tracker.record("tenant-1", 1000);
464
465        // Next check should be exhausted
466        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        // Record 500 tokens (50% threshold)
481        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        // Record 300 more tokens (80% threshold)
486        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        // Record 200 more tokens (95% + 100% threshold, but 100% not in thresholds)
491        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        // No more alerts
496        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); // 10% burst
518
519        let tracker = TokenBudgetTracker::new(config, "test-route");
520
521        // Use 1050 tokens (5% over limit, within burst)
522        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        // Use up budget
543        tracker.record("tenant-1", 1000);
544
545        // Should still be allowed (soft)
546        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        // Use 300 tokens (700 unused)
569        tracker.record("tenant-1", 300);
570
571        // Reset with rollover
572        tracker.reset_period("tenant-1");
573
574        // Should have 700 tokens carried over
575        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        // Force evictions with many distinct tenants
615        for i in 0..50 {
616            tracker.record(&format!("tenant-{i}"), 1);
617        }
618
619        // A fresh tenant must still get correct budget accounting
620        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}