diff --git a/lib/crates/fabro-model/src/billing.rs b/lib/crates/fabro-model/src/billing.rs index 29cad2335..f06cad06a 100644 --- a/lib/crates/fabro-model/src/billing.rs +++ b/lib/crates/fabro-model/src/billing.rs @@ -358,6 +358,19 @@ impl BilledTokenCounts { } } + /// Returns the five disjoint per-call token buckets, dropping the derived + /// `total_tokens` sum and the optional `total_usd_micros` cost. + #[must_use] + pub fn token_counts(&self) -> TokenCounts { + TokenCounts { + input_tokens: self.input_tokens, + output_tokens: self.output_tokens, + reasoning_tokens: self.reasoning_tokens, + cache_read_tokens: self.cache_read_tokens, + cache_write_tokens: self.cache_write_tokens, + } + } + pub fn add_counts(&mut self, source: &Self) { self.input_tokens += source.input_tokens; self.output_tokens += source.output_tokens; @@ -435,6 +448,27 @@ impl Catalog { self.provider(&model_ref.provider) .and_then(|provider| ModelBillingFacts::for_policy(provider.billing_policy, tokens)) } + + /// Price a partial token sample for `model` using catalog pricing. + /// + /// Returns `None` when the provider has no billing policy, the model is + /// unknown, or the pricing algorithm cannot produce a result for the given + /// tokens. Used by read-side rollups so in-flight stages can show an + /// exact cost for the tokens consumed so far. + #[must_use] + pub fn price_tokens(&self, model: &ModelRef, tokens: &TokenCounts) -> Option { + let facts = self.billing_facts_for(model, tokens)?; + let input = ModelBillingInput { + usage: ModelUsage { + model: model.clone(), + tokens: tokens.clone(), + }, + facts, + }; + self.pricing_for(model) + .and_then(|pricing| pricing.bill(&input)) + .map(|amount| amount.0) + } } fn costs_for_speed( diff --git a/lib/crates/fabro-server/src/server.rs b/lib/crates/fabro-server/src/server.rs index b5258fa65..731c3f926 100644 --- a/lib/crates/fabro-server/src/server.rs +++ b/lib/crates/fabro-server/src/server.rs @@ -3238,7 +3238,7 @@ async fn execute_run_in_process(state: Arc, run_id: RunId) { .expect("aggregate_billing lock poisoned"); accumulate_billing_rollup( &mut agg, - &fabro_workflow::billing_rollup_from_projection(projection), + &fabro_workflow::billing_rollup_from_projection(projection, None), ); } } @@ -3523,7 +3523,7 @@ async fn execute_run_subprocess(state: Arc, run_id: RunId) { .expect("aggregate_billing lock poisoned"); accumulate_billing_rollup( &mut agg, - &fabro_workflow::billing_rollup_from_projection(&final_state), + &fabro_workflow::billing_rollup_from_projection(&final_state, None), ); } diff --git a/lib/crates/fabro-server/src/server/handler/billing.rs b/lib/crates/fabro-server/src/server/handler/billing.rs index 6eee9bbfb..cbaf7cc3e 100644 --- a/lib/crates/fabro-server/src/server/handler/billing.rs +++ b/lib/crates/fabro-server/src/server/handler/billing.rs @@ -79,7 +79,8 @@ async fn get_run_billing( }; let projection = cached.projection; - let rollup = fabro_workflow::billing_rollup_from_projection(&projection); + let catalog = state.catalog(); + let rollup = fabro_workflow::billing_rollup_from_projection(&projection, Some(&catalog)); let by_model = rollup .by_model .iter() diff --git a/lib/crates/fabro-workflow/src/billing_rollup.rs b/lib/crates/fabro-workflow/src/billing_rollup.rs index a91b06061..a796258a3 100644 --- a/lib/crates/fabro-workflow/src/billing_rollup.rs +++ b/lib/crates/fabro-workflow/src/billing_rollup.rs @@ -1,6 +1,30 @@ +use std::borrow::Cow; use std::collections::HashMap; -use fabro_types::{BilledTokenCounts, ModelRef, RunProjection}; +use fabro_model::Catalog; +use fabro_types::{BilledTokenCounts, ModelRef, RunProjection, StageProjection}; + +fn stage_usage_with_cost<'a>( + catalog: Option<&Catalog>, + stage: &'a StageProjection, +) -> Cow<'a, BilledTokenCounts> { + let Some(catalog) = catalog else { + return Cow::Borrowed(&stage.usage); + }; + let Some(model) = stage.model.as_ref() else { + return Cow::Borrowed(&stage.usage); + }; + if stage.usage.total_usd_micros.is_some() { + return Cow::Borrowed(&stage.usage); + } + + let Some(total_usd_micros) = catalog.price_tokens(model, &stage.usage.token_counts()) else { + return Cow::Borrowed(&stage.usage); + }; + let mut usage = stage.usage.clone(); + usage.total_usd_micros = Some(total_usd_micros); + Cow::Owned(usage) +} #[derive(Debug, Clone, PartialEq)] pub struct ProjectionBillingStage { @@ -34,7 +58,10 @@ impl ProjectionBillingRollup { } #[must_use] -pub fn billing_rollup_from_projection(projection: &RunProjection) -> ProjectionBillingRollup { +pub fn billing_rollup_from_projection( + projection: &RunProjection, + catalog: Option<&Catalog>, +) -> ProjectionBillingRollup { let mut stage_indices = HashMap::::new(); let mut stages = Vec::::new(); let mut by_model = HashMap::::new(); @@ -46,7 +73,9 @@ pub fn billing_rollup_from_projection(projection: &RunProjection) -> ProjectionB if is_boundary_stage(projection, stage_id.node_id()) { continue; } - if stage.completion.is_none() && stage.duration_ms.is_none() && stage.usage.is_zero() { + let usage = stage_usage_with_cost(catalog, stage); + let usage = usage.as_ref(); + if stage.completion.is_none() && stage.duration_ms.is_none() && usage.is_zero() { continue; } @@ -68,10 +97,10 @@ pub fn billing_rollup_from_projection(projection: &RunProjection) -> ProjectionB runtime_ms = runtime_ms.saturating_add(duration_ms); } - if !stage.usage.is_zero() { + if !usage.is_zero() { billed_visit_count += 1; - row.billing.add_counts(&stage.usage); - totals.add_counts(&stage.usage); + row.billing.add_counts(usage); + totals.add_counts(usage); if let Some(model) = &stage.model { row.model = Some(model.clone()); @@ -84,7 +113,7 @@ pub fn billing_rollup_from_projection(projection: &RunProjection) -> ProjectionB billing: BilledTokenCounts::default(), }); model_entry.stages += 1; - model_entry.billing.add_counts(&stage.usage); + model_entry.billing.add_counts(usage); } } } @@ -126,6 +155,7 @@ fn is_boundary_stage(projection: &RunProjection, node_id: &str) -> bool { mod tests { use std::collections::HashMap; + use fabro_model::{Catalog, ModelRef, ProviderId}; use fabro_types::{ AttrValue, BilledModelUsage, BilledTokenCounts, Graph, Node, RunProjection, RunSpec, StageCompletion, StageOutcome, WorkflowSettings, first_event_seq, fixtures, @@ -190,7 +220,7 @@ mod tests { timestamp: chrono::Utc::now(), }); - let rollup = billing_rollup_from_projection(&projection); + let rollup = billing_rollup_from_projection(&projection, None); assert_eq!(rollup.stages.len(), 1); assert_eq!(rollup.stages[0].node_id, "verify"); @@ -233,7 +263,7 @@ mod tests { timestamp: chrono::Utc::now(), }); - let rollup = billing_rollup_from_projection(&projection); + let rollup = billing_rollup_from_projection(&projection, None); assert_eq!(rollup.stages.len(), 1); assert_eq!(rollup.stages[0].node_id, "build"); @@ -266,12 +296,48 @@ mod tests { timestamp: chrono::Utc::now(), }); - let rollup = billing_rollup_from_projection(&projection); + let rollup = billing_rollup_from_projection(&projection, None); assert_eq!(rollup.stages.len(), 0); assert_eq!(rollup.runtime_ms, 0); } + #[test] + fn rollup_prices_in_flight_stage_usage_using_catalog() { + let mut projection = test_projection(); + let model = ModelRef { + provider: ProviderId::openai(), + model_id: "gpt-5.4".to_string(), + speed: None, + }; + let stage = projection.stage_entry("agent", 1, first_event_seq(1)); + stage.started_at = Some(chrono::Utc::now()); + stage.usage = BilledTokenCounts { + input_tokens: 500_000, + output_tokens: 125_000, + total_tokens: 625_000, + ..BilledTokenCounts::default() + }; + stage.model = Some(model.clone()); + + let priced = billing_rollup_from_projection(&projection, Some(Catalog::builtin())); + let unpriced = billing_rollup_from_projection(&projection, None); + + assert_eq!(priced.stages.len(), 1); + assert_eq!(priced.stages[0].node_id, "agent"); + let stage_cost = priced.stages[0].billing.total_usd_micros; + assert!( + stage_cost.is_some_and(|cost| cost > 0), + "expected priced stage cost, got {stage_cost:?}" + ); + assert_eq!(priced.totals.total_usd_micros, stage_cost); + assert_eq!(priced.by_model.len(), 1); + assert_eq!(priced.by_model[0].billing.total_usd_micros, stage_cost); + assert_eq!(unpriced.stages.len(), 1); + assert_eq!(unpriced.stages[0].billing.total_usd_micros, None); + assert_eq!(unpriced.totals.total_usd_micros, None); + } + fn run_spec_with_boundary_nodes() -> RunSpec { let mut graph = Graph::new("test"); graph.nodes.insert("start".to_string(), { diff --git a/lib/crates/fabro-workflow/src/pipeline/finalize.rs b/lib/crates/fabro-workflow/src/pipeline/finalize.rs index 5345bfb4a..a7f4619db 100644 --- a/lib/crates/fabro-workflow/src/pipeline/finalize.rs +++ b/lib/crates/fabro-workflow/src/pipeline/finalize.rs @@ -82,7 +82,7 @@ pub(crate) async fn build_conclusion_from_store( .unwrap_or_default(); let projection_billing = projection .as_ref() - .map(billing_rollup_from_projection) + .map(|projection| billing_rollup_from_projection(projection, None)) .unwrap_or_default(); let checkpoint = projection .as_ref() @@ -438,7 +438,7 @@ async fn compute_final_patch( } pub(crate) fn billing_from_projection(projection: &RunProjection) -> Option { - billing_rollup_from_projection(projection).billing_if_present() + billing_rollup_from_projection(projection, None).billing_if_present() } pub(crate) fn build_terminal_event( @@ -554,7 +554,7 @@ pub async fn finalize(executed: Executed, options: &FinalizeOptions) -> Result