use serde::{Deserialize, Serialize};
use crate::json::Json;
pub use super::pricing::{
CacheReadAccounting, ModelPricing, PricingCatalog, PricingCatalogError, PricingConfig,
PricingResolver, PricingSource, PricingSourceConfig, PricingUnit, PromptCachePricing,
TokenPricingRates, active_pricing_resolver, attach_estimated_cost,
attach_estimated_cost_for_provider, estimate_cost, estimate_cost_for_provider,
estimate_cost_with_catalog, estimate_cost_with_provider, infer_model_provider,
pricing_for_model, pricing_for_provider, reset_active_pricing_resolver,
set_active_pricing_resolver,
};
use super::request::MessageContent;
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct AnnotatedLlmResponse {
#[serde(skip_serializing_if = "Option::is_none")]
pub id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub model: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub message: Option<MessageContent>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tool_calls: Option<Vec<ResponseToolCall>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub finish_reason: Option<FinishReason>,
#[serde(skip_serializing_if = "Option::is_none")]
pub usage: Option<Usage>,
#[serde(skip_serializing_if = "Option::is_none")]
pub api_specific: Option<ApiSpecificResponse>,
#[serde(flatten)]
pub extra: serde_json::Map<String, Json>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, Default)]
pub struct Usage {
#[serde(skip_serializing_if = "Option::is_none")]
pub prompt_tokens: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub completion_tokens: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub total_tokens: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub cache_read_tokens: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub cache_write_tokens: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub cost: Option<CostEstimate>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum CostSource {
ModelPricing,
ProviderReported,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct CostEstimate {
#[serde(skip_serializing_if = "Option::is_none")]
pub total: Option<f64>,
#[serde(default = "default_cost_currency")]
pub currency: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub input: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub output: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub cache_read: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub cache_write: Option<f64>,
pub source: CostSource,
#[serde(skip_serializing_if = "Option::is_none")]
pub pricing_provider: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub pricing_model: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub pricing_as_of: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub pricing_source: Option<String>,
}
impl CostEstimate {
#[must_use]
pub fn total_or_component_sum(&self) -> Option<f64> {
self.total.or_else(|| {
let (has_component, total) =
[self.input, self.output, self.cache_read, self.cache_write]
.into_iter()
.flatten()
.fold((false, 0.0), |(_, total), value| (true, total + value));
has_component.then_some(total)
})
}
#[must_use]
pub fn total_for_currency(&self, currency: &str) -> Option<f64> {
self.currency
.eq_ignore_ascii_case(currency)
.then_some(self.total)
.flatten()
}
#[must_use]
pub fn total_or_component_sum_for_currency(&self, currency: &str) -> Option<f64> {
self.currency
.eq_ignore_ascii_case(currency)
.then(|| self.total_or_component_sum())
.flatten()
}
}
#[derive(Debug, Clone, Default, Deserialize)]
pub(crate) struct RawUsageCost {
pub total: Option<f64>,
pub input: Option<f64>,
pub output: Option<f64>,
pub cache_read: Option<f64>,
pub cache_write: Option<f64>,
pub currency: Option<String>,
pub pricing_provider: Option<String>,
pub pricing_model: Option<String>,
pub pricing_as_of: Option<String>,
pub pricing_source: Option<String>,
}
pub(crate) fn provider_reported_cost(
provider_total_cost: Option<f64>,
cost: Option<RawUsageCost>,
) -> Option<CostEstimate> {
let cost = cost.unwrap_or_default();
let provider_total_uses_default_currency = provider_total_cost.is_some();
let nested_currency_is_default = cost
.currency
.as_deref()
.is_none_or(|currency| currency.eq_ignore_ascii_case("USD"));
let keep_component_costs = !provider_total_uses_default_currency || nested_currency_is_default;
let input = keep_component_costs.then_some(cost.input).flatten();
let output = keep_component_costs.then_some(cost.output).flatten();
let cache_read = keep_component_costs.then_some(cost.cache_read).flatten();
let cache_write = keep_component_costs.then_some(cost.cache_write).flatten();
let has_currency_native_amount = cost.total.is_some()
|| cost.input.is_some()
|| cost.output.is_some()
|| cost.cache_read.is_some()
|| cost.cache_write.is_some();
let component_total = [input, output, cache_read, cache_write]
.into_iter()
.flatten()
.sum();
let has_component_cost =
input.is_some() || output.is_some() || cache_read.is_some() || cache_write.is_some();
let total = provider_total_cost
.or(cost.total)
.or_else(|| has_component_cost.then_some(component_total));
if total.is_none()
&& input.is_none()
&& output.is_none()
&& cache_read.is_none()
&& cache_write.is_none()
{
return None;
}
Some(CostEstimate {
total,
currency: if provider_total_uses_default_currency {
default_cost_currency()
} else if has_currency_native_amount {
cost.currency.unwrap_or_else(default_cost_currency)
} else {
default_cost_currency()
},
input,
output,
cache_read,
cache_write,
source: CostSource::ProviderReported,
pricing_provider: cost.pricing_provider,
pricing_model: cost.pricing_model,
pricing_as_of: cost.pricing_as_of,
pricing_source: cost.pricing_source,
})
}
fn default_cost_currency() -> String {
"USD".into()
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum FinishReason {
Complete,
Length,
ToolUse,
ContentFilter,
Unknown(String),
}
impl FinishReason {
#[must_use]
pub fn is_complete(&self) -> bool {
matches!(self, FinishReason::Complete)
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ResponseToolCall {
pub id: String,
pub name: String,
pub arguments: Json,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "api")]
pub enum ApiSpecificResponse {
#[serde(rename = "openai_chat")]
OpenAIChat {
#[serde(skip_serializing_if = "Option::is_none")]
logprobs: Option<Json>,
#[serde(skip_serializing_if = "Option::is_none")]
system_fingerprint: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
service_tier: Option<String>,
},
#[serde(rename = "openai_responses")]
OpenAIResponses {
#[serde(skip_serializing_if = "Option::is_none")]
output_items: Option<Vec<Json>>,
#[serde(skip_serializing_if = "Option::is_none")]
status: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
incomplete_details: Option<Json>,
#[serde(skip_serializing_if = "Option::is_none")]
previous_response_id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
store: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
service_tier: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
truncation: Option<Json>,
#[serde(skip_serializing_if = "Option::is_none")]
reasoning: Option<Json>,
#[serde(skip_serializing_if = "Option::is_none")]
input_tokens_details: Option<Json>,
#[serde(skip_serializing_if = "Option::is_none")]
output_tokens_details: Option<Json>,
},
#[serde(rename = "anthropic_messages")]
AnthropicMessages {
#[serde(skip_serializing_if = "Option::is_none")]
object_type: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
role: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
stop_reason: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
stop_sequence: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
service_tier: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
container: Option<Json>,
#[serde(skip_serializing_if = "Option::is_none")]
content_blocks: Option<Vec<Json>>,
},
#[serde(rename = "custom")]
Custom {
api_name: String,
data: Json,
},
}
impl AnnotatedLlmResponse {
#[must_use]
pub fn response_text(&self) -> Option<&str> {
match self.message.as_ref()? {
MessageContent::Text(s) => Some(s.as_str()),
MessageContent::Parts(parts) => parts.iter().find_map(|p| match p {
super::request::ContentPart::Text { text } => Some(text.as_str()),
super::request::ContentPart::ImageUrl { .. } => None,
}),
}
}
#[must_use]
pub fn has_tool_calls(&self) -> bool {
self.tool_calls
.as_ref()
.is_some_and(|calls| !calls.is_empty())
}
}
#[cfg(test)]
#[path = "../../tests/unit/codec/response_tests.rs"]
mod tests;