1use std::collections::{BTreeSet, HashMap};
4use std::sync::Mutex;
5
6use crate::endpoints::ResolvedEndpoint;
7
8#[derive(Debug, Default, Clone)]
10pub struct EndpointUsage {
11 pub calls: u64,
13 pub input_tokens: u64,
15 pub output_tokens: u64,
17 pub cached_input_tokens: u64,
19 pub cache_creation_input_tokens: u64,
21 pub cost: f64,
23 pub models: BTreeSet<String>,
25}
26
27impl EndpointUsage {
28 const fn tokens(&self) -> u64 {
29 self.input_tokens + self.output_tokens
30 }
31}
32
33#[derive(Debug, Default, Clone)]
35pub struct UsageSnapshot {
36 pub endpoints: HashMap<String, EndpointUsage>,
38 pub total_cost: f64,
40 pub total_calls: u64,
42 pub total_tokens: u64,
44 pub total_input_tokens: u64,
46 pub total_output_tokens: u64,
48 pub total_cached_input_tokens: u64,
50 pub total_cache_creation_input_tokens: u64,
52 pub models: Vec<String>,
54}
55
56impl UsageSnapshot {
57 #[must_use]
59 pub fn from_endpoints(endpoints: &HashMap<String, EndpointUsage>) -> Self {
60 let total_cost = endpoints.values().map(|u| u.cost).sum();
61 let total_calls = endpoints.values().map(|u| u.calls).sum();
62 let total_tokens = endpoints.values().map(EndpointUsage::tokens).sum();
63 let total_input_tokens = endpoints.values().map(|u| u.input_tokens).sum();
64 let total_output_tokens = endpoints.values().map(|u| u.output_tokens).sum();
65 let total_cached_input_tokens = endpoints.values().map(|u| u.cached_input_tokens).sum();
66 let total_cache_creation_input_tokens = endpoints
67 .values()
68 .map(|u| u.cache_creation_input_tokens)
69 .sum();
70 let models: Vec<String> = endpoints
71 .values()
72 .flat_map(|u| u.models.iter().cloned())
73 .collect::<BTreeSet<_>>()
74 .into_iter()
75 .collect();
76 Self {
77 endpoints: endpoints.clone(),
78 total_cost,
79 total_calls,
80 total_tokens,
81 total_input_tokens,
82 total_output_tokens,
83 total_cached_input_tokens,
84 total_cache_creation_input_tokens,
85 models,
86 }
87 }
88}
89
90pub struct UsageTracker {
92 inner: Mutex<UsageInner>,
93}
94
95struct UsageInner {
96 per_endpoint: HashMap<String, EndpointUsage>,
98 global: UsageSnapshot,
100 per_test: Vec<(String, UsageSnapshot)>,
102}
103
104impl UsageTracker {
105 #[must_use]
107 pub fn new() -> Self {
108 Self {
109 inner: Mutex::new(UsageInner {
110 per_endpoint: HashMap::new(),
111 global: UsageSnapshot::default(),
112 per_test: Vec::new(),
113 }),
114 }
115 }
116
117 #[allow(clippy::significant_drop_tightening)]
124 pub fn record_llm_call(
125 &self,
126 endpoint_name: &str,
127 endpoint: &ResolvedEndpoint,
128 model: &str,
129 usage: &LlmUsage,
130 ) {
131 let cost = calculate_llm_cost(endpoint, usage.prompt_tokens, usage.completion_tokens);
132 let mut inner = self.inner.lock().unwrap();
133 let eu = inner
134 .per_endpoint
135 .entry(endpoint_name.to_owned())
136 .or_default();
137 eu.calls += 1;
138 eu.input_tokens += usage.prompt_tokens;
139 eu.output_tokens += usage.completion_tokens;
140 eu.cached_input_tokens += usage.cached_input_tokens;
141 eu.cache_creation_input_tokens += usage.cache_creation_input_tokens;
142 eu.cost += cost;
143 if !model.is_empty() {
144 eu.models.insert(model.to_owned());
145 }
146 }
147
148 #[allow(clippy::significant_drop_tightening)]
154 pub fn record_flat_call(&self, endpoint_name: &str, endpoint: &ResolvedEndpoint) {
155 let mut inner = self.inner.lock().unwrap();
156 let eu = inner
157 .per_endpoint
158 .entry(endpoint_name.to_owned())
159 .or_default();
160 eu.calls += 1;
161 eu.cost += endpoint.per_call_price;
162 }
163
164 #[must_use]
170 pub fn current_test_snapshot(&self) -> UsageSnapshot {
171 let inner = self.inner.lock().unwrap();
172 UsageSnapshot::from_endpoints(&inner.per_endpoint)
173 }
174
175 #[must_use]
181 pub fn global_snapshot(&self) -> UsageSnapshot {
182 let inner = self.inner.lock().unwrap();
183 inner.global.clone()
184 }
185
186 #[must_use]
192 pub fn per_test_snapshots(&self) -> Vec<(String, UsageSnapshot)> {
193 let inner = self.inner.lock().unwrap();
194 inner.per_test.clone()
195 }
196
197 pub fn reset_per_test(&self) {
203 let mut inner = self.inner.lock().unwrap();
204 inner.per_endpoint.clear();
205 }
206
207 pub fn commit_test(&self, test_name: &str) {
214 let mut inner = self.inner.lock().unwrap();
215 let snapshot = UsageSnapshot::from_endpoints(&inner.per_endpoint);
216 let ep_snapshot = inner.per_endpoint.clone();
218 for (ep_name, ep_usage) in &ep_snapshot {
219 let ge = inner.global.endpoints.entry(ep_name.clone()).or_default();
220 ge.calls += ep_usage.calls;
221 ge.input_tokens += ep_usage.input_tokens;
222 ge.output_tokens += ep_usage.output_tokens;
223 ge.cached_input_tokens += ep_usage.cached_input_tokens;
224 ge.cache_creation_input_tokens += ep_usage.cache_creation_input_tokens;
225 ge.cost += ep_usage.cost;
226 ge.models.extend(ep_usage.models.iter().cloned());
227 }
228 inner.global.total_cost += snapshot.total_cost;
229 inner.global.total_calls += snapshot.total_calls;
230 inner.global.total_tokens += snapshot.total_tokens;
231 inner.global.total_input_tokens += snapshot.total_input_tokens;
232 inner.global.total_output_tokens += snapshot.total_output_tokens;
233 inner.global.total_cached_input_tokens += snapshot.total_cached_input_tokens;
234 inner.global.total_cache_creation_input_tokens +=
235 snapshot.total_cache_creation_input_tokens;
236 inner.global.models = inner
237 .global
238 .endpoints
239 .values()
240 .flat_map(|u| u.models.iter().cloned())
241 .collect::<BTreeSet<_>>()
242 .into_iter()
243 .collect();
244 inner.per_test.push((test_name.to_owned(), snapshot));
245 }
246}
247
248impl Default for UsageTracker {
249 fn default() -> Self {
250 Self::new()
251 }
252}
253
254#[allow(clippy::cast_precision_loss, clippy::suboptimal_flops)]
256#[must_use]
257pub fn calculate_llm_cost(
258 endpoint: &ResolvedEndpoint,
259 input_tokens: u64,
260 output_tokens: u64,
261) -> f64 {
262 let input_cost = (input_tokens as f64 / 1_000_000.0) * endpoint.input_price_per_1m;
263 let output_cost = (output_tokens as f64 / 1_000_000.0) * endpoint.output_price_per_1m;
264 input_cost + output_cost
265}
266
267#[derive(Debug, Default, Clone, Copy)]
269pub struct LlmUsage {
270 pub prompt_tokens: u64,
272 pub completion_tokens: u64,
274 pub total_tokens: u64,
276 pub cached_input_tokens: u64,
278 pub cache_creation_input_tokens: u64,
280}
281
282#[derive(Debug, Clone)]
284pub struct LlmResponse {
285 pub content: String,
287 pub usage: LlmUsage,
289}
290
291#[must_use]
307pub fn extract_usage(value: &serde_json::Value) -> LlmUsage {
308 let usage = &value["usage"];
309 LlmUsage {
310 prompt_tokens: usage["prompt_tokens"].as_u64().unwrap_or(0),
311 completion_tokens: usage["completion_tokens"].as_u64().unwrap_or(0),
312 total_tokens: usage["total_tokens"].as_u64().unwrap_or(0),
313 cached_input_tokens: usage["prompt_tokens_details"]["cached_tokens"]
314 .as_u64()
315 .or_else(|| usage["cache_read_input_tokens"].as_u64())
316 .or_else(|| usage["prompt_cache_hit_tokens"].as_u64())
317 .unwrap_or(0),
318 cache_creation_input_tokens: usage["prompt_tokens_details"]["cache_write_tokens"]
319 .as_u64()
320 .or_else(|| usage["cache_creation_input_tokens"].as_u64())
321 .unwrap_or(0),
322 }
323}
324
325#[cfg(test)]
326mod tests {
327 use crate::costs::{calculate_llm_cost, LlmUsage, UsageTracker};
328 use crate::endpoints::ResolvedEndpoint;
329 use crate::scenario::EndpointType;
330
331 fn make_endpoint(
332 name: &str,
333 input_price: f64,
334 output_price: f64,
335 per_call: f64,
336 ) -> ResolvedEndpoint {
337 ResolvedEndpoint {
338 name: name.to_owned(),
339 endpoint_type: EndpointType::Llm,
340 url: String::new(),
341 model: None,
342 api_key: None,
343 headers: std::collections::HashMap::new(),
344 command: None,
345 args: vec![],
346 vision: false,
347 input_price_per_1m: input_price,
348 output_price_per_1m: output_price,
349 per_call_price: per_call,
350 max_attempts: 3,
351 fallbacks: vec![],
352 provider: crate::scenario::Provider::Openai,
353 deployment: None,
354 api_version: None,
355 auth: crate::scenario::AuthConfig::default(),
356 header_commands: std::collections::HashMap::new(),
357 aws: crate::scenario::AwsConfig::default(),
358 }
359 }
360
361 fn usage(prompt: u64, completion: u64, cached: u64) -> LlmUsage {
362 LlmUsage {
363 prompt_tokens: prompt,
364 completion_tokens: completion,
365 total_tokens: prompt + completion,
366 cached_input_tokens: cached,
367 cache_creation_input_tokens: 0,
368 }
369 }
370
371 #[test]
372 fn test_calculate_llm_cost() {
373 let ep = make_endpoint("test", 0.15, 0.60, 0.0);
374 let cost = calculate_llm_cost(&ep, 1_000_000, 500_000);
376 assert!((cost - 0.45).abs() < 0.001);
377 }
378
379 #[test]
380 fn test_calculate_zero_cost() {
381 let ep = make_endpoint("free", 0.0, 0.0, 0.0);
382 let cost = calculate_llm_cost(&ep, 1_000_000, 1_000_000);
383 assert!((cost - 0.0).abs() < f64::EPSILON);
384 }
385
386 #[test]
387 fn test_usage_tracker_record_llm() {
388 let tracker = UsageTracker::new();
389 let ep = make_endpoint("gpt4", 2.50, 10.0, 0.0);
390 tracker.record_llm_call("gpt4", &ep, "gpt-4o", &usage(1000, 500, 200));
391
392 let snap = tracker.current_test_snapshot();
393 assert_eq!(snap.total_calls, 1);
394 assert_eq!(snap.total_tokens, 1500);
395 assert_eq!(snap.total_input_tokens, 1000);
396 assert_eq!(snap.total_output_tokens, 500);
397 assert_eq!(snap.total_cached_input_tokens, 200);
398 assert_eq!(snap.models, vec!["gpt-4o".to_owned()]);
399 assert!(
400 snap.total_cost > 0.0,
401 "expected cost > 0, got {}",
402 snap.total_cost
403 );
404
405 let ep_usage = snap.endpoints.get("gpt4").unwrap();
406 assert_eq!(ep_usage.calls, 1);
407 assert_eq!(ep_usage.input_tokens, 1000);
408 assert_eq!(ep_usage.output_tokens, 500);
409 assert_eq!(ep_usage.cached_input_tokens, 200);
410 assert!(ep_usage.models.contains("gpt-4o"));
411 }
412
413 #[test]
414 fn test_usage_tracker_record_flat() {
415 let tracker = UsageTracker::new();
416 let ep = make_endpoint("agent", 0.0, 0.0, 0.01);
417 tracker.record_flat_call("agent", &ep);
418 tracker.record_flat_call("agent", &ep);
419
420 let snap = tracker.current_test_snapshot();
421 assert_eq!(snap.total_calls, 2);
422 assert!((snap.total_cost - 0.02).abs() < f64::EPSILON);
423 }
424
425 #[test]
426 fn test_usage_tracker_multiple_endpoints() {
427 let tracker = UsageTracker::new();
428 let fast = make_endpoint("fast", 0.15, 0.60, 0.0);
429 let slow = make_endpoint("slow", 2.50, 10.0, 0.0);
430
431 tracker.record_llm_call("fast", &fast, "fast-model", &usage(100, 50, 0));
432 tracker.record_llm_call("slow", &slow, "slow-model", &usage(200, 100, 0));
433
434 let snap = tracker.current_test_snapshot();
435 assert_eq!(snap.total_calls, 2);
436 assert_eq!(snap.endpoints.len(), 2);
437 assert_eq!(
438 snap.models,
439 vec!["fast-model".to_owned(), "slow-model".to_owned()]
440 );
441 }
442
443 #[test]
444 fn test_usage_tracker_reset_and_commit() {
445 let tracker = UsageTracker::new();
446 let ep = make_endpoint("test", 0.15, 0.60, 0.0);
447
448 tracker.record_llm_call("test", &ep, "m1", &usage(100, 50, 0));
449 tracker.commit_test("test1");
450 tracker.reset_per_test();
451
452 tracker.record_llm_call("test", &ep, "m2", &usage(200, 100, 0));
453 tracker.commit_test("test2");
454
455 let global = tracker.global_snapshot();
456 assert_eq!(global.total_calls, 2);
457 assert_eq!(global.total_tokens, 450);
458 assert_eq!(global.models, vec!["m1".to_owned(), "m2".to_owned()]);
459
460 let per_test = tracker.per_test_snapshots();
461 assert_eq!(per_test.len(), 2);
462 assert_eq!(per_test[0].0, "test1");
463 assert_eq!(per_test[1].0, "test2");
464 }
465}