use super::auth::Auth;
use super::response::{CompletionResponse, StreamChunk, Usage};
use super::{CostInfo, CostResolution};
use crate::generator::CompletionParameters;
use crate::message::Message;
use std::future::Future;
use std::pin::Pin;
pub const AUDIO_RATE_FALLBACK_MULTIPLE: f64 = 13.0;
#[derive(Debug, Clone, Default, PartialEq)]
#[non_exhaustive]
pub struct TokenPrice {
pub input_per_mtok: f64,
pub output_per_mtok: f64,
pub cache_read_per_mtok: Option<f64>,
pub cache_write_per_mtok: Option<f64>,
pub audio_per_mtok: Option<f64>,
pub image_per_mtok: Option<f64>,
}
impl TokenPrice {
pub fn new(input_per_mtok: f64, output_per_mtok: f64) -> Self {
Self {
input_per_mtok,
output_per_mtok,
cache_read_per_mtok: None,
cache_write_per_mtok: None,
audio_per_mtok: None,
image_per_mtok: None,
}
}
pub fn with_cache_rates(mut self, read_per_mtok: f64, write_per_mtok: f64) -> Self {
self.cache_read_per_mtok = Some(read_per_mtok);
self.cache_write_per_mtok = Some(write_per_mtok);
self
}
pub fn with_media_rates(
mut self,
audio_per_mtok: Option<f64>,
image_per_mtok: Option<f64>,
) -> Self {
self.audio_per_mtok = audio_per_mtok;
self.image_per_mtok = image_per_mtok;
self
}
pub fn audio_rate(&self) -> f64 {
self.audio_per_mtok
.unwrap_or(self.input_per_mtok * AUDIO_RATE_FALLBACK_MULTIPLE)
}
pub fn image_rate(&self) -> f64 {
self.image_per_mtok.unwrap_or(self.input_per_mtok)
}
pub fn cost_of(&self, usage: &Usage) -> f64 {
let read_rate = self.cache_read_per_mtok.unwrap_or(self.input_per_mtok);
let write_rate = self.cache_write_per_mtok.unwrap_or(self.input_per_mtok);
(usage.uncached_input_tokens as f64 * self.input_per_mtok
+ usage.cache_read_tokens as f64 * read_rate
+ usage.cache_write_tokens as f64 * write_rate
+ usage.completion_tokens as f64 * self.output_per_mtok)
/ 1_000_000.0
}
}
#[derive(Debug, Clone)]
pub struct CostOutcome {
pub resolution: CostResolution,
pub usd: f64,
pub usage: Usage,
}
impl CostOutcome {
pub fn resolved(usd: f64, usage: Usage) -> Self {
Self {
resolution: CostResolution::Resolved,
usd,
usage,
}
}
pub fn unpriced(usage: Usage) -> Self {
Self {
resolution: CostResolution::Unpriced,
usd: 0.0,
usage,
}
}
pub fn unknown() -> Self {
Self {
resolution: CostResolution::Unknown,
usd: 0.0,
usage: Usage::default(),
}
}
pub fn into_cost_info(
self,
model: impl Into<String>,
response_id: impl Into<String>,
) -> CostInfo {
CostInfo {
cost: self.usd,
prompt_tokens: self.usage.prompt_tokens(),
completion_tokens: self.usage.completion_tokens,
total_tokens: self.usage.total_tokens(),
cache_read_tokens: self.usage.cache_read_tokens,
cache_write_tokens: self.usage.cache_write_tokens,
reasoning_tokens: self.usage.reasoning_tokens,
model: model.into(),
response_id: response_id.into(),
resolution: self.resolution,
}
}
}
pub struct PostStreamCtx<'a> {
pub client: reqwest_middleware::ClientWithMiddleware,
pub base_url: &'a str,
pub generation_id: &'a str,
pub auth: &'a Auth,
pub price: Option<&'a TokenPrice>,
}
pub type CostFuture<'a> = Pin<Box<dyn Future<Output = CostOutcome> + Send + 'a>>;
#[derive(Debug, Clone)]
pub struct AppIdentity {
pub url: String,
pub title: String,
}
pub trait Provider: Send + Sync + std::fmt::Debug {
fn openrouter_slug(&self) -> Option<&'static str> {
None
}
fn endpoint_url(&self, base_url: &str) -> String {
format!("{}/chat/completions", base_url.trim_end_matches('/'))
}
fn auth_headers(&self, auth: &Auth) -> crate::error::Result<Vec<(String, String)>> {
super::openai_wire::openai_auth_headers(auth)
}
fn build_request(
&self,
model: &str,
messages: &[Message],
params: &CompletionParameters,
stream: bool,
include_usage: bool,
) -> crate::error::Result<serde_json::Value> {
super::openai_wire::openai_build_request(
model,
messages,
params,
stream,
include_usage,
self,
)
}
fn openai_token_limit_field(&self) -> &'static str {
"max_completion_tokens"
}
fn openai_request_usage(&self, _body: &mut serde_json::Value, _stream: bool) {}
fn max_cache_breakpoints(&self) -> usize {
usize::MAX
}
fn openai_messages_value(&self, model: &str, messages: &[Message]) -> Vec<serde_json::Value> {
let _ = model;
super::openai_wire::messages_to_payload(messages, self.wire_keeps_estimation_metadata())
}
fn openai_tools_value(&self, tools: &[crate::tools::ToolDefinition]) -> serde_json::Value {
serde_json::Value::Array(
tools
.iter()
.map(super::openai_wire::tool_definition_value)
.collect(),
)
}
fn openai_tool_choice_value(&self, choice: &crate::tools::ToolChoice) -> serde_json::Value {
super::openai_wire::tool_choice_value(choice)
}
fn parse_response(&self, raw: serde_json::Value) -> crate::error::Result<CompletionResponse> {
super::openai_wire::parse_openai_response(raw, self)
}
fn parse_chunk(&self, data: &str) -> Option<crate::error::Result<StreamChunk>> {
super::openai_wire::parse_openai_chunk(data, self)
}
fn parse_usage(&self, raw: &serde_json::Value) -> Option<Usage> {
super::openai_wire::parse_openai_usage_field(raw)
}
fn parse_response_media(
&self,
message: &serde_json::Value,
) -> crate::error::Result<Vec<crate::message::Media>> {
super::openai_wire::parse_openai_response_images(message)
}
fn emits_stream_usage(&self, requested: bool) -> bool {
requested
}
fn wire_keeps_estimation_metadata(&self) -> bool {
false
}
fn attribution_headers(&self, _app: Option<&AppIdentity>) -> Vec<(String, String)> {
Vec::new()
}
fn cost_of(&self, usage: Usage, price: Option<&TokenPrice>) -> CostOutcome;
fn resolve_post_stream<'a>(&'a self, _ctx: PostStreamCtx<'a>) -> CostFuture<'a> {
Box::pin(async { CostOutcome::unknown() })
}
}
pub(crate) fn kept_cache_breakpoints(
messages: &[Message],
max: usize,
) -> std::collections::HashSet<usize> {
let marked: Vec<usize> = messages
.iter()
.enumerate()
.filter(|(_, m)| m.cache_breakpoint)
.map(|(i, _)| i)
.collect();
if marked.len() > max {
tracing::warn!(
"this provider allows at most {} cache breakpoints per request; {} were marked, keeping the last {}",
max,
marked.len(),
max
);
}
marked.iter().rev().take(max).copied().collect()
}
pub(crate) fn price_or_unpriced(usage: Usage, price: Option<&TokenPrice>) -> CostOutcome {
match price {
Some(p) => CostOutcome::resolved(p.cost_of(&usage), usage),
None => CostOutcome::unpriced(usage),
}
}
#[cfg(test)]
mod tests {
use super::*;
fn usage(prompt: u32, completion: u32) -> Usage {
Usage {
uncached_input_tokens: prompt,
completion_tokens: completion,
..Default::default()
}
}
#[test]
fn token_price_costs_prompt_and_completion_per_mtok() {
let price = TokenPrice::new(3.0, 15.0); let u = usage(1_000_000, 1_000_000);
assert!((price.cost_of(&u) - 18.0).abs() < 1e-9);
}
#[test]
fn token_price_bills_cache_read_and_write_at_their_own_rates() {
let price = TokenPrice::new(3.0, 15.0).with_cache_rates(0.3, 3.75);
let u = Usage {
uncached_input_tokens: 200_000,
cache_read_tokens: 800_000,
cache_write_tokens: 100_000,
..Default::default()
};
assert!(
(price.cost_of(&u) - 1.215).abs() < 1e-9,
"got {}",
price.cost_of(&u)
);
}
#[test]
fn cache_rates_fall_back_to_input_rate_when_unset() {
let price = TokenPrice::new(2.0, 0.0);
let u = Usage {
uncached_input_tokens: 0,
cache_read_tokens: 1_000_000,
cache_write_tokens: 1_000_000,
..Default::default()
};
assert!((price.cost_of(&u) - 4.0).abs() < 1e-9);
}
}