mirror of
https://github.com/fabro-sh/fabro.git
synced 2026-10-10 03:30:59 +00:00
Preserve provider-reported workflow costs
This commit is contained in:
parent
f8ed856959
commit
7f25689fb6
13 changed files with 211 additions and 16 deletions
|
|
@ -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,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,
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
}
|
||||
|
|
|
|||
|
|
@ -528,6 +528,8 @@ mod tests {
|
|||
speed: None,
|
||||
},
|
||||
usage: TokenCounts::default(),
|
||||
cost_usd: None,
|
||||
cost_source: None,
|
||||
tool_call_count: 0,
|
||||
context_window: None,
|
||||
})
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
},
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()),
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
},
|
||||
|
|
|
|||
|
|
@ -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<UsdMicros>, cost: Option<UsdMicros>) {
|
||||
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]
|
||||
|
|
|
|||
|
|
@ -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<UsdMicros>,
|
||||
) -> Result<BilledModelUsage, Error> {
|
||||
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<Speed>) -> 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");
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue