use jiff::Timestamp;
use std::collections::BTreeMap;
use serde::{Deserialize, Serialize};
use super::{CostSource, RunCost, TokenUsage};
use crate::agent::AgentName;
use crate::flight::{ItineraryId, RunId};
#[derive(Debug, Clone, Copy, Default, PartialEq, Deserialize, Serialize)]
#[serde(deny_unknown_fields)]
pub struct ModelRates {
#[serde(default)]
pub input_usd: f64,
#[serde(default)]
pub output_usd: f64,
#[serde(default)]
pub cache_read_usd: f64,
#[serde(default)]
pub cache_write_usd: f64,
}
impl ModelRates {
#[must_use]
pub fn cost_of(&self, usage: TokenUsage) -> Option<f64> {
let rates = [
(self.input_usd, usage.input),
(self.output_usd, usage.output),
(self.cache_read_usd, usage.cache_read),
(self.cache_write_usd, usage.cache_write),
];
let mut total = 0.0;
for (rate, tokens) in rates {
if !rate.is_finite() || rate < 0.0 {
return None;
}
#[allow(clippy::cast_precision_loss)]
let tokens = tokens as f64;
total += rate * tokens / 1_000_000.0;
}
total.is_finite().then_some(total)
}
}
#[derive(Debug, Clone, Default, Deserialize, Serialize)]
#[serde(transparent)]
pub struct RateCard {
models: BTreeMap<String, ModelRates>,
}
impl RateCard {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn with(mut self, model: impl Into<String>, rates: ModelRates) -> Self {
self.models.insert(model.into(), rates);
self
}
#[must_use]
pub fn rates_for(&self, model: &str) -> Option<&ModelRates> {
self.models.get(model)
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.models.is_empty()
}
#[must_use]
pub fn estimate(&self, model: &str, usage: TokenUsage) -> Option<f64> {
self.rates_for(model)?.cost_of(usage)
}
#[must_use]
pub fn price(
&self,
run: RunId,
itinerary: ItineraryId,
agent: AgentName,
model: Option<String>,
usage: TokenUsage,
) -> RunCost {
let estimate = model
.as_deref()
.filter(|_| !usage.is_empty())
.and_then(|model| self.estimate(model, usage));
let (usd, source) = match estimate {
Some(usd) => (usd, CostSource::RateCard),
None => (0.0, CostSource::Unreported),
};
RunCost {
run,
itinerary,
agent,
pipeline: None,
model,
usage,
usd,
source,
at: Timestamp::now(),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn opus() -> ModelRates {
ModelRates {
input_usd: 5.0,
output_usd: 25.0,
cache_read_usd: 0.5,
cache_write_usd: 6.25,
}
}
fn card() -> RateCard {
RateCard::new().with("claude-opus-5", opus())
}
fn usage() -> TokenUsage {
TokenUsage {
input: 1_000_000,
output: 1_000_000,
cache_read: 1_000_000,
cache_write: 1_000_000,
}
}
fn priced(card: &RateCard, model: Option<&str>, usage: TokenUsage) -> RunCost {
card.price(
RunId::generate(),
ItineraryId::generate(),
"analyst".into(),
model.map(ToOwned::to_owned),
usage,
)
}
#[test]
fn a_million_of_each_costs_the_sum_of_the_rates() {
let cost = card().estimate("claude-opus-5", usage()).expect("priced");
assert!((cost - 36.75).abs() < 1e-9, "got {cost}");
}
#[test]
fn cached_tokens_are_priced_apart_from_fresh_input() {
let fresh = card()
.estimate(
"claude-opus-5",
TokenUsage {
input: 1_000_000,
..TokenUsage::default()
},
)
.expect("priced");
let cached = card()
.estimate(
"claude-opus-5",
TokenUsage {
cache_read: 1_000_000,
..TokenUsage::default()
},
)
.expect("priced");
assert!((fresh - 5.0).abs() < 1e-9);
assert!((cached - 0.5).abs() < 1e-9);
assert!(cached < fresh);
}
#[test]
fn an_estimate_is_labelled_as_an_estimate() {
let cost = priced(&card(), Some("claude-opus-5"), usage());
assert_eq!(cost.source, CostSource::RateCard);
assert!(
!cost.source.is_measured(),
"an estimate must never count as measured"
);
}
#[test]
fn an_unknown_model_is_unreported_rather_than_free() {
let cost = priced(&card(), Some("some-new-model"), usage());
assert_eq!(cost.source, CostSource::Unreported);
assert!((cost.usd - 0.0).abs() < f64::EPSILON);
}
#[test]
fn a_run_with_no_model_cannot_be_priced() {
assert_eq!(
priced(&card(), None, usage()).source,
CostSource::Unreported
);
}
#[test]
fn a_run_with_no_tokens_cannot_be_priced() {
assert_eq!(
priced(&card(), Some("claude-opus-5"), TokenUsage::default()).source,
CostSource::Unreported
);
}
#[test]
fn a_nonsense_rate_is_refused_rather_than_propagated() {
for bad in [f64::NAN, f64::INFINITY, -1.0] {
let card = RateCard::new().with(
"broken",
ModelRates {
output_usd: bad,
..ModelRates::default()
},
);
assert_eq!(
card.estimate("broken", usage()),
None,
"{bad} must not produce a price"
);
}
}
#[test]
fn rates_parse_from_configuration() {
let card: RateCard = toml::from_str(
r"
[claude-opus-5]
input_usd = 5.0
output_usd = 25.0
cache_read_usd = 0.5
cache_write_usd = 6.25
",
)
.expect("a rate card parses");
assert!(!card.is_empty());
assert_eq!(card.rates_for("claude-opus-5"), Some(&opus()));
}
#[test]
fn an_unknown_rate_field_is_rejected() {
let error = toml::from_str::<RateCard>(
r"
[claude-opus-5]
output_used = 25.0
",
);
assert!(error.is_err());
}
}