use std::collections::HashMap;
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Price {
pub input_per_1m: u64,
pub output_per_1m: u64,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
pub struct TokenUsage {
pub prompt_tokens: usize,
pub completion_tokens: usize,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub model: Option<String>,
}
impl TokenUsage {
pub fn total(&self) -> usize {
self.prompt_tokens.saturating_add(self.completion_tokens)
}
}
#[derive(Debug, Clone, Default)]
pub struct PriceBook {
rates: HashMap<String, Price>,
}
impl PriceBook {
pub fn default_set() -> Self {
let mut book = Self::default();
book.set(
"claude-3-5-sonnet",
Price {
input_per_1m: 3,
output_per_1m: 15,
},
);
book.set(
"claude-3-5-haiku",
Price {
input_per_1m: 80,
output_per_1m: 400,
},
);
book.set(
"gpt-4o-mini",
Price {
input_per_1m: 15,
output_per_1m: 60,
},
);
book.set(
"gpt-4o",
Price {
input_per_1m: 250,
output_per_1m: 1000,
},
);
book.set(
"deepseek-chat",
Price {
input_per_1m: 27,
output_per_1m: 110,
},
);
book
}
pub fn set(&mut self, model: impl Into<String>, price: Price) {
self.rates.insert(model.into(), price);
}
pub fn get(&self, model: &str) -> Option<Price> {
self.rates.get(model).copied()
}
pub fn estimate(&self, usage: &TokenUsage) -> Option<f64> {
let model = usage.model.as_deref()?;
self.estimate_cost(usage.prompt_tokens, usage.completion_tokens, model)
}
pub fn estimate_cost(
&self,
prompt_tokens: usize,
completion_tokens: usize,
model: &str,
) -> Option<f64> {
let p = self.rates.get(model)?;
let usd = prompt_tokens as f64 / 1e6 * (p.input_per_1m as f64 / 100.0)
+ completion_tokens as f64 / 1e6 * (p.output_per_1m as f64 / 100.0);
Some(usd)
}
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct OverallCost {
pub prompt_tokens: usize,
pub completion_tokens: usize,
pub total_tokens: usize,
pub cost_usd: Option<f64>,
}
impl OverallCost {
pub fn accumulate(&mut self, usage: &TokenUsage, book: &PriceBook) {
self.prompt_tokens = self.prompt_tokens.saturating_add(usage.prompt_tokens);
self.completion_tokens = self
.completion_tokens
.saturating_add(usage.completion_tokens);
self.total_tokens = self.total_tokens.saturating_add(usage.total());
if let Some(usd) = book.estimate(usage) {
self.cost_usd = Some(self.cost_usd.unwrap_or(0.0) + usd);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn cheap_book() -> PriceBook {
let mut b = PriceBook::default();
b.set(
"test-model",
Price {
input_per_1m: 300,
output_per_1m: 1500,
},
); b
}
#[test]
fn estimate_is_pure_and_per_1m_scaled() {
let b = cheap_book();
let usd = b.estimate_cost(1_000_000, 1_000_000, "test-model").unwrap();
assert!((usd - 18.0).abs() < 1e-9, "got {usd}");
let small = b.estimate_cost(1000, 1000, "test-model").unwrap();
assert!((small - 0.018).abs() < 1e-9, "got {small}");
}
#[test]
fn unknown_model_yields_none_not_a_guess() {
let b = cheap_book();
assert!(b.estimate_cost(1, 1, "mystery-model").is_none());
assert!(b.get("nonexistent").is_none());
}
#[test]
fn exact_match_not_substring() {
let mut b = cheap_book();
b.set(
"gpt-4o",
Price {
input_per_1m: 250,
output_per_1m: 1000,
},
);
assert!(b.get("gpt-4o-mini").is_none());
assert!(b.get("gpt-4o").is_some());
}
#[test]
fn accumulate_totals_and_priced_usd() {
let b = cheap_book();
let mut cost = OverallCost::default();
cost.accumulate(
&TokenUsage {
prompt_tokens: 1000,
completion_tokens: 1000,
model: Some("test-model".into()),
},
&b,
);
cost.accumulate(
&TokenUsage {
prompt_tokens: 500,
completion_tokens: 0,
model: None,
},
&b,
);
assert_eq!(cost.prompt_tokens, 1500);
assert_eq!(cost.completion_tokens, 1000);
assert_eq!(cost.total_tokens, 2500);
assert!(
(cost.cost_usd.unwrap() - 0.018).abs() < 1e-9,
"got {:?}",
cost.cost_usd
);
}
#[test]
fn accumulate_without_any_priced_usage_keeps_cost_none() {
let b = cheap_book();
let mut cost = OverallCost::default();
cost.accumulate(
&TokenUsage {
prompt_tokens: 10,
completion_tokens: 10,
model: None,
},
&b,
);
assert_eq!(cost.total_tokens, 20);
assert!(cost.cost_usd.is_none());
}
#[test]
fn default_set_has_common_models() {
let b = PriceBook::default_set();
assert!(b.get("gpt-4o-mini").is_some());
assert!(b.get("deepseek-chat").is_some());
let usd = b
.estimate_cost(1_000_000, 1_000_000, "gpt-4o-mini")
.unwrap();
assert!((usd - 0.75).abs() < 1e-9, "got {usd}");
}
}