use std::collections::BTreeMap;
use crate::provider::Usage;
use crate::state::ProviderCall;
pub const MICROS_PER_UNIT: u64 = 1_000_000;
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct Price {
pub input: u64,
pub output: u64,
pub cache_read: u64,
pub cache_write: u64,
pub per_server_tool_request: u64,
}
impl Price {
pub const ZERO: Self = Self {
input: 0,
output: 0,
cache_read: 0,
cache_write: 0,
per_server_tool_request: 0,
};
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct PriceTable {
as_of: String,
prices: BTreeMap<String, Price>,
}
impl PriceTable {
pub fn new(as_of: impl Into<String>) -> Self {
Self {
as_of: as_of.into(),
prices: BTreeMap::new(),
}
}
#[must_use]
pub fn with(mut self, model: impl Into<String>, price: Price) -> Self {
self.prices.insert(model.into(), price);
self
}
pub fn as_of(&self) -> &str {
&self.as_of
}
pub fn price(&self, model: &str) -> Option<Price> {
self.prices.get(model).copied()
}
pub fn cost_micros(&self, model: &str, usage: &Usage) -> Option<u64> {
let p = self.price(model)?;
let fresh_input = usage
.prompt_tokens
.saturating_sub(usage.cache_read_tokens)
.saturating_sub(usage.cache_write_tokens);
let per_million = |tokens: u64, price: u64| tokens as u128 * price as u128;
let mtok = per_million(fresh_input, p.input)
+ per_million(usage.completion_tokens, p.output)
+ per_million(usage.cache_read_tokens, p.cache_read)
+ per_million(usage.cache_write_tokens, p.cache_write);
let requests = usage.server_tool_requests as u128 * p.per_server_tool_request as u128;
let micros = (mtok + 500_000) / 1_000_000 + requests;
Some(micros.min(u64::MAX as u128) as u64)
}
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct Spend {
pub key: String,
pub calls: u64,
pub usage: Usage,
pub cost_micros: u64,
pub unpriced_calls: u64,
}
pub(crate) fn group(key: impl Into<String>, calls: &[&ProviderCall], prices: &PriceTable) -> Spend {
let mut spend = Spend {
key: key.into(),
calls: calls.len() as u64,
..Default::default()
};
for call in calls {
let Some(usage) = call.usage else {
spend.unpriced_calls += 1;
continue;
};
spend.usage.prompt_tokens += usage.prompt_tokens;
spend.usage.completion_tokens += usage.completion_tokens;
spend.usage.total_tokens += usage.total_tokens;
spend.usage.cache_read_tokens += usage.cache_read_tokens;
spend.usage.cache_write_tokens += usage.cache_write_tokens;
spend.usage.reasoning_tokens += usage.reasoning_tokens;
spend.usage.server_tool_requests += usage.server_tool_requests;
match call
.model
.as_deref()
.and_then(|m| prices.cost_micros(m, &usage))
{
Some(micros) => spend.cost_micros += micros,
None => spend.unpriced_calls += 1,
}
}
spend
}
#[cfg(test)]
mod tests {
use super::*;
fn table() -> PriceTable {
PriceTable::new("2026-07-29").with(
"m",
Price {
input: 3_000_000,
output: 15_000_000,
cache_read: 300_000,
cache_write: 3_750_000,
per_server_tool_request: 10_000,
},
)
}
#[test]
fn a_hand_computed_million_token_figure_comes_out_exact() {
let usage = Usage {
prompt_tokens: 1_000_000,
completion_tokens: 1_000_000,
total_tokens: 2_000_000,
cache_read_tokens: 500_000,
cache_write_tokens: 100_000,
reasoning_tokens: 200_000,
server_tool_requests: 3,
};
assert_eq!(
table().cost_micros("m", &usage),
Some(1_200_000 + 15_000_000 + 150_000 + 375_000 + 30_000)
);
}
#[test]
fn an_unpriced_model_is_unknown_rather_than_free() {
let usage = Usage {
prompt_tokens: 10,
completion_tokens: 10,
total_tokens: 20,
..Default::default()
};
assert_eq!(table().cost_micros("not-in-the-table", &usage), None);
assert_eq!(table().cost_micros("m", &usage), Some(180));
}
#[test]
fn cache_figures_larger_than_the_prompt_saturate_rather_than_underflow() {
let usage = Usage {
prompt_tokens: 10,
cache_read_tokens: 900,
total_tokens: 10,
..Default::default()
};
assert_eq!(table().cost_micros("m", &usage), Some(270));
}
#[test]
fn rounding_is_once_at_the_end_and_half_up() {
let half = PriceTable::new("x").with(
"m",
Price {
input: 1,
..Price::ZERO
},
);
let usage = Usage {
prompt_tokens: 500_000,
total_tokens: 500_000,
..Default::default()
};
assert_eq!(half.cost_micros("m", &usage), Some(1));
let under = Usage {
prompt_tokens: 499_999,
total_tokens: 499_999,
..Default::default()
};
assert_eq!(half.cost_micros("m", &under), Some(0));
}
}