Saturate billing cost accumulation

This commit is contained in:
Bryan Helmkamp 2026-07-23 07:53:01 -04:00
parent 4d2712b8db
commit 08cef66ec8
No known key found for this signature in database
3 changed files with 113 additions and 30 deletions

View file

@ -68,13 +68,13 @@ impl std::ops::Add for UsdMicros {
type Output = Self;
fn add(self, rhs: Self) -> Self::Output {
Self(self.0 + rhs.0)
Self(self.0.saturating_add(rhs.0))
}
}
impl std::ops::AddAssign for UsdMicros {
fn add_assign(&mut self, rhs: Self) {
self.0 += rhs.0;
*self = *self + rhs;
}
}
@ -84,6 +84,12 @@ impl std::iter::Sum for UsdMicros {
}
}
fn accumulate_optional_usd_micros(total: &mut Option<i64>, cost: Option<i64>) {
let mut typed_total = (*total).map(UsdMicros);
UsdMicros::accumulate(&mut typed_total, cost.map(UsdMicros));
*total = typed_total.map(|value| value.0);
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub struct PricePerMTok {
pub usd_micros: i64,
@ -376,25 +382,21 @@ impl BilledTokenCounts {
#[must_use]
pub fn from_billed_usage(billed: &[BilledModelUsage]) -> Self {
let mut tokens = TokenCounts::default();
let mut total_usd_micros = 0_i64;
let mut has_total = false;
let mut total_usd_micros = None;
for entry in billed {
tokens += entry.input.usage.tokens.clone();
if let Some(value) = entry.total_usd_micros {
total_usd_micros += value;
has_total = true;
}
accumulate_optional_usd_micros(&mut total_usd_micros, entry.total_usd_micros);
}
Self {
input_tokens: tokens.input_tokens,
output_tokens: tokens.output_tokens,
total_tokens: tokens.total_tokens(),
reasoning_tokens: tokens.reasoning_tokens,
cache_read_tokens: tokens.cache_read_tokens,
input_tokens: tokens.input_tokens,
output_tokens: tokens.output_tokens,
total_tokens: tokens.total_tokens(),
reasoning_tokens: tokens.reasoning_tokens,
cache_read_tokens: tokens.cache_read_tokens,
cache_write_tokens: tokens.cache_write_tokens,
total_usd_micros: has_total.then_some(total_usd_micros),
total_usd_micros,
}
}
@ -418,9 +420,7 @@ impl BilledTokenCounts {
self.reasoning_tokens += source.reasoning_tokens;
self.cache_read_tokens += source.cache_read_tokens;
self.cache_write_tokens += source.cache_write_tokens;
if let Some(value) = source.total_usd_micros {
*self.total_usd_micros.get_or_insert(0) += value;
}
accumulate_optional_usd_micros(&mut self.total_usd_micros, source.total_usd_micros);
}
pub fn add_billed_usage(&mut self, usage: &BilledModelUsage) {
@ -431,9 +431,7 @@ impl BilledTokenCounts {
self.cache_read_tokens += tokens.cache_read_tokens;
self.cache_write_tokens += tokens.cache_write_tokens;
self.total_tokens += tokens.total_tokens();
if let Some(value) = usage.total_usd_micros {
*self.total_usd_micros.get_or_insert(0) += value;
}
accumulate_optional_usd_micros(&mut self.total_usd_micros, usage.total_usd_micros);
}
pub fn replace_with_billed_usage(&mut self, usage: &BilledModelUsage) {
@ -779,6 +777,33 @@ mod tests {
assert_eq!(total, Some(UsdMicros(100_000)));
}
#[test]
fn usd_micros_arithmetic_saturates_at_i64_bounds() {
assert_eq!(UsdMicros(i64::MAX) + UsdMicros(1), UsdMicros(i64::MAX));
let mut minimum = UsdMicros(i64::MIN);
minimum += UsdMicros(-1);
assert_eq!(minimum, UsdMicros(i64::MIN));
assert_eq!(
[UsdMicros(i64::MAX), UsdMicros(1)]
.into_iter()
.sum::<UsdMicros>(),
UsdMicros(i64::MAX)
);
}
#[test]
fn usd_micros_accumulate_saturates_at_i64_bounds() {
let mut maximum = Some(UsdMicros(i64::MAX));
UsdMicros::accumulate(&mut maximum, Some(UsdMicros(1)));
assert_eq!(maximum, Some(UsdMicros(i64::MAX)));
let mut minimum = Some(UsdMicros(i64::MIN));
UsdMicros::accumulate(&mut minimum, Some(UsdMicros(-1)));
assert_eq!(minimum, Some(UsdMicros(i64::MIN)));
}
#[test]
fn model_billing_policy_override_changes_the_billing_algorithm() {
let catalog = catalog_from_toml(
@ -912,6 +937,31 @@ cache_input_cost_per_mtok = 0.3
assert_eq!(counts.total_usd_micros, Some(150));
}
#[test]
fn billed_token_counts_cost_rollups_saturate() {
let billed = [
billed_usage(0, 0, Some(i64::MAX)),
billed_usage(0, 0, Some(1)),
];
assert_eq!(
BilledTokenCounts::from_billed_usage(&billed).total_usd_micros,
Some(i64::MAX)
);
let mut counts = BilledTokenCounts {
total_usd_micros: Some(i64::MAX),
..BilledTokenCounts::default()
};
counts.add_counts(&BilledTokenCounts {
total_usd_micros: Some(1),
..BilledTokenCounts::default()
});
assert_eq!(counts.total_usd_micros, Some(i64::MAX));
counts.add_billed_usage(&billed_usage(0, 0, Some(1)));
assert_eq!(counts.total_usd_micros, Some(i64::MAX));
}
#[test]
fn billed_token_counts_replace_with_billed_usage_discards_previous_values() {
let mut counts = BilledTokenCounts {

View file

@ -718,16 +718,7 @@ async fn get_aggregate_billing(
agg.by_model
.values()
.fold(BilledTokenCounts::default(), |mut acc, totals| {
let billing = &totals.billing;
acc.input_tokens += billing.input_tokens;
acc.output_tokens += billing.output_tokens;
acc.reasoning_tokens += billing.reasoning_tokens;
acc.cache_read_tokens += billing.cache_read_tokens;
acc.cache_write_tokens += billing.cache_write_tokens;
acc.total_tokens += billing.total_tokens;
if let Some(value) = billing.total_usd_micros {
*acc.total_usd_micros.get_or_insert(0) += value;
}
acc.add_counts(&totals.billing);
acc
});
let response = AggregateBilling {

View file

@ -13551,6 +13551,48 @@ async fn get_aggregate_billing_returns_provider_model_speed_identity() {
assert_eq!(fast["billing"]["input_tokens"], 20);
}
#[tokio::test]
async fn get_aggregate_billing_saturates_total_cost_across_models() {
let state = test_app_state();
{
let mut agg = state
.aggregate_billing
.lock()
.expect("aggregate billing lock");
for (model_id, total_usd_micros) in [("maximum", i64::MAX), ("one", 1)] {
agg.by_model.insert(
ModelRef {
provider: ProviderId::openai(),
model_id: model_id.to_string(),
speed: None,
},
ModelBillingTotals {
stages: 1,
billing: BilledTokenCounts {
total_usd_micros: Some(total_usd_micros),
..BilledTokenCounts::default()
},
},
);
}
}
let app = crate::test_support::build_test_router(Arc::clone(&state));
let response = app
.oneshot(
Request::builder()
.method("GET")
.uri(api("/billing"))
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
let body = response_json!(response, StatusCode::OK).await;
assert_eq!(body["totals"]["total_usd_micros"].as_i64(), Some(i64::MAX));
}
#[test]
fn aggregate_billing_counts_projection_rollup_usage_visits() {
let mut accumulator = BillingAccumulator::default();