diff --git a/docs/public/api-reference/fabro-api.yaml b/docs/public/api-reference/fabro-api.yaml index 46b1b81cf..a767a7262 100644 --- a/docs/public/api-reference/fabro-api.yaml +++ b/docs/public/api-reference/fabro-api.yaml @@ -6160,7 +6160,7 @@ components: example: 3501.0 BillingStageRef: - description: Reference to a billing stage. + description: Reference to a workflow node in a billing stage row. type: object required: - id @@ -6616,7 +6616,7 @@ components: # ── Billing Schemas ────────────────────────────────────────────────── RunBillingStage: - description: Token counts and billed totals for a single stage within a run. + description: Token counts and billed totals for one workflow node within a run. Rows are grouped by node; billing and runtime sum every visit of that node. type: object required: - stage @@ -6627,7 +6627,7 @@ components: stage: $ref: "#/components/schemas/BillingStageRef" model: - description: Model used for this stage; null for non-LLM stages. + description: Latest usage-bearing visit model for this node; null when no visit used an LLM model. oneOf: - $ref: "#/components/schemas/ModelReference" - type: "null" @@ -6635,7 +6635,7 @@ components: $ref: "#/components/schemas/BilledTokenCounts" runtime_secs: type: number - description: Wall-clock runtime in seconds. + description: Wall-clock runtime in seconds, summed across every visit of this node. example: 154.0 RunBillingTotals: @@ -6696,7 +6696,7 @@ components: $ref: "#/components/schemas/ModelReference" stages: type: integer - description: Number of stages that used this model. + description: Number of usage-bearing stage visits that used this model. example: 2 billing: $ref: "#/components/schemas/BilledTokenCounts" @@ -6711,7 +6711,7 @@ components: properties: stages: type: array - description: Per-stage billing breakdown. + description: Per-node billing breakdown. Each row sums billing and runtime across all visits of that node. items: $ref: "#/components/schemas/RunBillingStage" totals: diff --git a/lib/crates/fabro-server/src/demo/mod.rs b/lib/crates/fabro-server/src/demo/mod.rs index 127589a0e..bc527007d 100644 --- a/lib/crates/fabro-server/src/demo/mod.rs +++ b/lib/crates/fabro-server/src/demo/mod.rs @@ -784,9 +784,10 @@ mod runs { RunNamespace, RunPrepareSettings, RunSandboxSettings, }; use fabro_types::settings::{InterpString, ProjectNamespace, WorkflowNamespace}; - use fabro_types::{RunId, WorkflowSettings}; + use fabro_types::{RunId, StageId, WorkflowSettings}; use super::ts; + use crate::server::run_stage_from_stage_id; fn labels(entries: &[(&str, &str)]) -> HashMap { entries @@ -1180,50 +1181,37 @@ mod runs { } pub(super) fn stages() -> Vec { - fn visit(n: u32) -> std::num::NonZeroU32 { - std::num::NonZeroU32::new(n).expect("visit is 1-based") - } vec![ - RunStage { - id: "detect-drift@1".into(), - name: "Detect Drift".into(), - status: StageState::Succeeded, - duration_secs: Some(72.0), - node_id: "detect".into(), - visit: visit(1), - }, - RunStage { - id: "propose-changes@1".into(), - name: "Propose Changes".into(), - status: StageState::Succeeded, - duration_secs: Some(154.0), - node_id: "propose".into(), - visit: visit(1), - }, - RunStage { - id: "review-changes@1".into(), - name: "Review Changes".into(), - status: StageState::Succeeded, - duration_secs: Some(45.0), - node_id: "review".into(), - visit: visit(1), - }, - RunStage { - id: "apply-changes@1".into(), - name: "Apply Changes".into(), - status: StageState::Succeeded, - duration_secs: Some(118.0), - node_id: "apply".into(), - visit: visit(1), - }, - RunStage { - id: "apply-changes@2".into(), - name: "Apply Changes".into(), - status: StageState::Running, - duration_secs: None, - node_id: "apply".into(), - visit: visit(2), - }, + run_stage_from_stage_id( + &StageId::new("detect-drift", 1), + "Detect Drift", + StageState::Succeeded, + Some(72.0), + ), + run_stage_from_stage_id( + &StageId::new("propose-changes", 1), + "Propose Changes", + StageState::Succeeded, + Some(154.0), + ), + run_stage_from_stage_id( + &StageId::new("review-changes", 1), + "Review Changes", + StageState::Succeeded, + Some(45.0), + ), + run_stage_from_stage_id( + &StageId::new("apply-changes", 1), + "Apply Changes", + StageState::Succeeded, + Some(118.0), + ), + run_stage_from_stage_id( + &StageId::new("apply-changes", 2), + "Apply Changes", + StageState::Running, + None, + ), ] } diff --git a/lib/crates/fabro-server/src/server.rs b/lib/crates/fabro-server/src/server.rs index 9fdee7e5f..26fbdd8d6 100644 --- a/lib/crates/fabro-server/src/server.rs +++ b/lib/crates/fabro-server/src/server.rs @@ -54,7 +54,7 @@ use fabro_llm::types::{ ContentPart, FinishReason, Message as LlmMessage, Request as LlmRequest, Role, ToolChoice, ToolDefinition, }; -use fabro_model::{BilledModelUsage, BilledTokenCounts, Catalog, ModelTestMode, Provider}; +use fabro_model::{BilledTokenCounts, Catalog, ModelTestMode, Provider}; use fabro_redact::redact_jsonl_line; use fabro_sandbox::daytona::{self, DaytonaSandbox}; use fabro_sandbox::reconnect::reconnect; @@ -536,17 +536,48 @@ pub(crate) struct ResolvedAppStateSettings { pub(crate) manifest_run_settings: std::result::Result, } -fn accumulate_model_billing(entry: &mut ModelBillingTotals, usage: &BilledModelUsage) { - let tokens = usage.tokens(); - entry.stages += 1; - entry.billing.input_tokens += tokens.input_tokens; - entry.billing.output_tokens += tokens.output_tokens; - entry.billing.reasoning_tokens += tokens.reasoning_tokens; - entry.billing.cache_read_tokens += tokens.cache_read_tokens; - entry.billing.cache_write_tokens += tokens.cache_write_tokens; - entry.billing.total_tokens += tokens.total_tokens(); - if let Some(value) = usage.total_usd_micros { - *entry.billing.total_usd_micros.get_or_insert(0) += value; +fn accumulate_billed_token_counts(target: &mut BilledTokenCounts, source: &BilledTokenCounts) { + target.input_tokens += source.input_tokens; + target.output_tokens += source.output_tokens; + target.reasoning_tokens += source.reasoning_tokens; + target.cache_read_tokens += source.cache_read_tokens; + target.cache_write_tokens += source.cache_write_tokens; + target.total_tokens += source.total_tokens; + if let Some(value) = source.total_usd_micros { + *target.total_usd_micros.get_or_insert(0) += value; + } +} + +fn accumulate_billing_rollup( + accumulator: &mut BillingAccumulator, + rollup: &fabro_workflow::ProjectionBillingRollup, +) { + accumulator.total_runs += 1; + accumulator.total_runtime_secs += rollup.runtime_ms as f64 / 1000.0; + for model in &rollup.by_model { + let entry = accumulator + .by_model + .entry(model.model_id.clone()) + .or_default(); + entry.stages += model.stages; + accumulate_billed_token_counts(&mut entry.billing, &model.billing); + } +} + +pub(crate) fn run_stage_from_stage_id( + stage_id: &StageId, + name: impl Into, + status: StageState, + duration_secs: Option, +) -> RunStage { + RunStage { + id: stage_id.to_string(), + name: name.into(), + status, + duration_secs, + node_id: stage_id.node_id().to_string(), + visit: std::num::NonZeroU32::new(stage_id.visit()) + .expect("StageId stores a non-zero visit"), } } @@ -2776,9 +2807,9 @@ async fn execute_run_in_process(state: Arc, run_id: RunId) { } } - // Save final checkpoint - let checkpoint = match run_store.state().await { - Ok(state) => state.checkpoint, + // Save final projection + let final_projection = match run_store.state().await { + Ok(state) => Some(state), Err(err) => { tracing::warn!(run_id = %run_id, error = %err, "Failed to load run state from store"); None @@ -2786,32 +2817,17 @@ async fn execute_run_in_process(state: Arc, run_id: RunId) { }; // Accumulate aggregate usage after execution completes. - if let Some(ref cp) = checkpoint { - let stage_durations = match run_store.list_events().await { - Ok(events) => fabro_workflow::total_stage_duration_by_node(&events), - Err(err) => { - tracing::warn!(run_id = %run_id, error = %err, "Failed to load run events from store"); - HashMap::default() - } - }; - let mut agg = state - .aggregate_billing - .lock() - .expect("aggregate_billing lock poisoned"); - agg.total_runs += 1; - let mut run_runtime: f64 = 0.0; - for (node_id, outcome) in &cp.node_outcomes { - if let Some(usage) = &outcome.usage { - let entry = agg - .by_model - .entry(usage.model_id().to_string()) - .or_default(); - accumulate_model_billing(entry, usage); - } - let duration_ms = stage_durations.get(node_id).copied().unwrap_or(0); - run_runtime += duration_ms as f64 / 1000.0; + if let Some(ref projection) = final_projection { + if projection.checkpoint.is_some() { + let mut agg = state + .aggregate_billing + .lock() + .expect("aggregate_billing lock poisoned"); + accumulate_billing_rollup( + &mut agg, + &fabro_workflow::billing_rollup_from_projection(projection), + ); } - agg.total_runtime_secs += run_runtime; } let mut runs = state.runs.lock().expect("runs lock poisoned"); @@ -2860,7 +2876,9 @@ async fn execute_run_in_process(state: Arc, run_id: RunId) { }; } } - managed_run.checkpoint = checkpoint; + managed_run.checkpoint = final_projection + .as_ref() + .and_then(|projection| projection.checkpoint.clone()); managed_run.run_dir = Some(run_dir); clear_live_run_state(managed_run); } @@ -3103,32 +3121,15 @@ async fn execute_run_subprocess(state: Arc, run_id: RunId) { } }; - if let Some(ref checkpoint) = final_state.checkpoint { - let stage_durations = match run_store.list_events().await { - Ok(events) => fabro_workflow::total_stage_duration_by_node(&events), - Err(err) => { - tracing::warn!(run_id = %run_id, error = %err, "Failed to load run events from store"); - HashMap::default() - } - }; + if final_state.checkpoint.is_some() { let mut agg = state .aggregate_billing .lock() .expect("aggregate_billing lock poisoned"); - agg.total_runs += 1; - let mut run_runtime: f64 = 0.0; - for (node_id, outcome) in &checkpoint.node_outcomes { - if let Some(usage) = &outcome.usage { - let entry = agg - .by_model - .entry(usage.model_id().to_string()) - .or_default(); - accumulate_model_billing(entry, usage); - } - let duration_ms = stage_durations.get(node_id).copied().unwrap_or(0); - run_runtime += duration_ms as f64 / 1000.0; - } - agg.total_runtime_secs += run_runtime; + accumulate_billing_rollup( + &mut agg, + &fabro_workflow::billing_rollup_from_projection(&final_state), + ); } let mut runs = state.runs.lock().expect("runs lock poisoned"); diff --git a/lib/crates/fabro-server/src/server/handler/billing.rs b/lib/crates/fabro-server/src/server/handler/billing.rs index 8482bfce0..193ebd78d 100644 --- a/lib/crates/fabro-server/src/server/handler/billing.rs +++ b/lib/crates/fabro-server/src/server/handler/billing.rs @@ -1,15 +1,13 @@ -use std::collections::HashSet; -use std::num::NonZeroU32; use std::sync::Arc; use fabro_store::RunProjectionReducer; use fabro_types::{EventBody, RunProjection, StageId}; use super::super::{ - ApiError, AppState, BilledTokenCounts, BillingByModel, BillingStageRef, EventEnvelope, HashMap, - IntoResponse, Json, ListResponse, ModelBillingTotals, ModelReference, PaginationParams, Path, - Query, RequiredUser, Response, Router, RunBilling, RunBillingStage, RunBillingTotals, RunId, - RunStage, StageState, State, StatusCode, accumulate_model_billing, get, parse_run_id_path, + ApiError, AppState, BillingByModel, BillingStageRef, EventEnvelope, HashMap, IntoResponse, + Json, ListResponse, ModelReference, PaginationParams, Path, Query, RequiredUser, Response, + Router, RunBilling, RunBillingStage, RunBillingTotals, RunId, StageState, State, StatusCode, + get, parse_run_id_path, run_stage_from_stage_id, }; pub(super) fn routes() -> Router> { @@ -85,15 +83,6 @@ async fn list_run_stages( let mut stages = Vec::new(); for (stage_id, stage_projection) in projection.iter_stages() { - let node_id = stage_id.node_id().to_string(); - let Some(visit) = NonZeroU32::new(stage_id.visit()) else { - tracing::warn!( - run_id = %id, - stage_id = %stage_id, - "Skipping stage with non-positive visit", - ); - continue; - }; // Prefer the latest lifecycle event; fall back to the projection's // stored completion (e.g. for runs recovered from snapshot only). let status = lifecycle_states.get(stage_id).copied().unwrap_or_else(|| { @@ -102,14 +91,12 @@ async fn list_run_stages( .as_ref() .map_or(StageState::Pending, |c| StageState::from(c.outcome)) }); - stages.push(RunStage { - id: stage_id.to_string(), - name: node_id.clone(), + stages.push(run_stage_from_stage_id( + stage_id, + stage_id.node_id().to_string(), status, - duration_secs: stage_durations.get(stage_id).map(|ms| *ms as f64 / 1000.0), - node_id, - visit, - }); + stage_durations.get(stage_id).map(|ms| *ms as f64 / 1000.0), + )); } (StatusCode::OK, Json(ListResponse::new(stages))).into_response() @@ -127,100 +114,39 @@ async fn get_run_billing( } }; - let checkpoint = match run_store.state().await { - Ok(state) => state.checkpoint, + let projection = match run_store.state().await { + Ok(state) => state, Err(err) => { return ApiError::new(StatusCode::INTERNAL_SERVER_ERROR, err.to_string()) .into_response(); } }; - - let Some(checkpoint) = checkpoint else { - let empty = RunBilling { - by_model: Vec::new(), - stages: Vec::new(), - totals: RunBillingTotals { - cache_read_tokens: 0, - cache_write_tokens: 0, - input_tokens: 0, - output_tokens: 0, - reasoning_tokens: 0, - runtime_secs: 0.0, - total_tokens: 0, - total_usd_micros: None, + let rollup = fabro_workflow::billing_rollup_from_projection(&projection); + let by_model = rollup + .by_model + .iter() + .map(|model| BillingByModel { + billing: model.billing.clone(), + model: ModelReference { + id: model.model_id.clone(), }, - }; - return (StatusCode::OK, Json(empty)).into_response(); - }; - - let stage_durations = match run_store.list_events().await { - Ok(events) => fabro_workflow::total_stage_duration_by_node(&events), - Err(err) => { - return ApiError::new(StatusCode::INTERNAL_SERVER_ERROR, err.to_string()) - .into_response(); - } - }; - - let mut by_model_totals = HashMap::::new(); - let mut billed_usages = Vec::new(); - let mut runtime_secs = 0.0_f64; - let mut stages = Vec::new(); - - // `completed_nodes` records every visit (one entry per re-entry of a - // looped node), but billing is per-node: the duration helper already sums - // across visits, and `node_outcomes` only stores the latest visit's usage. - // Dedup so we emit one row per node and don't multiply the sum by visit - // count. - let mut seen_nodes = HashSet::new(); - for node_id in &checkpoint.completed_nodes { - if !seen_nodes.insert(node_id.as_str()) { - continue; - } - let duration_ms = stage_durations.get(node_id).copied().unwrap_or(0); - runtime_secs += duration_ms as f64 / 1000.0; - - let usage = checkpoint - .node_outcomes - .get(node_id) - .and_then(|outcome| outcome.usage.as_ref()); - - let (billing, model) = if let Some(usage) = usage { - billed_usages.push(usage.clone()); - let tokens = usage.tokens(); - let billing = BilledTokenCounts { - cache_read_tokens: tokens.cache_read_tokens, - cache_write_tokens: tokens.cache_write_tokens, - input_tokens: tokens.input_tokens, - output_tokens: tokens.output_tokens, - reasoning_tokens: tokens.reasoning_tokens, - total_tokens: tokens.total_tokens(), - total_usd_micros: usage.total_usd_micros, - }; - let model_id = usage.model_id().to_string(); - accumulate_model_billing(by_model_totals.entry(model_id.clone()).or_default(), usage); - (billing, Some(ModelReference { id: model_id })) - } else { - (BilledTokenCounts::default(), None) - }; - - stages.push(RunBillingStage { - billing, - model, - runtime_secs: duration_ms as f64 / 1000.0, - stage: BillingStageRef { - id: node_id.clone(), - name: node_id.clone(), + stages: model.stages, + }) + .collect::>(); + let stages = rollup + .stages + .iter() + .map(|stage| RunBillingStage { + billing: stage.billing.clone(), + model: stage + .model_id + .as_ref() + .map(|id| ModelReference { id: id.clone() }), + runtime_secs: stage.duration_ms as f64 / 1000.0, + stage: BillingStageRef { + id: stage.node_id.clone(), + name: stage.node_id.clone(), }, - }); - } - - let totals = BilledTokenCounts::from_billed_usage(&billed_usages); - let by_model = by_model_totals - .into_iter() - .map(|(model, totals)| BillingByModel { - billing: totals.billing, - model: ModelReference { id: model }, - stages: totals.stages, }) .collect::>(); @@ -228,14 +154,14 @@ async fn get_run_billing( by_model, stages, totals: RunBillingTotals { - cache_read_tokens: totals.cache_read_tokens, - cache_write_tokens: totals.cache_write_tokens, - input_tokens: totals.input_tokens, - output_tokens: totals.output_tokens, - reasoning_tokens: totals.reasoning_tokens, - runtime_secs, - total_tokens: totals.total_tokens, - total_usd_micros: totals.total_usd_micros, + cache_read_tokens: rollup.totals.cache_read_tokens, + cache_write_tokens: rollup.totals.cache_write_tokens, + input_tokens: rollup.totals.input_tokens, + output_tokens: rollup.totals.output_tokens, + reasoning_tokens: rollup.totals.reasoning_tokens, + runtime_secs: rollup.runtime_ms as f64 / 1000.0, + total_tokens: rollup.totals.total_tokens, + total_usd_micros: rollup.totals.total_usd_micros, }, }; diff --git a/lib/crates/fabro-server/src/server/tests.rs b/lib/crates/fabro-server/src/server/tests.rs index 12984c3a3..08196dbdd 100644 --- a/lib/crates/fabro-server/src/server/tests.rs +++ b/lib/crates/fabro-server/src/server/tests.rs @@ -2320,6 +2320,32 @@ fn stage_entry<'a>(body: &'a serde_json::Value, id: &str) -> &'a serde_json::Val .unwrap_or_else(|| panic!("stage {id} not found in {body:#?}")) } +fn test_billed_usage( + model_id: &str, + input_tokens: i64, + output_tokens: i64, +) -> fabro_model::BilledModelUsage { + serde_json::from_value(json!({ + "input": { + "usage": { + "model": { + "provider": "openai", + "model_id": model_id + }, + "tokens": { + "input_tokens": input_tokens, + "output_tokens": output_tokens + } + }, + "facts": { + "provider": "open_ai" + } + }, + "total_usd_micros": input_tokens + output_tokens + })) + .unwrap() +} + #[tokio::test] async fn list_run_stages_distinguishes_visits() { let state = test_app_state_with_isolated_storage(); @@ -2573,6 +2599,142 @@ async fn run_billing_dedups_retried_nodes_and_sums_their_durations() { ); } +#[tokio::test] +async fn run_billing_sums_usage_across_retry_visits_and_uses_latest_model() { + let state = test_app_state_with_isolated_storage(); + let app = crate::test_support::build_test_router(Arc::clone(&state)); + let run_id = RunId::new(); + let failed_usage = test_billed_usage("gpt-old", 100, 10); + let success_usage = test_billed_usage("gpt-new", 200, 20); + + create_durable_run_with_events(&state, run_id, &[ + workflow_event::Event::RunSubmitted { + definition_blob: None, + }, + workflow_event::Event::RunStarting, + workflow_event::Event::RunRunning, + ]) + .await; + + append_scoped_stage_event( + &state, + run_id, + "verify", + 1, + &workflow_event::Event::StageFailed { + node_id: "verify".to_string(), + name: "Verify".to_string(), + index: 1, + failure: FailureDetail::new("try again", FailureCategory::TransientInfra), + will_retry: true, + duration_ms: 1200, + billing: Some(failed_usage), + actor: None, + }, + ) + .await; + append_scoped_stage_event( + &state, + run_id, + "verify", + 2, + &workflow_event::Event::StageCompleted { + node_id: "verify".to_string(), + name: "Verify".to_string(), + index: 1, + duration_ms: 800, + status: "succeeded".to_string(), + preferred_label: None, + suggested_next_ids: Vec::new(), + billing: Some(success_usage.clone()), + failure: None, + notes: None, + files_touched: Vec::new(), + context_updates: None, + jump_to_node: None, + context_values: None, + node_visits: None, + loop_failure_signatures: None, + restart_failure_signatures: None, + response: None, + attempt: 2, + max_attempts: 2, + }, + ) + .await; + + let mut latest_outcome: Outcome> = Outcome::success(); + latest_outcome.usage = Some(success_usage); + latest_outcome.duration_ms = Some(800); + let run_store = state.store.open_run(&run_id).await.unwrap(); + workflow_event::append_event( + &run_store, + &run_id, + &workflow_event::Event::CheckpointCompleted { + node_id: "verify".to_string(), + status: "running".to_string(), + current_node: "verify".to_string(), + completed_nodes: vec!["verify".to_string(), "verify".to_string()], + node_retries: std::collections::BTreeMap::from([("verify".to_string(), 2)]), + context_values: std::collections::BTreeMap::new(), + node_outcomes: std::collections::BTreeMap::from([( + "verify".to_string(), + latest_outcome, + )]), + next_node_id: None, + git_commit_sha: None, + loop_failure_signatures: std::collections::BTreeMap::new(), + restart_failure_signatures: std::collections::BTreeMap::new(), + node_visits: std::collections::BTreeMap::from([("verify".to_string(), 2usize)]), + diff: None, + }, + ) + .await + .unwrap(); + + let response = app + .clone() + .oneshot( + Request::builder() + .method("GET") + .uri(api(&format!("/runs/{run_id}/billing"))) + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + let body = response_json!(response, StatusCode::OK).await; + + let stages = body["stages"].as_array().unwrap(); + assert_eq!(stages.len(), 1); + assert_eq!(stages[0]["stage"]["id"], "verify"); + assert_eq!(stages[0]["model"]["id"], "gpt-new"); + assert_eq!(stages[0]["billing"]["input_tokens"], 300); + assert_eq!(stages[0]["billing"]["output_tokens"], 30); + assert_eq!(stages[0]["billing"]["total_usd_micros"], 330); + assert!((stages[0]["runtime_secs"].as_f64().unwrap() - 2.0).abs() < f64::EPSILON); + + assert_eq!(body["totals"]["input_tokens"], 300); + assert_eq!(body["totals"]["output_tokens"], 30); + assert_eq!(body["totals"]["total_usd_micros"], 330); + assert!((body["totals"]["runtime_secs"].as_f64().unwrap() - 2.0).abs() < f64::EPSILON); + + let by_model = body["by_model"].as_array().unwrap(); + assert_eq!(by_model.len(), 2); + let old_model = by_model + .iter() + .find(|entry| entry["model"]["id"] == "gpt-old") + .unwrap(); + let new_model = by_model + .iter() + .find(|entry| entry["model"]["id"] == "gpt-new") + .unwrap(); + assert_eq!(old_model["stages"], 1); + assert_eq!(old_model["billing"]["input_tokens"], 100); + assert_eq!(new_model["stages"], 1); + assert_eq!(new_model["billing"]["input_tokens"], 200); +} + #[tokio::test] async fn list_run_stages_shows_retrying_after_failed_event() { let state = test_app_state_with_isolated_storage(); @@ -6282,6 +6444,62 @@ async fn get_aggregate_billing_returns_zeros_initially() { assert!(body["by_model"].as_array().unwrap().is_empty()); } +#[test] +fn aggregate_billing_counts_projection_rollup_usage_visits() { + let mut accumulator = BillingAccumulator::default(); + let rollup = fabro_workflow::ProjectionBillingRollup { + stages: Vec::new(), + totals: BilledTokenCounts { + input_tokens: 300, + output_tokens: 30, + total_tokens: 330, + reasoning_tokens: 0, + cache_read_tokens: 0, + cache_write_tokens: 0, + total_usd_micros: Some(330), + }, + by_model: vec![ + fabro_workflow::ProjectionBillingByModel { + model_id: "gpt-old".to_string(), + stages: 1, + billing: BilledTokenCounts { + input_tokens: 100, + output_tokens: 10, + total_tokens: 110, + reasoning_tokens: 0, + cache_read_tokens: 0, + cache_write_tokens: 0, + total_usd_micros: Some(110), + }, + }, + fabro_workflow::ProjectionBillingByModel { + model_id: "gpt-new".to_string(), + stages: 1, + billing: BilledTokenCounts { + input_tokens: 200, + output_tokens: 20, + total_tokens: 220, + reasoning_tokens: 0, + cache_read_tokens: 0, + cache_write_tokens: 0, + total_usd_micros: Some(220), + }, + }, + ], + runtime_ms: 2000, + billed_visit_count: 2, + }; + + accumulate_billing_rollup(&mut accumulator, &rollup); + + assert_eq!(accumulator.total_runs, 1); + assert_eq!(accumulator.total_runtime_secs, 2.0); + assert_eq!(accumulator.by_model["gpt-old"].stages, 1); + assert_eq!(accumulator.by_model["gpt-old"].billing.input_tokens, 100); + assert_eq!(accumulator.by_model["gpt-new"].stages, 1); + assert_eq!(accumulator.by_model["gpt-new"].billing.input_tokens, 200); +} + #[tokio::test] async fn post_runs_returns_submitted_status() { let state = test_app_state(); diff --git a/lib/crates/fabro-store/src/artifact_store.rs b/lib/crates/fabro-store/src/artifact_store.rs index 4cd4f85e1..5be00d792 100644 --- a/lib/crates/fabro-store/src/artifact_store.rs +++ b/lib/crates/fabro-store/src/artifact_store.rs @@ -297,8 +297,13 @@ fn decode_artifact_location( )) })?; let (retry, filename) = decode_retry_and_filename(location, &mut parts)?; + let stage_id = StageId::try_new(node_id, visit).map_err(|err| { + Error::Other(format!( + "artifact location {location} has an invalid stage id: {err}" + )) + })?; Ok(NodeArtifact { - node: StageId::new(node_id, visit), + node: stage_id, retry, filename, size, diff --git a/lib/crates/fabro-store/src/run_state.rs b/lib/crates/fabro-store/src/run_state.rs index cd19243d4..af467f557 100644 --- a/lib/crates/fabro-store/src/run_state.rs +++ b/lib/crates/fabro-store/src/run_state.rs @@ -9,8 +9,8 @@ use fabro_types::run_event::{ use fabro_types::{ BilledModelUsage, Checkpoint, Conclusion, EventBody, FailureSignature, InterviewQuestionRecord, Outcome, PendingInterviewRecord, PullRequestRecord, RunControlAction, RunEvent, RunId, - RunProjection, RunSpec, RunStatus, RunSummary, SandboxRecord, StageCompletion, StageOutcome, - StageProjection, StartRecord, TerminalStatus, first_event_seq, + RunProjection, RunSpec, RunStatus, RunSummary, SandboxRecord, StageCompletion, StageId, + StageOutcome, StageProjection, StartRecord, TerminalStatus, first_event_seq, }; use fabro_util::error::render_with_causes; use serde_json::Value; @@ -297,27 +297,28 @@ impl RunProjectionReducer for RunProjection { ); } EventBody::StagePrompt(props) => { - let Some(stage) = stage_at_visit(self, stored, props.visit, event.seq) else { + let Some(stage) = stage_at_stored_or_visit(self, stored, props.visit, event.seq) + else { return Ok(()); }; stage.prompt = Some(props.text.clone()); stage.provider_used = provider_used_from_prompt(props); } EventBody::PromptCompleted(props) => { - let Some(stage) = stage_at_current_visit(self, stored, event.seq) else { + let Some(stage) = stage_at_stored_or_current_visit(self, stored, event.seq) else { return Ok(()); }; stage.response = Some(props.response.clone()); } EventBody::StageCompleted(props) => { - let Some(node_id) = stored.node_id.as_deref() else { - return Ok(()); - }; - let visit = stage_visit(node_id, props.node_visits.as_ref(), self).unwrap_or(1); let response = props.response.clone(); let outcome = stage_outcome_from_props(props); let completion = stage_completion_from_outcome(&outcome, ts); - let stage = self.stage_entry(node_id, visit, first_event_seq(event.seq)); + let Some(stage) = + stage_at_completed_visit(self, stored, props.node_visits.as_ref(), event.seq) + else { + return Ok(()); + }; stage.response = response; stage.completion = Some(completion); stage.duration_ms = Some(props.duration_ms); @@ -325,12 +326,12 @@ impl RunProjectionReducer for RunProjection { } EventBody::StageFailed(props) => { let failure_reason = props.failure.as_ref().map(|detail| detail.message.clone()); - let Some(stage) = stage_at_current_visit(self, stored, event.seq) else { + let Some(stage) = stage_at_stored_or_current_visit(self, stored, event.seq) else { return Ok(()); }; stage.completion = Some(StageCompletion { outcome: StageOutcome::Failed { - retry_requested: false, + retry_requested: props.will_retry, }, notes: None, failure_reason, @@ -340,13 +341,15 @@ impl RunProjectionReducer for RunProjection { stage.usage.clone_from(&props.billing); } EventBody::AgentSessionStarted(props) => { - let Some(stage) = stage_at_visit(self, stored, props.visit, event.seq) else { + let Some(stage) = stage_at_stored_or_visit(self, stored, props.visit, event.seq) + else { return Ok(()); }; stage.provider_used = Some(provider_used_from_agent_session_started(props)); } EventBody::AgentCliStarted(props) => { - let Some(stage) = stage_at_visit(self, stored, props.visit, event.seq) else { + let Some(stage) = stage_at_stored_or_visit(self, stored, props.visit, event.seq) + else { return Ok(()); }; stage.provider_used = Some(provider_used_from_agent_cli_started(props)); @@ -355,7 +358,7 @@ impl RunProjectionReducer for RunProjection { let script_invocation = serde_json::to_value(props).map_err(|err| { Error::InvalidEvent(format!("invalid command.started payload: {err}")) })?; - let Some(stage) = stage_at_current_visit(self, stored, event.seq) else { + let Some(stage) = stage_at_stored_or_current_visit(self, stored, event.seq) else { return Ok(()); }; stage.script_invocation = Some(script_invocation); @@ -364,7 +367,7 @@ impl RunProjectionReducer for RunProjection { let script_timing = serde_json::to_value(props).map_err(|err| { Error::InvalidEvent(format!("invalid command.completed payload: {err}")) })?; - let Some(stage) = stage_at_current_visit(self, stored, event.seq) else { + let Some(stage) = stage_at_stored_or_current_visit(self, stored, event.seq) else { return Ok(()); }; stage.stdout = Some(props.stdout.clone()); @@ -380,7 +383,7 @@ impl RunProjectionReducer for RunProjection { let parallel_results = serde_json::to_value(&props.results).map_err(|err| { Error::InvalidEvent(format!("invalid parallel.completed payload: {err}")) })?; - let Some(stage) = stage_at_current_visit(self, stored, event.seq) else { + let Some(stage) = stage_at_stored_or_current_visit(self, stored, event.seq) else { return Ok(()); }; stage.parallel_results = Some(parallel_results); @@ -398,6 +401,9 @@ fn stage_at_visit<'a>( visit: u32, seq: u32, ) -> Option<&'a mut StageProjection> { + if visit == 0 { + return None; + } let node_id = stored.node_id.as_deref()?; Some(state.stage_entry(node_id, visit, first_event_seq(seq))) } @@ -412,6 +418,51 @@ fn stage_at_current_visit<'a>( Some(state.stage_entry(node_id, visit, first_event_seq(seq))) } +fn stage_at_stored_stage_id<'a>( + state: &'a mut RunProjection, + stage_id: &StageId, + seq: u32, +) -> &'a mut StageProjection { + state.stage_entry(stage_id.node_id(), stage_id.visit(), first_event_seq(seq)) +} + +fn stage_at_stored_or_visit<'a>( + state: &'a mut RunProjection, + stored: &RunEvent, + visit: u32, + seq: u32, +) -> Option<&'a mut StageProjection> { + if let Some(stage_id) = stored.stage_id.as_ref() { + return Some(stage_at_stored_stage_id(state, stage_id, seq)); + } + stage_at_visit(state, stored, visit, seq) +} + +fn stage_at_stored_or_current_visit<'a>( + state: &'a mut RunProjection, + stored: &RunEvent, + seq: u32, +) -> Option<&'a mut StageProjection> { + if let Some(stage_id) = stored.stage_id.as_ref() { + return Some(stage_at_stored_stage_id(state, stage_id, seq)); + } + stage_at_current_visit(state, stored, seq) +} + +fn stage_at_completed_visit<'a>( + state: &'a mut RunProjection, + stored: &RunEvent, + node_visits: Option<&BTreeMap>, + seq: u32, +) -> Option<&'a mut StageProjection> { + if let Some(stage_id) = stored.stage_id.as_ref() { + return Some(stage_at_stored_stage_id(state, stage_id, seq)); + } + let node_id = stored.node_id.as_deref()?; + let visit = stage_visit(node_id, node_visits, state).unwrap_or(1); + Some(state.stage_entry(node_id, visit, first_event_seq(seq))) +} + pub(crate) fn build_summary(state: &RunProjection, run_id: &RunId) -> RunSummary { let workflow_name = state.spec.as_ref().map(|spec| { if spec.graph.name.is_empty() { @@ -533,6 +584,7 @@ fn stage_visit( node_visits .and_then(|visits| visits.get(node_id).copied()) .and_then(|visit| u32::try_from(visit).ok()) + .filter(|visit| *visit > 0) .or_else(|| state.current_visit_for(node_id)) } @@ -617,7 +669,8 @@ mod tests { use fabro_types::run_event::run::RunFailedProps; use fabro_types::run_event::{ CheckpointCompletedProps, InterviewCompletedProps, InterviewOption, InterviewStartedProps, - RunControlEffectProps, StageCompletedProps, StagePromptProps, StageStartedProps, + RunControlEffectProps, StageCompletedProps, StageFailedProps, StagePromptProps, + StageStartedProps, }; use fabro_types::{ BilledModelUsage, BlockedReason, Checkpoint, EventBody, FailureReason, Outcome, @@ -1015,6 +1068,86 @@ mod tests { assert_eq!(second_stage.usage.as_ref(), Some(&second_usage)); } + #[test] + fn stage_completed_prefers_stored_stage_id_over_legacy_node_visits() { + let mut state = RunProjection::default(); + let usage = test_usage("gpt-5.2", 300, 30); + let scoped_stage_id = StageId::new("build", 2); + + state + .apply_event(&test_stage_event( + 3, + EventBody::StageCompleted(StageCompletedProps { + index: 0, + duration_ms: 333, + status: StageOutcome::Succeeded, + preferred_label: None, + suggested_next_ids: Vec::new(), + billing: Some(usage.clone()), + failure: None, + notes: None, + files_touched: Vec::new(), + context_updates: None, + jump_to_node: None, + context_values: None, + node_visits: Some(BTreeMap::from([("build".to_string(), 1usize)])), + loop_failure_signatures: None, + restart_failure_signatures: None, + response: Some("done".to_string()), + attempt: 1, + max_attempts: 1, + }), + scoped_stage_id.clone(), + )) + .unwrap(); + + assert!( + state.stage(&StageId::new("build", 1)).is_none(), + "legacy node_visits must not override stored stage_id" + ); + let stage = state.stage(&scoped_stage_id).unwrap(); + assert_eq!(stage.duration_ms, Some(333)); + assert_eq!(stage.usage.as_ref(), Some(&usage)); + assert_eq!(stage.response.as_deref(), Some("done")); + } + + #[test] + fn stage_failed_prefers_stored_stage_id_and_preserves_retry_request() { + let mut state = RunProjection::default(); + let usage = test_usage("gpt-5.2", 400, 40); + let scoped_stage_id = StageId::new("build", 2); + + state + .apply_event(&test_stage_event( + 3, + EventBody::StageFailed(StageFailedProps { + index: 0, + failure: Some(fabro_types::FailureDetail::new( + "try again", + fabro_types::FailureCategory::TransientInfra, + )), + will_retry: true, + duration_ms: 444, + billing: Some(usage.clone()), + }), + scoped_stage_id.clone(), + )) + .unwrap(); + + assert!( + state.stage(&StageId::new("build", 1)).is_none(), + "current-visit fallback must not override stored stage_id" + ); + let stage = state.stage(&scoped_stage_id).unwrap(); + assert_eq!(stage.duration_ms, Some(444)); + assert_eq!(stage.usage.as_ref(), Some(&usage)); + let completion = stage.completion.as_ref().unwrap(); + assert_eq!(completion.outcome, StageOutcome::Failed { + retry_requested: true, + }); + assert_eq!(completion.failure_reason.as_deref(), Some("try again")); + } + #[test] fn checkpoint_completed_creates_projection_entry_for_skipped_stage() { let mut state = RunProjection::default(); diff --git a/lib/crates/fabro-types/src/lib.rs b/lib/crates/fabro-types/src/lib.rs index c04702c1e..22675855c 100644 --- a/lib/crates/fabro-types/src/lib.rs +++ b/lib/crates/fabro-types/src/lib.rs @@ -80,7 +80,7 @@ pub use run_summary::RunSummary; pub use sandbox_record::SandboxRecord; pub use secret::{SecretMetadata, SecretType}; pub use stage_completion::StageCompletion; -pub use stage_id::{ParallelBranchId, StageId}; +pub use stage_id::{InvalidStageVisit, ParallelBranchId, StageId}; pub use start::StartRecord; pub use status::{ BlockedReason, FailureReason, InvalidTransition, ParseFailureReasonError, diff --git a/lib/crates/fabro-types/src/outcome.rs b/lib/crates/fabro-types/src/outcome.rs index af34cd10b..8ef0c67dc 100644 --- a/lib/crates/fabro-types/src/outcome.rs +++ b/lib/crates/fabro-types/src/outcome.rs @@ -138,12 +138,32 @@ impl From for StageState { match outcome { StageOutcome::Succeeded => Self::Succeeded, StageOutcome::PartiallySucceeded => Self::PartiallySucceeded, - StageOutcome::Failed { .. } => Self::Failed, + StageOutcome::Failed { + retry_requested: true, + } => Self::Retrying, + StageOutcome::Failed { + retry_requested: false, + } => Self::Failed, StageOutcome::Skipped => Self::Skipped, } } } +#[cfg(test)] +mod stage_state_tests { + use super::{StageOutcome, StageState}; + + #[test] + fn retry_requested_failure_projects_as_retrying() { + assert_eq!( + StageState::from(StageOutcome::Failed { + retry_requested: true, + }), + StageState::Retrying + ); + } +} + #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)] #[serde(rename_all = "snake_case")] pub enum FailureCategory { @@ -340,9 +360,10 @@ mod tests { StageState::from(StageOutcome::Failed { retry_requested: true, }), - StageState::Failed + StageState::Retrying ); assert!(StageState::Cancelled.is_terminal()); + assert!(!StageState::Retrying.is_terminal()); assert!(!StageState::Running.is_terminal()); } } diff --git a/lib/crates/fabro-types/src/run_projection.rs b/lib/crates/fabro-types/src/run_projection.rs index 712ceaa7f..87321d348 100644 --- a/lib/crates/fabro-types/src/run_projection.rs +++ b/lib/crates/fabro-types/src/run_projection.rs @@ -112,14 +112,24 @@ impl RunProjection { /// here once instead of asking each caller to remember. pub fn iter_stages(&self) -> impl Iterator { let mut entries: Vec<(&StageId, &StageProjection)> = self.stages.iter().collect(); - entries.sort_by_key(|(_, stage)| stage.first_event_seq); + entries.sort_by(|(left_id, left_stage), (right_id, right_stage)| { + left_stage + .first_event_seq + .cmp(&right_stage.first_event_seq) + .then_with(|| left_id.cmp(right_id)) + }); entries.into_iter() } /// Mutable counterpart of [`iter_stages`]. Same chronological ordering. pub fn iter_stages_mut(&mut self) -> impl Iterator { let mut entries: Vec<(&StageId, &mut StageProjection)> = self.stages.iter_mut().collect(); - entries.sort_by_key(|(_, stage)| stage.first_event_seq); + entries.sort_by(|(left_id, left_stage), (right_id, right_stage)| { + left_stage + .first_event_seq + .cmp(&right_stage.first_event_seq) + .then_with(|| left_id.cmp(right_id)) + }); entries.into_iter() } @@ -256,4 +266,36 @@ mod iter_stages_tests { .collect(); assert_eq!(order, vec!["a", "b", "c"]); } + + #[test] + fn iter_stages_tie_breaks_same_first_event_seq_by_stage_id() { + for _ in 0..128 { + let mut p = RunProjection::default(); + p.stage_entry("verify", 2, seq(10)); + p.stage_entry("build", 1, seq(10)); + p.stage_entry("verify", 1, seq(10)); + + let order: Vec = p + .iter_stages() + .map(|(stage_id, _)| stage_id.to_string()) + .collect(); + assert_eq!(order, vec!["build@1", "verify@1", "verify@2"]); + } + } + + #[test] + fn iter_stages_mut_tie_breaks_same_first_event_seq_by_stage_id() { + for _ in 0..128 { + let mut p = RunProjection::default(); + p.stage_entry("verify", 2, seq(10)); + p.stage_entry("build", 1, seq(10)); + p.stage_entry("verify", 1, seq(10)); + + let order: Vec = p + .iter_stages_mut() + .map(|(stage_id, _)| stage_id.to_string()) + .collect(); + assert_eq!(order, vec!["build@1", "verify@1", "verify@2"]); + } + } } diff --git a/lib/crates/fabro-types/src/stage_id.rs b/lib/crates/fabro-types/src/stage_id.rs index baae58864..91711cb79 100644 --- a/lib/crates/fabro-types/src/stage_id.rs +++ b/lib/crates/fabro-types/src/stage_id.rs @@ -1,4 +1,5 @@ use std::fmt; +use std::num::NonZeroU32; use std::str::FromStr; use serde::de::Error as _; @@ -7,16 +8,21 @@ use serde::{Deserialize, Deserializer, Serialize, Serializer}; #[derive(Clone, Debug, PartialEq, Eq, Hash, PartialOrd, Ord)] pub struct StageId { node_id: String, - visit: u32, + visit: NonZeroU32, } impl StageId { #[must_use] pub fn new(node_id: impl Into, visit: u32) -> Self { - Self { + Self::try_new(node_id, visit).expect("stage id visit must be greater than zero") + } + + pub fn try_new(node_id: impl Into, visit: u32) -> Result { + let visit = NonZeroU32::new(visit).ok_or(InvalidStageVisit)?; + Ok(Self { node_id: node_id.into(), visit, - } + }) } #[must_use] @@ -26,7 +32,7 @@ impl StageId { #[must_use] pub fn visit(&self) -> u32 { - self.visit + self.visit.get() } } @@ -47,6 +53,17 @@ impl fmt::Display for ParseStageIdError { impl std::error::Error for ParseStageIdError {} +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct InvalidStageVisit; + +impl fmt::Display for InvalidStageVisit { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str("stage id visit must be greater than zero") + } +} + +impl std::error::Error for InvalidStageVisit {} + impl FromStr for StageId { type Err = ParseStageIdError; @@ -67,7 +84,7 @@ impl FromStr for StageId { let visit = visit .parse() .map_err(|err| ParseStageIdError(format!("invalid stage id visit: {err}")))?; - Ok(Self::new(node_id, visit)) + Self::try_new(node_id, visit).map_err(|err| ParseStageIdError(err.to_string())) } } @@ -224,6 +241,18 @@ mod tests { assert!(err.to_string().starts_with("invalid stage id visit:")); } + #[test] + fn parse_rejects_zero_visit() { + let err = "code@0".parse::().unwrap_err(); + assert_eq!(err.to_string(), "stage id visit must be greater than zero"); + } + + #[test] + fn try_new_rejects_zero_visit() { + let err = StageId::try_new("code", 0).unwrap_err(); + assert_eq!(err.to_string(), "stage id visit must be greater than zero"); + } + #[test] fn parse_rejects_empty_node_id() { let err = "@3".parse::().unwrap_err(); diff --git a/lib/crates/fabro-workflow/src/billing_rollup.rs b/lib/crates/fabro-workflow/src/billing_rollup.rs new file mode 100644 index 000000000..5ee06da81 --- /dev/null +++ b/lib/crates/fabro-workflow/src/billing_rollup.rs @@ -0,0 +1,213 @@ +use std::collections::{BTreeMap, HashMap}; + +use fabro_types::{BilledModelUsage, BilledTokenCounts, RunProjection}; + +#[derive(Debug, Clone, PartialEq)] +pub struct ProjectionBillingStage { + pub node_id: String, + pub billing: BilledTokenCounts, + pub duration_ms: u64, + pub model_id: Option, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct ProjectionBillingByModel { + pub model_id: String, + pub stages: i64, + pub billing: BilledTokenCounts, +} + +#[derive(Debug, Clone, Default, PartialEq)] +pub struct ProjectionBillingRollup { + pub stages: Vec, + pub totals: BilledTokenCounts, + pub by_model: Vec, + pub runtime_ms: u64, + pub billed_visit_count: usize, +} + +impl ProjectionBillingRollup { + #[must_use] + pub fn billing_if_present(&self) -> Option { + (self.billed_visit_count > 0).then(|| self.totals.clone()) + } +} + +#[must_use] +pub fn billing_rollup_from_projection(projection: &RunProjection) -> ProjectionBillingRollup { + let mut stage_indices = HashMap::::new(); + let mut stages = Vec::::new(); + let mut by_model = BTreeMap::::new(); + let mut totals = BilledTokenCounts::default(); + let mut runtime_ms = 0_u64; + let mut billed_visit_count = 0_usize; + + for (stage_id, stage) in projection.iter_stages() { + if stage.completion.is_none() && stage.duration_ms.is_none() && stage.usage.is_none() { + continue; + } + + let node_id = stage_id.node_id(); + let index = *stage_indices.entry(node_id.to_string()).or_insert_with(|| { + let index = stages.len(); + stages.push(ProjectionBillingStage { + node_id: node_id.to_string(), + billing: BilledTokenCounts::default(), + duration_ms: 0, + model_id: None, + }); + index + }); + let row = &mut stages[index]; + + if let Some(duration_ms) = stage.duration_ms { + row.duration_ms = row.duration_ms.saturating_add(duration_ms); + runtime_ms = runtime_ms.saturating_add(duration_ms); + } + + if let Some(usage) = stage.usage.as_ref() { + billed_visit_count += 1; + row.model_id = Some(usage.model_id().to_string()); + accumulate_usage(&mut row.billing, usage); + accumulate_usage(&mut totals, usage); + + let model_id = usage.model_id().to_string(); + let model_entry = + by_model + .entry(model_id.clone()) + .or_insert_with(|| ProjectionBillingByModel { + model_id, + stages: 0, + billing: BilledTokenCounts::default(), + }); + model_entry.stages += 1; + accumulate_usage(&mut model_entry.billing, usage); + } + } + + ProjectionBillingRollup { + stages, + totals, + by_model: by_model.into_values().collect(), + runtime_ms, + billed_visit_count, + } +} + +fn accumulate_usage(counts: &mut BilledTokenCounts, usage: &BilledModelUsage) { + let tokens = usage.tokens(); + counts.input_tokens += tokens.input_tokens; + counts.output_tokens += tokens.output_tokens; + counts.reasoning_tokens += tokens.reasoning_tokens; + counts.cache_read_tokens += tokens.cache_read_tokens; + counts.cache_write_tokens += tokens.cache_write_tokens; + counts.total_tokens += tokens.total_tokens(); + if let Some(value) = usage.total_usd_micros { + *counts.total_usd_micros.get_or_insert(0) += value; + } +} + +#[cfg(test)] +mod tests { + use fabro_types::{BilledModelUsage, RunProjection, StageOutcome, first_event_seq}; + use serde_json::json; + + use super::billing_rollup_from_projection; + + fn test_usage(model_id: &str, input_tokens: i64, output_tokens: i64) -> BilledModelUsage { + serde_json::from_value(json!({ + "input": { + "usage": { + "model": { + "provider": "openai", + "model_id": model_id + }, + "tokens": { + "input_tokens": input_tokens, + "output_tokens": output_tokens + } + }, + "facts": { + "provider": "open_ai" + } + }, + "total_usd_micros": input_tokens + output_tokens + })) + .unwrap() + } + + #[test] + fn rollup_groups_stage_rows_by_node_and_sums_retry_visit_usage() { + let mut projection = RunProjection::default(); + let failed_usage = test_usage("gpt-old", 100, 10); + let success_usage = test_usage("gpt-new", 200, 20); + let first = projection.stage_entry("verify", 1, first_event_seq(1)); + first.duration_ms = Some(1200); + first.usage = Some(failed_usage); + first.completion = Some(fabro_types::StageCompletion { + outcome: StageOutcome::Failed { + retry_requested: true, + }, + notes: None, + failure_reason: Some("try again".to_string()), + timestamp: chrono::Utc::now(), + }); + let second = projection.stage_entry("verify", 2, first_event_seq(2)); + second.duration_ms = Some(800); + second.usage = Some(success_usage); + second.completion = Some(fabro_types::StageCompletion { + outcome: StageOutcome::Succeeded, + notes: None, + failure_reason: None, + timestamp: chrono::Utc::now(), + }); + + let rollup = billing_rollup_from_projection(&projection); + + assert_eq!(rollup.stages.len(), 1); + assert_eq!(rollup.stages[0].node_id, "verify"); + assert_eq!(rollup.stages[0].model_id.as_deref(), Some("gpt-new")); + assert_eq!(rollup.stages[0].duration_ms, 2000); + assert_eq!(rollup.stages[0].billing.input_tokens, 300); + assert_eq!(rollup.stages[0].billing.output_tokens, 30); + assert_eq!(rollup.stages[0].billing.total_usd_micros, Some(330)); + + assert_eq!(rollup.runtime_ms, 2000); + assert_eq!(rollup.totals.input_tokens, 300); + assert_eq!(rollup.totals.output_tokens, 30); + assert_eq!(rollup.totals.total_usd_micros, Some(330)); + assert_eq!(rollup.billed_visit_count, 2); + + assert_eq!(rollup.by_model.len(), 2); + assert_eq!(rollup.by_model[0].model_id, "gpt-new"); + assert_eq!(rollup.by_model[0].stages, 1); + assert_eq!(rollup.by_model[0].billing.input_tokens, 200); + assert_eq!(rollup.by_model[1].model_id, "gpt-old"); + assert_eq!(rollup.by_model[1].stages, 1); + assert_eq!(rollup.by_model[1].billing.input_tokens, 100); + } + + #[test] + fn rollup_includes_completed_non_llm_stage_rows_with_zero_billing() { + let mut projection = RunProjection::default(); + let stage = projection.stage_entry("start", 1, first_event_seq(1)); + stage.duration_ms = Some(25); + stage.completion = Some(fabro_types::StageCompletion { + outcome: StageOutcome::Succeeded, + notes: None, + failure_reason: None, + timestamp: chrono::Utc::now(), + }); + + let rollup = billing_rollup_from_projection(&projection); + + assert_eq!(rollup.stages.len(), 1); + assert_eq!(rollup.stages[0].node_id, "start"); + assert_eq!(rollup.stages[0].duration_ms, 25); + assert!(rollup.stages[0].model_id.is_none()); + assert_eq!(rollup.stages[0].billing.input_tokens, 0); + assert_eq!(rollup.runtime_ms, 25); + assert!(rollup.by_model.is_empty()); + assert!(rollup.billing_if_present().is_none()); + } +} diff --git a/lib/crates/fabro-workflow/src/lib.rs b/lib/crates/fabro-workflow/src/lib.rs index fe814118d..5478e611d 100644 --- a/lib/crates/fabro-workflow/src/lib.rs +++ b/lib/crates/fabro-workflow/src/lib.rs @@ -270,6 +270,7 @@ mod duration_tests { pub mod artifact; pub mod artifact_snapshot; pub mod artifact_upload; +pub mod billing_rollup; pub mod command_log; pub(crate) mod condition; pub mod context; @@ -298,6 +299,10 @@ pub mod run_control; pub(crate) mod run_dir; pub mod run_lookup; +pub use billing_rollup::{ + ProjectionBillingByModel, ProjectionBillingRollup, ProjectionBillingStage, + billing_rollup_from_projection, +}; pub use error::{Error, FailureCategory, FailureSignature, FailureSignatureExt, Result}; pub use manifest_path::ManifestPath; pub mod run_materialization; diff --git a/lib/crates/fabro-workflow/src/pipeline/finalize.rs b/lib/crates/fabro-workflow/src/pipeline/finalize.rs index 997de0f4e..927d2d833 100644 --- a/lib/crates/fabro-workflow/src/pipeline/finalize.rs +++ b/lib/crates/fabro-workflow/src/pipeline/finalize.rs @@ -19,6 +19,7 @@ use crate::run_status::{FailureReason, RunStatus, SuccessReason}; use crate::runtime_store::RunStoreHandle; use crate::sandbox_git::git_diff_with_timeout; use crate::services::RunServices; +use crate::{ProjectionBillingRollup, billing_rollup_from_projection}; pub fn classify_engine_result( engine_result: &Result, @@ -68,22 +69,22 @@ pub(crate) async fn build_conclusion_from_store( run_duration_ms: u64, final_git_commit_sha: Option, ) -> Conclusion { - let (state_result, events_result) = tokio::join!(run_store.state(), run_store.list_events()); - let projection = state_result.ok(); + let projection = run_store.state().await.ok(); let projection_order = projection .as_ref() .map(stage_projection_order) .unwrap_or_default(); + let projection_billing = projection + .as_ref() + .map(billing_rollup_from_projection) + .unwrap_or_default(); let checkpoint = projection .as_ref() .and_then(|state| state.checkpoint.as_ref()); - let stage_durations = events_result - .map(|events| crate::latest_stage_duration_by_node(&events)) - .unwrap_or_default(); build_conclusion_from_parts( checkpoint, - &stage_durations, + &projection_billing, &projection_order, status, failure_reason, @@ -94,7 +95,7 @@ pub(crate) async fn build_conclusion_from_store( fn build_conclusion_from_parts( checkpoint: Option<&Checkpoint>, - stage_durations: &HashMap, + projection_billing: &ProjectionBillingRollup, projection_order: &HashMap, status: StageOutcome, failure_reason: Option, @@ -105,6 +106,11 @@ fn build_conclusion_from_parts( // while the other checkpoint maps are keyed by node_id. Dedupe to one row // per node so the stages table matches the deduped billing total. let (stages, total_retries) = if let Some(cp) = checkpoint { + let billing_by_node = projection_billing + .stages + .iter() + .map(|stage| (stage.node_id.as_str(), stage)) + .collect::>(); let mut stage_rows = Vec::new(); let mut seen = std::collections::HashSet::new(); let mut retries_sum: u32 = 0; @@ -130,7 +136,6 @@ fn build_conclusion_from_parts( } for (original_checkpoint_order, node_id) in stage_order { - let outcome = cp.node_outcomes.get(node_id); let retries = cp .node_retries .get(node_id) @@ -138,14 +143,13 @@ fn build_conclusion_from_parts( .unwrap_or(1) .saturating_sub(1); retries_sum += retries; + let billing = billing_by_node.get(node_id); let summary = StageSummary { stage_id: node_id.to_string(), stage_label: node_id.to_string(), - duration_ms: stage_durations.get(node_id).copied().unwrap_or(0), - billing_usd_micros: outcome - .and_then(|o| o.usage.as_ref()) - .and_then(|usage| usage.total_usd_micros), + duration_ms: billing.map_or(0, |stage| stage.duration_ms), + billing_usd_micros: billing.and_then(|stage| stage.billing.total_usd_micros), retries, }; stage_rows.push(( @@ -176,7 +180,7 @@ fn build_conclusion_from_parts( failure_reason, final_git_commit_sha, stages, - billing: checkpoint.and_then(billing_from_checkpoint), + billing: projection_billing.billing_if_present(), total_retries, } } @@ -391,15 +395,8 @@ async fn compute_final_patch( } } -/// Iterates `node_outcomes.values()` rather than `completed_nodes` to avoid -/// over-counting the last visit's usage on looping workflows. -pub(crate) fn billing_from_checkpoint(cp: &Checkpoint) -> Option { - let usage: Vec<_> = cp - .node_outcomes - .values() - .filter_map(|o| o.usage.clone()) - .collect(); - (!usage.is_empty()).then(|| BilledTokenCounts::from_billed_usage(&usage)) +pub(crate) fn billing_from_projection(projection: &RunProjection) -> Option { + billing_rollup_from_projection(projection).billing_if_present() } pub(crate) fn build_terminal_event( @@ -503,7 +500,6 @@ pub async fn finalize(retroed: Retroed, options: &FinalizeOptions) -> Result Result BilledModelUsage { + serde_json::from_value(serde_json::json!({ + "input": { + "usage": { + "model": { + "provider": "openai", + "model_id": model_id + }, + "tokens": { + "input_tokens": input_tokens, + "output_tokens": output_tokens + } + }, + "facts": { + "provider": "open_ai" + } + }, + "total_usd_micros": input_tokens + output_tokens + })) + .unwrap() + } + #[test] fn conclusion_stage_order_follows_projection_first_event_order() { let mut projection = RunProjection::default(); @@ -753,7 +776,7 @@ mod tests { let conclusion = build_conclusion_from_parts( Some(&checkpoint), - &HashMap::new(), + &ProjectionBillingRollup::default(), &projection_order, StageOutcome::Succeeded, None, @@ -788,7 +811,7 @@ mod tests { let conclusion = build_conclusion_from_parts( Some(&checkpoint), - &HashMap::new(), + &ProjectionBillingRollup::default(), &projection_order, StageOutcome::Succeeded, None, @@ -804,6 +827,66 @@ mod tests { assert_eq!(stage_ids, vec!["skipped", "finished"]); } + #[test] + fn conclusion_billing_sums_retry_visit_usage_from_projection() { + let mut projection = RunProjection::default(); + let failed_usage = test_usage("gpt-old", 100, 10); + let success_usage = test_usage("gpt-new", 200, 20); + let failed = projection.stage_entry("verify", 1, first_event_seq(1)); + failed.duration_ms = Some(1200); + failed.usage = Some(failed_usage); + failed.completion = Some(StageCompletion { + outcome: StageOutcome::Failed { + retry_requested: true, + }, + notes: None, + failure_reason: Some("try again".to_string()), + timestamp: chrono::Utc::now(), + }); + let succeeded = projection.stage_entry("verify", 2, first_event_seq(2)); + succeeded.duration_ms = Some(800); + succeeded.usage = Some(success_usage.clone()); + succeeded.completion = Some(StageCompletion { + outcome: StageOutcome::Succeeded, + notes: None, + failure_reason: None, + timestamp: chrono::Utc::now(), + }); + + let projection_order = stage_projection_order(&projection); + let projection_billing = billing_rollup_from_projection(&projection); + let mut latest_outcome = Outcome::success(); + latest_outcome.usage = Some(success_usage); + latest_outcome.duration_ms = Some(800); + let mut checkpoint = checkpoint_with( + vec!["verify", "verify"], + HashMap::from([("verify".to_string(), latest_outcome)]), + ); + checkpoint.node_retries.insert("verify".to_string(), 2); + + let conclusion = build_conclusion_from_parts( + Some(&checkpoint), + &projection_billing, + &projection_order, + StageOutcome::Succeeded, + None, + 10, + None, + ); + + assert_eq!(conclusion.billing.as_ref().unwrap().input_tokens, 300); + assert_eq!(conclusion.billing.as_ref().unwrap().output_tokens, 30); + assert_eq!( + conclusion.billing.as_ref().unwrap().total_usd_micros, + Some(330) + ); + assert_eq!(conclusion.stages.len(), 1); + assert_eq!(conclusion.stages[0].stage_id, "verify"); + assert_eq!(conclusion.stages[0].duration_ms, 2000); + assert_eq!(conclusion.stages[0].billing_usd_micros, Some(330)); + assert_eq!(conclusion.stages[0].retries, 1); + } + fn test_services( run_store: RunStoreHandle, emitter: Arc, diff --git a/lib/crates/fabro-workflow/src/pipeline/mod.rs b/lib/crates/fabro-workflow/src/pipeline/mod.rs index 71aefbef9..284c32a1f 100644 --- a/lib/crates/fabro-workflow/src/pipeline/mod.rs +++ b/lib/crates/fabro-workflow/src/pipeline/mod.rs @@ -12,7 +12,7 @@ mod validate; pub use execute::execute; pub use fabro_types::PullRequestRecord; pub(crate) use finalize::{ - billing_from_checkpoint, build_conclusion_from_store, build_terminal_event, + billing_from_projection, build_conclusion_from_store, build_terminal_event, }; pub use finalize::{classify_engine_result, finalize, write_finalize_commit}; pub use initialize::initialize; diff --git a/lib/crates/fabro-workflow/src/test_support.rs b/lib/crates/fabro-workflow/src/test_support.rs index 12aa40f10..0bf6377a6 100644 --- a/lib/crates/fabro-workflow/src/test_support.rs +++ b/lib/crates/fabro-workflow/src/test_support.rs @@ -17,7 +17,7 @@ use crate::handler::HandlerRegistry; use crate::outcome::Outcome; use crate::pipeline; use crate::pipeline::types::{Executed, Initialized}; -use crate::pipeline::{billing_from_checkpoint, build_terminal_event}; +use crate::pipeline::{billing_from_projection, build_terminal_event}; use crate::records::Checkpoint; use crate::run_metadata::RunMetadataRuntime; use crate::run_options::RunOptions; @@ -36,10 +36,7 @@ async fn execute_and_emit_terminal(initialized: InitializedState) -> Executed { let executed = Box::pin(pipeline::execute(initialized.initialized)).await; initialized.store_logger.flush().await; let state = executed.engine.run.run_store.state().await.ok(); - let billing = state - .as_ref() - .and_then(|s| s.checkpoint.as_ref()) - .and_then(billing_from_checkpoint); + let billing = state.as_ref().and_then(billing_from_projection); let event = build_terminal_event( &executed.outcome, executed.duration_ms, diff --git a/lib/packages/fabro-api-client/src/models/billing-by-model.ts b/lib/packages/fabro-api-client/src/models/billing-by-model.ts index d8d802575..74f074401 100644 --- a/lib/packages/fabro-api-client/src/models/billing-by-model.ts +++ b/lib/packages/fabro-api-client/src/models/billing-by-model.ts @@ -26,7 +26,7 @@ import type { ModelReference } from './model-reference'; export interface BillingByModel { 'model': ModelReference; /** - * Number of stages that used this model. + * Number of usage-bearing stage visits that used this model. */ 'stages': number; 'billing': BilledTokenCounts; diff --git a/lib/packages/fabro-api-client/src/models/billing-stage-ref.ts b/lib/packages/fabro-api-client/src/models/billing-stage-ref.ts index bcc94a656..7b87ef631 100644 --- a/lib/packages/fabro-api-client/src/models/billing-stage-ref.ts +++ b/lib/packages/fabro-api-client/src/models/billing-stage-ref.ts @@ -15,7 +15,7 @@ /** - * Reference to a billing stage. + * Reference to a workflow node in a billing stage row. */ export interface BillingStageRef { /** diff --git a/lib/packages/fabro-api-client/src/models/run-billing-stage.ts b/lib/packages/fabro-api-client/src/models/run-billing-stage.ts index b6d807d8f..5d376b7cc 100644 --- a/lib/packages/fabro-api-client/src/models/run-billing-stage.ts +++ b/lib/packages/fabro-api-client/src/models/run-billing-stage.ts @@ -24,14 +24,14 @@ import type { BillingStageRef } from './billing-stage-ref'; import type { ModelReference } from './model-reference'; /** - * Token counts and billed totals for a single stage within a run. + * Token counts and billed totals for one workflow node within a run. Rows are grouped by node; billing and runtime sum every visit of that node. */ export interface RunBillingStage { 'stage': BillingStageRef; 'model': ModelReference | null; 'billing': BilledTokenCounts; /** - * Wall-clock runtime in seconds. + * Wall-clock runtime in seconds, summed across every visit of this node. */ 'runtime_secs': number; } diff --git a/lib/packages/fabro-api-client/src/models/run-billing.ts b/lib/packages/fabro-api-client/src/models/run-billing.ts index 88544b1cc..09cfed9e3 100644 --- a/lib/packages/fabro-api-client/src/models/run-billing.ts +++ b/lib/packages/fabro-api-client/src/models/run-billing.ts @@ -28,7 +28,7 @@ import type { RunBillingTotals } from './run-billing-totals'; */ export interface RunBilling { /** - * Per-stage billing breakdown. + * Per-node billing breakdown. Each row sums billing and runtime across all visits of that node. */ 'stages': Array; 'totals': RunBillingTotals;