diff --git a/lib/crates/fabro-agent/src/session.rs b/lib/crates/fabro-agent/src/session.rs index 2bd840bda..93b2b5c7b 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,9 @@ impl Session { &usage, )); *usage_accumulator += usage.clone(); + if let Some(cost_usd) = response.cost_usd { + *cost_accumulator.get_or_insert_default() += UsdMicros::from_usd(cost_usd); + } self.history.push(Message::Assistant { content: text.clone(), @@ -1729,6 +1755,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 +2304,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 +4056,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-model/src/billing.rs b/lib/crates/fabro-model/src/billing.rs index d75992d7a..9754e8908 100644 --- a/lib/crates/fabro-model/src/billing.rs +++ b/lib/crates/fabro-model/src/billing.rs @@ -47,6 +47,15 @@ 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(), + )) + } +} + impl std::ops::Add for UsdMicros { type Output = Self; @@ -76,7 +85,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, } } 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 8304bf1ec..5b9d36270 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/tests.rs b/lib/crates/fabro-server/src/server/tests.rs index f7fed5606..7314693dd 100644 --- a/lib/crates/fabro-server/src/server/tests.rs +++ b/lib/crates/fabro-server/src/server/tests.rs @@ -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), }, diff --git a/lib/crates/fabro-store/src/run_state.rs b/lib/crates/fabro-store/src/run_state.rs index ec67aa4f0..4f8d55dd0 100644 --- a/lib/crates/fabro-store/src/run_state.rs +++ b/lib/crates/fabro-store/src/run_state.rs @@ -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, diff --git a/lib/crates/fabro-types/src/run_event/agent.rs b/lib/crates/fabro-types/src/run_event/agent.rs index 64b2ede3d..d61ccfb14 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 e76558f3e..cb1c38b81 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-workflow/src/event/convert.rs b/lib/crates/fabro-workflow/src/event/convert.rs index 5f0fc9959..b00ba0251 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 mut billing = billed_token_counts_from_llm(usage); + billing.total_usd_micros = cost_usd.map(|cost| UsdMicros::from_usd(cost).0); 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 c1f27a7a3..feb5cc4f5 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; @@ -40,7 +40,7 @@ use crate::context::WorkflowContext; use crate::context::keys::Fidelity; use crate::error::Error; use crate::event::{Emitter, Event, StageScope}; -use crate::outcome::billed_model_usage_from_llm; +use crate::outcome::billed_model_usage_from_llm_with_cost; use crate::services::FabroRunToolServices; use crate::steering_hub::{ActiveControlHandle, SteeringHub}; @@ -601,6 +601,12 @@ struct OneShotCompletion { model: ModelRef, } +fn add_cost(total: &mut Option, cost: Option) { + if let Some(cost) = cost { + *total.get_or_insert_default() += cost; + } +} + impl AgentApiBackend { #[must_use] pub fn new( @@ -1058,6 +1064,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 +1100,10 @@ impl CodergenBackend for AgentApiBackend { inference_duration = inference_duration.saturating_add(inference_start.elapsed()); let completion = completion_result?; total_usage += completion.response.usage.clone(); + add_cost( + &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 { @@ -1116,10 +1127,11 @@ impl CodergenBackend for AgentApiBackend { continue; } - let stage_usage = billed_model_usage_from_llm( + let stage_usage = billed_model_usage_from_llm_with_cost( self.catalog.as_ref(), &completion.model, &total_usage, + total_cost, )?; return Ok(CodergenResult::Text { @@ -1214,6 +1226,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 +1289,7 @@ impl CodergenBackend for AgentApiBackend { tool_duration = tool_duration.saturating_add(timing.tool); if process_result.is_ok() { total_usage += session.last_input_usage(); + add_cost(&mut total_cost, session.last_input_cost()); } process_result } @@ -1411,6 +1425,7 @@ impl CodergenBackend for AgentApiBackend { match process_result { Ok(()) => { total_usage += session.last_input_usage(); + add_cost(&mut total_cost, session.last_input_cost()); succeeded = true; break; } @@ -1482,6 +1497,7 @@ impl CodergenBackend for AgentApiBackend { match repair_result { Ok(()) => { total_usage += session.last_input_usage(); + add_cost(&mut total_cost, session.last_input_cost()); repair_attempts += 1; response = last_assistant_response(&session); } @@ -1509,7 +1525,7 @@ impl CodergenBackend for AgentApiBackend { } let billing_controls = self.resolve_effective_request_controls(node)?; - let stage_usage = billed_model_usage_from_llm( + let stage_usage = billed_model_usage_from_llm_with_cost( self.catalog.as_ref(), &ModelRef { provider: session.provider_id(), @@ -1517,6 +1533,7 @@ impl CodergenBackend for AgentApiBackend { speed: billing_controls.speed, }, &total_usage, + total_cost, )?; // Collect files_touched from the shared tracking state. @@ -2778,7 +2795,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 +2810,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 +2851,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 532d26860..2b84f7bc5 100644 --- a/lib/crates/fabro-workflow/src/outcome.rs +++ b/lib/crates/fabro-workflow/src/outcome.rs @@ -3,7 +3,7 @@ pub use fabro_core::outcome::{ }; use fabro_llm::types::TokenCounts as LlmTokenCounts; use fabro_model::{ - BilledTokenCounts, Catalog, ModelBillingInput, ModelRef, ModelUsage, TokenCounts, + BilledTokenCounts, Catalog, ModelBillingInput, ModelRef, ModelUsage, TokenCounts, UsdMicros, }; pub use fabro_types::BilledModelUsage; @@ -39,6 +39,19 @@ pub fn billed_model_usage_from_llm( }) } +pub fn billed_model_usage_from_llm_with_cost( + catalog: &Catalog, + model: &ModelRef, + usage: &LlmTokenCounts, + total_cost: Option, +) -> Result { + let mut billed = billed_model_usage_from_llm(catalog, model, usage)?; + if let Some(total_cost) = total_cost { + billed.total_usd_micros = Some(total_cost.0); + } + Ok(billed) +} + #[must_use] pub fn billed_token_counts_from_llm(usage: &LlmTokenCounts) -> BilledTokenCounts { let tokens = token_counts_from_llm_usage(usage); @@ -149,9 +162,9 @@ 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}; + use super::{OutcomeExt, billed_model_usage_from_llm, billed_model_usage_from_llm_with_cost}; fn model_ref(provider: ProviderId, model_id: &str, speed: Option) -> ModelRef { ModelRef { @@ -182,6 +195,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_with_cost( + Catalog::builtin(), + &model_ref(ProviderId::openai(), "gpt-5.4", None), + &usage, + Some(UsdMicros(125_000)), + ) + .unwrap(); + + 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");