use rust_decimal::prelude::*;
use rust_decimal::Decimal;
use serde::{Deserialize, Serialize};
use crate::budget::Budget;
const EWMA_ALPHA_NUM: u64 = 3;
const EWMA_ALPHA_DEN: u64 = 10;
#[derive(Debug, Clone)]
pub struct ForecastInputs {
pub cost_so_far_cents: u64,
pub tool_calls_so_far: u32,
pub wall_time_ms_so_far: u64,
pub steps_completed: u32,
pub step_cost_cents: u64,
pub prior_ewma_step_cost_cents: Option<u64>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum BreachKind {
CostCents,
ToolCalls,
WallTime,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct ForecastSnapshot {
pub projected_cost_cents: u64,
pub ewma_cost_cents: u64,
pub ewma_step_cost_cents: u64,
pub budget_breach_projected: bool,
pub breach_kind: Option<BreachKind>,
}
pub fn compute_forecast(inputs: ForecastInputs, budget: &Budget) -> ForecastSnapshot {
let projected_cost_cents = linear_projection(&inputs, budget);
let ewma_step_cost_cents = update_ewma(&inputs);
let ewma_cost_cents = ewma_projection(&inputs, budget, ewma_step_cost_cents);
let breach_kind = projected_breach(projected_cost_cents, ewma_cost_cents, &inputs, budget);
ForecastSnapshot {
projected_cost_cents,
ewma_cost_cents,
ewma_step_cost_cents,
budget_breach_projected: breach_kind.is_some(),
breach_kind,
}
}
fn linear_projection(inputs: &ForecastInputs, budget: &Budget) -> u64 {
let cost = Decimal::from(inputs.cost_so_far_cents);
if cost.is_zero() {
return inputs.cost_so_far_cents;
}
let mut max_fraction: Option<Decimal> = None;
if let Some(cap) = budget.max_cost_cents {
if let Some(f) = fraction(inputs.cost_so_far_cents, cap) {
max_fraction = max_decimal(max_fraction, f);
}
}
if let Some(cap) = budget.max_tool_calls {
if let Some(f) = fraction(inputs.tool_calls_so_far as u64, cap as u64) {
max_fraction = max_decimal(max_fraction, f);
}
}
if let Some(cap) = budget.max_wall_time_ms {
if let Some(f) = fraction(inputs.wall_time_ms_so_far, cap) {
max_fraction = max_decimal(max_fraction, f);
}
}
match max_fraction {
Some(frac) if frac > Decimal::ZERO => {
let projected = cost / frac;
projected
.round_dp_with_strategy(0, rust_decimal::RoundingStrategy::AwayFromZero)
.to_u64()
.unwrap_or(u64::MAX)
}
_ => inputs.cost_so_far_cents,
}
}
fn update_ewma(inputs: &ForecastInputs) -> u64 {
let step = Decimal::from(inputs.step_cost_cents);
let alpha_num = Decimal::from(EWMA_ALPHA_NUM);
let alpha_den = Decimal::from(EWMA_ALPHA_DEN);
match inputs.prior_ewma_step_cost_cents {
None => inputs.step_cost_cents,
Some(prior) => {
let prior = Decimal::from(prior);
let weighted = alpha_num * step + (alpha_den - alpha_num) * prior;
(weighted / alpha_den)
.round_dp_with_strategy(0, rust_decimal::RoundingStrategy::AwayFromZero)
.to_u64()
.unwrap_or(prior.to_u64().unwrap_or(0))
}
}
}
fn ewma_projection(inputs: &ForecastInputs, budget: &Budget, ewma_step: u64) -> u64 {
let remaining_steps = match budget.max_tool_calls {
Some(cap) if cap > inputs.tool_calls_so_far => (cap - inputs.tool_calls_so_far) as u64,
_ => return linear_projection(inputs, budget),
};
inputs
.cost_so_far_cents
.saturating_add(ewma_step.saturating_mul(remaining_steps))
}
fn projected_breach(
projected_cost_cents: u64,
ewma_cost_cents: u64,
inputs: &ForecastInputs,
budget: &Budget,
) -> Option<BreachKind> {
if let Some(cap) = budget.max_wall_time_ms {
if inputs.wall_time_ms_so_far >= cap {
return Some(BreachKind::WallTime);
}
}
if let Some(cap) = budget.max_cost_cents {
if projected_cost_cents > cap || ewma_cost_cents > cap {
return Some(BreachKind::CostCents);
}
}
if let (Some(cap), Some(rate)) = (
budget.max_tool_calls,
per_ms_rate(inputs.tool_calls_so_far as u64, inputs.wall_time_ms_so_far),
) {
if let Some(remaining_ms) = remaining_wall_ms(inputs, budget) {
let projected_calls = inputs.tool_calls_so_far as u64
+ (rate * Decimal::from(remaining_ms))
.round_dp_with_strategy(0, rust_decimal::RoundingStrategy::AwayFromZero)
.to_u64()
.unwrap_or(0);
if projected_calls > cap as u64 {
return Some(BreachKind::ToolCalls);
}
}
}
None
}
fn fraction(used: u64, cap: u64) -> Option<Decimal> {
if cap == 0 {
return None;
}
Some(Decimal::from(used) / Decimal::from(cap))
}
fn per_ms_rate(used: u64, elapsed_ms: u64) -> Option<Decimal> {
if elapsed_ms == 0 {
return None;
}
Some(Decimal::from(used) / Decimal::from(elapsed_ms))
}
fn remaining_wall_ms(inputs: &ForecastInputs, budget: &Budget) -> Option<u64> {
let cap = budget.max_wall_time_ms?;
Some(cap.saturating_sub(inputs.wall_time_ms_so_far))
}
fn max_decimal(current: Option<Decimal>, candidate: Decimal) -> Option<Decimal> {
match current {
Some(c) if c >= candidate => Some(c),
_ => Some(candidate),
}
}
#[cfg(test)]
mod tests {
use super::*;
fn small_budget() -> Budget {
Budget {
max_input_tokens: None,
max_output_tokens: None,
max_total_tokens: None,
max_tool_calls: Some(10),
max_wall_time_ms: Some(60_000),
max_cost_cents: Some(500),
}
}
#[test]
fn linear_projects_pro_rata_against_walltime() {
let inputs = ForecastInputs {
cost_so_far_cents: 100,
tool_calls_so_far: 1,
wall_time_ms_so_far: 10_000,
steps_completed: 1,
step_cost_cents: 100,
prior_ewma_step_cost_cents: None,
};
let snap = compute_forecast(inputs, &small_budget());
assert_eq!(snap.projected_cost_cents, 500);
}
#[test]
fn linear_zero_cost_stays_zero() {
let inputs = ForecastInputs {
cost_so_far_cents: 0,
tool_calls_so_far: 0,
wall_time_ms_so_far: 0,
steps_completed: 0,
step_cost_cents: 0,
prior_ewma_step_cost_cents: None,
};
let snap = compute_forecast(inputs, &small_budget());
assert_eq!(snap.projected_cost_cents, 0);
assert!(!snap.budget_breach_projected);
}
#[test]
fn linear_no_budget_returns_cost_so_far() {
let inputs = ForecastInputs {
cost_so_far_cents: 250,
tool_calls_so_far: 3,
wall_time_ms_so_far: 5_000,
steps_completed: 3,
step_cost_cents: 80,
prior_ewma_step_cost_cents: None,
};
let budget = Budget {
max_input_tokens: None,
max_output_tokens: None,
max_total_tokens: None,
max_tool_calls: None,
max_wall_time_ms: None,
max_cost_cents: None,
};
let snap = compute_forecast(inputs, &budget);
assert_eq!(snap.projected_cost_cents, 250);
assert!(!snap.budget_breach_projected);
}
#[test]
fn ewma_seeds_from_first_step() {
let inputs = ForecastInputs {
cost_so_far_cents: 50,
tool_calls_so_far: 1,
wall_time_ms_so_far: 5_000,
steps_completed: 1,
step_cost_cents: 50,
prior_ewma_step_cost_cents: None,
};
let snap = compute_forecast(inputs, &small_budget());
assert_eq!(snap.ewma_step_cost_cents, 50);
}
#[test]
fn ewma_converges_under_constant_step_cost() {
let mut ewma = 30u64;
for _ in 0..10 {
let inputs = ForecastInputs {
cost_so_far_cents: 30 * 10,
tool_calls_so_far: 3,
wall_time_ms_so_far: 10_000,
steps_completed: 3,
step_cost_cents: 30,
prior_ewma_step_cost_cents: Some(ewma),
};
let snap = compute_forecast(inputs, &small_budget());
ewma = snap.ewma_step_cost_cents;
}
assert_eq!(ewma, 30, "EWMA must converge to the constant step cost");
}
#[test]
fn ewma_dampens_single_spike() {
let inputs = ForecastInputs {
cost_so_far_cents: 50,
tool_calls_so_far: 5,
wall_time_ms_so_far: 5_000,
steps_completed: 5,
step_cost_cents: 100,
prior_ewma_step_cost_cents: Some(10),
};
let snap = compute_forecast(inputs, &small_budget());
assert_eq!(snap.ewma_step_cost_cents, 37);
}
#[test]
fn no_breach_when_well_under_budget() {
let inputs = ForecastInputs {
cost_so_far_cents: 25,
tool_calls_so_far: 1,
wall_time_ms_so_far: 30_000,
steps_completed: 1,
step_cost_cents: 25,
prior_ewma_step_cost_cents: None,
};
let snap = compute_forecast(inputs, &small_budget());
assert!(!snap.budget_breach_projected);
assert!(snap.breach_kind.is_none());
}
#[test]
fn breach_flagged_when_projection_exceeds_cost_cap() {
let inputs = ForecastInputs {
cost_so_far_cents: 350,
tool_calls_so_far: 6,
wall_time_ms_so_far: 10_000,
steps_completed: 6,
step_cost_cents: 50,
prior_ewma_step_cost_cents: Some(50),
};
let snap = compute_forecast(inputs, &small_budget());
assert!(snap.budget_breach_projected);
assert_eq!(snap.breach_kind, Some(BreachKind::CostCents));
}
#[test]
fn breach_flagged_when_walltime_already_exceeded() {
let inputs = ForecastInputs {
cost_so_far_cents: 100,
tool_calls_so_far: 1,
wall_time_ms_so_far: 65_000,
steps_completed: 1,
step_cost_cents: 100,
prior_ewma_step_cost_cents: None,
};
let snap = compute_forecast(inputs, &small_budget());
assert!(snap.budget_breach_projected);
assert_eq!(snap.breach_kind, Some(BreachKind::WallTime));
}
#[test]
fn breach_flagged_for_runaway_tool_calls() {
let inputs = ForecastInputs {
cost_so_far_cents: 50,
tool_calls_so_far: 4,
wall_time_ms_so_far: 10_000,
steps_completed: 4,
step_cost_cents: 10,
prior_ewma_step_cost_cents: Some(10),
};
let snap = compute_forecast(inputs, &small_budget());
assert!(snap.budget_breach_projected);
assert_eq!(snap.breach_kind, Some(BreachKind::ToolCalls));
}
#[test]
fn snapshot_is_serde_round_trip() {
let snap = ForecastSnapshot {
projected_cost_cents: 750,
ewma_cost_cents: 720,
ewma_step_cost_cents: 60,
budget_breach_projected: true,
breach_kind: Some(BreachKind::CostCents),
};
let json = serde_json::to_string(&snap).unwrap();
let back: ForecastSnapshot = serde_json::from_str(&json).unwrap();
assert_eq!(snap, back);
}
}