use crate::outcome::priced_input_usd;
use crate::registry::UnifiedRegistry;
use crate::InferenceResult;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum InferenceBilling {
Metered,
Subscription,
Unknown,
}
impl InferenceBilling {
pub const fn as_str(self) -> &'static str {
match self {
Self::Metered => "metered",
Self::Subscription => "subscription",
Self::Unknown => "unknown",
}
}
pub fn from_wire(value: &str) -> Option<Self> {
match value {
"metered" => Some(Self::Metered),
"subscription" => Some(Self::Subscription),
"unknown" => Some(Self::Unknown),
_ => None,
}
}
pub const fn combine(self, other: Self) -> Self {
match (self, other) {
(Self::Subscription, _) | (_, Self::Subscription) => Self::Subscription,
(Self::Unknown, _) | (_, Self::Unknown) => Self::Unknown,
_ => Self::Metered,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub enum PricedInferenceCall {
Metered(f64),
Subscription,
Unknown,
}
impl PricedInferenceCall {
pub const fn billing(self) -> InferenceBilling {
match self {
Self::Metered(_) => InferenceBilling::Metered,
Self::Subscription => InferenceBilling::Subscription,
Self::Unknown => InferenceBilling::Unknown,
}
}
pub const fn cost_usd(self) -> Option<f64> {
match self {
Self::Metered(cost_usd) => Some(cost_usd),
Self::Subscription | Self::Unknown => None,
}
}
}
pub fn price_inference_call(
registry: &UnifiedRegistry,
result: &InferenceResult,
) -> PricedInferenceCall {
let Some(schema) = registry
.get(&result.model_used)
.or_else(|| registry.find_by_name(&result.model_used))
else {
return PricedInferenceCall::Unknown;
};
if schema.tags.iter().any(|tag| tag == "subscription") {
return PricedInferenceCall::Subscription;
}
let Some(usage) = result.usage.as_ref() else {
return PricedInferenceCall::Unknown;
};
let (Some(input_per_mtok), Some(output_per_mtok)) =
(schema.cost.input_per_mtok, schema.cost.output_per_mtok)
else {
return PricedInferenceCall::Unknown;
};
let cost_usd = priced_input_usd(
usage.prompt_tokens,
usage.cache_read_input_tokens,
usage.cache_creation_input_tokens,
input_per_mtok,
schema.cache_rates(),
) + (usage.completion_tokens as f64 / 1_000_000.0) * output_per_mtok;
PricedInferenceCall::Metered(cost_usd)
}
#[cfg(test)]
mod tests {
use super::{price_inference_call, InferenceBilling};
use crate::registry::UnifiedRegistry;
use crate::schema::{
ApiProtocol, CostModel, ModelCapability, ModelSchema, ModelSource, PerformanceEnvelope,
TrustTier,
};
use crate::{InferenceResult, TokenUsage};
fn schema(id: &str, name: &str, tags: &[&str]) -> ModelSchema {
ModelSchema {
id: id.into(),
name: name.into(),
provider: "test".into(),
family: "test".into(),
version: "test".into(),
capabilities: vec![ModelCapability::Generate],
context_length: 4096,
max_output_tokens: None,
param_count: String::new(),
quantization: None,
performance: PerformanceEnvelope::default(),
cost: CostModel {
input_per_mtok: Some(2.0),
output_per_mtok: Some(10.0),
..Default::default()
},
source: ModelSource::RemoteApi {
endpoint: "https://example.invalid/v1".into(),
api_key_env: "TEST_API_KEY".into(),
api_key_envs: vec![],
api_version: None,
protocol: ApiProtocol::Anthropic,
},
tags: tags.iter().map(|tag| (*tag).to_string()).collect(),
supported_params: vec![],
public_benchmarks: vec![],
trust_tier: TrustTier::Curated,
deprecated: false,
available: false,
weights_ready: false,
}
}
fn result(model_used: &str, usage: Option<TokenUsage>) -> InferenceResult {
serde_json::from_value(serde_json::json!({
"text": "ok",
"tool_calls": [],
"trace_id": "trace",
"model_used": model_used,
"latency_ms": 0,
"usage": usage,
}))
.expect("inference result fixture")
}
fn usage() -> TokenUsage {
TokenUsage {
prompt_tokens: 1_000,
completion_tokens: 500,
total_tokens: 1_800,
context_window: 4_096,
cache_read_input_tokens: 200,
cache_creation_input_tokens: 100,
}
}
#[test]
fn prices_all_token_buckets_after_exact_id_or_name_resolution() {
let directory = tempfile::tempdir().unwrap();
let mut registry = UnifiedRegistry::new_empty(directory.path().to_path_buf());
registry.register(schema("test/metered:exact", "Metered Name", &[]));
for serving_model in ["test/metered:exact", "Metered Name"] {
let priced = price_inference_call(®istry, &result(serving_model, Some(usage())));
assert_eq!(priced.billing(), InferenceBilling::Metered);
assert_eq!(priced.cost_usd(), Some(0.00729));
}
}
#[test]
fn an_exact_subscription_tag_overrides_usage_and_catalog_dollar_fields() {
let directory = tempfile::tempdir().unwrap();
let mut registry = UnifiedRegistry::new_empty(directory.path().to_path_buf());
registry.register(schema(
"test/subscription:exact",
"Subscription Name",
&["subscription"],
));
for usage in [Some(usage()), None] {
let priced = price_inference_call(®istry, &result("test/subscription:exact", usage));
assert_eq!(priced.billing(), InferenceBilling::Subscription);
assert_eq!(priced.cost_usd(), None);
}
}
#[test]
fn unpriceable_calls_are_unknown_not_free() {
let directory = tempfile::tempdir().unwrap();
let mut registry = UnifiedRegistry::new_empty(directory.path().to_path_buf());
let mut missing_rate = schema("test/unpriced", "Unpriced", &[]);
missing_rate.cost.output_per_mtok = None;
registry.register(missing_rate);
for call in [
result("missing/model", Some(usage())),
result("test/unpriced", Some(usage())),
result("test/unpriced", None),
] {
let priced = price_inference_call(®istry, &call);
assert_eq!(priced.billing(), InferenceBilling::Unknown);
assert_eq!(priced.cost_usd(), None);
}
}
}