use std::collections::HashMap;
use std::sync::Mutex;
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct UsageStats {
pub requests: u64,
pub input_tokens: u64,
pub output_tokens: u64,
pub cache_creation_tokens: u64,
pub cache_read_tokens: u64,
}
impl UsageStats {
fn add(&mut self, other: &UsageStats) {
self.requests += other.requests;
self.input_tokens += other.input_tokens;
self.output_tokens += other.output_tokens;
self.cache_creation_tokens += other.cache_creation_tokens;
self.cache_read_tokens += other.cache_read_tokens;
}
pub(crate) fn merge_max(&mut self, other: &UsageStats) {
self.input_tokens = self.input_tokens.max(other.input_tokens);
self.output_tokens = self.output_tokens.max(other.output_tokens);
self.cache_creation_tokens = self.cache_creation_tokens.max(other.cache_creation_tokens);
self.cache_read_tokens = self.cache_read_tokens.max(other.cache_read_tokens);
}
}
pub struct UsageTracker {
stats: Mutex<HashMap<String, UsageStats>>,
notifier: Mutex<Option<tokio::sync::watch::Sender<()>>>,
}
impl Default for UsageTracker {
fn default() -> Self {
Self {
stats: Mutex::new(HashMap::new()),
notifier: Mutex::new(None),
}
}
}
impl UsageTracker {
pub fn new() -> Self {
Self::default()
}
pub fn set_notifier(&self, sender: tokio::sync::watch::Sender<()>) {
*self.notifier.lock().unwrap() = Some(sender);
}
pub fn record(&self, label: &str, delta: UsageStats) {
let mut stats = self.stats.lock().unwrap();
let entry = stats.entry(label.to_string()).or_default();
entry.add(&UsageStats {
requests: 1,
..delta
});
if let Some(tx) = self.notifier.lock().unwrap().as_ref() {
let _ = tx.send(());
}
}
pub fn snapshot(&self) -> HashMap<String, UsageStats> {
self.stats.lock().unwrap().clone()
}
}
pub fn parse_usage(value: &serde_json::Value) -> Option<UsageStats> {
let usage_obj = value
.get("message")
.and_then(|m| m.get("usage"))
.or_else(|| value.get("usage"));
if let Some(usage_obj) = usage_obj
&& let Some(stats) =
parse_anthropic_usage(usage_obj).or_else(|| parse_openai_usage(usage_obj))
{
return Some(stats);
}
parse_ollama_usage(value)
}
fn parse_anthropic_usage(usage: &serde_json::Value) -> Option<UsageStats> {
let get = |key: &str| usage.get(key).and_then(|v| v.as_u64()).unwrap_or(0);
if usage.get("input_tokens").is_none() && usage.get("output_tokens").is_none() {
return None;
}
Some(UsageStats {
requests: 0,
input_tokens: get("input_tokens"),
output_tokens: get("output_tokens"),
cache_creation_tokens: get("cache_creation_input_tokens"),
cache_read_tokens: get("cache_read_input_tokens"),
})
}
fn parse_openai_usage(usage: &serde_json::Value) -> Option<UsageStats> {
let get = |key: &str| usage.get(key).and_then(|v| v.as_u64()).unwrap_or(0);
if usage.get("prompt_tokens").is_none() && usage.get("completion_tokens").is_none() {
return None;
}
let cache_read_tokens = usage
.get("prompt_tokens_details")
.and_then(|d| d.get("cached_tokens"))
.and_then(|v| v.as_u64())
.unwrap_or(0);
Some(UsageStats {
requests: 0,
input_tokens: get("prompt_tokens"),
output_tokens: get("completion_tokens"),
cache_creation_tokens: 0,
cache_read_tokens,
})
}
fn parse_ollama_usage(body: &serde_json::Value) -> Option<UsageStats> {
let prompt_eval_count = body.get("prompt_eval_count").and_then(|v| v.as_u64());
let eval_count = body.get("eval_count").and_then(|v| v.as_u64());
if prompt_eval_count.is_none() && eval_count.is_none() {
return None;
}
Some(UsageStats {
requests: 0,
input_tokens: prompt_eval_count.unwrap_or(0),
output_tokens: eval_count.unwrap_or(0),
cache_creation_tokens: 0,
cache_read_tokens: 0,
})
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn tracker_records_and_accumulates_by_label() {
let tracker = UsageTracker::new();
tracker.record(
"chat",
UsageStats {
input_tokens: 10,
output_tokens: 5,
..Default::default()
},
);
tracker.record(
"chat",
UsageStats {
input_tokens: 3,
output_tokens: 7,
..Default::default()
},
);
tracker.record(
"other",
UsageStats {
input_tokens: 100,
..Default::default()
},
);
let snapshot = tracker.snapshot();
let chat = snapshot.get("chat").unwrap();
assert_eq!(chat.requests, 2);
assert_eq!(chat.input_tokens, 13);
assert_eq!(chat.output_tokens, 12);
assert_eq!(snapshot.get("other").unwrap().input_tokens, 100);
}
#[test]
fn merge_max_takes_running_totals_not_sums() {
let mut running = UsageStats {
input_tokens: 10,
output_tokens: 1,
..Default::default()
};
running.merge_max(&UsageStats {
input_tokens: 10,
output_tokens: 4,
..Default::default()
});
assert_eq!(running.input_tokens, 10);
assert_eq!(running.output_tokens, 4);
}
#[test]
fn parse_usage_anthropic_non_streaming() {
let body = json!({
"usage": {
"input_tokens": 12,
"output_tokens": 34,
"cache_creation_input_tokens": 5,
"cache_read_input_tokens": 6
}
});
let usage = parse_usage(&body).unwrap();
assert_eq!(usage.input_tokens, 12);
assert_eq!(usage.output_tokens, 34);
assert_eq!(usage.cache_creation_tokens, 5);
assert_eq!(usage.cache_read_tokens, 6);
}
#[test]
fn parse_usage_missing_usage_returns_none() {
let body = json!({"content": []});
assert!(parse_usage(&body).is_none());
}
#[test]
fn parse_usage_anthropic_message_start() {
let event = json!({
"type": "message_start",
"message": {
"usage": {
"input_tokens": 100,
"output_tokens": 1,
"cache_read_input_tokens": 20
}
}
});
let usage = parse_usage(&event).unwrap();
assert_eq!(usage.input_tokens, 100);
assert_eq!(usage.cache_read_tokens, 20);
}
#[test]
fn parse_usage_anthropic_message_delta() {
let event = json!({
"type": "message_delta",
"usage": {
"output_tokens": 42
}
});
let usage = parse_usage(&event).unwrap();
assert_eq!(usage.output_tokens, 42);
}
#[test]
fn parse_usage_openai_with_cached_tokens() {
let body = json!({
"usage": {
"prompt_tokens": 50,
"completion_tokens": 25,
"prompt_tokens_details": {"cached_tokens": 10}
}
});
let usage = parse_usage(&body).unwrap();
assert_eq!(usage.input_tokens, 50);
assert_eq!(usage.output_tokens, 25);
assert_eq!(usage.cache_read_tokens, 10);
}
#[test]
fn parse_usage_ollama_done() {
let body = json!({
"done": true,
"prompt_eval_count": 8,
"eval_count": 16
});
let usage = parse_usage(&body).unwrap();
assert_eq!(usage.input_tokens, 8);
assert_eq!(usage.output_tokens, 16);
}
#[test]
fn parse_usage_ollama_non_done_line_without_counts_is_none() {
let body = json!({"done": false, "response": "partial"});
assert!(parse_usage(&body).is_none());
}
}