diff --git a/lib/crates/fabro-model/src/billing.rs b/lib/crates/fabro-model/src/billing.rs index 0df9fc824..baa592cb5 100644 --- a/lib/crates/fabro-model/src/billing.rs +++ b/lib/crates/fabro-model/src/billing.rs @@ -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, cost: Option) { + 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(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 { diff --git a/lib/crates/fabro-server/src/server/handler/system.rs b/lib/crates/fabro-server/src/server/handler/system.rs index ba387b1ab..d26e445a6 100644 --- a/lib/crates/fabro-server/src/server/handler/system.rs +++ b/lib/crates/fabro-server/src/server/handler/system.rs @@ -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 { diff --git a/lib/crates/fabro-server/src/server/tests.rs b/lib/crates/fabro-server/src/server/tests.rs index 7314693dd..3d412d5a4 100644 --- a/lib/crates/fabro-server/src/server/tests.rs +++ b/lib/crates/fabro-server/src/server/tests.rs @@ -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();