use std::time::Duration;
use salvor_core::{Budget, BudgetKind};
use serde_json::Value;
#[derive(Clone, Debug, Default, PartialEq)]
pub struct Budgets {
pub max_steps: Option<u64>,
pub max_tokens: Option<u64>,
pub max_cost_usd: Option<f64>,
pub max_wall_time: Option<Duration>,
}
impl Budgets {
#[must_use]
pub fn any_declared(&self) -> bool {
self.max_steps.is_some()
|| self.max_tokens.is_some()
|| self.max_cost_usd.is_some()
|| self.max_wall_time.is_some()
}
#[must_use]
pub fn first_crossing(
&self,
extensions: &BudgetExtensions,
pricing: Option<&Pricing>,
observations: &BudgetObservations,
) -> Option<(Budget, f64)> {
if let Some(max_steps) = self.max_steps {
let limit = to_f64(max_steps.saturating_add(extensions.steps));
let observed = to_f64(observations.steps);
if observed >= limit {
return Some((
Budget {
kind: BudgetKind::Steps,
limit,
},
observed,
));
}
}
if let Some(max_tokens) = self.max_tokens {
let limit = to_f64(max_tokens.saturating_add(extensions.tokens));
let observed = to_f64(
observations
.input_tokens
.saturating_add(observations.output_tokens),
);
if observed >= limit {
return Some((
Budget {
kind: BudgetKind::Tokens,
limit,
},
observed,
));
}
}
if let (Some(max_cost), Some(pricing)) = (self.max_cost_usd, pricing) {
let limit = max_cost + extensions.cost_usd;
let observed = pricing.cost_usd(observations.input_tokens, observations.output_tokens);
if observed >= limit {
return Some((
Budget {
kind: BudgetKind::CostUsd,
limit,
},
observed,
));
}
}
if let Some(max_wall) = self.max_wall_time {
let limit = max_wall.as_secs_f64() + extensions.wall_time_seconds;
let observed = observations.elapsed_seconds;
if observed >= limit {
return Some((
Budget {
kind: BudgetKind::WallTime,
limit,
},
observed,
));
}
}
None
}
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct Pricing {
pub input_per_mtok: f64,
pub output_per_mtok: f64,
}
impl Pricing {
#[must_use]
pub fn cost_usd(&self, input_tokens: u64, output_tokens: u64) -> f64 {
to_f64(input_tokens) / 1_000_000.0 * self.input_per_mtok
+ to_f64(output_tokens) / 1_000_000.0 * self.output_per_mtok
}
}
#[derive(Clone, Copy, Debug, Default, PartialEq)]
pub struct BudgetObservations {
pub steps: u64,
pub input_tokens: u64,
pub output_tokens: u64,
pub elapsed_seconds: f64,
}
#[derive(Clone, Copy, Debug, Default, PartialEq)]
pub struct BudgetExtensions {
pub steps: u64,
pub tokens: u64,
pub cost_usd: f64,
pub wall_time_seconds: f64,
}
impl BudgetExtensions {
pub fn absorb(&mut self, resume_input: &Value) {
let Some(extend) = resume_input.get("extend").and_then(Value::as_object) else {
return;
};
if let Some(steps) = extend.get("steps").and_then(Value::as_u64) {
self.steps = self.steps.saturating_add(steps);
}
if let Some(tokens) = extend.get("tokens").and_then(Value::as_u64) {
self.tokens = self.tokens.saturating_add(tokens);
}
if let Some(cost) = extend.get("cost_usd").and_then(Value::as_f64) {
self.cost_usd += cost;
}
if let Some(seconds) = extend.get("wall_time_seconds").and_then(Value::as_f64) {
self.wall_time_seconds += seconds;
}
}
}
pub fn validate_extension_input(input: &Value) -> Result<(), String> {
let Some(top) = input.as_object() else {
return Err("a budget-crossing resume input must be a JSON object".to_owned());
};
for key in top.keys() {
if key != "extend" {
return Err(format!(
"unexpected top-level key `{key}`; a budget-crossing resume input may only carry `extend`"
));
}
}
let Some(extend) = top.get("extend") else {
return Ok(());
};
let Some(extend) = extend.as_object() else {
return Err("`extend` must be a JSON object".to_owned());
};
for (key, value) in extend {
match key.as_str() {
"steps" | "tokens" => {
if value.as_u64().is_none() {
return Err(format!("`extend.{key}` must be an unsigned integer"));
}
}
"cost_usd" | "wall_time_seconds" => {
if value.as_f64().is_none() {
return Err(format!("`extend.{key}` must be a number"));
}
}
other => {
return Err(format!(
"unknown key `extend.{other}`; expected steps, tokens, cost_usd, or wall_time_seconds"
));
}
}
}
Ok(())
}
#[allow(clippy::cast_precision_loss)]
fn to_f64(count: u64) -> f64 {
count as f64
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn crossing_fires_at_limit_and_extensions_raise_it() {
let budgets = Budgets {
max_steps: Some(2),
..Budgets::default()
};
let mut extensions = BudgetExtensions::default();
let observations = BudgetObservations {
steps: 2,
..BudgetObservations::default()
};
let (budget, observed) = budgets
.first_crossing(&extensions, None, &observations)
.expect("steps crossing fires at the limit");
assert_eq!(budget.kind, BudgetKind::Steps);
assert_eq!(budget.limit, 2.0);
assert_eq!(observed, 2.0);
extensions.absorb(&json!({"extend": {"steps": 3}}));
assert_eq!(
budgets.first_crossing(&extensions, None, &observations),
None,
"the extension raises the effective limit past the observation"
);
}
#[test]
fn check_order_is_steps_first() {
let budgets = Budgets {
max_steps: Some(1),
max_tokens: Some(10),
..Budgets::default()
};
let observations = BudgetObservations {
steps: 1,
input_tokens: 100,
output_tokens: 100,
..BudgetObservations::default()
};
let (budget, _) = budgets
.first_crossing(&BudgetExtensions::default(), None, &observations)
.expect("a crossing fires");
assert_eq!(budget.kind, BudgetKind::Steps);
}
#[test]
fn cost_crossing_uses_pricing() {
let budgets = Budgets {
max_cost_usd: Some(1.0),
..Budgets::default()
};
let pricing = Pricing {
input_per_mtok: 3.0,
output_per_mtok: 15.0,
};
let observations = BudgetObservations {
input_tokens: 200_000,
output_tokens: 40_000,
..BudgetObservations::default()
};
let (budget, observed) = budgets
.first_crossing(&BudgetExtensions::default(), Some(&pricing), &observations)
.expect("cost crossing fires");
assert_eq!(budget.kind, BudgetKind::CostUsd);
assert!((observed - 1.2).abs() < 1e-12);
}
#[test]
fn extension_validation_rejects_wrong_shapes() {
assert!(validate_extension_input(&json!({})).is_ok());
assert!(validate_extension_input(&json!({"extend": {"steps": 2}})).is_ok());
assert!(
validate_extension_input(&json!({
"extend": {"steps": 1, "tokens": 2, "cost_usd": 0.5, "wall_time_seconds": 60}
}))
.is_ok()
);
assert!(validate_extension_input(&json!("more please")).is_err());
assert!(validate_extension_input(&json!({"other": 1})).is_err());
assert!(validate_extension_input(&json!({"extend": 5})).is_err());
assert!(validate_extension_input(&json!({"extend": {"stepz": 1}})).is_err());
assert!(validate_extension_input(&json!({"extend": {"steps": -1}})).is_err());
assert!(validate_extension_input(&json!({"extend": {"cost_usd": "1"}})).is_err());
}
}