use std::{
collections::{BTreeMap, HashMap},
ops::AddAssign,
path::Path,
sync::Arc,
};
use chrono::{Duration, NaiveDate, NaiveDateTime, TimeZone, Timelike};
use rho_providers::model::{ModelMetadata, ModelUsage};
use rusqlite::{Connection, OpenFlags};
use super::{migrations::SCHEMA_VERSION, pricing::catalog_cost_usd_micros, UsageLedgerError};
use crate::sqlite_support::BUSY_TIMEOUT;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum SpendRange {
Today,
Last7Days,
Last30Days,
AllTime,
}
impl SpendRange {
pub(crate) const ALL: [Self; 4] = [
Self::Today,
Self::Last7Days,
Self::Last30Days,
Self::AllTime,
];
fn index(self) -> usize {
match self {
Self::Today => 0,
Self::Last7Days => 1,
Self::Last30Days => 2,
Self::AllTime => 3,
}
}
pub(crate) fn cycled(self, step: isize) -> Self {
let len = Self::ALL.len() as isize;
Self::ALL[(self.index() as isize + step).rem_euclid(len) as usize]
}
fn first_day(self, today: NaiveDate) -> Option<NaiveDate> {
match self {
Self::AllTime => None,
Self::Last30Days => Some(today - Duration::days(29)),
Self::Last7Days => Some(today - Duration::days(6)),
Self::Today => Some(today),
}
}
fn timeline_unit(self) -> TimelineUnit {
match self {
Self::Today => TimelineUnit::Hour,
Self::AllTime | Self::Last30Days | Self::Last7Days => TimelineUnit::Day,
}
}
}
#[derive(Clone, Debug, PartialEq)]
pub(crate) enum RoutePricing {
Local,
Catalog(Arc<ModelMetadata>),
Unknown,
}
pub(crate) fn catalog_route_pricing(provider: &str, model: &str) -> RoutePricing {
use rho_providers::{
model::models_dev::cached_model_metadata,
provider::{provider_descriptor, ProviderId},
};
if provider_descriptor(provider).is_some_and(|descriptor| descriptor.id == ProviderId::Ollama) {
return RoutePricing::Local;
}
match cached_model_metadata(provider, model) {
Some(metadata) if metadata.cost_default.is_some() => {
RoutePricing::Catalog(Arc::new(metadata))
}
_ => RoutePricing::Unknown,
}
}
fn model_leaf(model: &str) -> &str {
model.rsplit('/').next().unwrap_or(model)
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub(crate) struct SpendReports([SpendReport; 4]);
impl SpendReports {
pub(crate) fn get(&self, range: SpendRange) -> &SpendReport {
&self.0[range.index()]
}
pub(crate) fn empty(now: NaiveDateTime) -> Self {
match Self::build(now, None, |_| Ok::<_, std::convert::Infallible>(())) {
Ok(reports) => reports,
}
}
fn build<E>(
now: NaiveDateTime,
first_day: Option<NaiveDate>,
feed: impl FnOnce(&mut dyn FnMut(&PricedRequest<'_>)) -> Result<(), E>,
) -> Result<Self, E> {
let mut builders = SpendRange::ALL.map(|range| ReportBuilder::new(range, now, first_day));
feed(&mut |request| {
for builder in &mut builders {
builder.add(request);
}
})?;
Ok(Self(builders.map(ReportBuilder::finish)))
}
}
pub(crate) fn load_spend_reports<Tz: TimeZone>(
path: &Path,
tz: &Tz,
now: NaiveDateTime,
mut price_route: impl FnMut(&str, &str) -> RoutePricing,
) -> Result<SpendReports, UsageLedgerError> {
if !path.is_file() {
return Ok(SpendReports::empty(now));
}
let connection = Connection::open_with_flags(
path,
OpenFlags::SQLITE_OPEN_READ_ONLY | OpenFlags::SQLITE_OPEN_NO_MUTEX,
)?;
connection.busy_timeout(BUSY_TIMEOUT)?;
let version: i64 = connection.pragma_query_value(None, "user_version", |row| row.get(0))?;
if version > SCHEMA_VERSION {
return Err(UsageLedgerError::UnsupportedSchema {
found: version,
supported: SCHEMA_VERSION,
});
}
if version < 1 {
return Ok(SpendReports::empty(now));
}
let pricing = resolve_route_pricing(&connection, &mut price_route)?;
let local_date = |ms: i64| {
chrono::DateTime::from_timestamp_millis(ms).map(|utc| utc.with_timezone(tz).naive_local())
};
let first_ms: Option<i64> =
connection.query_row("SELECT MIN(occurred_at_ms) FROM usage_events", [], |row| {
row.get(0)
})?;
let first_day = first_ms.and_then(local_date).map(|at| at.date());
SpendReports::build(now, first_day, |add| {
let mut statement = connection.prepare(
"SELECT occurred_at_ms, provider, model, purpose, input_tokens,
output_tokens, cache_read_tokens, cache_write_tokens, total_tokens,
cost_usd_micros
FROM usage_events",
)?;
let mut rows = statement.query([])?;
while let Some(row) = rows.next()? {
let Some(at) = local_date(row.get(0)?) else {
continue;
};
let count = |index: usize| -> rusqlite::Result<Option<u64>> {
Ok(row
.get::<_, Option<i64>>(index)?
.and_then(|value| u64::try_from(value).ok()))
};
let text =
|index: usize| -> rusqlite::Result<&str> { Ok(row.get_ref(index)?.as_str()?) };
let usage = ModelUsage {
input_tokens: count(4)?,
output_tokens: count(5)?,
cache_read_tokens: count(6)?,
cache_write_tokens: count(7)?,
total_tokens: count(8)?,
..ModelUsage::default()
};
let (provider, model) = (text(1)?, text(2)?);
let cost = match count(9)? {
Some(micros) => RowCost::Actual(micros),
None => match pricing.get(provider).and_then(|models| models.get(model)) {
Some(RoutePricing::Local) => RowCost::Local,
Some(RoutePricing::Catalog(metadata)) => {
catalog_cost_usd_micros(&usage, metadata)
.map_or(RowCost::Unpriced, RowCost::Computed)
}
Some(RoutePricing::Unknown) | None => RowCost::Unpriced,
},
};
add(&PricedRequest {
at,
provider,
model: model_leaf(model),
purpose: text(3)?,
totals: SpendTotals::request(row_tokens(&usage), cost),
});
}
Ok::<_, UsageLedgerError>(())
})
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum RowCost {
Actual(u64),
Computed(u64),
Local,
Unpriced,
}
struct PricedRequest<'a> {
at: NaiveDateTime,
provider: &'a str,
model: &'a str,
purpose: &'a str,
totals: SpendTotals,
}
struct ReportBuilder {
start: Option<NaiveDateTime>,
end: NaiveDateTime,
unit: TimelineUnit,
totals: SpendTotals,
providers: HashMap<String, SpendTotals>,
models: HashMap<String, SpendTotals>,
purposes: HashMap<String, SpendTotals>,
buckets: Vec<u64>,
}
impl ReportBuilder {
fn new(range: SpendRange, now: NaiveDateTime, first_day: Option<NaiveDate>) -> Self {
let today = now.date();
let first_day = range.first_day(today).or(first_day);
let unit = range.timeline_unit();
let bucket_count = match (unit, first_day) {
(TimelineUnit::Hour, _) => 24,
(TimelineUnit::Day, Some(first)) => (today - first).num_days().max(0) as usize + 1,
(TimelineUnit::Day, None) => 0,
};
Self {
start: first_day.map(midnight),
end: midnight(today + Duration::days(1)),
unit,
totals: SpendTotals::default(),
providers: HashMap::new(),
models: HashMap::new(),
purposes: HashMap::new(),
buckets: vec![0; bucket_count],
}
}
fn add(&mut self, request: &PricedRequest<'_>) {
let Some(start) = self.start else {
return;
};
if request.at < start || request.at >= self.end {
return;
}
let totals = request.totals;
self.totals += totals;
for (groups, name) in [
(&mut self.providers, request.provider),
(&mut self.models, request.model),
(&mut self.purposes, request.purpose),
] {
match groups.get_mut(name) {
Some(group) => *group += totals,
None => {
groups.insert(name.to_owned(), totals);
}
}
}
let index = match self.unit {
TimelineUnit::Hour => request.at.hour() as usize,
TimelineUnit::Day => (request.at.date() - start.date()).num_days() as usize,
};
let last = self.buckets.len() - 1;
self.buckets[index.min(last)] += totals.equivalent_usd_micros();
}
fn finish(self) -> SpendReport {
SpendReport {
totals: self.totals,
providers: ranked(self.providers),
models: ranked(self.models),
purposes: ranked(self.purposes),
timeline: Timeline {
unit: self.unit,
start: self.start.unwrap_or(self.end),
equivalent_usd_micros: self.buckets,
},
}
}
}
fn midnight(day: NaiveDate) -> NaiveDateTime {
day.and_hms_opt(0, 0, 0).expect("midnight is valid")
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub(crate) struct SpendTotals {
pub(crate) requests: u64,
pub(crate) tokens: u64,
pub(crate) actual_usd_micros: u64,
pub(crate) computed_usd_micros: u64,
pub(crate) local_requests: u64,
pub(crate) unpriced_requests: u64,
}
impl SpendTotals {
fn request(tokens: u64, cost: RowCost) -> Self {
let mut totals = Self {
requests: 1,
tokens,
..Self::default()
};
match cost {
RowCost::Actual(micros) => totals.actual_usd_micros = micros,
RowCost::Computed(micros) => totals.computed_usd_micros = micros,
RowCost::Local => totals.local_requests = 1,
RowCost::Unpriced => totals.unpriced_requests = 1,
}
totals
}
pub(crate) fn equivalent_usd_micros(&self) -> u64 {
self.actual_usd_micros
.saturating_add(self.computed_usd_micros)
}
pub(crate) fn valuation(&self) -> Valuation {
if self.requests > 0 && self.local_requests == self.requests {
Valuation::Local
} else if self.local_requests + self.unpriced_requests == self.requests {
Valuation::Unpriced
} else {
Valuation::Usd(self.equivalent_usd_micros())
}
}
}
impl AddAssign for SpendTotals {
fn add_assign(&mut self, other: Self) {
self.requests += other.requests;
self.tokens = self.tokens.saturating_add(other.tokens);
self.actual_usd_micros = self
.actual_usd_micros
.saturating_add(other.actual_usd_micros);
self.computed_usd_micros = self
.computed_usd_micros
.saturating_add(other.computed_usd_micros);
self.local_requests += other.local_requests;
self.unpriced_requests += other.unpriced_requests;
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum Valuation {
Local,
Unpriced,
Usd(u64),
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub(crate) struct SpendGroup {
pub(crate) name: String,
pub(crate) totals: SpendTotals,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum TimelineUnit {
Hour,
Day,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub(crate) struct Timeline {
pub(crate) unit: TimelineUnit,
pub(crate) start: NaiveDateTime,
pub(crate) equivalent_usd_micros: Vec<u64>,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub(crate) struct SpendReport {
pub(crate) totals: SpendTotals,
pub(crate) providers: Vec<SpendGroup>,
pub(crate) models: Vec<SpendGroup>,
pub(crate) purposes: Vec<SpendGroup>,
pub(crate) timeline: Timeline,
}
fn ranked(groups: HashMap<String, SpendTotals>) -> Vec<SpendGroup> {
let mut groups: Vec<_> = groups
.into_iter()
.map(|(name, totals)| SpendGroup { name, totals })
.collect();
groups.sort_by(|left, right| {
right
.totals
.equivalent_usd_micros()
.cmp(&left.totals.equivalent_usd_micros())
.then(right.totals.requests.cmp(&left.totals.requests))
.then_with(|| left.name.cmp(&right.name))
});
groups
}
fn row_tokens(usage: &ModelUsage) -> u64 {
usage.total_tokens.unwrap_or_else(|| {
usage
.inclusive_prompt_tokens()
.unwrap_or_default()
.saturating_add(usage.output_tokens.unwrap_or_default())
})
}
fn resolve_route_pricing(
connection: &Connection,
price_route: &mut impl FnMut(&str, &str) -> RoutePricing,
) -> Result<HashMap<String, HashMap<String, RoutePricing>>, UsageLedgerError> {
let mut statement = connection
.prepare("SELECT DISTINCT provider, model FROM usage_events ORDER BY provider, model")?;
let routes = statement
.query_map([], |row| {
Ok((row.get::<_, String>(0)?, row.get::<_, String>(1)?))
})?
.collect::<Result<Vec<_>, _>>()?;
let resolved: Vec<_> = routes
.into_iter()
.map(|(provider, model)| {
let pricing = price_route(&provider, &model);
(provider, model, pricing)
})
.collect();
let mut leaf_prices: BTreeMap<&str, &Arc<ModelMetadata>> = BTreeMap::new();
for (_, model, pricing) in &resolved {
if let RoutePricing::Catalog(metadata) = pricing {
leaf_prices.entry(model_leaf(model)).or_insert(metadata);
}
}
let borrowed: Vec<_> = resolved
.iter()
.map(|(_, model, pricing)| match pricing {
RoutePricing::Unknown => leaf_prices
.get(model_leaf(model))
.map(|metadata| RoutePricing::Catalog(Arc::clone(metadata))),
RoutePricing::Local | RoutePricing::Catalog(_) => None,
})
.collect();
let mut pricing: HashMap<String, HashMap<String, RoutePricing>> = HashMap::new();
for ((provider, model, own), borrowed) in resolved.into_iter().zip(borrowed) {
pricing
.entry(provider)
.or_default()
.insert(model, borrowed.unwrap_or(own));
}
Ok(pricing)
}
#[cfg(test)]
#[path = "report_tests.rs"]
mod tests;