use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use serde::{Deserialize, Serialize};
use crate::types::{ModelParams, UsageReport};
const UNLIMITED: u64 = u64::MAX;
const COST_SCALE: f64 = 1_000_000.0;
#[derive(Debug, Clone)]
pub struct TokenLedger {
inner: Arc<TokenLedgerInner>,
}
#[derive(Debug)]
struct TokenLedgerInner {
total_input: AtomicU64,
total_output: AtomicU64,
total_cache_read: AtomicU64,
total_cache_write: AtomicU64,
total_cost_micros: AtomicU64,
rolling_tokens: AtomicU64,
budget_max_tokens: AtomicU64,
}
impl Default for TokenLedger {
fn default() -> Self {
Self::new()
}
}
impl TokenLedger {
#[must_use]
pub fn new() -> Self {
Self {
inner: Arc::new(TokenLedgerInner {
total_input: AtomicU64::new(0),
total_output: AtomicU64::new(0),
total_cache_read: AtomicU64::new(0),
total_cache_write: AtomicU64::new(0),
total_cost_micros: AtomicU64::new(0),
rolling_tokens: AtomicU64::new(0),
budget_max_tokens: AtomicU64::new(UNLIMITED),
}),
}
}
#[must_use]
pub fn with_budget(max_tokens: Option<u64>) -> Self {
let ledger = Self::new();
ledger.set_budget(max_tokens);
ledger
}
pub fn record(&self, usage: &UsageReport) {
let input = fold(usage.input);
let output = fold(usage.output);
let cache_read = fold(usage.cache_read);
let cache_write = fold(usage.cache_write);
self.inner.total_input.fetch_add(input, Ordering::Relaxed);
self.inner.total_output.fetch_add(output, Ordering::Relaxed);
self.inner
.total_cache_read
.fetch_add(cache_read, Ordering::Relaxed);
self.inner
.total_cache_write
.fetch_add(cache_write, Ordering::Relaxed);
self.inner
.rolling_tokens
.fetch_add(input + output, Ordering::Relaxed);
#[allow(clippy::cast_possible_truncation, clippy::cast_sign_loss)]
let cost_micros = (cost_of(usage) * COST_SCALE).round() as u64;
self.inner
.total_cost_micros
.fetch_add(cost_micros, Ordering::Relaxed);
}
#[must_use]
pub fn total(&self) -> u64 {
self.inner.total_input.load(Ordering::Relaxed)
+ self.inner.total_output.load(Ordering::Relaxed)
+ self.inner.total_cache_read.load(Ordering::Relaxed)
+ self.inner.total_cache_write.load(Ordering::Relaxed)
}
#[must_use]
pub fn input(&self) -> u64 {
self.inner.total_input.load(Ordering::Relaxed)
}
#[must_use]
pub fn output(&self) -> u64 {
self.inner.total_output.load(Ordering::Relaxed)
}
#[must_use]
pub fn cost(&self) -> f64 {
#[allow(clippy::cast_precision_loss)]
let micros = self.inner.total_cost_micros.load(Ordering::Relaxed) as f64;
micros / COST_SCALE
}
#[must_use]
pub fn rolling(&self) -> u64 {
self.inner.rolling_tokens.load(Ordering::Relaxed)
}
pub fn reset_rolling(&self) {
self.inner.rolling_tokens.store(0, Ordering::Relaxed);
}
pub fn set_budget(&self, max_tokens: Option<u64>) {
let v = max_tokens.unwrap_or(UNLIMITED);
self.inner.budget_max_tokens.store(v, Ordering::Relaxed);
}
#[must_use]
pub fn budget(&self) -> Option<u64> {
match self.inner.budget_max_tokens.load(Ordering::Relaxed) {
UNLIMITED => None,
v => Some(v),
}
}
#[must_use]
pub fn remaining(&self) -> Option<u64> {
self.budget().map(|b| b.saturating_sub(self.total()))
}
#[must_use]
pub fn is_exhausted(&self) -> bool {
match self.budget() {
None => false,
Some(b) => self.total() >= b,
}
}
#[must_use]
pub fn snapshot(&self) -> TokenSnapshot {
TokenSnapshot {
total: self.total(),
input: self.input(),
output: self.output(),
rolling: self.rolling(),
budget: self.budget(),
remaining: self.remaining(),
cost: self.cost(),
}
}
#[must_use]
pub fn effective_max_tokens(&self, params: &ModelParams) -> usize {
match self.remaining() {
None => params.max_tokens,
#[allow(clippy::cast_possible_truncation)]
Some(rem) => params.max_tokens.min(rem as usize),
}
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct TokenSnapshot {
pub total: u64,
pub input: u64,
pub output: u64,
pub rolling: u64,
pub budget: Option<u64>,
pub remaining: Option<u64>,
pub cost: f64,
}
fn fold(v: f64) -> u64 {
#[allow(clippy::cast_possible_truncation, clippy::cast_sign_loss)]
let result = v.max(0.0).round() as u64;
result
}
fn cost_of(usage: &UsageReport) -> f64 {
let c = &usage.cost;
(c.input + c.output + c.cache_read + c.cache_write).max(0.0)
}