use super::estimate_cost_usd;
use crate::model_profile_data::types::{Speed, SpeedConfig, SpeedValue};
pub(super) fn speed(value: Speed, name: &str) -> SpeedValue {
SpeedValue {
value,
name: name.into(),
cost_multiplier: None,
}
}
pub(super) fn priced_speed(value: Speed, name: &str, multiplier: f64) -> SpeedValue {
SpeedValue {
cost_multiplier: Some(multiplier),
..speed(value, name)
}
}
pub(super) fn speed_flex_priority() -> SpeedConfig {
SpeedConfig {
values: vec![
speed(Speed::Flex, "Flex"),
speed(Speed::Default, "Standard"),
speed(Speed::Priority, "Fast"),
],
default: Speed::Default,
}
}
pub(super) fn speed_flex_only() -> SpeedConfig {
SpeedConfig {
values: vec![
speed(Speed::Flex, "Flex"),
speed(Speed::Default, "Standard"),
],
default: Speed::Default,
}
}
pub(super) fn speed_priority_only() -> SpeedConfig {
SpeedConfig {
values: vec![
speed(Speed::Default, "Standard"),
speed(Speed::Priority, "Fast"),
],
default: Speed::Default,
}
}
pub fn estimate_cost_usd_for_speed(
provider_type: &str,
model_id: &str,
input_tokens: u32,
output_tokens: u32,
cache_read_tokens: u32,
cache_creation_tokens: u32,
served_tier: Option<&str>,
) -> Option<f64> {
let standard = estimate_cost_usd(
provider_type,
model_id,
input_tokens,
output_tokens,
cache_read_tokens,
cache_creation_tokens,
)?;
let multiplier = served_tier
.and_then(|tier| {
super::get_model_profile(provider_type, model_id)?
.speed?
.values
.into_iter()
.find(|value| value.value.matches_tier(tier))?
.cost_multiplier
})
.unwrap_or(1.0);
Some(standard * multiplier)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn fast_and_priority_are_one_tier() {
for speed in [Speed::Fast, Speed::Priority] {
assert!(speed.matches_tier("fast"));
assert!(speed.matches_tier("priority"));
assert!(!speed.matches_tier("ultrafast"));
}
assert!(Speed::Ultrafast.matches_tier("ultrafast"));
assert!(!Speed::Default.matches_tier("flex"));
}
#[test]
fn served_tier_scales_the_standard_estimate() {
let tokens = (100_000, 100_000, 0, 0);
let at = |model: &str, tier: Option<&str>| {
estimate_cost_usd_for_speed("openai", model, tokens.0, tokens.1, 0, 0, tier).unwrap()
};
let standard = 1.0 + 5.0;
let close = |a: f64, b: f64| (a - b).abs() < 1e-9;
assert!(close(at("gpt-6-astra", None), standard));
assert!(close(at("gpt-6-astra", Some("default")), standard));
assert!(close(at("gpt-6-astra", Some("flex")), standard * 0.5));
assert!(close(at("gpt-6-astra", Some("priority")), standard * 2.0));
assert!(close(at("gpt-6-astra", Some("fast")), standard * 2.0));
assert!(close(at("gpt-6-astra", Some("ultrafast")), standard * 6.0));
assert!(close(at("gpt-6.1-sol", Some("ultrafast")), 0.2 + 1.0));
assert!(close(at("gpt-6.1-sol", Some("fast")), (0.2 + 1.0) * 2.0));
let gpt55 = estimate_cost_usd("openai", "gpt-5.5", tokens.0, tokens.1, 0, 0);
assert_eq!(
estimate_cost_usd_for_speed(
"openai",
"gpt-5.5",
tokens.0,
tokens.1,
0,
0,
Some("priority")
),
gpt55
);
}
#[test]
fn tier_multiplier_applies_on_top_of_the_long_context_tier() {
let cost = estimate_cost_usd_for_speed(
"openai",
"gpt-6-astra",
300_000,
0,
0,
0,
Some("ultrafast"),
)
.unwrap();
assert!((cost - 0.3 * 20.0 * 6.0).abs() < 1e-9, "{cost}");
}
}