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