use std::collections::HashMap;
use std::sync::RwLock;
use std::sync::atomic::{AtomicU64, Ordering};
use openclaw_core::types::TokenUsage;
pub struct UsageTracker {
totals: RwLock<HashMap<String, ModelUsage>>,
}
#[derive(Debug, Default)]
pub struct ModelUsage {
input_tokens: AtomicU64,
output_tokens: AtomicU64,
request_count: AtomicU64,
}
impl UsageTracker {
#[must_use]
pub fn new() -> Self {
Self {
totals: RwLock::new(HashMap::new()),
}
}
pub fn record(&self, model: &str, usage: &TokenUsage) {
let mut totals = self.totals.write().unwrap();
let entry = totals.entry(model.to_string()).or_default();
entry
.input_tokens
.fetch_add(usage.input_tokens, Ordering::Relaxed);
entry
.output_tokens
.fetch_add(usage.output_tokens, Ordering::Relaxed);
entry.request_count.fetch_add(1, Ordering::Relaxed);
}
#[must_use]
pub fn get_usage(&self, model: &str) -> Option<TokenUsageSummary> {
let totals = self.totals.read().unwrap();
totals.get(model).map(|u| TokenUsageSummary {
input_tokens: u.input_tokens.load(Ordering::Relaxed),
output_tokens: u.output_tokens.load(Ordering::Relaxed),
request_count: u.request_count.load(Ordering::Relaxed),
})
}
#[must_use]
pub fn total_usage(&self) -> TokenUsageSummary {
let totals = self.totals.read().unwrap();
let mut summary = TokenUsageSummary::default();
for usage in totals.values() {
summary.input_tokens += usage.input_tokens.load(Ordering::Relaxed);
summary.output_tokens += usage.output_tokens.load(Ordering::Relaxed);
summary.request_count += usage.request_count.load(Ordering::Relaxed);
}
summary
}
pub fn reset(&self) {
let mut totals = self.totals.write().unwrap();
totals.clear();
}
}
impl Default for UsageTracker {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Clone, Default)]
pub struct TokenUsageSummary {
pub input_tokens: u64,
pub output_tokens: u64,
pub request_count: u64,
}
impl TokenUsageSummary {
#[must_use]
pub const fn total_tokens(&self) -> u64 {
self.input_tokens + self.output_tokens
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_usage_tracking() {
let tracker = UsageTracker::new();
tracker.record(
"claude-3-5-sonnet",
&TokenUsage {
input_tokens: 100,
output_tokens: 50,
cache_read_tokens: None,
cache_write_tokens: None,
},
);
tracker.record(
"claude-3-5-sonnet",
&TokenUsage {
input_tokens: 200,
output_tokens: 100,
cache_read_tokens: None,
cache_write_tokens: None,
},
);
let usage = tracker.get_usage("claude-3-5-sonnet").unwrap();
assert_eq!(usage.input_tokens, 300);
assert_eq!(usage.output_tokens, 150);
assert_eq!(usage.request_count, 2);
}
#[test]
fn test_total_usage() {
let tracker = UsageTracker::new();
tracker.record(
"model1",
&TokenUsage {
input_tokens: 100,
output_tokens: 50,
cache_read_tokens: None,
cache_write_tokens: None,
},
);
tracker.record(
"model2",
&TokenUsage {
input_tokens: 200,
output_tokens: 100,
cache_read_tokens: None,
cache_write_tokens: None,
},
);
let total = tracker.total_usage();
assert_eq!(total.input_tokens, 300);
assert_eq!(total.output_tokens, 150);
assert_eq!(total.total_tokens(), 450);
}
}