Skip to main content

llm_browser_testkit/
costs.rs

1//! Cost calculation, usage tracking, and pricing logic.
2
3use std::collections::{BTreeSet, HashMap};
4use std::sync::Mutex;
5
6use crate::endpoints::ResolvedEndpoint;
7
8/// Accumulated usage for a single endpoint.
9#[derive(Debug, Default, Clone)]
10pub struct EndpointUsage {
11    /// Number of calls made.
12    pub calls: u64,
13    /// Total input tokens consumed.
14    pub input_tokens: u64,
15    /// Total output tokens consumed.
16    pub output_tokens: u64,
17    /// Input tokens served from the provider's prompt cache.
18    pub cached_input_tokens: u64,
19    /// Input tokens written to the provider's prompt cache (cache creation).
20    pub cache_creation_input_tokens: u64,
21    /// Accumulated cost in USD.
22    pub cost: f64,
23    /// Model names observed on this endpoint, sorted and deduplicated.
24    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/// Aggregated usage across all endpoints for a test or scenario run.
34#[derive(Debug, Default, Clone)]
35pub struct UsageSnapshot {
36    /// Per-endpoint usage.
37    pub endpoints: HashMap<String, EndpointUsage>,
38    /// Total cost across all endpoints.
39    pub total_cost: f64,
40    /// Total calls across all endpoints.
41    pub total_calls: u64,
42    /// Total tokens across all endpoints.
43    pub total_tokens: u64,
44    /// Total input (prompt) tokens across all endpoints.
45    pub total_input_tokens: u64,
46    /// Total output (completion) tokens across all endpoints.
47    pub total_output_tokens: u64,
48    /// Total input tokens served from provider prompt caches.
49    pub total_cached_input_tokens: u64,
50    /// Total input tokens written to provider prompt caches.
51    pub total_cache_creation_input_tokens: u64,
52    /// Model names observed across all endpoints, sorted and deduplicated.
53    pub models: Vec<String>,
54}
55
56impl UsageSnapshot {
57    /// Creates a snapshot from per-endpoint usage data.
58    #[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
90/// Thread-safe usage tracker for the test runner.
91pub struct UsageTracker {
92    inner: Mutex<UsageInner>,
93}
94
95struct UsageInner {
96    /// Per-endpoint usage for the current test.
97    per_endpoint: HashMap<String, EndpointUsage>,
98    /// Aggregated usage across all completed tests.
99    global: UsageSnapshot,
100    /// Per-test snapshots keyed by test name.
101    per_test: Vec<(String, UsageSnapshot)>,
102}
103
104impl UsageTracker {
105    /// Creates a new empty usage tracker.
106    #[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    /// Records a completed LLM call, adding usage, cost, and the answering
118    /// model to the endpoint's accumulator.
119    ///
120    /// # Panics
121    ///
122    /// Panics if the mutex is poisoned.
123    #[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    /// Records a flat-cost call (MCP tool, agent task).
149    ///
150    /// # Panics
151    ///
152    /// Panics if the mutex is poisoned.
153    #[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    /// Reads current usage without locking for the full snapshot.
165    ///
166    /// # Panics
167    ///
168    /// Panics if the mutex is poisoned.
169    #[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    /// Reads the global aggregated snapshot.
176    ///
177    /// # Panics
178    ///
179    /// Panics if the mutex is poisoned.
180    #[must_use]
181    pub fn global_snapshot(&self) -> UsageSnapshot {
182        let inner = self.inner.lock().unwrap();
183        inner.global.clone()
184    }
185
186    /// Reads per-test snapshots.
187    ///
188    /// # Panics
189    ///
190    /// Panics if the mutex is poisoned.
191    #[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    /// Resets the per-test accumulator. Call at the start of each test.
198    ///
199    /// # Panics
200    ///
201    /// Panics if the mutex is poisoned.
202    pub fn reset_per_test(&self) {
203        let mut inner = self.inner.lock().unwrap();
204        inner.per_endpoint.clear();
205    }
206
207    /// Commits the current test's usage to the global accumulator and stores
208    /// it as a per-test snapshot.
209    ///
210    /// # Panics
211    ///
212    /// Panics if the mutex is poisoned.
213    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        // Merge into global
217        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/// Calculates the cost of an LLM call based on token pricing.
255#[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/// Usage info extracted from an LLM API response.
268#[derive(Debug, Default, Clone, Copy)]
269pub struct LlmUsage {
270    /// Number of prompt / input tokens.
271    pub prompt_tokens: u64,
272    /// Number of completion / output tokens.
273    pub completion_tokens: u64,
274    /// Total tokens used.
275    pub total_tokens: u64,
276    /// Input tokens served from the provider's prompt cache (cache hit).
277    pub cached_input_tokens: u64,
278    /// Input tokens written to the provider's prompt cache (cache creation).
279    pub cache_creation_input_tokens: u64,
280}
281
282/// Result of an LLM chat call including usage data.
283#[derive(Debug, Clone)]
284pub struct LlmResponse {
285    /// The message content from the LLM.
286    pub content: String,
287    /// Token usage from the API response.
288    pub usage: LlmUsage,
289}
290
291/// Extracts token usage from an OpenAI-compatible API response JSON.
292///
293/// Cache reads are read from `usage.prompt_tokens_details.cached_tokens`
294/// (`OpenAI`, `Azure`, `OpenRouter`, Groq, xAI, …), falling back to the
295/// `Anthropic`-compatible `usage.cache_read_input_tokens` and `DeepSeek`
296/// `usage.prompt_cache_hit_tokens` spellings. Cache writes are read from
297/// `usage.prompt_tokens_details.cache_write_tokens` (`OpenRouter`, and
298/// `OpenAI`/`Azure` on GPT-5.6+), falling back to the `Anthropic`-compatible
299/// `usage.cache_creation_input_tokens`. Providers that do not report prompt
300/// caching yield `0`.
301///
302/// Note the provider semantics differ: for `OpenAI`-compatible responses
303/// `cached_tokens`/`cache_write_tokens` are subsets of `prompt_tokens`,
304/// whereas for `Anthropic`/Bedrock-style responses the cache counters are
305/// reported in addition to `inputTokens`.
306#[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        // 1M input tokens = $0.15, 500K output = $0.30
375        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}