diff --git a/docs/public/examples/definition-of-done.mdx b/docs/public/examples/definition-of-done.mdx index 5b1ef9091..6291293b1 100644 --- a/docs/public/examples/definition-of-done.mdx +++ b/docs/public/examples/definition-of-done.mdx @@ -262,7 +262,7 @@ Otherwise set preferred_next_label to \"more_work_needed\"." **Three-category triage** — failures are classified as IMPLEMENTABLE (can fix now), STRUCTURAL (needs architecture work), or DEFERRED (needs external resources). This prevents the agent from wasting cycles on items it can't address in a code-only pass. -**Batched fixes** — the `fix_batch` node tackles up to 5 failures per iteration. The self-loop (`fix_batch -> fix_batch` with `loop_restart=true`) allows it to keep going when more fixes remain, while `goal_gate=true` ensures the workflow only succeeds if fixes were actually applied. +**Batched fixes** — the `fix_batch` node tackles up to 5 failures per iteration. The self-loop (`fix_batch -> fix_batch` with `loop_restart=true`) allows it to keep going when more fixes remain, while `goal_gate=true` ensures the workflow only succeeds if fixes were actually applied. Because `loop_restart` begins each round with a fresh, empty context, every iteration re-derives the remaining work from the repository state rather than from accumulated conversation history — see [Failures — Loop restart edges](/execution/failures#loop-restart-edges). **Build gate** — after each fix batch, a script node runs `cargo build` and `cargo test`. If the build breaks, a dedicated `build_fix` node diagnoses and repairs compilation errors before retrying. diff --git a/docs/public/execution/failures.mdx b/docs/public/execution/failures.mdx index 1caa84184..a74eefa9f 100644 --- a/docs/public/execution/failures.mdx +++ b/docs/public/execution/failures.mdx @@ -232,7 +232,11 @@ Failure signature counts are never reset on success. This is intentional — it ### Loop restart edges -Edges marked with `loop_restart=true` trigger a special restart of the workflow from the target node. These have an additional guard: only `transient_infra` failures may cross a `loop_restart` edge. If the failure class is anything else, the run is terminated: +Taking an edge marked with `loop_restart=true` restarts the workflow from the edge's target node. A restart is more than a jump: the completed-stage history, per-node outcomes, and retry counts are cleared, and the run context is replaced with a **fresh, empty context** — the target node starts over as if the run had just begun there, with no preamble of prior stages. Node visit counts are the one thing preserved, so `max_visits` and `max_node_visits` still bound how many times a restart loop can run. + +A **successful** outcome may take a `loop_restart` edge freely. This is the "start another round from a clean slate" pattern — for example, a self-loop that begins a fresh batch of work and re-derives its remaining work from the repository state rather than from accumulated context. + +A **failed** outcome faces an additional guard: only `transient_infra` failures may cross a `loop_restart` edge. If the failure class is anything else, the run is terminated: ``` loop_restart blocked: failure_class=deterministic (requires transient_infra) diff --git a/docs/public/reference/dot-language.mdx b/docs/public/reference/dot-language.mdx index eb162718a..f01fbf2ed 100644 --- a/docs/public/reference/dot-language.mdx +++ b/docs/public/reference/dot-language.mdx @@ -284,7 +284,7 @@ audit [ | `weight` | Integer | Priority for tiebreaking (higher wins, default: 0) | | `fidelity` | String | Override fidelity level for this transition | | `thread_id` | String | Override thread ID for this transition | -| `loop_restart` | Boolean | Mark this edge as a loop restart point | +| `loop_restart` | Boolean | Restart the workflow from this edge's target when taken: stage history and retry counts clear and the context resets to empty (visit counts are kept). Failed outcomes may only take it for `transient_infra` failures — see [Failures](/execution/failures#loop-restart-edges) | | `freeform` | Boolean | When `true` on a human-gate edge, accept free-text input instead of fixed choices | ## Condition expressions diff --git a/lib/crates/fabro-agent/src/session.rs b/lib/crates/fabro-agent/src/session.rs index dfc754d1f..34986a29d 100644 --- a/lib/crates/fabro-agent/src/session.rs +++ b/lib/crates/fabro-agent/src/session.rs @@ -15,7 +15,7 @@ use fabro_llm::{Error as LlmError, retry}; use fabro_mcp::config::{McpServerSettings, McpTransport}; use fabro_mcp::connection_manager::McpConnectionManager; use fabro_mcp::http_transport; -use fabro_model::{AgentProfileKind, Catalog, ModelRef, Speed}; +use fabro_model::{AgentProfileKind, Catalog, ModelRef, Speed, UsdMicros}; use fabro_types::{ AgentToolSummary, PermissionLevel, Principal, SessionMessage, SessionRecord, StageContextWindowProjection, SteeringMessage, @@ -354,6 +354,7 @@ pub struct Session { completion_coordinator: Option>, last_input_timing: SessionInputTiming, last_input_usage: TokenCounts, + last_input_cost: Option, } impl Session { @@ -393,6 +394,7 @@ impl Session { completion_coordinator: None, last_input_timing: SessionInputTiming::default(), last_input_usage: TokenCounts::default(), + last_input_cost: None, } } @@ -1222,6 +1224,11 @@ impl Session { self.last_input_usage.clone() } + #[must_use] + pub const fn last_input_cost(&self) -> Option { + self.last_input_cost + } + /// Process an input. The inference/tool timing accumulated during the call /// is available via [`Self::last_input_timing`] after this returns, even on /// error. @@ -1232,8 +1239,10 @@ impl Session { ) -> Result<(), Error> { let mut timing = SessionInputTiming::default(); let mut usage = TokenCounts::default(); + let mut cost = None; self.last_input_timing = timing; self.last_input_usage = TokenCounts::default(); + self.last_input_cost = None; if self.state == SessionState::Closed { return Err(Error::SessionClosed); } @@ -1258,7 +1267,13 @@ impl Session { // Process the initial input, then drain any followups let mut result = self - .run_single_input(input, &agent_tool_runtime, &mut timing, &mut usage) + .run_single_input( + input, + &agent_tool_runtime, + &mut timing, + &mut usage, + &mut cost, + ) .await; if result.is_ok() { @@ -1270,7 +1285,13 @@ impl Session { .pop_front(); let Some(followup) = followup else { break }; result = self - .run_single_input(&followup, &agent_tool_runtime, &mut timing, &mut usage) + .run_single_input( + &followup, + &agent_tool_runtime, + &mut timing, + &mut usage, + &mut cost, + ) .await; if result.is_err() { break; @@ -1290,6 +1311,7 @@ impl Session { self.last_input_timing = timing; self.last_input_usage = usage; + self.last_input_cost = cost; result } @@ -1299,6 +1321,7 @@ impl Session { agent_tool_runtime: &AgentToolRuntime, timing: &mut SessionInputTiming, usage_accumulator: &mut TokenCounts, + cost_accumulator: &mut Option, ) -> Result<(), Error> { const STREAM_CONSUME_RETRIES: usize = 3; @@ -1704,6 +1727,7 @@ impl Session { &usage, )); *usage_accumulator += usage.clone(); + UsdMicros::accumulate(cost_accumulator, response.cost_usd.map(UsdMicros::from_usd)); self.history.push(Message::Assistant { content: text.clone(), @@ -1729,6 +1753,8 @@ impl Session { text: text.clone(), model, usage: response.usage.clone(), + cost_usd: response.cost_usd, + cost_source: response.cost_source, tool_call_count: tool_calls.len(), context_window, }); @@ -2276,6 +2302,25 @@ mod tests { } } + #[tokio::test] + async fn last_input_cost_sums_each_response_in_a_multi_turn_input() { + let mut registry = ToolRegistry::new(); + registry.register(make_echo_tool()); + + let responses = vec![ + response_with_cost( + tool_call_response("echo", "call_1", serde_json::json!({"text": "hello"})), + 0.04, + ), + response_with_cost(text_response("Done!"), 0.06), + ]; + + let mut session = make_session_with_tools(responses, registry).await; + session.process_input("Use echo tool").await.unwrap(); + + assert_eq!(session.last_input_cost(), Some(UsdMicros(100_000))); + } + #[tokio::test] async fn last_input_timing_reports_inference_and_tool_per_call() { let mut registry = ToolRegistry::new(); @@ -4009,6 +4054,12 @@ mod tests { response } + fn response_with_cost(mut response: Response, cost_usd: f64) -> Response { + response.cost_usd = Some(cost_usd); + response.cost_source = Some(fabro_model::CostSource::Authoritative); + response + } + fn response_with_input_tokens(response: Response, input_tokens: i64) -> Response { response_with_usage(response, TokenCounts { input_tokens, diff --git a/lib/crates/fabro-agent/src/types.rs b/lib/crates/fabro-agent/src/types.rs index d6907880f..94cb6808e 100644 --- a/lib/crates/fabro-agent/src/types.rs +++ b/lib/crates/fabro-agent/src/types.rs @@ -3,7 +3,7 @@ use std::time::SystemTime; use chrono::{DateTime, Utc}; use fabro_llm::Error as LlmError; use fabro_llm::types::{ContentPart, ThinkingData, TokenCounts, ToolCall, ToolResult}; -use fabro_model::ModelRef; +use fabro_model::{CostSource, ModelRef}; use fabro_types::{SessionMessage, StageContextWindowProjection}; use serde::de::DeserializeOwned; use serde::{Deserialize, Serialize}; @@ -245,6 +245,12 @@ pub enum AgentEvent { text: String, model: ModelRef, usage: TokenCounts, + /// USD cost reported or estimated for this individual response. + #[serde(default, skip_serializing_if = "Option::is_none")] + cost_usd: Option, + /// Provenance of `cost_usd`. + #[serde(default, skip_serializing_if = "Option::is_none")] + cost_source: Option, tool_call_count: usize, #[serde(default, skip_serializing_if = "Option::is_none")] context_window: Option, @@ -880,12 +886,16 @@ mod tests { speed: None, }, usage: usage.clone(), + cost_usd: Some(0.125), + cost_source: Some(CostSource::Authoritative), tool_call_count: 2, context_window: None, }; match &event { AgentEvent::AssistantMessage { usage, + cost_usd, + cost_source, tool_call_count, .. } => { @@ -893,6 +903,8 @@ mod tests { assert_eq!(usage.input_tokens, 100); assert_eq!(usage.cache_read_tokens, 80); assert_eq!(usage.reasoning_tokens, 20); + assert_eq!(*cost_usd, Some(0.125)); + assert_eq!(*cost_source, Some(CostSource::Authoritative)); } _ => panic!("expected AssistantMessage"), } diff --git a/lib/crates/fabro-cli/src/commands/run/run_progress/mod.rs b/lib/crates/fabro-cli/src/commands/run/run_progress/mod.rs index 75d7fc0ff..6e1a6227b 100644 --- a/lib/crates/fabro-cli/src/commands/run/run_progress/mod.rs +++ b/lib/crates/fabro-cli/src/commands/run/run_progress/mod.rs @@ -528,6 +528,8 @@ mod tests { speed: None, }, usage: TokenCounts::default(), + cost_usd: None, + cost_source: None, tool_call_count: 0, context_window: None, }) diff --git a/lib/crates/fabro-cli/tests/it/cmd/upgrade.rs b/lib/crates/fabro-cli/tests/it/cmd/upgrade.rs index 7a3a35576..b75ebb779 100644 --- a/lib/crates/fabro-cli/tests/it/cmd/upgrade.rs +++ b/lib/crates/fabro-cli/tests/it/cmd/upgrade.rs @@ -49,6 +49,12 @@ fn brew_command(context: &TestContext, formula: &str, version: &str) -> Command } } cmd.env(EnvVars::NO_COLOR, "1"); + // Unlike context.command(), this command inherits the developer's + // environment, and inherited FORCE_COLOR/CLICOLOR_FORCE override NO_COLOR + // in the CLI's color detection — breaking these snapshots. + cmd.env_remove(EnvVars::FORCE_COLOR); + cmd.env_remove(EnvVars::CLICOLOR_FORCE); + cmd.env_remove(EnvVars::CLICOLOR); cmd.env(EnvVars::HOME, &context.home_dir); cmd.env(EnvVars::FABRO_NO_UPGRADE_CHECK, "true") .env(EnvVars::FABRO_HTTP_PROXY_POLICY, "disabled") diff --git a/lib/crates/fabro-model/src/billing.rs b/lib/crates/fabro-model/src/billing.rs index 0f571bfb3..6e4de19f0 100644 --- a/lib/crates/fabro-model/src/billing.rs +++ b/lib/crates/fabro-model/src/billing.rs @@ -47,17 +47,34 @@ fn saturating_rounded_f64_to_i64(value: f64) -> i64 { #[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Default, Serialize, Deserialize)] pub struct UsdMicros(pub i64); +impl UsdMicros { + #[must_use] + pub fn from_usd(usd: f64) -> Self { + Self(saturating_rounded_f64_to_i64( + (usd * USD_MICROS_PER_USD_F64).round(), + )) + } + + /// Folds a cost into a running total that stays `None` until a cost is + /// observed (`None` means "no provider data", not $0). + pub fn accumulate(total: &mut Option, cost: Option) { + if let Some(cost) = cost { + *total.get_or_insert_default() += cost; + } + } +} + 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; } } @@ -67,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, @@ -76,7 +99,7 @@ impl PricePerMTok { #[must_use] pub fn from_usd(usd: f64) -> Self { Self { - usd_micros: saturating_rounded_f64_to_i64((usd * USD_MICROS_PER_USD_F64).round()), + usd_micros: UsdMicros::from_usd(usd).0, } } @@ -328,6 +351,16 @@ impl BilledModelUsage { pub fn tokens(&self) -> &TokenCounts { &self.input.usage.tokens } + + /// Overrides the billed total with a provider-reported cost; `None` leaves + /// the catalog estimate in place. + #[must_use] + pub fn with_reported_cost(mut self, cost: Option) -> Self { + if let Some(cost) = cost { + self.total_usd_micros = Some(cost.0); + } + self + } } #[derive(Debug, Clone, PartialEq, Eq, Default, Serialize, Deserialize)] @@ -349,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, } } @@ -391,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) { @@ -404,15 +431,23 @@ 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) { *self = Self::from_billed_usage(std::slice::from_ref(usage)); } + /// Overrides the billed total with a provider-reported cost; `None` leaves + /// any existing estimate in place. + #[must_use] + pub fn with_reported_cost(mut self, cost: Option) -> Self { + if let Some(cost) = cost { + self.total_usd_micros = Some(cost.0); + } + self + } + #[must_use] pub fn is_zero(&self) -> bool { self.input_tokens == 0 @@ -726,6 +761,45 @@ mod tests { } } + #[test] + fn usd_micros_accumulate_keeps_none_until_a_cost_is_observed() { + let mut total = None; + UsdMicros::accumulate(&mut total, None); + assert_eq!(total, None); + + UsdMicros::accumulate(&mut total, Some(UsdMicros(40_000))); + UsdMicros::accumulate(&mut total, None); + UsdMicros::accumulate(&mut total, Some(UsdMicros(60_000))); + 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( @@ -859,6 +933,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/demo/mod.rs b/lib/crates/fabro-server/src/demo/mod.rs index 283435092..4b5bb204e 100644 --- a/lib/crates/fabro-server/src/demo/mod.rs +++ b/lib/crates/fabro-server/src/demo/mod.rs @@ -1485,6 +1485,7 @@ mod runs { speed: None, }, billing: BilledTokenCounts::default(), + cost_source: None, tool_call_count: 0, visit: 1, message: None, @@ -1554,6 +1555,7 @@ mod runs { speed: None, }, billing: BilledTokenCounts::default(), + cost_source: None, tool_call_count: 0, visit: 1, message: None, diff --git a/lib/crates/fabro-server/src/server/handler/pair.rs b/lib/crates/fabro-server/src/server/handler/pair.rs index 8367e34ad..67bc1ad05 100644 --- a/lib/crates/fabro-server/src/server/handler/pair.rs +++ b/lib/crates/fabro-server/src/server/handler/pair.rs @@ -887,6 +887,7 @@ mod tests { speed: None, }, billing: BilledTokenCounts::default(), + cost_source: None, tool_call_count: 0, visit: 1, message: None, @@ -919,6 +920,7 @@ mod tests { speed: None, }, billing: BilledTokenCounts::default(), + cost_source: None, tool_call_count: 0, visit: 1, message: None, 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 74095bbf5..9ecdabcae 100644 --- a/lib/crates/fabro-server/src/server/tests.rs +++ b/lib/crates/fabro-server/src/server/tests.rs @@ -4148,6 +4148,8 @@ fn context_window_event( speed: None, }, usage: TokenCounts::default(), + cost_usd: None, + cost_source: None, tool_call_count: 0, context_window: Some(context_window), }, @@ -13771,6 +13773,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(); diff --git a/lib/crates/fabro-static/src/env_vars.rs b/lib/crates/fabro-static/src/env_vars.rs index 4ed7585ca..b2921b045 100644 --- a/lib/crates/fabro-static/src/env_vars.rs +++ b/lib/crates/fabro-static/src/env_vars.rs @@ -114,6 +114,7 @@ impl EnvVars { pub const CI: &'static str = "CI"; pub const CLICOLOR: &'static str = "CLICOLOR"; pub const CLICOLOR_FORCE: &'static str = "CLICOLOR_FORCE"; + pub const FORCE_COLOR: &'static str = "FORCE_COLOR"; pub const HOME: &'static str = "HOME"; pub const KUBERNETES_SERVICE_HOST: &'static str = "KUBERNETES_SERVICE_HOST"; pub const LANG: &'static str = "LANG"; @@ -251,6 +252,7 @@ mod tests { EnvVars::CI, EnvVars::CLICOLOR, EnvVars::CLICOLOR_FORCE, + EnvVars::FORCE_COLOR, EnvVars::HOME, EnvVars::KUBERNETES_SERVICE_HOST, EnvVars::LANG, diff --git a/lib/crates/fabro-store/src/run_state.rs b/lib/crates/fabro-store/src/run_state.rs index 0d21ba098..19c6c6a3d 100644 --- a/lib/crates/fabro-store/src/run_state.rs +++ b/lib/crates/fabro-store/src/run_state.rs @@ -513,6 +513,38 @@ impl RunProjectionReducer for RunProjection { }; stage.parallel_results = Some(parallel_results); } + EventBody::ParallelBranchStarted(_) => { + // Branches bypass the engine's StageStarted/StageCompleted + // lifecycle. Seed started_at so the branch stage drives a live + // wall-clock timer while it runs (the entry is created Running). + let Some(stage) = stage_at_stored_or_current_visit(self, stored, event.seq) else { + return Ok(()); + }; + if stage.started_at.is_none() { + stage.started_at = Some(ts); + } + stage.state = StageState::Running; + } + EventBody::ParallelBranchCompleted(props) => { + // A branch never emits its own StageCompleted, so finalize it + // here; otherwise the stage spins Running forever after the run + // (and the fan-in) is done. + let outcome = + StageOutcome::from_str(&props.status).unwrap_or(StageOutcome::Failed { + retry_requested: false, + }); + let Some(stage) = stage_at_stored_or_current_visit(self, stored, event.seq) else { + return Ok(()); + }; + stage.completion = Some(StageCompletion { + outcome, + notes: None, + failure_reason: None, + timestamp: ts, + }); + stage.timing = Some(fabro_types::StageTiming::wall_only(props.duration_ms)); + stage.state = StageState::from(outcome); + } EventBody::TodoCreated(props) => { let Some(stage) = stage_at_stored_or_current_visit(self, stored, event.seq) else { return Ok(()); @@ -1260,8 +1292,9 @@ mod tests { AgentSubFailedProps, AgentSubSpawnedProps, AgentToolCategory, AgentToolSource, AgentToolStartedProps, AgentToolSummary, AgentToolsAvailableProps, CheckpointCompletedProps, InterviewCompletedProps, InterviewOption, InterviewStartedProps, - RunCompletedProps, RunControlEffectProps, StageCompletedProps, StageFailedProps, - StagePromptProps, StageRetryingProps, StageStartedProps, + ParallelBranchCompletedProps, ParallelBranchStartedProps, RunCompletedProps, + RunControlEffectProps, StageCompletedProps, StageFailedProps, StagePromptProps, + StageRetryingProps, StageStartedProps, }; use fabro_types::settings::run::{DockerfileSource, EnvironmentProvider}; use fabro_types::{ @@ -1993,6 +2026,86 @@ mod tests { assert_eq!(stage.prompt.as_deref(), Some("prompt")); } + #[test] + fn parallel_branch_completed_finalizes_branch_stage() { + // A parallel branch never runs through the engine's StageStarted/ + // StageCompleted lifecycle: its stage entry is created Running by the + // first branch-scoped event, and only ParallelBranchCompleted marks it + // terminal. Guards against branches spinning Running forever. + let mut state = initialized_projection(); + let branch = StageId::new("review_ux", 1); + let branch_started_at = test_dt("2026-04-07T12:00:00Z"); + + state + .apply_event(&test_stage_event_at( + 3, + "2026-04-07T12:00:00Z", + EventBody::ParallelBranchStarted(ParallelBranchStartedProps { index: 0 }), + branch.clone(), + )) + .unwrap(); + let stage = state.stage(&branch).unwrap(); + assert_eq!(stage.state, StageState::Running); + assert_eq!(stage.started_at, Some(branch_started_at)); + + state + .apply_event(&test_stage_event( + 4, + EventBody::ParallelBranchCompleted(ParallelBranchCompletedProps { + index: 0, + duration_ms: 1234, + status: "succeeded".to_string(), + head_sha: None, + }), + branch.clone(), + )) + .unwrap(); + + let stage = state.stage(&branch).unwrap(); + assert_eq!(stage.state, StageState::Succeeded); + assert_eq!(stage.timing.unwrap().wall_time_ms, 1234); + assert_eq!( + stage.completion.as_ref().unwrap().outcome, + StageOutcome::Succeeded + ); + } + + #[test] + fn parallel_branch_completed_folds_failed_status_as_failed() { + let mut state = initialized_projection(); + let branch = StageId::new("review_ux", 1); + + state + .apply_event(&test_stage_event( + 3, + EventBody::ParallelBranchStarted(ParallelBranchStartedProps { index: 0 }), + branch.clone(), + )) + .unwrap(); + state + .apply_event(&test_stage_event( + 4, + EventBody::ParallelBranchCompleted(ParallelBranchCompletedProps { + index: 0, + duration_ms: 500, + status: "failed".to_string(), + head_sha: None, + }), + branch.clone(), + )) + .unwrap(); + + let stage = state.stage(&branch).unwrap(); + assert_eq!(stage.state, StageState::Failed); + assert_eq!(stage.timing.unwrap().wall_time_ms, 500); + assert_eq!( + stage.completion.as_ref().unwrap().outcome, + StageOutcome::Failed { + retry_requested: false, + } + ); + } + fn start_stage(state: &mut RunProjection, stage_id: &StageId) { state .apply_event(&test_stage_event( @@ -3788,6 +3901,7 @@ mod tests { text: "assistant text".to_string(), model: billed_usage().model().clone(), billing, + cost_source: None, tool_call_count: 0, visit: 1, message: None, diff --git a/lib/crates/fabro-types/src/run_event/agent.rs b/lib/crates/fabro-types/src/run_event/agent.rs index ada131b09..138344f9f 100644 --- a/lib/crates/fabro-types/src/run_event/agent.rs +++ b/lib/crates/fabro-types/src/run_event/agent.rs @@ -1,4 +1,4 @@ -use fabro_model::{ReasoningEffort, Speed}; +use fabro_model::{CostSource, ReasoningEffort, Speed}; use serde::{Deserialize, Serialize}; use serde_json::Value; use strum::{Display, EnumString, IntoStaticStr}; @@ -121,6 +121,9 @@ pub struct AgentMessageProps { pub text: String, pub model: ModelRef, pub billing: BilledTokenCounts, + /// Provenance of the optional total in `billing`. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub cost_source: Option, pub tool_call_count: usize, pub visit: u32, /// Canonical replay-authoritative transcript message. Present on events @@ -403,6 +406,7 @@ mod tests { }); let props: AgentMessageProps = serde_json::from_value(v).unwrap(); assert_eq!(props.text, "hello"); + assert!(props.cost_source.is_none()); assert!(props.message.is_none()); assert!(props.context_window.is_none()); } @@ -416,6 +420,7 @@ mod tests { text: "ok".to_string(), model: sample_model_ref(), billing: BilledTokenCounts::default(), + cost_source: None, tool_call_count: 0, visit: 1, message: Some(msg.clone()), diff --git a/lib/crates/fabro-types/src/run_event/mod.rs b/lib/crates/fabro-types/src/run_event/mod.rs index bfb396663..4aed5b3af 100644 --- a/lib/crates/fabro-types/src/run_event/mod.rs +++ b/lib/crates/fabro-types/src/run_event/mod.rs @@ -2139,6 +2139,7 @@ mod tests { speed: None, }, billing: BilledTokenCounts::default(), + cost_source: None, tool_call_count: 0, visit: 1, message: None, @@ -2191,6 +2192,7 @@ mod tests { speed: None, }, billing: BilledTokenCounts::default(), + cost_source: None, tool_call_count: 0, visit: 1, message: None, diff --git a/lib/crates/fabro-validate/src/rules/inert_attribute.rs b/lib/crates/fabro-validate/src/rules/inert_attribute.rs new file mode 100644 index 000000000..fef9c5454 --- /dev/null +++ b/lib/crates/fabro-validate/src/rules/inert_attribute.rs @@ -0,0 +1,239 @@ +use fabro_graphviz::graph::{self, Graph}; + +use crate::{Diagnostic, LintRule, Severity}; + +pub(super) fn rule() -> Box { + Box::new(Rule) +} + +/// Attributes that only specific handler types read, paired with the handler +/// types that consume them. On every other node type the attribute is inert: +/// accepted by the parser and read by nothing at runtime. +/// +/// Attributes read by several handlers (`timeout`), resolved for every node +/// (`fidelity`, `retry_policy`, `max_visits`, `goal_gate`), or injectable via +/// model stylesheets (`model`, `provider`, `reasoning_effort`, `speed`, +/// `backend`) are deliberately not listed. +const HANDLER_SPECIFIC_ATTRS: &[(&str, &[&str])] = &[ + ("script", &["command"]), + ("language", &["command"]), + ("duration", &["wait"]), + ("join_policy", &["parallel"]), + ("max_parallel", &["parallel"]), + ("output_schema", &["agent", "prompt"]), + ("prompt", &["agent", "prompt", "parallel.fan_in"]), +]; + +struct Rule; + +impl LintRule for Rule { + fn name(&self) -> &'static str { + "inert_attribute" + } + + fn apply(&self, graph: &Graph) -> Vec { + let mut diagnostics = Vec::new(); + for node in graph.nodes.values() { + // An unknown shape or type is covered by the type_known rule; a + // node this rule cannot classify is skipped rather than guessed at. + let Some(handler) = node.handler_type() else { + continue; + }; + if !graph::is_known_handler_type(handler) { + continue; + } + for (attr, consumers) in HANDLER_SPECIFIC_ATTRS { + if !node.attrs.contains_key(*attr) { + continue; + } + if consumers.contains(&handler) { + continue; + } + diagnostics.push(Diagnostic { + rule: self.name().to_string(), + severity: Severity::Warning, + message: format!( + "Node '{}' (type '{handler}') sets '{attr}', which is only read by {} nodes and has no effect here", + node.id, + consumers.join(", "), + ), + node_id: Some(node.id.clone()), + edge: None, + fix: Some(format!( + "Remove '{attr}' or change the node to a type that reads it ({})", + consumers.join(", "), + )), + ..Diagnostic::default() + }); + } + } + diagnostics + } +} + +#[cfg(test)] +mod tests { + use fabro_graphviz::graph::{AttrValue, Node}; + + use super::Rule; + use crate::rules::test_support::minimal_graph; + use crate::{LintRule, Severity}; + + fn node_with_attr(id: &str, shape: &str, attr: &str, value: &str) -> Node { + let mut node = Node::new(id); + node.attrs + .insert("shape".to_string(), AttrValue::String(shape.to_string())); + node.attrs + .insert(attr.to_string(), AttrValue::String(value.to_string())); + node + } + + #[test] + fn warns_on_script_on_agent_node() { + let mut g = minimal_graph(); + g.nodes.insert( + "work".to_string(), + node_with_attr("work", "box", "script", "echo hi"), + ); + let d = Rule.apply(&g); + assert_eq!(d.len(), 1); + assert_eq!(d[0].severity, Severity::Warning); + assert!(d[0].message.contains("'script'")); + assert!(d[0].message.contains("command")); + assert_eq!(d[0].node_id.as_deref(), Some("work")); + } + + #[test] + fn warns_on_prompt_on_start_and_command_nodes() { + let mut g = minimal_graph(); + g.nodes + .get_mut("start") + .expect("minimal graph has start") + .attrs + .insert( + "prompt".to_string(), + AttrValue::String("do things".to_string()), + ); + g.nodes.insert( + "run".to_string(), + node_with_attr("run", "parallelogram", "prompt", "do things"), + ); + let d = Rule.apply(&g); + assert_eq!(d.len(), 2); + assert!(d.iter().all(|d| d.message.contains("'prompt'"))); + } + + #[test] + fn warns_on_duration_on_command_node() { + let mut g = minimal_graph(); + g.nodes.insert( + "run".to_string(), + node_with_attr("run", "parallelogram", "duration", "30s"), + ); + let d = Rule.apply(&g); + assert_eq!(d.len(), 1); + assert!(d[0].message.contains("'duration'")); + assert!(d[0].message.contains("wait")); + } + + #[test] + fn warns_on_parallel_attrs_on_agent_node() { + let mut g = minimal_graph(); + let mut node = Node::new("work"); + node.attrs.insert( + "join_policy".to_string(), + AttrValue::String("wait_all".to_string()), + ); + node.attrs + .insert("max_parallel".to_string(), AttrValue::Integer(4)); + g.nodes.insert("work".to_string(), node); + let d = Rule.apply(&g); + assert_eq!(d.len(), 2); + } + + #[test] + fn warns_on_output_schema_on_command_node() { + let mut g = minimal_graph(); + g.nodes.insert( + "run".to_string(), + node_with_attr("run", "parallelogram", "output_schema", "routing"), + ); + let d = Rule.apply(&g); + assert_eq!(d.len(), 1); + assert!(d[0].message.contains("'output_schema'")); + } + + #[test] + fn accepts_attrs_on_their_own_handler_types() { + let mut g = minimal_graph(); + g.nodes.insert( + "run".to_string(), + node_with_attr("run", "parallelogram", "script", "echo hi"), + ); + g.nodes.insert( + "pause".to_string(), + node_with_attr("pause", "insulator", "duration", "30s"), + ); + g.nodes.insert( + "work".to_string(), + node_with_attr("work", "box", "prompt", "do things"), + ); + g.nodes.insert( + "fork".to_string(), + node_with_attr("fork", "component", "join_policy", "wait_all"), + ); + g.nodes.insert( + "spec".to_string(), + node_with_attr("spec", "tab", "output_schema", "routing"), + ); + assert!(Rule.apply(&g).is_empty()); + } + + #[test] + fn accepts_prompt_on_shapeless_node_defaulting_to_agent() { + let mut g = minimal_graph(); + let mut node = Node::new("work"); + node.attrs.insert( + "prompt".to_string(), + AttrValue::String("do things".to_string()), + ); + g.nodes.insert("work".to_string(), node); + assert!(Rule.apply(&g).is_empty()); + } + + #[test] + fn accepts_prompt_on_fan_in_judge() { + let mut g = minimal_graph(); + g.nodes.insert( + "merge".to_string(), + node_with_attr("merge", "tripleoctagon", "prompt", "pick the best"), + ); + assert!(Rule.apply(&g).is_empty()); + } + + #[test] + fn ignores_unclassifiable_node_shapes() { + let mut g = minimal_graph(); + g.nodes.insert( + "odd".to_string(), + node_with_attr("odd", "doubleoctagon", "script", "echo hi"), + ); + assert!(Rule.apply(&g).is_empty()); + } + + #[test] + fn ignores_handler_specific_attrs_on_unrecognized_explicit_types() { + let mut g = minimal_graph(); + let mut node = Node::new("custom"); + node.attrs.insert( + "type".to_string(), + AttrValue::String("custom.handler".to_string()), + ); + node.attrs.insert( + "script".to_string(), + AttrValue::String("echo hi".to_string()), + ); + g.nodes.insert("custom".to_string(), node); + assert!(Rule.apply(&g).is_empty()); + } +} diff --git a/lib/crates/fabro-validate/src/rules/mod.rs b/lib/crates/fabro-validate/src/rules/mod.rs index 98ece95af..428004240 100644 --- a/lib/crates/fabro-validate/src/rules/mod.rs +++ b/lib/crates/fabro-validate/src/rules/mod.rs @@ -8,9 +8,11 @@ mod fidelity_valid; mod freeform_edge_count; mod goal_gate_has_retry; mod import_error; +mod inert_attribute; mod model_support; mod node_model_known; mod orphan_custom_outcome; +mod parallel_branch_inert_attribute; mod prompt_on_llm_nodes; mod random_selection_no_conditions; mod reachability; @@ -60,6 +62,8 @@ pub fn built_in_rules() -> Vec> { thread_id_requires_fidelity_full::rule(), selection_valid::rule(), random_selection_no_conditions::rule(), + inert_attribute::rule(), + parallel_branch_inert_attribute::rule(), ] } diff --git a/lib/crates/fabro-validate/src/rules/parallel_branch_inert_attribute.rs b/lib/crates/fabro-validate/src/rules/parallel_branch_inert_attribute.rs new file mode 100644 index 000000000..eb051281e --- /dev/null +++ b/lib/crates/fabro-validate/src/rules/parallel_branch_inert_attribute.rs @@ -0,0 +1,289 @@ +use std::collections::BTreeSet; + +use fabro_graphviz::graph::Graph; + +use crate::{Diagnostic, LintRule, Severity}; + +pub(super) fn rule() -> Box { + Box::new(Rule) +} + +/// Attributes that parallel branch execution does not resolve. Branch nodes +/// are dispatched with a snapshot of the context taken when the parallel node +/// started, so per-branch `fidelity` never changes what a branch sees, and +/// per-branch `thread_id` never replaces the thread inherited in that snapshot. +const BRANCH_IGNORED_ATTRS: &[&str] = &["fidelity", "thread_id"]; + +struct Rule; + +/// Renders one or more parallel-node ids as `'a'` or `'a', 'b'`. +fn quoted_list(ids: &[String]) -> String { + ids.iter() + .map(|id| format!("'{id}'")) + .collect::>() + .join(", ") +} + +fn fix_message(attr: &str, parallel_ids: &[String]) -> String { + match attr { + "fidelity" => { + if parallel_ids.len() == 1 { + format!( + "Set fidelity on the parallel node {} (or its incoming edge) to control what every branch sees", + quoted_list(parallel_ids), + ) + } else { + format!( + "Set fidelity on the parallel nodes {} (or their incoming edges) to control what every branch sees", + quoted_list(parallel_ids), + ) + } + } + "thread_id" => format!( + "Remove '{attr}': parallel branches inherit the thread resolved when the parallel node started" + ), + _ => format!("Remove '{attr}'"), + } +} + +impl LintRule for Rule { + fn name(&self) -> &'static str { + "parallel_branch_inert_attribute" + } + + fn apply(&self, graph: &Graph) -> Vec { + let parallel_ids: BTreeSet<&str> = graph + .nodes + .values() + .filter(|n| n.handler_type() == Some("parallel")) + .map(|n| n.id.as_str()) + .collect(); + if parallel_ids.is_empty() { + return Vec::new(); + } + + let mut diagnostics = Vec::new(); + + // Branch edges (parallel node -> branch target) carrying an attribute + // that branch dispatch never reads. + for edge in &graph.edges { + if !parallel_ids.contains(edge.from.as_str()) { + continue; + } + for attr in BRANCH_IGNORED_ATTRS { + if !edge.attrs.contains_key(*attr) { + continue; + } + diagnostics.push(Diagnostic { + rule: self.name().to_string(), + severity: Severity::Warning, + message: format!( + "Edge {} -> {} sets '{attr}', which is ignored on parallel branch edges: branches receive the context snapshot taken when '{}' started", + edge.from, edge.to, edge.from, + ), + node_id: None, + edge: Some((edge.from.clone(), edge.to.clone())), + fix: Some(fix_message(attr, std::slice::from_ref(&edge.from))), + ..Diagnostic::default() + }); + } + } + + // Branch target nodes carrying such an attribute — but only when every + // incoming edge comes from a parallel node. A node that is also + // reachable through a normal edge resolves the attribute on that path, + // so it is not inert there. + let branch_targets: BTreeSet<&str> = graph + .edges + .iter() + .filter(|e| parallel_ids.contains(e.from.as_str())) + .map(|e| e.to.as_str()) + .collect(); + for target in branch_targets { + let only_branch_entries = graph + .edges + .iter() + .filter(|e| e.to == target) + .all(|e| parallel_ids.contains(e.from.as_str())); + if !only_branch_entries { + continue; + } + let Some(node) = graph.nodes.get(target) else { + continue; + }; + let parents: Vec = graph + .edges + .iter() + .filter(|e| e.to == target && parallel_ids.contains(e.from.as_str())) + .map(|e| e.from.clone()) + .collect::>() + .into_iter() + .collect(); + for attr in BRANCH_IGNORED_ATTRS { + if !node.attrs.contains_key(*attr) { + continue; + } + diagnostics.push(Diagnostic { + rule: self.name().to_string(), + severity: Severity::Warning, + message: format!( + "Node '{}' sets '{attr}', but it only runs as a parallel branch (of {}), where '{attr}' is ignored: branches receive the context snapshot taken when the parallel node started", + node.id, + quoted_list(&parents), + ), + node_id: Some(node.id.clone()), + edge: None, + fix: Some(fix_message(attr, &parents)), + ..Diagnostic::default() + }); + } + } + + diagnostics + } +} + +#[cfg(test)] +mod tests { + use fabro_graphviz::graph::{AttrValue, Edge, Graph, Node}; + + use super::Rule; + use crate::rules::test_support::minimal_graph; + use crate::{LintRule, Severity}; + + fn shaped_node(id: &str, shape: &str) -> Node { + let mut node = Node::new(id); + node.attrs + .insert("shape".to_string(), AttrValue::String(shape.to_string())); + node + } + + /// start -> fork -> {branch_a, branch_b} -> merge -> exit + fn parallel_graph() -> Graph { + let mut g = minimal_graph(); + g.nodes + .insert("fork".to_string(), shaped_node("fork", "component")); + g.nodes + .insert("branch_a".to_string(), shaped_node("branch_a", "tab")); + g.nodes + .insert("branch_b".to_string(), shaped_node("branch_b", "tab")); + g.nodes + .insert("merge".to_string(), shaped_node("merge", "tripleoctagon")); + g.edges = vec![ + Edge::new("start", "fork"), + Edge::new("fork", "branch_a"), + Edge::new("fork", "branch_b"), + Edge::new("branch_a", "merge"), + Edge::new("branch_b", "merge"), + Edge::new("merge", "exit"), + ]; + g + } + + #[test] + fn warns_on_fidelity_on_branch_node() { + let mut g = parallel_graph(); + g.nodes + .get_mut("branch_a") + .expect("graph has branch_a") + .attrs + .insert( + "fidelity".to_string(), + AttrValue::String("truncate".to_string()), + ); + let d = Rule.apply(&g); + assert_eq!(d.len(), 1); + assert_eq!(d[0].severity, Severity::Warning); + assert_eq!(d[0].node_id.as_deref(), Some("branch_a")); + assert!(d[0].message.contains("'fidelity'")); + assert!(d[0].fix.as_deref().is_some_and(|f| f.contains("'fork'"))); + } + + #[test] + fn warns_on_thread_id_on_branch_edge() { + let mut g = parallel_graph(); + g.edges[1].attrs.insert( + "thread_id".to_string(), + AttrValue::String("impl".to_string()), + ); + let d = Rule.apply(&g); + assert_eq!(d.len(), 1); + assert_eq!( + d[0].edge, + Some(("fork".to_string(), "branch_a".to_string())) + ); + assert!(d[0].message.contains("'thread_id'")); + assert_eq!( + d[0].fix.as_deref(), + Some( + "Remove 'thread_id': parallel branches inherit the thread resolved when the parallel node started" + ) + ); + } + + #[test] + fn accepts_fidelity_on_the_parallel_node_itself() { + let mut g = parallel_graph(); + g.nodes + .get_mut("fork") + .expect("graph has fork") + .attrs + .insert( + "fidelity".to_string(), + AttrValue::String("truncate".to_string()), + ); + assert!(Rule.apply(&g).is_empty()); + } + + #[test] + fn accepts_fidelity_on_branch_node_also_reached_by_normal_edge() { + let mut g = parallel_graph(); + // branch_a is also a normal successor of merge, so fidelity resolves + // on that path and is not inert. + g.edges.push(Edge::new("merge", "branch_a")); + g.nodes + .get_mut("branch_a") + .expect("graph has branch_a") + .attrs + .insert( + "fidelity".to_string(), + AttrValue::String("truncate".to_string()), + ); + assert!(Rule.apply(&g).is_empty()); + } + + #[test] + fn names_every_parallel_parent_of_a_shared_branch_node() { + let mut g = parallel_graph(); + g.nodes + .insert("fork2".to_string(), shaped_node("fork2", "component")); + g.edges.push(Edge::new("start", "fork2")); + g.edges.push(Edge::new("fork2", "branch_a")); + g.nodes + .get_mut("branch_a") + .expect("graph has branch_a") + .attrs + .insert( + "fidelity".to_string(), + AttrValue::String("truncate".to_string()), + ); + let d = Rule.apply(&g); + assert_eq!(d.len(), 1); + assert!(d[0].message.contains("'fork', 'fork2'")); + let fix = d[0].fix.as_deref().expect("diagnostic has a fix"); + assert!(fix.contains("'fork', 'fork2'")); + assert!(fix.contains("parallel nodes")); + } + + #[test] + fn accepts_graph_without_parallel_nodes() { + let mut g = minimal_graph(); + let mut node = shaped_node("work", "tab"); + node.attrs.insert( + "fidelity".to_string(), + AttrValue::String("truncate".to_string()), + ); + g.nodes.insert("work".to_string(), node); + assert!(Rule.apply(&g).is_empty()); + } +} diff --git a/lib/crates/fabro-workflow/src/event/convert.rs b/lib/crates/fabro-workflow/src/event/convert.rs index 8aa10c200..b2ba1d715 100644 --- a/lib/crates/fabro-workflow/src/event/convert.rs +++ b/lib/crates/fabro-workflow/src/event/convert.rs @@ -3,6 +3,7 @@ use ::fabro_types::{ }; use chrono::Utc; use fabro_agent::{AgentEvent, SandboxEvent, SkillActivationSource}; +use fabro_model::UsdMicros; use uuid::Uuid; use super::Event; @@ -611,14 +612,18 @@ fn event_body_from_event(event: &Event) -> EventBody { text, model, usage, + cost_usd, + cost_source, tool_call_count, context_window, } => { - let billing = billed_token_counts_from_llm(usage); + let billing = billed_token_counts_from_llm(usage) + .with_reported_cost(cost_usd.map(UsdMicros::from_usd)); EventBody::AgentMessage(fabro_types::AgentMessageProps { text: text.clone(), model: model.clone(), billing, + cost_source: *cost_source, tool_call_count: *tool_call_count, visit: *visit, message: None, @@ -2116,6 +2121,8 @@ mod tests { speed: None, }, usage: LlmTokenCounts::default(), + cost_usd: None, + cost_source: None, tool_call_count: 0, context_window: None, }, @@ -2148,6 +2155,8 @@ mod tests { output_tokens: 34, ..LlmTokenCounts::default() }, + cost_usd: None, + cost_source: None, tool_call_count: 0, context_window: None, }, @@ -2166,6 +2175,43 @@ mod tests { assert_eq!(message.billing.total_usd_micros, None); } + #[test] + fn agent_assistant_message_preserves_provider_cost() { + let stored = to_run_event(&fixtures::RUN_1, &Event::Agent { + stage: "code".to_string(), + visit: 1, + event: AgentEvent::AssistantMessage { + text: "ok".to_string(), + model: ModelRef { + provider: ProviderId::new("openrouter"), + model_id: "openai/gpt-5.4".to_string(), + speed: None, + }, + usage: LlmTokenCounts { + input_tokens: 12, + output_tokens: 34, + ..LlmTokenCounts::default() + }, + cost_usd: Some(0.125), + cost_source: Some(fabro_model::CostSource::Authoritative), + tool_call_count: 0, + context_window: None, + }, + session_id: Some("ses_agent".to_string()), + parent_session_id: None, + tool_call_id: None, + }); + + let EventBody::AgentMessage(message) = stored.body else { + panic!("expected agent message body"); + }; + assert_eq!(message.billing.total_usd_micros, Some(125_000)); + assert_eq!( + message.cost_source, + Some(fabro_model::CostSource::Authoritative) + ); + } + #[test] fn agent_assistant_message_copies_context_window_to_props() { let context_window = ::fabro_types::StageContextWindowProjection { @@ -2196,6 +2242,8 @@ mod tests { speed: None, }, usage: LlmTokenCounts::default(), + cost_usd: None, + cost_source: None, tool_call_count: 0, context_window: Some(context_window), }, diff --git a/lib/crates/fabro-workflow/src/handler/llm/api.rs b/lib/crates/fabro-workflow/src/handler/llm/api.rs index 9cb5469bb..573542d9f 100644 --- a/lib/crates/fabro-workflow/src/handler/llm/api.rs +++ b/lib/crates/fabro-workflow/src/handler/llm/api.rs @@ -20,7 +20,7 @@ use fabro_llm::types::{ use fabro_mcp::config::McpServerSettings; #[cfg(test)] use fabro_model::catalog::LlmCatalogSettings; -use fabro_model::{AgentProfileKind, Catalog, FallbackTarget, ModelRef, ProviderId}; +use fabro_model::{AgentProfileKind, Catalog, FallbackTarget, ModelRef, ProviderId, UsdMicros}; use fabro_types::settings::run::RunModelControls; use fabro_types::{PermissionLevel, RunId, SessionCapability, StageId, StageTiming}; use serde::de::DeserializeOwned; @@ -1060,6 +1060,7 @@ impl CodergenBackend for AgentApiBackend { .map(structured_output::prompt_response_format); let mut repair_attempts = 0_i64; let mut total_usage = TokenCounts::default(); + let mut total_cost = None; let mut inference_duration = Duration::ZERO; loop { @@ -1095,6 +1096,10 @@ impl CodergenBackend for AgentApiBackend { inference_duration = inference_duration.saturating_add(inference_start.elapsed()); let completion = completion_result?; total_usage += completion.response.usage.clone(); + UsdMicros::accumulate( + &mut total_cost, + completion.response.cost_usd.map(UsdMicros::from_usd), + ); let response_text = completion.response.text(); let validation_error = if let Some(schema) = &output_schema { @@ -1122,7 +1127,8 @@ impl CodergenBackend for AgentApiBackend { self.catalog.as_ref(), &completion.model, &total_usage, - )?; + )? + .with_reported_cost(total_cost); return Ok(CodergenResult::Text { text: response_text, @@ -1216,6 +1222,7 @@ impl CodergenBackend for AgentApiBackend { ); let mut total_usage = TokenCounts::default(); + let mut total_cost = None; let mut inference_duration = Duration::ZERO; let mut tool_duration = Duration::ZERO; @@ -1278,6 +1285,7 @@ impl CodergenBackend for AgentApiBackend { tool_duration = tool_duration.saturating_add(timing.tool); if process_result.is_ok() { total_usage += session.last_input_usage(); + UsdMicros::accumulate(&mut total_cost, session.last_input_cost()); } process_result } @@ -1413,6 +1421,7 @@ impl CodergenBackend for AgentApiBackend { match process_result { Ok(()) => { total_usage += session.last_input_usage(); + UsdMicros::accumulate(&mut total_cost, session.last_input_cost()); succeeded = true; break; } @@ -1484,6 +1493,7 @@ impl CodergenBackend for AgentApiBackend { match repair_result { Ok(()) => { total_usage += session.last_input_usage(); + UsdMicros::accumulate(&mut total_cost, session.last_input_cost()); repair_attempts += 1; response = last_assistant_response(&session); } @@ -1519,7 +1529,8 @@ impl CodergenBackend for AgentApiBackend { speed: billing_controls.speed, }, &total_usage, - )?; + )? + .with_reported_cost(total_cost); // Collect files_touched from the shared tracking state. let (files_touched, last_file_touched) = file_tracking_snapshot(&file_tracking); @@ -2780,7 +2791,11 @@ reasoning = false .body_excludes(r#""role":"assistant""#); then.status(200) .header("content-type", "application/json") - .json_body(chat_completion_response("not json", 10, 1)); + .json_body({ + let mut response = chat_completion_response("not json", 10, 1); + response["usage"]["cost"] = serde_json::json!(0.04); + response + }); }); let repair = server.mock(|when, then| { when.method(POST) @@ -2791,7 +2806,11 @@ reasoning = false .body_includes("output_schema"); then.status(200) .header("content-type", "application/json") - .json_body(chat_completion_response(r#"{"passed":true}"#, 11, 2)); + .json_body({ + let mut response = chat_completion_response(r#"{"passed":true}"#, 11, 2); + response["usage"]["cost"] = serde_json::json!(0.06); + response + }); }); let backend = mock_api_backend(&server); let mut node = Node::new("audit"); @@ -2828,6 +2847,7 @@ reasoning = false let usage = usage.expect("usage should be aggregated"); assert_eq!(usage.tokens().input_tokens, 21); assert_eq!(usage.tokens().output_tokens, 3); + assert_eq!(usage.total_usd_micros, Some(100_000)); } #[tokio::test] diff --git a/lib/crates/fabro-workflow/src/outcome.rs b/lib/crates/fabro-workflow/src/outcome.rs index 7d3c7096c..330969de1 100644 --- a/lib/crates/fabro-workflow/src/outcome.rs +++ b/lib/crates/fabro-workflow/src/outcome.rs @@ -149,7 +149,7 @@ fn token_counts_from_llm_usage(usage: &LlmTokenCounts) -> TokenCounts { mod tests { use fabro_llm::types::TokenCounts; use fabro_model::catalog::LlmCatalogSettings; - use fabro_model::{Catalog, ModelRef, ProviderId, Speed}; + use fabro_model::{Catalog, ModelRef, ProviderId, Speed, UsdMicros}; use super::{OutcomeExt, billed_model_usage_from_llm}; @@ -182,6 +182,24 @@ mod tests { assert_eq!(billed.tokens().reasoning_tokens, 25_000); } + #[test] + fn response_cost_overrides_catalog_estimate() { + let usage = TokenCounts { + input_tokens: 11, + output_tokens: 7, + ..TokenCounts::default() + }; + let billed = billed_model_usage_from_llm( + Catalog::builtin(), + &model_ref(ProviderId::openai(), "gpt-5.4", None), + &usage, + ) + .unwrap() + .with_reported_cost(Some(UsdMicros(125_000))); + + assert_eq!(billed.total_usd_micros, Some(125_000)); + } + #[test] fn retry_classify_marks_failed_outcome_with_retry_request() { let outcome = crate::outcome::Outcome::retry_classify("timeout"); diff --git a/lib/crates/fabro-workflow/tests/it/integration.rs b/lib/crates/fabro-workflow/tests/it/integration.rs index 086b443af..2b817434a 100644 --- a/lib/crates/fabro-workflow/tests/it/integration.rs +++ b/lib/crates/fabro-workflow/tests/it/integration.rs @@ -2319,6 +2319,121 @@ reasoning = false ); } +#[tokio::test] +async fn workflow_persists_authoritative_openrouter_cost_for_agent_stage() { + use fabro_auth::EnvCredentialSource; + use fabro_workflow::steering_hub::SteeringHub; + use httpmock::Method::POST; + use httpmock::MockServer; + + const AUTHORITATIVE_COST_USD: f64 = 0.125; + const AUTHORITATIVE_COST_USD_MICROS: i64 = 125_000; + + let server = MockServer::start_async().await; + let text_chunk = serde_json::json!({ + "id": "chatcmpl_authoritative_cost", + "model": "openai/gpt-5.4", + "choices": [{ + "delta": {"content": "done"}, + "finish_reason": null + }] + }); + let usage_chunk = serde_json::json!({ + "id": "chatcmpl_authoritative_cost", + "model": "openai/gpt-5.4", + "choices": [], + "usage": { + "prompt_tokens": 11, + "completion_tokens": 7, + "total_tokens": 18, + "cost": AUTHORITATIVE_COST_USD + } + }); + let response = format!("data: {text_chunk}\n\ndata: {usage_chunk}\n\ndata: [DONE]\n\n"); + let completion_mock = server + .mock_async(|when, then| { + when.method(POST) + .path("/chat/completions") + .body_includes(r#""stream":true"#) + .body_includes("Report completion"); + then.status(200) + .header("content-type", "text/event-stream") + .body(response); + }) + .await; + + let settings: LlmCatalogSettings = toml::from_str(&format!( + r#" +[providers.openrouter] +enabled = true +base_url = "{}" +"#, + server.base_url() + )) + .expect("test catalog should parse"); + let catalog = Arc::new(Catalog::from_builtin_with_overrides(&settings).unwrap()); + let source = Arc::new(EnvCredentialSource::with_env_lookup(Arc::new(|name| { + (name == "OPENROUTER_API_KEY").then(|| "sk-test".to_string()) + }))); + let backend = AgentApiBackend::new_with_catalog( + "openai/gpt-5.4".to_string(), + ProviderId::from("openrouter"), + Vec::new(), + source, + Arc::new(SteeringHub::new(Arc::new(Emitter::default()))), + catalog, + ); + + let mut graph = make_graph_with_start_exit("AuthoritativeOpenRouterCost"); + let mut work = Node::new("work"); + work.attrs.insert( + "prompt".to_string(), + AttrValue::String("Report completion".to_string()), + ); + graph.nodes.insert("work".to_string(), work); + graph.edges.push(Edge::new("start", "work")); + graph.edges.push(Edge::new("work", "exit")); + + let mut registry = HandlerRegistry::new(Box::new(AgentHandler::new(Some(Box::new(backend))))); + registry.register("start", Box::new(StartHandler)); + registry.register("exit", Box::new(ExitHandler)); + + let dir = tempfile::tempdir().unwrap(); + let engine = WorkflowRunner::new(registry, Arc::new(Emitter::default()), local_env()); + let run_options = RunOptions { + settings: WorkflowSettings::default(), + run_dir: dir.path().to_path_buf(), + cancel_token: CancellationToken::new(), + run_id: test_run_id("authoritative-openrouter-cost"), + labels: std::collections::HashMap::new(), + workflow_slug: None, + github_app: None, + base_branch: None, + display_base_sha: None, + pre_run_git: None, + fork_source_ref: None, + git: None, + }; + + let (outcome, state) = engine + .run_with_state(&graph, &run_options) + .await + .expect("workflow execution should complete"); + assert_eq!(outcome.status, StageOutcome::Succeeded); + assert_eq!(completion_mock.calls_async().await, 1); + + let work = state + .stage(&fabro_types::StageId::new("work", 1)) + .expect("agent stage should be projected"); + assert_eq!(work.usage.input_tokens, 11); + assert_eq!(work.usage.output_tokens, 7); + assert_eq!( + work.usage.total_usd_micros, + Some(AUTHORITATIVE_COST_USD_MICROS), + "provider-reported usage.cost should override the catalog estimate" + ); +} + // --------------------------------------------------------------------------- // 12. Parallel fan-out / fan-in integration test (Gap #14) // ---------------------------------------------------------------------------