use chrono::NaiveDate;
use serde::{Deserialize, Serialize};
use std::collections::{BTreeMap, BTreeSet};
use std::path::{Path, PathBuf};
use crate::error::{AppError, io_err, json_err};
use crate::io::load;
use crate::model::UsageEntry;
use crate::tokens::local_date;
const PRICING_VERSION: u32 = 1;
#[derive(Debug, Default, Clone, Copy, PartialEq, Serialize, Deserialize)]
pub struct PricingFields {
#[serde(default)]
pub input: Option<f64>,
#[serde(default)]
pub output: Option<f64>,
#[serde(default)]
pub cache_read: Option<f64>,
#[serde(default)]
pub cache_write: Option<f64>,
}
impl PricingFields {
fn is_empty(&self) -> bool {
self.input.is_none()
&& self.output.is_none()
&& self.cache_read.is_none()
&& self.cache_write.is_none()
}
fn overlay(&self, base: &PricingFields) -> PricingFields {
PricingFields {
input: self.input.or(base.input),
output: self.output.or(base.output),
cache_read: self.cache_read.or(base.cache_read),
cache_write: self.cache_write.or(base.cache_write),
}
}
}
#[derive(Debug, Clone, PartialEq, Default, Serialize, Deserialize)]
pub struct PeakPricing {
#[serde(default)]
pub hours: Vec<[u8; 2]>,
#[serde(default)]
pub utc_offset: i8,
#[serde(flatten)]
pub fields: PricingFields,
}
#[derive(Debug, Clone, PartialEq, Default, Serialize, Deserialize)]
pub struct LongContextPricing {
pub above: u64,
#[serde(flatten)]
pub fields: PricingFields,
}
#[derive(Debug, Clone, PartialEq, Default, Serialize, Deserialize)]
pub struct PricingVersion {
#[serde(default)]
pub since: Option<NaiveDate>,
#[serde(default, skip_serializing_if = "std::ops::Not::not")]
pub force: bool,
#[serde(flatten)]
pub fields: PricingFields,
#[serde(default)]
pub long_context: Option<LongContextPricing>,
#[serde(default)]
pub peak: Option<PeakPricing>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct PricingFile {
#[serde(default = "default_version")]
pub version: u32,
#[serde(default)]
pub models: BTreeMap<String, Vec<PricingVersion>>,
}
fn default_version() -> u32 {
PRICING_VERSION
}
impl Default for PricingFile {
fn default() -> Self {
Self {
version: PRICING_VERSION,
models: BTreeMap::new(),
}
}
}
pub fn pricing_path() -> Result<PathBuf, AppError> {
let root = match std::env::var_os("XDG_CONFIG_HOME") {
Some(v) if !v.is_empty() => PathBuf::from(v),
_ => load::home_dir()?.join(".config"),
};
Ok(root.join("tokrs").join("pricing.json"))
}
pub fn load_pricing(path: &Path) -> Result<PricingFile, AppError> {
if !path.is_file() {
return Ok(PricingFile::default());
}
let raw = std::fs::read(path).map_err(|e| io_err("read", path, e))?;
let file: PricingFile = serde_json::from_slice(&raw).map_err(|e| json_err(path, e))?;
if file.version != PRICING_VERSION {
let e = std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!("unsupported pricing version {}", file.version),
);
return Err(json_err(path, serde_json::Error::io(e)));
}
Ok(file)
}
pub fn sync_models(
table: &mut PricingFile,
path: &Path,
entries: &[UsageEntry],
) -> Result<usize, AppError> {
let models: BTreeSet<&str> = entries.iter().map(|e| e.model.as_str()).collect();
let mut added = 0usize;
for model in models {
if model.is_empty() || model == "unknown" {
continue;
}
if find_versions(table, model).is_none() {
table
.models
.insert(model.to_string(), vec![PricingVersion::default()]);
added += 1;
}
}
if added == 0 {
return Ok(0);
}
if let Some(dir) = path.parent() {
std::fs::create_dir_all(dir).map_err(|e| io_err("create", dir, e))?;
}
let json = serde_json::to_string_pretty(table).map_err(|e| {
io_err(
"serialize",
path,
std::io::Error::other(format!("pricing table: {e}")),
)
})?;
let tmp = path.with_extension("json.tmp");
std::fs::write(&tmp, &json).map_err(|e| io_err("write", &tmp, e))?;
std::fs::rename(&tmp, path).map_err(|e| io_err("rename", path, e))?;
Ok(added)
}
pub fn resolve(entries: &mut [UsageEntry], table: &PricingFile) {
for entry in entries {
let version = match_version(table, entry);
let estimated = version.and_then(|v| estimate(entry, v));
let forced = version.is_some_and(|v| v.force) && estimated.is_some();
entry.cost_usd = if forced {
estimated
} else {
entry.self_cost_usd.or(estimated)
};
}
}
fn match_version<'a>(table: &'a PricingFile, entry: &UsageEntry) -> Option<&'a PricingVersion> {
let versions = find_versions(table, &entry.model)?;
let date = local_date(entry.created_at);
versions
.iter()
.filter(|v| v.since.is_none_or(|since| date >= since))
.max_by_key(|v| v.since.unwrap_or(NaiveDate::MIN))
}
fn estimate(entry: &UsageEntry, version: &PricingVersion) -> Option<f64> {
let mut fields = version.fields;
if fields.is_empty() {
return None;
}
if let Some(peak) = &version.peak
&& in_peak_hours(peak, entry.created_at)
{
fields = peak.fields.overlay(&fields);
}
if let Some(long) = &version.long_context {
let context = entry.input_tokens + entry.cache_read_tokens + entry.cache_creation_tokens;
if context >= long.above {
fields = long.fields.overlay(&fields);
}
}
let price = |v: Option<f64>| v.unwrap_or(0.0).max(0.0);
let cost = entry.input_tokens as f64 * price(fields.input)
+ entry.output_tokens as f64 * price(fields.output)
+ entry.cache_read_tokens as f64 * price(fields.cache_read)
+ entry.cache_creation_tokens as f64 * price(fields.cache_write);
Some(cost / 1e6)
}
fn find_versions<'a>(table: &'a PricingFile, model: &str) -> Option<&'a [PricingVersion]> {
if let Some(v) = table.models.get(model) {
return Some(v);
}
table
.models
.iter()
.filter(|(key, _)| {
!key.is_empty()
&& model.len() > key.len()
&& model.starts_with(key.as_str())
&& !model.as_bytes()[key.len()].is_ascii_alphanumeric()
})
.max_by_key(|(key, _)| key.len())
.map(|(_, versions)| versions.as_slice())
}
fn in_peak_hours(peak: &PeakPricing, created_at: i64) -> bool {
let local = created_at + i64::from(peak.utc_offset) * 3600;
let hour = (local.rem_euclid(86_400) / 3_600) as u8;
peak.hours.iter().any(|&[start, end]| {
if start <= end {
(start..end).contains(&hour)
} else {
hour >= start || hour < end
}
})
}
#[cfg(test)]
#[path = "tests/prince_test.rs"]
mod tests;