use std::cmp::Reverse;
use std::collections::HashMap;
use std::sync::{Mutex, MutexGuard, OnceLock};
use super::token_envelope::TokenEnvelope;
#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]
pub(crate) struct SessionUsage {
pub session_id: String,
pub per_model: HashMap<String, ModelUsageEntry>,
pub per_provider: HashMap<String, ProviderUsageEntry>,
pub total_requests: usize,
pub total_tokens: usize,
pub total_saved: usize,
pub total_cost_usd: f64,
}
#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]
pub(crate) struct ModelUsageEntry {
pub requests: usize,
pub input_tokens: usize,
pub output_tokens: usize,
pub cache_read_tokens: usize,
pub tokens_saved: usize,
pub cost_usd: f64,
}
#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]
pub(crate) struct ProviderUsageEntry {
pub requests: usize,
pub total_tokens: usize,
pub tokens_saved: usize,
pub cost_usd: f64,
}
#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]
pub(crate) struct CompressionOverview {
pub total_tokens: usize,
pub total_saved: usize,
pub avg_compression_ratio: f64,
pub best_model: Option<String>,
pub best_ratio: f64,
}
static NORMALIZER: OnceLock<Mutex<SessionUsage>> = OnceLock::new();
fn usage_guard() -> MutexGuard<'static, SessionUsage> {
NORMALIZER
.get_or_init(|| Mutex::new(SessionUsage::default()))
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
}
fn compression_ratio(input_tokens: usize, tokens_saved: usize) -> f64 {
let original = input_tokens.saturating_add(tokens_saved);
if original == 0 {
0.0
} else {
tokens_saved as f64 / original as f64
}
}
pub(crate) fn record_envelope(envelope: &TokenEnvelope) {
let total_tokens = envelope.total_tokens();
let cost = envelope.cost_usd.unwrap_or_default();
let mut usage = usage_guard();
usage.total_requests = usage.total_requests.saturating_add(1);
usage.total_tokens = usage.total_tokens.saturating_add(total_tokens);
usage.total_saved = usage.total_saved.saturating_add(envelope.tokens_saved);
usage.total_cost_usd += cost;
let model = usage.per_model.entry(envelope.model.clone()).or_default();
model.requests = model.requests.saturating_add(1);
model.input_tokens = model.input_tokens.saturating_add(envelope.input_tokens);
model.output_tokens = model.output_tokens.saturating_add(envelope.output_tokens);
model.cache_read_tokens = model
.cache_read_tokens
.saturating_add(envelope.cache_read_tokens);
model.tokens_saved = model.tokens_saved.saturating_add(envelope.tokens_saved);
model.cost_usd += cost;
let provider = usage
.per_provider
.entry(format!("{:?}", envelope.provider))
.or_default();
provider.requests = provider.requests.saturating_add(1);
provider.total_tokens = provider.total_tokens.saturating_add(total_tokens);
provider.tokens_saved = provider.tokens_saved.saturating_add(envelope.tokens_saved);
provider.cost_usd += cost;
}
pub(crate) fn session_usage() -> SessionUsage {
usage_guard().clone()
}
pub(crate) fn model_breakdown() -> Vec<(String, ModelUsageEntry)> {
let mut entries = usage_guard()
.per_model
.clone()
.into_iter()
.collect::<Vec<_>>();
entries.sort_unstable_by_key(|(_, entry)| {
Reverse(entry.input_tokens.saturating_add(entry.output_tokens))
});
entries
}
#[must_use]
pub(crate) fn provider_breakdown() -> Vec<(String, ProviderUsageEntry)> {
let mut entries = usage_guard()
.per_provider
.clone()
.into_iter()
.collect::<Vec<_>>();
entries.sort_unstable_by_key(|(_, entry)| Reverse(entry.total_tokens));
entries
}
#[must_use]
pub(crate) fn compression_overview() -> CompressionOverview {
let usage = usage_guard();
let total_input = usage
.per_model
.values()
.map(|entry| entry.input_tokens)
.sum();
let best = usage.per_model.iter().fold(None, |best, (model, entry)| {
let ratio = compression_ratio(entry.input_tokens, entry.tokens_saved);
match best {
Some((_, best_ratio)) if ratio <= best_ratio => best,
_ => Some((model.clone(), ratio)),
}
});
let (best_model, best_ratio) = best.map_or((None, 0.0), |(model, ratio)| (Some(model), ratio));
CompressionOverview {
total_tokens: usage.total_tokens,
total_saved: usage.total_saved,
avg_compression_ratio: compression_ratio(total_input, usage.total_saved),
best_model,
best_ratio,
}
}
pub(crate) fn reset_usage() {
*usage_guard() = SessionUsage::default();
}
#[cfg(test)]
mod tests {
use super::{compression_overview as overview, model_breakdown, provider_breakdown};
use super::{record_envelope, reset_usage, session_usage};
use crate::core::context_kernel::token_envelope::{ProviderKind, TokenEnvelope};
use std::sync::{Mutex, MutexGuard};
static TEST_LOCK: Mutex<()> = Mutex::new(());
fn setup() -> MutexGuard<'static, ()> {
let guard = TEST_LOCK
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
reset_usage();
record_envelope(&envelope("low", ProviderKind::Anthropic, 20, 0));
record_envelope(&envelope("high", ProviderKind::Gemini, 60, 60));
record_envelope(&envelope("low", ProviderKind::OpenAi, 20, 0));
guard
}
fn envelope(model: &str, provider: ProviderKind, input: usize, saved: usize) -> TokenEnvelope {
TokenEnvelope {
model: model.to_owned(),
provider,
input_tokens: input,
tokens_saved: saved,
..TokenEnvelope::default()
}
}
#[test]
fn record_updates_totals() {
let _guard = setup();
let usage = session_usage();
assert_eq!(
(usage.total_requests, usage.total_tokens, usage.total_saved),
(3, 100, 60)
);
}
#[test]
fn per_model_breakdown() {
let _guard = setup();
let entries = model_breakdown();
assert_eq!((entries.len(), entries[0].0.as_str()), (2, "high"));
}
#[test]
fn per_provider_breakdown() {
let _guard = setup();
let entries = provider_breakdown();
assert_eq!((entries.len(), entries[0].0.as_str()), (3, "Gemini"));
assert!(entries[0].1.total_tokens > entries[1].1.total_tokens);
}
#[test]
fn compression_overview() {
let _guard = setup();
let summary = overview();
assert!((summary.avg_compression_ratio - 0.375).abs() < f64::EPSILON);
assert_eq!(summary.best_model.as_deref(), Some("high"));
assert!((summary.best_ratio - 0.5).abs() < f64::EPSILON);
}
#[test]
fn reset_clears() {
let _guard = setup();
reset_usage();
let usage = session_usage();
assert_eq!((usage.total_requests, usage.total_tokens), (0, 0));
assert!(usage.per_model.is_empty());
}
}