use crate::provider::{ModelCapabilities, TokenMeasurement, TokenMeasurementSource};
use serde::{Deserialize, Serialize};
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct BudgetPolicy {
pub reserved_output_tokens: u64,
pub extra_safety_margin_tokens: u64,
}
impl Default for BudgetPolicy {
fn default() -> Self {
Self {
reserved_output_tokens: 16_384,
extra_safety_margin_tokens: 1_024,
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum BudgetVerdict {
Within {
available_tokens: u64,
},
Exceeded {
measured_tokens: u64,
available_tokens: u64,
source: TokenMeasurementSource,
},
}
pub fn evaluate(
capabilities: &ModelCapabilities,
measurement: &TokenMeasurement,
policy: &BudgetPolicy,
) -> BudgetVerdict {
let margin = measurement
.safety_margin_tokens
.saturating_add(policy.extra_safety_margin_tokens);
let reservations = policy.reserved_output_tokens.saturating_add(margin);
let window = capabilities.context_tokens as u64;
let available_tokens = window.saturating_sub(reservations);
if measurement.input_tokens.saturating_add(reservations) <= window {
BudgetVerdict::Within { available_tokens }
} else {
BudgetVerdict::Exceeded {
measured_tokens: measurement.input_tokens,
available_tokens,
source: measurement.source,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn caps(context_tokens: u32) -> ModelCapabilities {
ModelCapabilities {
context_tokens,
max_output_tokens: 8_192,
supports_tools: true,
supports_images: true,
supports_reasoning: true,
}
}
fn policy() -> BudgetPolicy {
BudgetPolicy {
reserved_output_tokens: 10_000,
extra_safety_margin_tokens: 1_000,
}
}
fn measurement(tokens: u64, margin: u64) -> TokenMeasurement {
TokenMeasurement {
input_tokens: tokens,
source: TokenMeasurementSource::Heuristic,
safety_margin_tokens: margin,
}
}
#[test]
fn within_budget_when_formula_holds() {
let verdict = evaluate(&caps(200_000), &measurement(100_000, 500), &policy());
match verdict {
BudgetVerdict::Within { available_tokens } => {
assert_eq!(available_tokens, 200_000 - 10_000 - 1_000 - 500);
}
other => panic!("expected within, got {other:?}"),
}
}
#[test]
fn exceeds_when_input_plus_reservations_overflow_window() {
let verdict = evaluate(&caps(200_000), &measurement(189_001, 0), &policy());
match verdict {
BudgetVerdict::Exceeded {
measured_tokens,
available_tokens,
source,
} => {
assert_eq!(measured_tokens, 189_001);
assert_eq!(available_tokens, 189_000);
assert_eq!(source, TokenMeasurementSource::Heuristic);
}
other => panic!("expected exceeded, got {other:?}"),
}
}
#[test]
fn saturates_when_reservations_exceed_window() {
let verdict = evaluate(&caps(1_000), &measurement(0, 0), &policy());
assert!(matches!(verdict, BudgetVerdict::Exceeded { .. }));
}
#[test]
fn within_at_exact_boundary_is_inclusive() {
let verdict = evaluate(&caps(200_000), &measurement(189_000, 0), &policy());
assert!(matches!(
verdict,
BudgetVerdict::Within {
available_tokens: 189_000
}
));
}
}