use std::collections::BTreeMap;
use serde::{Deserialize, Deserializer};
use crate::config::LlmCfg;
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct Usage {
pub prompt_tokens: u64,
pub completion_tokens: u64,
pub cached_tokens: u64,
pub reasoning_tokens: u64,
}
#[derive(Deserialize)]
struct UsageWire {
#[serde(default)]
prompt_tokens: u64,
#[serde(default)]
completion_tokens: u64,
#[serde(default)]
prompt_tokens_details: Option<PromptTokensDetails>,
#[serde(default)]
completion_tokens_details: Option<CompletionTokensDetails>,
}
#[derive(Deserialize)]
struct PromptTokensDetails {
#[serde(default)]
cached_tokens: u64,
}
#[derive(Deserialize)]
struct CompletionTokensDetails {
#[serde(default)]
reasoning_tokens: u64,
}
impl<'de> Deserialize<'de> for Usage {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
let w = UsageWire::deserialize(deserializer)?;
let cached = w
.prompt_tokens_details
.map(|d| d.cached_tokens)
.unwrap_or(0)
.min(w.prompt_tokens);
let reasoning = w
.completion_tokens_details
.map(|d| d.reasoning_tokens)
.unwrap_or(0)
.min(w.completion_tokens);
Ok(Usage {
prompt_tokens: w.prompt_tokens,
completion_tokens: w.completion_tokens,
cached_tokens: cached,
reasoning_tokens: reasoning,
})
}
}
impl Usage {
fn is_empty(&self) -> bool {
self.prompt_tokens == 0 && self.completion_tokens == 0
}
pub fn total_tokens(&self) -> u64 {
self.prompt_tokens.saturating_add(self.completion_tokens)
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum UnpricedPolicy {
Count,
Block,
}
impl UnpricedPolicy {
pub fn parse(s: &str) -> anyhow::Result<UnpricedPolicy> {
match s.trim().to_ascii_lowercase().as_str() {
"count" | "" => Ok(UnpricedPolicy::Count),
"block" | "reject" | "deny" => Ok(UnpricedPolicy::Block),
other => {
anyhow::bail!("invalid llm.on_unpriced_model {other:?} (expected count|block)")
}
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
struct ModelRate {
input_micros_per_m: u64,
output_micros_per_m: u64,
cached_micros_per_m: u64,
reasoning_micros_per_m: u64,
}
#[derive(Clone, Debug)]
pub struct LlmRuntime {
pub enabled: bool,
pub api_style: String,
pub unpriced: UnpricedPolicy,
prices: BTreeMap<String, ModelRate>,
}
impl LlmRuntime {
pub fn build(cfg: &LlmCfg) -> Self {
let prices = cfg
.models
.iter()
.map(|(name, p)| {
let input = usd_per_m_to_micros(p.input_per_1m);
let output = usd_per_m_to_micros(p.output_per_1m);
let rate = ModelRate {
input_micros_per_m: input,
output_micros_per_m: output,
cached_micros_per_m: if p.cached_per_1m > 0.0 {
usd_per_m_to_micros(p.cached_per_1m)
} else {
input
},
reasoning_micros_per_m: if p.reasoning_per_1m > 0.0 {
usd_per_m_to_micros(p.reasoning_per_1m)
} else {
output
},
};
(name.clone(), rate)
})
.collect();
LlmRuntime {
enabled: cfg.enabled,
api_style: if cfg.api_style.trim().is_empty() {
"openai".to_string()
} else {
cfg.api_style.trim().to_ascii_lowercase()
},
unpriced: UnpricedPolicy::parse(&cfg.on_unpriced_model)
.unwrap_or(UnpricedPolicy::Count),
prices,
}
}
pub fn disabled() -> Self {
LlmRuntime {
enabled: false,
api_style: "openai".to_string(),
unpriced: UnpricedPolicy::Count,
prices: BTreeMap::new(),
}
}
pub fn has_price_book(&self) -> bool {
!self.prices.is_empty()
}
fn resolve_rate(&self, model: &str) -> Option<&ModelRate> {
if let Some(rate) = self.prices.get(model) {
return Some(rate);
}
strip_provider_prefix(model).and_then(|bare| self.prices.get(bare))
}
pub fn is_priced(&self, model: &str) -> bool {
self.resolve_rate(model).is_some()
}
pub fn reject_unpriced(&self, model: &str) -> bool {
self.unpriced == UnpricedPolicy::Block && self.has_price_book() && !self.is_priced(model)
}
pub fn cost_micros(&self, model: &str, usage: &Usage) -> Option<u64> {
let rate = self.resolve_rate(model)?;
let cached = usage.cached_tokens.min(usage.prompt_tokens);
let uncached_input = usage.prompt_tokens - cached;
let reasoning = usage.reasoning_tokens.min(usage.completion_tokens);
let base_output = usage.completion_tokens - reasoning;
let total = uncached_input as u128 * rate.input_micros_per_m as u128
+ cached as u128 * rate.cached_micros_per_m as u128
+ base_output as u128 * rate.output_micros_per_m as u128
+ reasoning as u128 * rate.reasoning_micros_per_m as u128;
Some((total / 1_000_000).min(u64::MAX as u128) as u64)
}
}
const PROVIDER_PREFIXES: &[&str] = &[
"openai/",
"anthropic/",
"azure/",
"azure_ai/",
"vertex_ai/",
"vertex/",
"bedrock/",
"gemini/",
"google/",
"mistral/",
"codestral/",
"cohere/",
"groq/",
"together_ai/",
"together/",
"fireworks_ai/",
"fireworks/",
"deepseek/",
"xai/",
"perplexity/",
"replicate/",
"anyscale/",
"deepinfra/",
"cloudflare/",
"watsonx/",
"sagemaker/",
"ollama_chat/",
"ollama/",
];
pub fn canonical_model(model: &str) -> &str {
strip_provider_prefix(model).unwrap_or(model)
}
fn strip_provider_prefix(model: &str) -> Option<&str> {
for p in PROVIDER_PREFIXES {
if model.len() > p.len() && model.as_bytes()[..p.len()].eq_ignore_ascii_case(p.as_bytes()) {
return Some(&model[p.len()..]);
}
}
None
}
fn usd_per_m_to_micros(usd: f64) -> u64 {
if !usd.is_finite() || usd <= 0.0 {
return 0;
}
(usd * 1_000_000.0).round() as u64
}
#[derive(Deserialize)]
struct ModelField {
model: Option<String>,
}
pub fn parse_request_model(body: &[u8]) -> Option<String> {
let parsed: ModelField = serde_json::from_slice(body).ok()?;
let model = parsed.model?;
(!model.trim().is_empty()).then_some(model)
}
#[derive(Deserialize)]
struct MaxTokensField {
max_tokens: Option<u64>,
max_completion_tokens: Option<u64>,
}
pub fn parse_request_max_tokens(body: &[u8]) -> Option<u64> {
let parsed: MaxTokensField = serde_json::from_slice(body).ok()?;
parsed.max_tokens.or(parsed.max_completion_tokens)
}
pub fn estimate_prompt_tokens(body_len: usize) -> u64 {
(body_len / 4) as u64
}
#[derive(Deserialize)]
struct UsageField {
usage: Option<Usage>,
}
pub fn parse_response_usage(body: &[u8]) -> Option<Usage> {
let parsed: UsageField = serde_json::from_slice(body).ok()?;
parsed.usage.filter(|u| !u.is_empty())
}
pub fn parse_sse_usage(bytes: &[u8]) -> Option<Usage> {
let text = std::str::from_utf8(bytes).ok()?;
let mut last = None;
for line in text.lines() {
let line = line.trim_start();
let Some(payload) = line.strip_prefix("data:") else {
continue;
};
let payload = payload.trim();
if payload.is_empty() || payload == "[DONE]" {
continue;
}
if let Ok(parsed) = serde_json::from_str::<UsageField>(payload) {
if let Some(u) = parsed.usage.filter(|u| !u.is_empty()) {
last = Some(u);
}
}
}
last
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::ModelPrice;
fn runtime() -> LlmRuntime {
let mut models = BTreeMap::new();
models.insert(
"gpt-4o".to_string(),
ModelPrice {
input_per_1m: 2.50,
output_per_1m: 10.00,
..Default::default()
},
);
LlmRuntime::build(&LlmCfg {
enabled: true,
api_style: "openai".into(),
models,
..Default::default()
})
}
#[test]
fn parses_request_model() {
let body = br#"{"model":"gpt-4o","messages":[{"role":"user","content":"hi"}]}"#;
assert_eq!(parse_request_model(body), Some("gpt-4o".to_string()));
assert_eq!(parse_request_model(b"not json"), None);
assert_eq!(parse_request_model(br#"{"messages":[]}"#), None);
assert_eq!(parse_request_model(br#"{"model":""}"#), None);
}
#[test]
fn parses_non_streaming_usage() {
let body = br#"{"id":"x","choices":[],"usage":{"prompt_tokens":12,"completion_tokens":34,"total_tokens":46}}"#;
assert_eq!(
parse_response_usage(body),
Some(Usage {
prompt_tokens: 12,
completion_tokens: 34,
..Default::default()
})
);
assert_eq!(parse_response_usage(br#"{"error":"nope"}"#), None);
assert_eq!(
parse_response_usage(br#"{"usage":{"prompt_tokens":0,"completion_tokens":0}}"#),
None
);
}
#[test]
fn parses_terminal_sse_usage() {
let stream = "data: {\"choices\":[{\"delta\":{\"content\":\"hi\"}}],\"usage\":null}\n\n\
data: {\"choices\":[],\"usage\":{\"prompt_tokens\":7,\"completion_tokens\":5,\"total_tokens\":12}}\n\n\
data: [DONE]\n\n";
assert_eq!(
parse_sse_usage(stream.as_bytes()),
Some(Usage {
prompt_tokens: 7,
completion_tokens: 5,
..Default::default()
})
);
let no_usage = "data: {\"choices\":[{\"delta\":{\"content\":\"hi\"}}]}\n\ndata: [DONE]\n\n";
assert_eq!(parse_sse_usage(no_usage.as_bytes()), None);
}
#[test]
fn prices_known_model_and_fails_open_on_unknown() {
let rt = runtime();
let usage = Usage {
prompt_tokens: 1_000_000,
completion_tokens: 1_000_000,
..Default::default()
};
assert_eq!(rt.cost_micros("gpt-4o", &usage), Some(12_500_000));
assert_eq!(rt.cost_micros("mystery-model", &usage), None);
}
#[test]
fn cost_is_proportional_for_small_counts() {
let rt = runtime();
let usage = Usage {
prompt_tokens: 1_000,
completion_tokens: 0,
..Default::default()
};
assert_eq!(rt.cost_micros("gpt-4o", &usage), Some(2_500));
}
#[test]
fn parses_llm_toml_models_map() {
let toml = r#"
[llm]
enabled = true
api_style = "openai"
[llm.models."gpt-4o"]
input_per_1m = 2.5
output_per_1m = 10.0
"#;
let cfg: crate::config::Config = toml::from_str(toml).unwrap();
assert!(cfg.llm.enabled);
assert_eq!(cfg.llm.models.len(), 1);
let rt = LlmRuntime::build(&cfg.llm);
let usage = Usage {
prompt_tokens: 1_000_000,
completion_tokens: 0,
..Default::default()
};
assert_eq!(rt.cost_micros("gpt-4o", &usage), Some(2_500_000));
}
#[test]
fn negative_or_zero_price_is_free_not_an_error() {
assert_eq!(usd_per_m_to_micros(-1.0), 0);
assert_eq!(usd_per_m_to_micros(0.0), 0);
assert_eq!(usd_per_m_to_micros(0.5), 500_000);
}
#[test]
fn parses_cached_and_reasoning_detail_dims() {
let body = br#"{"usage":{"prompt_tokens":100,"completion_tokens":80,
"prompt_tokens_details":{"cached_tokens":40},
"completion_tokens_details":{"reasoning_tokens":30}}}"#;
assert_eq!(
parse_response_usage(body),
Some(Usage {
prompt_tokens: 100,
completion_tokens: 80,
cached_tokens: 40,
reasoning_tokens: 30,
})
);
let bad = br#"{"usage":{"prompt_tokens":10,"completion_tokens":5,
"prompt_tokens_details":{"cached_tokens":9999}}}"#;
assert_eq!(parse_response_usage(bad).unwrap().cached_tokens, 10);
}
#[test]
fn cached_reasoning_default_to_base_rate_so_totals_are_unchanged() {
let rt = runtime();
let plain = Usage {
prompt_tokens: 1_000_000,
completion_tokens: 1_000_000,
..Default::default()
};
let with_dims = Usage {
prompt_tokens: 1_000_000,
completion_tokens: 1_000_000,
cached_tokens: 500_000,
reasoning_tokens: 400_000,
};
assert_eq!(
rt.cost_micros("gpt-4o", &plain),
rt.cost_micros("gpt-4o", &with_dims)
);
}
#[test]
fn explicit_cached_reasoning_rates_are_applied() {
let mut models = BTreeMap::new();
models.insert(
"gpt-4o".to_string(),
ModelPrice {
input_per_1m: 2.50,
output_per_1m: 10.00,
cached_per_1m: 1.25, reasoning_per_1m: 20.00, },
);
let rt = LlmRuntime::build(&LlmCfg {
enabled: true,
models,
..Default::default()
});
let usage = Usage {
prompt_tokens: 1_000_000, completion_tokens: 1_000_000, cached_tokens: 400_000,
reasoning_tokens: 300_000,
};
assert_eq!(rt.cost_micros("gpt-4o", &usage), Some(15_000_000));
}
#[test]
fn unpriced_policy_block_only_bites_with_a_price_book() {
let count = runtime();
assert!(!count.reject_unpriced("mystery"));
let mut models = BTreeMap::new();
models.insert(
"gpt-4o".to_string(),
ModelPrice {
input_per_1m: 2.5,
output_per_1m: 10.0,
..Default::default()
},
);
let block = LlmRuntime::build(&LlmCfg {
enabled: true,
models,
on_unpriced_model: "block".into(),
..Default::default()
});
assert!(block.reject_unpriced("mystery"));
assert!(!block.reject_unpriced("gpt-4o"));
let block_no_book = LlmRuntime::build(&LlmCfg {
enabled: true,
on_unpriced_model: "block".into(),
..Default::default()
});
assert!(!block_no_book.reject_unpriced("anything"));
}
#[test]
fn unpriced_policy_parse_rejects_typos() {
assert_eq!(
UnpricedPolicy::parse("count").unwrap(),
UnpricedPolicy::Count
);
assert_eq!(
UnpricedPolicy::parse("block").unwrap(),
UnpricedPolicy::Block
);
assert_eq!(UnpricedPolicy::parse("").unwrap(), UnpricedPolicy::Count);
assert!(UnpricedPolicy::parse("banana").is_err());
}
#[test]
fn sse_usage_is_the_last_frame_never_the_sum_of_cumulative_chunks() {
let stream = "data: {\"choices\":[{\"delta\":{\"content\":\"a\"}}],\"usage\":{\"prompt_tokens\":100,\"completion_tokens\":10}}\n\n\
data: {\"choices\":[{\"delta\":{\"content\":\"b\"}}],\"usage\":{\"prompt_tokens\":100,\"completion_tokens\":25}}\n\n\
data: {\"choices\":[],\"usage\":{\"prompt_tokens\":100,\"completion_tokens\":60,\"total_tokens\":160}}\n\n\
data: [DONE]\n\n";
let u = parse_sse_usage(stream.as_bytes()).expect("terminal usage");
assert_eq!(
u.prompt_tokens, 100,
"prompt must be the last frame, not 300 (summed)"
);
assert_eq!(
u.completion_tokens, 60,
"completion must be the last frame, not 95 (summed)"
);
}
#[test]
fn cached_prompt_tokens_are_never_billed_at_the_output_rate() {
let mut models = BTreeMap::new();
models.insert(
"m".to_string(),
ModelPrice {
input_per_1m: 3.00,
output_per_1m: 60.00, cached_per_1m: 0.30, ..Default::default()
},
);
let rt = LlmRuntime::build(&LlmCfg {
enabled: true,
models,
..Default::default()
});
let usage = Usage {
prompt_tokens: 1_000_000,
completion_tokens: 0,
cached_tokens: 1_000_000,
reasoning_tokens: 0,
};
assert_eq!(rt.cost_micros("m", &usage), Some(300_000));
}
#[test]
fn metering_reads_the_usage_object_not_the_request_body_size() {
let rt = runtime();
let resp = br#"{"usage":{"prompt_tokens":50,"completion_tokens":10}}"#;
let usage = parse_response_usage(resp).expect("usage");
assert_eq!(usage.prompt_tokens, 50);
let huge_b64_body_len = 4_000_000usize;
assert!(estimate_prompt_tokens(huge_b64_body_len) > usage.prompt_tokens);
assert_eq!(rt.cost_micros("gpt-4o", &usage), Some(225));
}
#[test]
fn provider_prefixed_model_resolves_to_the_bare_price() {
let rt = runtime(); let usage = Usage {
prompt_tokens: 1_000,
completion_tokens: 0,
..Default::default()
};
assert!(rt.is_priced("openai/gpt-4o"));
assert_eq!(
rt.cost_micros("openai/gpt-4o", &usage),
rt.cost_micros("gpt-4o", &usage)
);
assert!(rt.is_priced("OpenAI/gpt-4o"));
}
#[test]
fn exact_prefixed_entry_wins_over_normalization() {
let mut models = BTreeMap::new();
models.insert(
"gpt-4o".to_string(),
ModelPrice {
input_per_1m: 2.50,
output_per_1m: 10.00,
..Default::default()
},
);
models.insert(
"openai/gpt-4o".to_string(),
ModelPrice {
input_per_1m: 99.0, output_per_1m: 99.0,
..Default::default()
},
);
let rt = LlmRuntime::build(&LlmCfg {
enabled: true,
models,
..Default::default()
});
let usage = Usage {
prompt_tokens: 1_000_000,
completion_tokens: 0,
..Default::default()
};
assert_eq!(rt.cost_micros("openai/gpt-4o", &usage), Some(99_000_000));
assert_eq!(rt.cost_micros("gpt-4o", &usage), Some(2_500_000));
}
#[test]
fn unknown_or_huggingface_style_prefix_is_left_unpriced() {
let rt = runtime();
assert!(!rt.is_priced("meta-llama/Llama-3-8b"));
assert_eq!(
rt.cost_micros("meta-llama/Llama-3-8b", &Usage::default()),
None
);
assert!(!rt.is_priced("openai/mystery-model"));
}
#[test]
fn canonical_model_strips_known_prefixes_for_attribution() {
assert_eq!(canonical_model("openai/gpt-4o"), "gpt-4o");
assert_eq!(canonical_model("azure/gpt-4o"), "gpt-4o");
assert_eq!(canonical_model("gpt-4o"), "gpt-4o"); assert_eq!(canonical_model("meta-llama/Llama-3"), "meta-llama/Llama-3");
}
#[test]
fn block_policy_does_not_reject_a_prefixed_priced_model() {
let mut models = BTreeMap::new();
models.insert(
"gpt-4o".to_string(),
ModelPrice {
input_per_1m: 2.5,
output_per_1m: 10.0,
..Default::default()
},
);
let rt = LlmRuntime::build(&LlmCfg {
enabled: true,
models,
on_unpriced_model: "block".into(),
..Default::default()
});
assert!(!rt.reject_unpriced("openai/gpt-4o"));
assert!(rt.reject_unpriced("openai/mystery-model"));
assert!(rt.reject_unpriced("mystery-model"));
}
}