use std::collections::HashMap;
use std::sync::{Mutex, OnceLock};
use serde::{Deserialize, Serialize};
use crate::core::gain::model_pricing::{ModelPricing, PricingMatchKind};
#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq)]
pub struct ModelUsage {
pub requests: u64,
pub input_tokens: u64,
pub output_tokens: u64,
pub cache_read_tokens: u64,
pub cache_write_tokens: u64,
pub reasoning_tokens: u64,
#[serde(default)]
pub counterfactual_requests: u64,
#[serde(default)]
pub counterfactual_input_tokens: u64,
#[serde(default)]
pub counterfactual_billed_tokens: u64,
}
impl ModelUsage {
fn add(&mut self, u: &super::usage::RealUsage) {
self.requests += 1;
self.input_tokens += u.input_tokens;
self.output_tokens += u.output_tokens;
self.cache_read_tokens += u.cache_read_tokens;
self.cache_write_tokens += u.cache_write_tokens;
self.reasoning_tokens += u.reasoning_tokens;
if let Some(counted) = u
.wire
.as_deref()
.and_then(|w| w.counterfactual.as_ref())
.and_then(super::counterfactual::CounterfactualSlot::get)
{
self.counterfactual_requests += 1;
self.counterfactual_input_tokens += counted;
self.counterfactual_billed_tokens +=
u.input_tokens + u.cache_read_tokens + u.cache_write_tokens;
}
}
fn billable_tokens(&self) -> u64 {
self.input_tokens + self.output_tokens + self.cache_read_tokens + self.cache_write_tokens
}
}
#[derive(Debug, Clone, Serialize, PartialEq)]
pub struct ModelSpend {
pub model: String,
pub requests: u64,
pub input_tokens: u64,
pub output_tokens: u64,
pub cache_read_tokens: u64,
pub cache_write_tokens: u64,
pub reasoning_tokens: u64,
pub cost_usd: f64,
pub pricing_estimated: bool,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq)]
pub struct CohortUsage {
pub requests: u64,
pub input_tokens: u64,
pub output_tokens: u64,
#[serde(default)]
pub sum_sq_output: u64,
}
impl CohortUsage {
fn add(&mut self, u: &super::usage::RealUsage) {
self.requests += 1;
self.input_tokens += u.input_tokens;
self.output_tokens += u.output_tokens;
self.sum_sq_output += u.output_tokens.saturating_mul(u.output_tokens);
}
#[must_use]
pub fn avg_output(&self) -> Option<f64> {
if self.requests == 0 {
None
} else {
#[allow(clippy::cast_precision_loss)]
Some(self.output_tokens as f64 / self.requests as f64)
}
}
#[must_use]
pub fn variance_output(&self) -> Option<f64> {
if self.requests < 2 {
return None;
}
#[allow(clippy::cast_precision_loss)]
let n = self.requests as f64;
#[allow(clippy::cast_precision_loss)]
let sum = self.output_tokens as f64;
#[allow(clippy::cast_precision_loss)]
let sum_sq = self.sum_sq_output as f64;
let var = (sum_sq - sum * sum / n) / (n - 1.0);
Some(var.max(0.0))
}
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct PersistedUsage {
pub ts: u64,
pub models: HashMap<String, ModelUsage>,
#[serde(default)]
pub cohorts: HashMap<String, CohortUsage>,
}
const MAX_TRACKED_MODELS: usize = 256;
const PROXY_USAGE_FILE: &str = "proxy_usage.json";
fn store() -> &'static Mutex<HashMap<String, ModelUsage>> {
static STORE: OnceLock<Mutex<HashMap<String, ModelUsage>>> = OnceLock::new();
STORE.get_or_init(|| Mutex::new(HashMap::new()))
}
fn cohort_store() -> &'static Mutex<HashMap<String, CohortUsage>> {
static STORE: OnceLock<Mutex<HashMap<String, CohortUsage>>> = OnceLock::new();
STORE.get_or_init(|| Mutex::new(HashMap::new()))
}
pub fn resume_from_disk() {
let Some(persisted) = load_persisted() else {
return;
};
let mut map = store()
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
for (model, usage) in persisted.models {
let acc = map.entry(model).or_default();
acc.requests += usage.requests;
acc.input_tokens += usage.input_tokens;
acc.output_tokens += usage.output_tokens;
acc.cache_read_tokens += usage.cache_read_tokens;
acc.cache_write_tokens += usage.cache_write_tokens;
acc.reasoning_tokens += usage.reasoning_tokens;
acc.counterfactual_requests += usage.counterfactual_requests;
acc.counterfactual_input_tokens += usage.counterfactual_input_tokens;
acc.counterfactual_billed_tokens += usage.counterfactual_billed_tokens;
}
drop(map);
let mut cohorts = cohort_store()
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
for (arm, usage) in persisted.cohorts {
let acc = cohorts.entry(arm).or_default();
acc.requests += usage.requests;
acc.input_tokens += usage.input_tokens;
acc.output_tokens += usage.output_tokens;
}
}
pub fn record(u: &super::usage::RealUsage) {
super::usage_sink::push(u);
if let Some(wire) = u.wire.as_deref()
&& (wire.person.is_some() || wire.project.is_some())
{
let pricing = crate::core::gain::model_pricing::ModelPricing::load();
let baseline = crate::core::config::Config::load().proxy.baseline.clone();
#[allow(clippy::cast_precision_loss)]
let cost_usd = if wire.is_local {
let billable =
u.input_tokens + u.output_tokens + u.cache_read_tokens + u.cache_write_tokens;
baseline.effective_local_shadow_rate() / 1_000_000.0 * billable as f64
} else {
pricing.quote(Some(&u.model)).cost.estimate_usd(
u.input_tokens,
u.output_tokens,
u.cache_write_tokens,
u.cache_read_tokens,
)
};
super::policy_gate::record_spend(wire.person.as_deref(), wire.project.as_deref(), cost_usd);
}
if let Some(wire) = u.wire.as_deref()
&& let Some(routed_from) = wire.routed_from.as_deref()
{
crate::core::savings_ledger::record_routing_event(routed_from, &u.model, u.input_tokens);
}
if u.cache_read_tokens > 0 {
let cost = crate::core::gain::model_pricing::ModelPricing::load()
.quote(Some(&u.model))
.cost;
#[allow(clippy::cast_precision_loss)]
let discount_usd =
(cost.input_per_m - cost.cache_read_per_m) / 1_000_000.0 * u.cache_read_tokens as f64;
crate::core::savings_ledger::record_caching_event(
&u.model,
u.cache_read_tokens,
discount_usd,
);
}
let key = normalize_key(&u.model);
{
let mut map = store()
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let key = if !map.contains_key(&key) && map.len() >= MAX_TRACKED_MODELS {
"unknown".to_string()
} else {
key
};
map.entry(key).or_default().add(u);
}
if let Some(arm) = u.cohort {
let mut cohorts = cohort_store()
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
cohorts.entry(arm.as_str().to_string()).or_default().add(u);
}
persist();
}
#[must_use]
pub fn cohort_snapshot() -> HashMap<String, CohortUsage> {
cohort_store()
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.clone()
}
#[must_use]
pub fn persisted_cohorts() -> HashMap<String, CohortUsage> {
load_persisted().map(|p| p.cohorts).unwrap_or_default()
}
fn normalize_key(model: &str) -> String {
let m = model.trim();
if m.is_empty() {
"unknown".to_string()
} else {
m.to_string()
}
}
pub fn snapshot() -> Vec<ModelSpend> {
let map = store()
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
price_models(&map)
}
pub fn total_cost_usd() -> f64 {
snapshot().iter().map(|m| m.cost_usd).sum()
}
#[derive(Debug, Clone, Serialize, PartialEq)]
pub struct VerifiedSavings {
pub requests: u64,
pub counterfactual_input_tokens: u64,
pub billed_input_tokens: u64,
pub verified_saved_tokens: i64,
}
pub fn verified_savings() -> Option<VerifiedSavings> {
let map = store()
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
verified_of(&map)
}
#[allow(clippy::cast_possible_wrap)]
fn verified_of(map: &HashMap<String, ModelUsage>) -> Option<VerifiedSavings> {
let (mut requests, mut counterfactual, mut billed) = (0u64, 0u64, 0u64);
for usage in map.values() {
requests += usage.counterfactual_requests;
counterfactual += usage.counterfactual_input_tokens;
billed += usage.counterfactual_billed_tokens;
}
(requests > 0).then(|| VerifiedSavings {
requests,
counterfactual_input_tokens: counterfactual,
billed_input_tokens: billed,
verified_saved_tokens: counterfactual as i64 - billed as i64,
})
}
pub fn price_models(map: &HashMap<String, ModelUsage>) -> Vec<ModelSpend> {
let pricing = ModelPricing::load();
let mut rows: Vec<ModelSpend> = map
.iter()
.map(|(model, usage)| price_one(&pricing, model, usage))
.collect();
rows.sort_by(|a, b| {
b.cost_usd
.partial_cmp(&a.cost_usd)
.unwrap_or(std::cmp::Ordering::Equal)
});
rows
}
fn price_one(pricing: &ModelPricing, model: &str, usage: &ModelUsage) -> ModelSpend {
let quote = pricing.quote(Some(model));
let cost_usd = quote.cost.estimate_usd(
usage.input_tokens,
usage.output_tokens,
usage.cache_write_tokens,
usage.cache_read_tokens,
);
ModelSpend {
model: model.to_string(),
requests: usage.requests,
input_tokens: usage.input_tokens,
output_tokens: usage.output_tokens,
cache_read_tokens: usage.cache_read_tokens,
cache_write_tokens: usage.cache_write_tokens,
reasoning_tokens: usage.reasoning_tokens,
cost_usd,
pricing_estimated: !matches!(quote.match_kind, PricingMatchKind::Exact),
}
}
fn usage_path() -> Option<std::path::PathBuf> {
crate::core::data_dir::lean_ctx_data_dir()
.ok()
.map(|d| d.join(PROXY_USAGE_FILE))
}
fn persist() {
let Some(path) = usage_path() else {
return;
};
let models = {
let map = store()
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
map.clone()
};
let cohorts = {
let map = cohort_store()
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
map.clone()
};
let payload = PersistedUsage {
ts: std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs(),
models,
cohorts,
};
let Ok(json) = serde_json::to_string(&payload) else {
return;
};
let tmp = path.with_extension("json.tmp");
if std::fs::write(&tmp, json).is_ok() {
let _ = std::fs::rename(&tmp, &path);
}
}
pub fn load_persisted() -> Option<PersistedUsage> {
let path = usage_path()?;
let data = std::fs::read_to_string(path).ok()?;
serde_json::from_str(&data).ok()
}
pub fn persisted_snapshot() -> Vec<ModelSpend> {
load_persisted()
.map(|p| price_models(&p.models))
.unwrap_or_default()
}
pub fn persisted_dominant_model() -> Option<String> {
let persisted = load_persisted()?;
persisted
.models
.iter()
.filter(|(m, _)| m.as_str() != "unknown" && !m.trim().is_empty())
.max_by_key(|(_, u)| u.billable_tokens())
.filter(|(_, u)| u.billable_tokens() > 0)
.map(|(m, _)| m.clone())
}
#[cfg(test)]
mod tests {
use super::*;
fn usage(
model: &str,
input: u64,
output: u64,
cache_read: u64,
) -> super::super::usage::RealUsage {
super::super::usage::RealUsage {
model: model.to_string(),
input_tokens: input,
output_tokens: output,
cache_read_tokens: cache_read,
..Default::default()
}
}
#[test]
fn counterfactual_pair_recorded_only_when_probe_answered() {
let slot = super::super::counterfactual::CounterfactualSlot::new();
slot.set(5_000);
let with_probe = super::super::usage::RealUsage {
wire: Some(Box::new(super::super::usage::WireContext {
counterfactual: Some(slot),
..Default::default()
})),
cache_write_tokens: 200,
..usage("claude-sonnet-4.5", 1_000, 0, 300)
};
let mut acc = ModelUsage::default();
acc.add(&with_probe);
assert_eq!(acc.counterfactual_requests, 1);
assert_eq!(acc.counterfactual_input_tokens, 5_000);
assert_eq!(acc.counterfactual_billed_tokens, 1_000 + 300 + 200);
let empty = super::super::usage::RealUsage {
wire: Some(Box::new(super::super::usage::WireContext {
counterfactual: Some(super::super::counterfactual::CounterfactualSlot::new()),
..Default::default()
})),
..usage("claude-sonnet-4.5", 1_000, 0, 0)
};
acc.add(&empty);
assert_eq!(acc.requests, 2);
assert_eq!(acc.counterfactual_requests, 1, "empty slot adds no pair");
acc.add(&usage("claude-sonnet-4.5", 10, 0, 0));
assert_eq!(acc.counterfactual_requests, 1);
}
#[test]
fn verified_of_aggregates_and_reports_signed_savings() {
let mut map = HashMap::new();
assert!(verified_of(&map).is_none(), "no coverage → None, not zeros");
map.insert(
"claude-sonnet-4.5".to_string(),
ModelUsage {
counterfactual_requests: 2,
counterfactual_input_tokens: 10_000,
counterfactual_billed_tokens: 6_000,
..Default::default()
},
);
map.insert(
"claude-haiku-4.5".to_string(),
ModelUsage {
counterfactual_requests: 1,
counterfactual_input_tokens: 1_000,
counterfactual_billed_tokens: 1_400,
..Default::default()
},
);
let v = verified_of(&map).expect("covered rows present");
assert_eq!(v.requests, 3);
assert_eq!(v.counterfactual_input_tokens, 11_000);
assert_eq!(v.billed_input_tokens, 7_400);
assert_eq!(v.verified_saved_tokens, 3_600);
map.get_mut("claude-sonnet-4.5")
.unwrap()
.counterfactual_input_tokens = 0;
map.get_mut("claude-sonnet-4.5")
.unwrap()
.counterfactual_billed_tokens = 0;
let negative = verified_of(&map).unwrap();
assert_eq!(
negative.verified_saved_tokens, -400,
"a net-negative verified saving is reported, never clamped"
);
}
#[test]
fn counterfactual_fields_roundtrip_and_default_for_legacy_files() {
let legacy = r#"{"ts":1,"models":{"m":{"requests":1,"input_tokens":10,"output_tokens":5,"cache_read_tokens":0,"cache_write_tokens":0,"reasoning_tokens":0}}}"#;
let p: PersistedUsage = serde_json::from_str(legacy).expect("legacy file loads");
assert_eq!(p.models["m"].counterfactual_requests, 0);
let mut p = PersistedUsage::default();
p.models.insert(
"m".into(),
ModelUsage {
counterfactual_requests: 4,
counterfactual_input_tokens: 9_999,
counterfactual_billed_tokens: 5_555,
..Default::default()
},
);
let back: PersistedUsage =
serde_json::from_str(&serde_json::to_string(&p).unwrap()).unwrap();
assert_eq!(back.models["m"].counterfactual_input_tokens, 9_999);
assert_eq!(back.models["m"].counterfactual_billed_tokens, 5_555);
}
#[test]
fn prices_known_model_with_cache_split() {
let mut map = HashMap::new();
let mut acc = ModelUsage::default();
acc.add(&usage("claude-sonnet-4.5", 1_000_000, 1_000_000, 1_000_000));
map.insert("claude-sonnet-4.5".to_string(), acc);
let rows = price_models(&map);
assert_eq!(rows.len(), 1);
let row = &rows[0];
assert!(
(row.cost_usd - 18.30).abs() < 1e-6,
"cost was {}",
row.cost_usd
);
assert!(!row.pricing_estimated, "exact model match");
assert_eq!(row.requests, 1);
}
#[test]
fn unknown_model_prices_with_fallback_and_is_estimated() {
let mut map = HashMap::new();
let mut acc = ModelUsage::default();
acc.add(&usage("some-novel-model-xyz", 1_000_000, 0, 0));
map.insert("some-novel-model-xyz".to_string(), acc);
let rows = price_models(&map);
assert!(rows[0].pricing_estimated, "fallback pricing is estimated");
assert!(rows[0].cost_usd > 0.0);
}
#[test]
fn dominant_model_picks_highest_token_real_model() {
let mut models = HashMap::new();
models.insert("claude-haiku-4.5".to_string(), {
let mut u = ModelUsage::default();
u.add(&usage("claude-haiku-4.5", 100, 100, 0));
u
});
models.insert("claude-opus-4.5".to_string(), {
let mut u = ModelUsage::default();
u.add(&usage("claude-opus-4.5", 10_000, 10_000, 0));
u
});
models.insert("unknown".to_string(), {
let mut u = ModelUsage::default();
u.add(&usage("unknown", 999_999, 0, 0));
u
});
let dominant = models
.iter()
.filter(|(m, _)| m.as_str() != "unknown")
.max_by_key(|(_, u)| u.billable_tokens())
.map(|(m, _)| m.clone());
assert_eq!(dominant.as_deref(), Some("claude-opus-4.5"));
}
#[test]
fn empty_model_buckets_as_unknown() {
assert_eq!(normalize_key(" "), "unknown");
assert_eq!(normalize_key(""), "unknown");
assert_eq!(normalize_key("gpt-5.4"), "gpt-5.4");
}
#[test]
fn cohort_avg_output_is_mean_per_turn() {
let mut c = CohortUsage::default();
assert_eq!(c.avg_output(), None, "no observations → None");
c.add(&usage("m", 10, 100, 0));
c.add(&usage("m", 10, 50, 0));
assert_eq!(c.requests, 2);
assert_eq!(c.output_tokens, 150);
assert!((c.avg_output().unwrap() - 75.0).abs() < f64::EPSILON);
}
#[test]
fn persisted_usage_without_cohorts_field_loads() {
let json = r#"{"ts":1,"models":{"gpt-5.4":{"requests":1,"input_tokens":10,"output_tokens":5,"cache_read_tokens":0,"cache_write_tokens":0,"reasoning_tokens":0}}}"#;
let p: PersistedUsage = serde_json::from_str(json).expect("loads legacy file");
assert_eq!(p.models.len(), 1);
assert!(p.cohorts.is_empty());
}
#[test]
fn persisted_usage_roundtrips_cohorts() {
let mut p = PersistedUsage::default();
p.cohorts.insert(
"control".into(),
CohortUsage {
requests: 3,
input_tokens: 30,
output_tokens: 300,
sum_sq_output: 30_000,
},
);
let json = serde_json::to_string(&p).unwrap();
let back: PersistedUsage = serde_json::from_str(&json).unwrap();
assert_eq!(back.cohorts.get("control").unwrap().output_tokens, 300);
}
}