Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
49 changes: 15 additions & 34 deletions rust/src/providers/mistral/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -9,9 +9,11 @@ use reqwest::Client;
use serde::Deserialize;
use std::collections::HashMap;

mod pricing;
mod subscription;
mod token_math;

use pricing::{MistralPrice, PriceIndex};
use subscription::{SubscriptionBudget, SubscriptionBudgets};

use crate::core::{
Expand Down Expand Up @@ -84,24 +86,21 @@ struct ModelUsageData {

#[derive(Debug, Deserialize)]
struct UsageEntry {
#[serde(rename = "event_type")]
event_type: Option<String>,
#[serde(rename = "billing_metric")]
billing_metric: Option<String>,
#[serde(rename = "billing_group")]
billing_group: Option<String>,
#[serde(rename = "api_zone")]
api_zone: Option<String>,
#[serde(rename = "service_tier")]
service_tier: Option<String>,
value: Option<i64>,
#[serde(rename = "value_paid")]
value_paid: Option<i64>,
}

#[derive(Debug, Deserialize)]
struct MistralPrice {
#[serde(rename = "billing_metric")]
billing_metric: Option<String>,
#[serde(rename = "billing_group")]
billing_group: Option<String>,
price: Option<String>,
}

#[derive(Debug)]
struct MistralUsageSummary {
total_cost: f64,
Expand Down Expand Up @@ -288,7 +287,7 @@ impl MistralProvider {
}

fn summarize_billing(billing: BillingResponse) -> Result<MistralUsageSummary, ProviderError> {
let prices = Self::build_price_index(billing.prices.unwrap_or_default());
let prices = PriceIndex::new(billing.prices.unwrap_or_default());
let mut total_cost = 0.0;
let mut total_tokens = TokenCounts::default();
let mut model_count = 0;
Expand Down Expand Up @@ -448,24 +447,9 @@ impl MistralProvider {
.with_description_as_detail()
}

fn build_price_index(prices: Vec<MistralPrice>) -> HashMap<String, f64> {
prices
.into_iter()
.filter_map(|price| {
let metric = price.billing_metric?;
let group = price.billing_group?;
let value = price.price?.parse::<f64>().ok()?;
if !value.is_finite() {
return None;
}
Some((format!("{metric}::{group}"), value))
})
.collect()
}

fn aggregate_model(
data: &ModelUsageData,
prices: &HashMap<String, f64>,
prices: &PriceIndex,
mode: AggregationMode,
) -> Result<ModelAggregation, ProviderError> {
let mut tokens = TokenCounts::default();
Expand All @@ -481,11 +465,11 @@ impl MistralProvider {
if matches!(mode, AggregationMode::CostAndTokens) {
tokens.add_lane(entry.value.or(entry.value_paid).unwrap_or(0), kind)?;
}
if let (Some(metric), Some(group)) = (&entry.billing_metric, &entry.billing_group) {
// Upstream 0.70.0 (#4076): priced by event type, API zone and
// service tier, not by metric and group alone.
if let Some(unit_price) = prices.unit_price(entry) {
let billed_units = entry.value_paid.or(entry.value).unwrap_or(0);
let entry_cost = (billed_units as f64)
* prices.get(&format!("{metric}::{group}")).unwrap_or(&0.0);
Self::accumulate_finite_cost(entry_cost, &mut cost);
Self::accumulate_finite_cost(billed_units as f64 * unit_price, &mut cost);
}
}
}
Expand All @@ -495,10 +479,7 @@ impl MistralProvider {
})
}

fn aggregate_cost(
data: &ModelUsageData,
prices: &HashMap<String, f64>,
) -> Result<f64, ProviderError> {
fn aggregate_cost(data: &ModelUsageData, prices: &PriceIndex) -> Result<f64, ProviderError> {
match Self::aggregate_model(data, prices, AggregationMode::CostOnly)? {
ModelAggregation::Cost(cost) => Ok(cost),
ModelAggregation::CostAndTokens { .. } => {
Expand Down
Loading