Merge pull request #594 from fabro-sh/preserve-provider-costs

Preserve provider-reported workflow costs
This commit is contained in:
Bryan Helmkamp 2026-07-23 08:10:18 -04:00 • committed by GitHub
commit 580ee156b2
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
15 changed files with 455 additions and 43 deletions

View file

@ -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<Arc<dyn CompletionCoordinator>>,
last_input_timing: SessionInputTiming,
last_input_usage: TokenCounts,
last_input_cost: Option<UsdMicros>,
}
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<UsdMicros> {
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<UsdMicros>,
) -> 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,

View file

@ -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<f64>,
/// Provenance of `cost_usd`.
#[serde(default, skip_serializing_if = "Option::is_none")]
cost_source: Option<CostSource>,
tool_call_count: usize,
#[serde(default, skip_serializing_if = "Option::is_none")]
context_window: Option<StageContextWindowProjection>,
@ -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"),
}

View file

@ -528,6 +528,8 @@ mod tests {
speed: None,
},
usage: TokenCounts::default(),
cost_usd: None,
cost_source: None,
tool_call_count: 0,
context_window: None,
})

View file

@ -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<Self>, cost: Option<Self>) {
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<i64>, cost: Option<i64>) {
let mut typed_total = (*total).map(UsdMicros);
UsdMicros::accumulate(&mut typed_total, cost.map(UsdMicros));
*total = typed_total.map(|value| value.0);
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub struct PricePerMTok {
pub usd_micros: i64,
@ -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<UsdMicros>) -> 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<UsdMicros>) -> 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
@ -730,6 +765,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>(),
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(
@ -863,6 +937,31 @@ cache_input_cost_per_mtok = 0.3
assert_eq!(counts.total_usd_micros, Some(150));
}
#[test]
fn billed_token_counts_cost_rollups_saturate() {
let billed = [
billed_usage(0, 0, Some(i64::MAX)),
billed_usage(0, 0, Some(1)),
];
assert_eq!(
BilledTokenCounts::from_billed_usage(&billed).total_usd_micros,
Some(i64::MAX)
);
let mut counts = BilledTokenCounts {
total_usd_micros: Some(i64::MAX),
..BilledTokenCounts::default()
};
counts.add_counts(&BilledTokenCounts {
total_usd_micros: Some(1),
..BilledTokenCounts::default()
});
assert_eq!(counts.total_usd_micros, Some(i64::MAX));
counts.add_billed_usage(&billed_usage(0, 0, Some(1)));
assert_eq!(counts.total_usd_micros, Some(i64::MAX));
}
#[test]
fn billed_token_counts_replace_with_billed_usage_discards_previous_values() {
let mut counts = BilledTokenCounts {

View file

@ -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,

View file

@ -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,

View file

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

View file

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

View file

@ -3788,6 +3788,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,

View file

@ -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<CostSource>,
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()),

View file

@ -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,

View file

@ -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),
},

View file

@ -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;
@ -1058,6 +1058,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 {
@ -1093,6 +1094,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 {
@ -1120,7 +1125,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,
@ -1214,6 +1220,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;
@ -1276,6 +1283,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
}
@ -1411,6 +1419,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;
}
@ -1482,6 +1491,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);
}
@ -1517,7 +1527,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);
@ -2778,7 +2789,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)
@ -2789,7 +2804,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");
@ -2826,6 +2845,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]

View file

@ -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");

View file

@ -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)
// ---------------------------------------------------------------------------