use std::collections::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 cost: f64,
}
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,
}
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();
Self {
endpoints: endpoints.clone(),
total_cost,
total_calls,
total_tokens,
}
}
}
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,
input_tokens: u64,
output_tokens: u64,
) {
let cost = calculate_llm_cost(endpoint, input_tokens, output_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 += input_tokens;
eu.output_tokens += output_tokens;
eu.cost += cost;
}
#[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.cost += ep_usage.cost;
}
inner.global.total_cost += snapshot.total_cost;
inner.global.total_calls += snapshot.total_calls;
inner.global.total_tokens += snapshot.total_tokens;
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,
}
#[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),
}
}
#[cfg(test)]
mod tests {
use crate::costs::{calculate_llm_cost, 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![],
input_price_per_1m: input_price,
output_price_per_1m: output_price,
per_call_price: per_call,
}
}
#[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, 1000, 500);
let snap = tracker.current_test_snapshot();
assert_eq!(snap.total_calls, 1);
assert_eq!(snap.total_tokens, 1500);
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);
}
#[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, 100, 50);
tracker.record_llm_call("slow", &slow, 200, 100);
let snap = tracker.current_test_snapshot();
assert_eq!(snap.total_calls, 2);
assert_eq!(snap.endpoints.len(), 2);
}
#[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, 100, 50);
tracker.commit_test("test1");
tracker.reset_per_test();
tracker.record_llm_call("test", &ep, 200, 100);
tracker.commit_test("test2");
let global = tracker.global_snapshot();
assert_eq!(global.total_calls, 2);
assert_eq!(global.total_tokens, 450);
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");
}
}