use std::collections::{BTreeSet, HashMap};
use std::sync::Mutex;
use crate::endpoints::ResolvedEndpoint;
#[derive(Debug, Default, Clone)]
pub struct EndpointUsage {
pub calls: u64,
pub input_tokens: u64,
pub output_tokens: u64,
pub cached_input_tokens: u64,
pub cache_creation_input_tokens: u64,
pub cost: f64,
pub models: BTreeSet<String>,
}
impl EndpointUsage {
const fn tokens(&self) -> u64 {
self.input_tokens + self.output_tokens
}
}
#[derive(Debug, Default, Clone)]
pub struct UsageSnapshot {
pub endpoints: HashMap<String, EndpointUsage>,
pub total_cost: f64,
pub total_calls: u64,
pub total_tokens: u64,
pub total_input_tokens: u64,
pub total_output_tokens: u64,
pub total_cached_input_tokens: u64,
pub total_cache_creation_input_tokens: u64,
pub models: Vec<String>,
}
impl UsageSnapshot {
#[must_use]
pub fn from_endpoints(endpoints: &HashMap<String, EndpointUsage>) -> Self {
let total_cost = endpoints.values().map(|u| u.cost).sum();
let total_calls = endpoints.values().map(|u| u.calls).sum();
let total_tokens = endpoints.values().map(EndpointUsage::tokens).sum();
let total_input_tokens = endpoints.values().map(|u| u.input_tokens).sum();
let total_output_tokens = endpoints.values().map(|u| u.output_tokens).sum();
let total_cached_input_tokens = endpoints.values().map(|u| u.cached_input_tokens).sum();
let total_cache_creation_input_tokens = endpoints
.values()
.map(|u| u.cache_creation_input_tokens)
.sum();
let models: Vec<String> = endpoints
.values()
.flat_map(|u| u.models.iter().cloned())
.collect::<BTreeSet<_>>()
.into_iter()
.collect();
Self {
endpoints: endpoints.clone(),
total_cost,
total_calls,
total_tokens,
total_input_tokens,
total_output_tokens,
total_cached_input_tokens,
total_cache_creation_input_tokens,
models,
}
}
}
pub struct UsageTracker {
inner: Mutex<UsageInner>,
}
struct UsageInner {
per_endpoint: HashMap<String, EndpointUsage>,
global: UsageSnapshot,
per_test: Vec<(String, UsageSnapshot)>,
}
impl UsageTracker {
#[must_use]
pub fn new() -> Self {
Self {
inner: Mutex::new(UsageInner {
per_endpoint: HashMap::new(),
global: UsageSnapshot::default(),
per_test: Vec::new(),
}),
}
}
#[allow(clippy::significant_drop_tightening)]
pub fn record_llm_call(
&self,
endpoint_name: &str,
endpoint: &ResolvedEndpoint,
model: &str,
usage: &LlmUsage,
) {
let cost = calculate_llm_cost(endpoint, usage.prompt_tokens, usage.completion_tokens);
let mut inner = self.inner.lock().unwrap();
let eu = inner
.per_endpoint
.entry(endpoint_name.to_owned())
.or_default();
eu.calls += 1;
eu.input_tokens += usage.prompt_tokens;
eu.output_tokens += usage.completion_tokens;
eu.cached_input_tokens += usage.cached_input_tokens;
eu.cache_creation_input_tokens += usage.cache_creation_input_tokens;
eu.cost += cost;
if !model.is_empty() {
eu.models.insert(model.to_owned());
}
}
#[allow(clippy::significant_drop_tightening)]
pub fn record_flat_call(&self, endpoint_name: &str, endpoint: &ResolvedEndpoint) {
let mut inner = self.inner.lock().unwrap();
let eu = inner
.per_endpoint
.entry(endpoint_name.to_owned())
.or_default();
eu.calls += 1;
eu.cost += endpoint.per_call_price;
}
#[must_use]
pub fn current_test_snapshot(&self) -> UsageSnapshot {
let inner = self.inner.lock().unwrap();
UsageSnapshot::from_endpoints(&inner.per_endpoint)
}
#[must_use]
pub fn global_snapshot(&self) -> UsageSnapshot {
let inner = self.inner.lock().unwrap();
inner.global.clone()
}
#[must_use]
pub fn per_test_snapshots(&self) -> Vec<(String, UsageSnapshot)> {
let inner = self.inner.lock().unwrap();
inner.per_test.clone()
}
pub fn reset_per_test(&self) {
let mut inner = self.inner.lock().unwrap();
inner.per_endpoint.clear();
}
pub fn commit_test(&self, test_name: &str) {
let mut inner = self.inner.lock().unwrap();
let snapshot = UsageSnapshot::from_endpoints(&inner.per_endpoint);
let ep_snapshot = inner.per_endpoint.clone();
for (ep_name, ep_usage) in &ep_snapshot {
let ge = inner.global.endpoints.entry(ep_name.clone()).or_default();
ge.calls += ep_usage.calls;
ge.input_tokens += ep_usage.input_tokens;
ge.output_tokens += ep_usage.output_tokens;
ge.cached_input_tokens += ep_usage.cached_input_tokens;
ge.cache_creation_input_tokens += ep_usage.cache_creation_input_tokens;
ge.cost += ep_usage.cost;
ge.models.extend(ep_usage.models.iter().cloned());
}
inner.global.total_cost += snapshot.total_cost;
inner.global.total_calls += snapshot.total_calls;
inner.global.total_tokens += snapshot.total_tokens;
inner.global.total_input_tokens += snapshot.total_input_tokens;
inner.global.total_output_tokens += snapshot.total_output_tokens;
inner.global.total_cached_input_tokens += snapshot.total_cached_input_tokens;
inner.global.total_cache_creation_input_tokens +=
snapshot.total_cache_creation_input_tokens;
inner.global.models = inner
.global
.endpoints
.values()
.flat_map(|u| u.models.iter().cloned())
.collect::<BTreeSet<_>>()
.into_iter()
.collect();
inner.per_test.push((test_name.to_owned(), snapshot));
}
}
impl Default for UsageTracker {
fn default() -> Self {
Self::new()
}
}
#[allow(clippy::cast_precision_loss, clippy::suboptimal_flops)]
#[must_use]
pub fn calculate_llm_cost(
endpoint: &ResolvedEndpoint,
input_tokens: u64,
output_tokens: u64,
) -> f64 {
let input_cost = (input_tokens as f64 / 1_000_000.0) * endpoint.input_price_per_1m;
let output_cost = (output_tokens as f64 / 1_000_000.0) * endpoint.output_price_per_1m;
input_cost + output_cost
}
#[derive(Debug, Default, Clone, Copy)]
pub struct LlmUsage {
pub prompt_tokens: u64,
pub completion_tokens: u64,
pub total_tokens: u64,
pub cached_input_tokens: u64,
pub cache_creation_input_tokens: u64,
}
#[derive(Debug, Clone)]
pub struct LlmResponse {
pub content: String,
pub usage: LlmUsage,
}
#[must_use]
pub fn extract_usage(value: &serde_json::Value) -> LlmUsage {
let usage = &value["usage"];
LlmUsage {
prompt_tokens: usage["prompt_tokens"].as_u64().unwrap_or(0),
completion_tokens: usage["completion_tokens"].as_u64().unwrap_or(0),
total_tokens: usage["total_tokens"].as_u64().unwrap_or(0),
cached_input_tokens: usage["prompt_tokens_details"]["cached_tokens"]
.as_u64()
.or_else(|| usage["cache_read_input_tokens"].as_u64())
.or_else(|| usage["prompt_cache_hit_tokens"].as_u64())
.unwrap_or(0),
cache_creation_input_tokens: usage["prompt_tokens_details"]["cache_write_tokens"]
.as_u64()
.or_else(|| usage["cache_creation_input_tokens"].as_u64())
.unwrap_or(0),
}
}
#[cfg(test)]
mod tests {
use crate::costs::{calculate_llm_cost, LlmUsage, UsageTracker};
use crate::endpoints::ResolvedEndpoint;
use crate::scenario::EndpointType;
fn make_endpoint(
name: &str,
input_price: f64,
output_price: f64,
per_call: f64,
) -> ResolvedEndpoint {
ResolvedEndpoint {
name: name.to_owned(),
endpoint_type: EndpointType::Llm,
url: String::new(),
model: None,
api_key: None,
headers: std::collections::HashMap::new(),
command: None,
args: vec![],
vision: false,
input_price_per_1m: input_price,
output_price_per_1m: output_price,
per_call_price: per_call,
max_attempts: 3,
fallbacks: vec![],
provider: crate::scenario::Provider::Openai,
deployment: None,
api_version: None,
auth: crate::scenario::AuthConfig::default(),
header_commands: std::collections::HashMap::new(),
aws: crate::scenario::AwsConfig::default(),
}
}
fn usage(prompt: u64, completion: u64, cached: u64) -> LlmUsage {
LlmUsage {
prompt_tokens: prompt,
completion_tokens: completion,
total_tokens: prompt + completion,
cached_input_tokens: cached,
cache_creation_input_tokens: 0,
}
}
#[test]
fn test_calculate_llm_cost() {
let ep = make_endpoint("test", 0.15, 0.60, 0.0);
let cost = calculate_llm_cost(&ep, 1_000_000, 500_000);
assert!((cost - 0.45).abs() < 0.001);
}
#[test]
fn test_calculate_zero_cost() {
let ep = make_endpoint("free", 0.0, 0.0, 0.0);
let cost = calculate_llm_cost(&ep, 1_000_000, 1_000_000);
assert!((cost - 0.0).abs() < f64::EPSILON);
}
#[test]
fn test_usage_tracker_record_llm() {
let tracker = UsageTracker::new();
let ep = make_endpoint("gpt4", 2.50, 10.0, 0.0);
tracker.record_llm_call("gpt4", &ep, "gpt-4o", &usage(1000, 500, 200));
let snap = tracker.current_test_snapshot();
assert_eq!(snap.total_calls, 1);
assert_eq!(snap.total_tokens, 1500);
assert_eq!(snap.total_input_tokens, 1000);
assert_eq!(snap.total_output_tokens, 500);
assert_eq!(snap.total_cached_input_tokens, 200);
assert_eq!(snap.models, vec!["gpt-4o".to_owned()]);
assert!(
snap.total_cost > 0.0,
"expected cost > 0, got {}",
snap.total_cost
);
let ep_usage = snap.endpoints.get("gpt4").unwrap();
assert_eq!(ep_usage.calls, 1);
assert_eq!(ep_usage.input_tokens, 1000);
assert_eq!(ep_usage.output_tokens, 500);
assert_eq!(ep_usage.cached_input_tokens, 200);
assert!(ep_usage.models.contains("gpt-4o"));
}
#[test]
fn test_usage_tracker_record_flat() {
let tracker = UsageTracker::new();
let ep = make_endpoint("agent", 0.0, 0.0, 0.01);
tracker.record_flat_call("agent", &ep);
tracker.record_flat_call("agent", &ep);
let snap = tracker.current_test_snapshot();
assert_eq!(snap.total_calls, 2);
assert!((snap.total_cost - 0.02).abs() < f64::EPSILON);
}
#[test]
fn test_usage_tracker_multiple_endpoints() {
let tracker = UsageTracker::new();
let fast = make_endpoint("fast", 0.15, 0.60, 0.0);
let slow = make_endpoint("slow", 2.50, 10.0, 0.0);
tracker.record_llm_call("fast", &fast, "fast-model", &usage(100, 50, 0));
tracker.record_llm_call("slow", &slow, "slow-model", &usage(200, 100, 0));
let snap = tracker.current_test_snapshot();
assert_eq!(snap.total_calls, 2);
assert_eq!(snap.endpoints.len(), 2);
assert_eq!(
snap.models,
vec!["fast-model".to_owned(), "slow-model".to_owned()]
);
}
#[test]
fn test_usage_tracker_reset_and_commit() {
let tracker = UsageTracker::new();
let ep = make_endpoint("test", 0.15, 0.60, 0.0);
tracker.record_llm_call("test", &ep, "m1", &usage(100, 50, 0));
tracker.commit_test("test1");
tracker.reset_per_test();
tracker.record_llm_call("test", &ep, "m2", &usage(200, 100, 0));
tracker.commit_test("test2");
let global = tracker.global_snapshot();
assert_eq!(global.total_calls, 2);
assert_eq!(global.total_tokens, 450);
assert_eq!(global.models, vec!["m1".to_owned(), "m2".to_owned()]);
let per_test = tracker.per_test_snapshots();
assert_eq!(per_test.len(), 2);
assert_eq!(per_test[0].0, "test1");
assert_eq!(per_test[1].0, "test2");
}
}