1use std::collections::{BTreeSet, HashMap};
4use std::sync::Mutex;
5
6use crate::endpoints::ResolvedEndpoint;
7use crate::scenario::Provider;
8
9#[derive(Debug, Default, Clone)]
11pub struct EndpointUsage {
12 pub calls: u64,
14 pub input_tokens: u64,
16 pub output_tokens: u64,
18 pub cached_input_tokens: u64,
20 pub cache_creation_input_tokens: u64,
22 pub cost: f64,
24 pub models: BTreeSet<String>,
26}
27
28impl EndpointUsage {
29 const fn tokens(&self) -> u64 {
30 self.input_tokens + self.output_tokens
31 }
32}
33
34#[derive(Debug, Default, Clone)]
36pub struct UsageSnapshot {
37 pub endpoints: HashMap<String, EndpointUsage>,
39 pub total_cost: f64,
41 pub total_calls: u64,
43 pub total_tokens: u64,
45 pub total_input_tokens: u64,
47 pub total_output_tokens: u64,
49 pub total_cached_input_tokens: u64,
51 pub total_cache_creation_input_tokens: u64,
53 pub models: Vec<String>,
55}
56
57impl UsageSnapshot {
58 #[must_use]
60 pub fn from_endpoints(endpoints: &HashMap<String, EndpointUsage>) -> Self {
61 let total_cost = endpoints.values().map(|u| u.cost).sum();
62 let total_calls = endpoints.values().map(|u| u.calls).sum();
63 let total_tokens = endpoints.values().map(EndpointUsage::tokens).sum();
64 let total_input_tokens = endpoints.values().map(|u| u.input_tokens).sum();
65 let total_output_tokens = endpoints.values().map(|u| u.output_tokens).sum();
66 let total_cached_input_tokens = endpoints.values().map(|u| u.cached_input_tokens).sum();
67 let total_cache_creation_input_tokens = endpoints
68 .values()
69 .map(|u| u.cache_creation_input_tokens)
70 .sum();
71 let models: Vec<String> = endpoints
72 .values()
73 .flat_map(|u| u.models.iter().cloned())
74 .collect::<BTreeSet<_>>()
75 .into_iter()
76 .collect();
77 Self {
78 endpoints: endpoints.clone(),
79 total_cost,
80 total_calls,
81 total_tokens,
82 total_input_tokens,
83 total_output_tokens,
84 total_cached_input_tokens,
85 total_cache_creation_input_tokens,
86 models,
87 }
88 }
89}
90
91pub struct UsageTracker {
93 inner: Mutex<UsageInner>,
94}
95
96struct UsageInner {
97 per_endpoint: HashMap<String, EndpointUsage>,
99 global: UsageSnapshot,
101 per_test: Vec<(String, UsageSnapshot)>,
103}
104
105impl UsageTracker {
106 #[must_use]
108 pub fn new() -> Self {
109 Self {
110 inner: Mutex::new(UsageInner {
111 per_endpoint: HashMap::new(),
112 global: UsageSnapshot::default(),
113 per_test: Vec::new(),
114 }),
115 }
116 }
117
118 #[allow(clippy::significant_drop_tightening)]
125 pub fn record_llm_call(
126 &self,
127 endpoint_name: &str,
128 endpoint: &ResolvedEndpoint,
129 model: &str,
130 usage: &LlmUsage,
131 ) {
132 let cost = calculate_llm_cost(endpoint, usage);
133 let mut inner = self.inner.lock().unwrap();
134 let eu = inner
135 .per_endpoint
136 .entry(endpoint_name.to_owned())
137 .or_default();
138 eu.calls += 1;
139 eu.input_tokens += usage.prompt_tokens;
140 eu.output_tokens += usage.completion_tokens;
141 eu.cached_input_tokens += usage.cached_input_tokens;
142 eu.cache_creation_input_tokens += usage.cache_creation_input_tokens;
143 eu.cost += cost;
144 if !model.is_empty() {
145 eu.models.insert(model.to_owned());
146 }
147 }
148
149 #[allow(clippy::significant_drop_tightening)]
155 pub fn record_flat_call(&self, endpoint_name: &str, endpoint: &ResolvedEndpoint) {
156 let mut inner = self.inner.lock().unwrap();
157 let eu = inner
158 .per_endpoint
159 .entry(endpoint_name.to_owned())
160 .or_default();
161 eu.calls += 1;
162 eu.cost += endpoint.per_call_price;
163 }
164
165 #[must_use]
171 pub fn current_test_snapshot(&self) -> UsageSnapshot {
172 let inner = self.inner.lock().unwrap();
173 UsageSnapshot::from_endpoints(&inner.per_endpoint)
174 }
175
176 #[must_use]
182 pub fn global_snapshot(&self) -> UsageSnapshot {
183 let inner = self.inner.lock().unwrap();
184 inner.global.clone()
185 }
186
187 #[must_use]
193 pub fn per_test_snapshots(&self) -> Vec<(String, UsageSnapshot)> {
194 let inner = self.inner.lock().unwrap();
195 inner.per_test.clone()
196 }
197
198 pub fn reset_per_test(&self) {
204 let mut inner = self.inner.lock().unwrap();
205 inner.per_endpoint.clear();
206 }
207
208 pub fn commit_test(&self, test_name: &str) {
215 let mut inner = self.inner.lock().unwrap();
216 let snapshot = UsageSnapshot::from_endpoints(&inner.per_endpoint);
217 let ep_snapshot = inner.per_endpoint.clone();
219 for (ep_name, ep_usage) in &ep_snapshot {
220 let ge = inner.global.endpoints.entry(ep_name.clone()).or_default();
221 ge.calls += ep_usage.calls;
222 ge.input_tokens += ep_usage.input_tokens;
223 ge.output_tokens += ep_usage.output_tokens;
224 ge.cached_input_tokens += ep_usage.cached_input_tokens;
225 ge.cache_creation_input_tokens += ep_usage.cache_creation_input_tokens;
226 ge.cost += ep_usage.cost;
227 ge.models.extend(ep_usage.models.iter().cloned());
228 }
229 inner.global.total_cost += snapshot.total_cost;
230 inner.global.total_calls += snapshot.total_calls;
231 inner.global.total_tokens += snapshot.total_tokens;
232 inner.global.total_input_tokens += snapshot.total_input_tokens;
233 inner.global.total_output_tokens += snapshot.total_output_tokens;
234 inner.global.total_cached_input_tokens += snapshot.total_cached_input_tokens;
235 inner.global.total_cache_creation_input_tokens +=
236 snapshot.total_cache_creation_input_tokens;
237 inner.global.models = inner
238 .global
239 .endpoints
240 .values()
241 .flat_map(|u| u.models.iter().cloned())
242 .collect::<BTreeSet<_>>()
243 .into_iter()
244 .collect();
245 inner.per_test.push((test_name.to_owned(), snapshot));
246 }
247}
248
249impl Default for UsageTracker {
250 fn default() -> Self {
251 Self::new()
252 }
253}
254
255#[derive(Debug, Clone, Copy, PartialEq, Eq)]
257pub enum CacheAccounting {
258 Subset,
261 Additive,
264}
265
266#[must_use]
268pub const fn cache_accounting(provider: Provider) -> CacheAccounting {
269 match provider {
270 Provider::Openai | Provider::Azure => CacheAccounting::Subset,
271 Provider::Bedrock => CacheAccounting::Additive,
272 }
273}
274
275#[allow(clippy::cast_precision_loss, clippy::suboptimal_flops)]
283#[must_use]
284pub fn calculate_llm_cost(endpoint: &ResolvedEndpoint, usage: &LlmUsage) -> f64 {
285 let million = 1_000_000.0;
286 let (ordinary_input, cached, cache_write) =
287 if cache_accounting(endpoint.provider) == CacheAccounting::Subset {
288 (
289 usage
290 .prompt_tokens
291 .saturating_sub(usage.cached_input_tokens)
292 .saturating_sub(usage.cache_creation_input_tokens),
293 usage.cached_input_tokens,
294 usage.cache_creation_input_tokens,
295 )
296 } else {
297 (
298 usage.prompt_tokens,
299 usage.cached_input_tokens,
300 usage.cache_creation_input_tokens,
301 )
302 };
303 let input_cost = if endpoint.cache_pricing {
304 (ordinary_input as f64 / million) * endpoint.input_price_per_1m
305 + (cached as f64 / million) * endpoint.cached_input_price_per_1m
306 + (cache_write as f64 / million) * endpoint.cache_write_price_per_1m
307 } else {
308 ((ordinary_input + cached + cache_write) as f64 / million) * endpoint.input_price_per_1m
309 };
310 let output_cost = (usage.completion_tokens as f64 / million) * endpoint.output_price_per_1m;
311 input_cost + output_cost + endpoint.per_call_price
312}
313
314#[derive(Debug, Default, Clone, Copy)]
316pub struct LlmUsage {
317 pub prompt_tokens: u64,
319 pub completion_tokens: u64,
321 pub total_tokens: u64,
323 pub cached_input_tokens: u64,
325 pub cache_creation_input_tokens: u64,
327}
328
329#[derive(Debug, Clone)]
331pub struct LlmResponse {
332 pub content: String,
334 pub usage: LlmUsage,
336}
337
338#[must_use]
354pub fn extract_usage(value: &serde_json::Value) -> LlmUsage {
355 let usage = &value["usage"];
356 if usage.is_null() {
357 return extract_gemini_usage(value);
358 }
359 LlmUsage {
360 prompt_tokens: usage["prompt_tokens"].as_u64().unwrap_or(0),
361 completion_tokens: usage["completion_tokens"].as_u64().unwrap_or(0),
362 total_tokens: usage["total_tokens"].as_u64().unwrap_or(0),
363 cached_input_tokens: usage["prompt_tokens_details"]["cached_tokens"]
364 .as_u64()
365 .or_else(|| usage["cache_read_input_tokens"].as_u64())
366 .or_else(|| usage["prompt_cache_hit_tokens"].as_u64())
367 .unwrap_or(0),
368 cache_creation_input_tokens: usage["prompt_tokens_details"]["cache_write_tokens"]
369 .as_u64()
370 .or_else(|| usage["cache_creation_input_tokens"].as_u64())
371 .unwrap_or(0),
372 }
373}
374
375#[must_use]
381fn extract_gemini_usage(value: &serde_json::Value) -> LlmUsage {
382 let metadata = &value["usageMetadata"];
383 LlmUsage {
384 prompt_tokens: metadata["promptTokenCount"].as_u64().unwrap_or(0),
385 completion_tokens: metadata["candidatesTokenCount"].as_u64().unwrap_or(0),
386 total_tokens: metadata["totalTokenCount"].as_u64().unwrap_or(0),
387 cached_input_tokens: metadata["cachedContentTokenCount"].as_u64().unwrap_or(0),
388 cache_creation_input_tokens: 0,
389 }
390}
391
392#[cfg(test)]
393mod tests {
394 use crate::costs::{calculate_llm_cost, LlmUsage, UsageTracker};
395 use crate::endpoints::ResolvedEndpoint;
396 use crate::scenario::EndpointType;
397 use crate::scenario::Provider;
398
399 fn make_endpoint(
400 name: &str,
401 input_price: f64,
402 output_price: f64,
403 per_call: f64,
404 ) -> ResolvedEndpoint {
405 ResolvedEndpoint {
406 name: name.to_owned(),
407 endpoint_type: EndpointType::Llm,
408 url: String::new(),
409 model: None,
410 api_key: None,
411 headers: std::collections::HashMap::new(),
412 command: None,
413 args: vec![],
414 vision: false,
415 input_price_per_1m: input_price,
416 output_price_per_1m: output_price,
417 cached_input_price_per_1m: input_price * 0.1,
418 cache_write_price_per_1m: input_price * 1.25,
419 cache_pricing: true,
420 cache_markers: true,
421 per_call_price: per_call,
422 max_attempts: 3,
423 fallbacks: vec![],
424 provider: Provider::Openai,
425 deployment: None,
426 api_version: None,
427 auth: crate::scenario::AuthConfig::default(),
428 header_commands: std::collections::HashMap::new(),
429 aws: crate::scenario::AwsConfig::default(),
430 }
431 }
432
433 fn usage(prompt: u64, completion: u64, cached: u64) -> LlmUsage {
434 LlmUsage {
435 prompt_tokens: prompt,
436 completion_tokens: completion,
437 total_tokens: prompt + completion,
438 cached_input_tokens: cached,
439 cache_creation_input_tokens: 0,
440 }
441 }
442
443 #[test]
444 fn test_calculate_llm_cost() {
445 let ep = make_endpoint("test", 0.15, 0.60, 0.0);
446 let cost = calculate_llm_cost(&ep, &usage(1_000_000, 500_000, 0));
448 assert!((cost - 0.45).abs() < 0.001);
449 }
450
451 #[test]
452 fn test_calculate_zero_cost() {
453 let ep = make_endpoint("free", 0.0, 0.0, 0.0);
454 let cost = calculate_llm_cost(&ep, &usage(1_000_000, 1_000_000, 0));
455 assert!((cost - 0.0).abs() < f64::EPSILON);
456 }
457
458 #[test]
459 fn test_calculate_llm_cost_cache_subset() {
460 let ep = make_endpoint("gpt", 1.0, 0.0, 0.0);
462 let cost = calculate_llm_cost(&ep, &usage(1_000_000, 0, 1_000_000));
463 assert!((cost - 0.10).abs() < 0.001, "got {cost}");
465 }
466
467 #[test]
468 fn test_calculate_llm_cost_cache_additive_bedrock() {
469 let mut ep = make_endpoint("bedrock", 1.0, 0.0, 0.0);
471 ep.provider = Provider::Bedrock;
472 let cost = calculate_llm_cost(&ep, &usage(1_000_000, 0, 1_000_000));
473 assert!((cost - 1.10).abs() < 0.001, "got {cost}");
475 }
476
477 #[test]
478 fn test_calculate_llm_cost_cache_pricing_disabled() {
479 let mut ep = make_endpoint("bedrock", 1.0, 0.0, 0.0);
482 ep.provider = Provider::Bedrock;
483 ep.cache_pricing = false;
484 let cost = calculate_llm_cost(&ep, &usage(1_000_000, 0, 1_000_000));
485 assert!((cost - 2.0).abs() < 0.001, "got {cost}");
486 }
487
488 #[test]
489 fn test_usage_tracker_record_llm() {
490 let tracker = UsageTracker::new();
491 let ep = make_endpoint("gpt4", 2.50, 10.0, 0.0);
492 tracker.record_llm_call("gpt4", &ep, "gpt-4o", &usage(1000, 500, 200));
493
494 let snap = tracker.current_test_snapshot();
495 assert_eq!(snap.total_calls, 1);
496 assert_eq!(snap.total_tokens, 1500);
497 assert_eq!(snap.total_input_tokens, 1000);
498 assert_eq!(snap.total_output_tokens, 500);
499 assert_eq!(snap.total_cached_input_tokens, 200);
500 assert_eq!(snap.models, vec!["gpt-4o".to_owned()]);
501 assert!(
502 snap.total_cost > 0.0,
503 "expected cost > 0, got {}",
504 snap.total_cost
505 );
506
507 let ep_usage = snap.endpoints.get("gpt4").unwrap();
508 assert_eq!(ep_usage.calls, 1);
509 assert_eq!(ep_usage.input_tokens, 1000);
510 assert_eq!(ep_usage.output_tokens, 500);
511 assert_eq!(ep_usage.cached_input_tokens, 200);
512 assert!(ep_usage.models.contains("gpt-4o"));
513 }
514
515 #[test]
516 fn test_usage_tracker_record_flat() {
517 let tracker = UsageTracker::new();
518 let ep = make_endpoint("agent", 0.0, 0.0, 0.01);
519 tracker.record_flat_call("agent", &ep);
520 tracker.record_flat_call("agent", &ep);
521
522 let snap = tracker.current_test_snapshot();
523 assert_eq!(snap.total_calls, 2);
524 assert!((snap.total_cost - 0.02).abs() < f64::EPSILON);
525 }
526
527 #[test]
528 fn test_usage_tracker_multiple_endpoints() {
529 let tracker = UsageTracker::new();
530 let fast = make_endpoint("fast", 0.15, 0.60, 0.0);
531 let slow = make_endpoint("slow", 2.50, 10.0, 0.0);
532
533 tracker.record_llm_call("fast", &fast, "fast-model", &usage(100, 50, 0));
534 tracker.record_llm_call("slow", &slow, "slow-model", &usage(200, 100, 0));
535
536 let snap = tracker.current_test_snapshot();
537 assert_eq!(snap.total_calls, 2);
538 assert_eq!(snap.endpoints.len(), 2);
539 assert_eq!(
540 snap.models,
541 vec!["fast-model".to_owned(), "slow-model".to_owned()]
542 );
543 }
544
545 #[test]
546 fn test_usage_tracker_reset_and_commit() {
547 let tracker = UsageTracker::new();
548 let ep = make_endpoint("test", 0.15, 0.60, 0.0);
549
550 tracker.record_llm_call("test", &ep, "m1", &usage(100, 50, 0));
551 tracker.commit_test("test1");
552 tracker.reset_per_test();
553
554 tracker.record_llm_call("test", &ep, "m2", &usage(200, 100, 0));
555 tracker.commit_test("test2");
556
557 let global = tracker.global_snapshot();
558 assert_eq!(global.total_calls, 2);
559 assert_eq!(global.total_tokens, 450);
560 assert_eq!(global.models, vec!["m1".to_owned(), "m2".to_owned()]);
561
562 let per_test = tracker.per_test_snapshots();
563 assert_eq!(per_test.len(), 2);
564 assert_eq!(per_test[0].0, "test1");
565 assert_eq!(per_test[1].0, "test2");
566 }
567}