use crate::data::PROVIDERS;
use crate::{AiCost, Error, ModelPricing, Result};
use genai::ModelIden;
use genai::chat::Usage;
pub fn compute(provider_type: &str, model_name: &str, usage: &Usage) -> Result<AiCost> {
let model_name = normalize_model_name(model_name);
let provider_type = normalize_provider_type(provider_type);
let model = find_model_entry(provider_type, model_name)
.ok_or_else(|| Error::ModelNotFound(provider_type.to_string(), model_name.to_string()))?;
let prompt_tokens = usage.prompt_tokens.unwrap_or(0) as f64;
let (prompt_tokens_normal, prompt_cached_tokens, prompt_cache_creation_tokens) = match &usage.prompt_tokens_details
{
Some(details) => {
let cached = details.cached_tokens.unwrap_or(0) as f64;
let cache_creation_tokens = details.cache_creation_tokens.unwrap_or(0) as f64;
let normal = prompt_tokens - cached - cache_creation_tokens;
(normal, cached, cache_creation_tokens)
}
None => (prompt_tokens, 0.0, 0.0),
};
let price_prompt_normal = model.input_normal;
let price_prompt_cached = model.input_cached.unwrap_or(price_prompt_normal);
let price_prompt_cache_creation = 1.25 * price_prompt_normal;
let completion_tokens = usage.completion_tokens.unwrap_or(0) as f64;
let (completion_tokens_normal, completion_tokens_reasoning) = if let Some(reasoning_tokens) = usage
.completion_tokens_details
.as_ref()
.and_then(|v| v.reasoning_tokens.map(|v| v as f64))
{
(completion_tokens - reasoning_tokens, reasoning_tokens)
} else {
(completion_tokens, 0.)
};
let price_completion_normal = model.output_normal;
let price_completion_reasoning = model.output_reasoning.unwrap_or(price_completion_normal);
let cost_prompt_normal = prompt_tokens_normal * price_prompt_normal / 1_000_000.0;
let cost_prompt_cached = prompt_cached_tokens * price_prompt_cached / 1_000_000.0;
let cost_prompt_cache_creation = prompt_cache_creation_tokens * price_prompt_cache_creation / 1_000_000.0;
let cost_completion_normal = completion_tokens_normal * price_completion_normal / 1_000_000.0;
let cost_completion_reasoning = completion_tokens_reasoning * price_completion_reasoning / 1_000_000.0;
let input_normal = (cost_prompt_normal * 1_000_000.0).round() / 1_000_000.0;
let input_cache_read = (cost_prompt_cached * 1_000_000.0).round() / 1_000_000.0;
let input_cache_write = (cost_prompt_cache_creation * 1_000_000.0).round() / 1_000_000.0;
let input_total = ((input_normal + input_cache_read + input_cache_write) * 1_000_000.0).round() / 1_000_000.0;
let output_normal = (cost_completion_normal * 1_000_000.0).round() / 1_000_000.0;
let output_reasoning = (cost_completion_reasoning * 1_000_000.0).round() / 1_000_000.0;
let output_total = ((output_normal + output_reasoning) * 1_000_000.0).round() / 1_000_000.0;
let total = ((input_total + output_total) * 1_000_000.0).round() / 1_000_000.0;
let input_cache_saving = if prompt_cached_tokens > 0.0 {
let would_have_cost = prompt_cached_tokens * price_prompt_normal / 1_000_000.0;
let saving = (would_have_cost - cost_prompt_cached).max(0.0);
(saving * 1_000_000.0).round() / 1_000_000.0
} else {
0.0
};
Ok(AiCost {
total,
input_total,
input_normal,
input_cache_read,
input_cache_write,
output_total,
output_normal,
output_reasoning,
input_cache_saving,
})
}
pub fn compute_iden(model_iden: &ModelIden, usage: &Usage) -> Result<AiCost> {
let provider_type = model_iden.adapter_kind.as_lower_str();
let model_name = &*model_iden.model_name;
compute(provider_type, model_name, usage)
}
pub fn model_pricing(model_iden: &ModelIden) -> Option<ModelPricing> {
let provider_type = model_iden.adapter_kind.as_lower_str();
let model_name = &*model_iden.model_name;
let model_name = normalize_model_name(model_name);
let provider_type = normalize_provider_type(provider_type);
find_model_entry(provider_type, model_name).copied()
}
fn normalize_model_name(model_name: &str) -> &str {
match model_name.split_once("::") {
Some((_, after)) => after,
None => model_name,
}
}
fn normalize_provider_type(provider_type: &str) -> &str {
if provider_type == "openai_resp" {
"openai"
} else {
provider_type
}
}
fn find_model_entry(provider_type: &str, model_name: &str) -> Option<&'static ModelPricing> {
let provider = PROVIDERS.iter().find(|p| p.name == provider_type)?;
let mut model: Option<&ModelPricing> = None;
for m in provider.models.iter() {
if model_name.starts_with(m.name) {
let current_len = model.map(|m| m.name.len()).unwrap_or(0);
if current_len < m.name.len() {
model = Some(m)
}
}
}
model
}
#[cfg(test)]
mod tests {
use super::*;
use genai::chat::{PromptTokensDetails, Usage};
type TestResult = std::result::Result<(), Box<dyn std::error::Error>>;
#[test]
fn test_pricing_core_cost_simple() -> TestResult {
let usage = Usage {
prompt_tokens: Some(1000),
completion_tokens: Some(500),
prompt_tokens_details: None,
..Default::default()
};
let ai_cost = compute("openai", "gpt-4o", &usage)?;
let price = ai_cost.total;
let expected = 0.0075;
assert!((price - expected).abs() < f64::EPSILON);
assert!(ai_cost.input_cache_write == 0.0);
assert!(ai_cost.input_cache_saving == 0.0);
Ok(())
}
#[test]
fn test_pricing_core_cost_with_cached() -> TestResult {
let fx_prompt_normal_tokens = 1000;
let fx_completion_tokens = 500;
let fx_cached_tokens = 400;
let usage = Usage {
prompt_tokens: Some(fx_prompt_normal_tokens + fx_cached_tokens),
completion_tokens: Some(fx_completion_tokens),
prompt_tokens_details: Some(PromptTokensDetails {
cached_tokens: Some(fx_cached_tokens),
..Default::default()
}),
..Default::default()
};
let ai_cost = compute("openai", "gpt-4o-mini", &usage)?;
let price = ai_cost.total;
let cached = fx_cached_tokens as f64 * 0.075 / 1_000_000.0;
let prompt = fx_prompt_normal_tokens as f64 * 0.150 / 1_000_000.0;
let completion = fx_completion_tokens as f64 * 0.6 / 1_000_000.0;
let expected = cached + prompt + completion;
let expected = (expected * 1_000_000.0).round() / 1_000_000.0;
assert!((price - expected).abs() < f64::EPSILON);
assert!(ai_cost.input_cache_saving > 0.0);
assert!(ai_cost.input_cache_write == 0.0);
Ok(())
}
}