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;
7use crate::scenario::Provider;
8
9/// Accumulated usage for a single endpoint.
10#[derive(Debug, Default, Clone)]
11pub struct EndpointUsage {
12    /// Number of calls made.
13    pub calls: u64,
14    /// Total input tokens consumed.
15    pub input_tokens: u64,
16    /// Total output tokens consumed.
17    pub output_tokens: u64,
18    /// Input tokens served from the provider's prompt cache.
19    pub cached_input_tokens: u64,
20    /// Input tokens written to the provider's prompt cache (cache creation).
21    pub cache_creation_input_tokens: u64,
22    /// Accumulated cost in USD.
23    pub cost: f64,
24    /// Model names observed on this endpoint, sorted and deduplicated.
25    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/// Aggregated usage across all endpoints for a test or scenario run.
35#[derive(Debug, Default, Clone)]
36pub struct UsageSnapshot {
37    /// Per-endpoint usage.
38    pub endpoints: HashMap<String, EndpointUsage>,
39    /// Total cost across all endpoints.
40    pub total_cost: f64,
41    /// Total calls across all endpoints.
42    pub total_calls: u64,
43    /// Total tokens across all endpoints.
44    pub total_tokens: u64,
45    /// Total input (prompt) tokens across all endpoints.
46    pub total_input_tokens: u64,
47    /// Total output (completion) tokens across all endpoints.
48    pub total_output_tokens: u64,
49    /// Total input tokens served from provider prompt caches.
50    pub total_cached_input_tokens: u64,
51    /// Total input tokens written to provider prompt caches.
52    pub total_cache_creation_input_tokens: u64,
53    /// Model names observed across all endpoints, sorted and deduplicated.
54    pub models: Vec<String>,
55}
56
57impl UsageSnapshot {
58    /// Creates a snapshot from per-endpoint usage data.
59    #[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
91/// Thread-safe usage tracker for the test runner.
92pub struct UsageTracker {
93    inner: Mutex<UsageInner>,
94}
95
96struct UsageInner {
97    /// Per-endpoint usage for the current test.
98    per_endpoint: HashMap<String, EndpointUsage>,
99    /// Aggregated usage across all completed tests.
100    global: UsageSnapshot,
101    /// Per-test snapshots keyed by test name.
102    per_test: Vec<(String, UsageSnapshot)>,
103}
104
105impl UsageTracker {
106    /// Creates a new empty usage tracker.
107    #[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    /// Records a completed LLM call, adding usage, cost, and the answering
119    /// model to the endpoint's accumulator.
120    ///
121    /// # Panics
122    ///
123    /// Panics if the mutex is poisoned.
124    #[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    /// Records a flat-cost call (MCP tool, agent task).
150    ///
151    /// # Panics
152    ///
153    /// Panics if the mutex is poisoned.
154    #[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    /// Reads current usage without locking for the full snapshot.
166    ///
167    /// # Panics
168    ///
169    /// Panics if the mutex is poisoned.
170    #[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    /// Reads the global aggregated snapshot.
177    ///
178    /// # Panics
179    ///
180    /// Panics if the mutex is poisoned.
181    #[must_use]
182    pub fn global_snapshot(&self) -> UsageSnapshot {
183        let inner = self.inner.lock().unwrap();
184        inner.global.clone()
185    }
186
187    /// Reads per-test snapshots.
188    ///
189    /// # Panics
190    ///
191    /// Panics if the mutex is poisoned.
192    #[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    /// Resets the per-test accumulator. Call at the start of each test.
199    ///
200    /// # Panics
201    ///
202    /// Panics if the mutex is poisoned.
203    pub fn reset_per_test(&self) {
204        let mut inner = self.inner.lock().unwrap();
205        inner.per_endpoint.clear();
206    }
207
208    /// Commits the current test's usage to the global accumulator and stores
209    /// it as a per-test snapshot.
210    ///
211    /// # Panics
212    ///
213    /// Panics if the mutex is poisoned.
214    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        // Merge into global
218        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/// How a provider reports cache tokens relative to its input token count.
256#[derive(Debug, Clone, Copy, PartialEq, Eq)]
257pub enum CacheAccounting {
258    /// Cache tokens are a **subset** of the reported input tokens
259    /// (`OpenAI`, `Azure`, `OpenRouter`, Groq, xAI, `DeepSeek`, …).
260    Subset,
261    /// Cache tokens are reported **in addition** to the input tokens
262    /// (`Anthropic` Messages and AWS Bedrock Converse).
263    Additive,
264}
265
266/// Returns the cache accounting convention for a provider.
267#[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/// Calculates the cost of an LLM call based on token pricing.
276///
277/// When [`ResolvedEndpoint::cache_pricing`] is enabled (the default), cache
278/// reads and writes are billed at their cache rates; otherwise every prompt
279/// token is billed at the flat input price. Cache tokens are treated as a
280/// subset or as additive to the input count depending on the provider (see
281/// [`CacheAccounting`]).
282#[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/// Usage info extracted from an LLM API response.
315#[derive(Debug, Default, Clone, Copy)]
316pub struct LlmUsage {
317    /// Number of prompt / input tokens.
318    pub prompt_tokens: u64,
319    /// Number of completion / output tokens.
320    pub completion_tokens: u64,
321    /// Total tokens used.
322    pub total_tokens: u64,
323    /// Input tokens served from the provider's prompt cache (cache hit).
324    pub cached_input_tokens: u64,
325    /// Input tokens written to the provider's prompt cache (cache creation).
326    pub cache_creation_input_tokens: u64,
327}
328
329/// Result of an LLM chat call including usage data.
330#[derive(Debug, Clone)]
331pub struct LlmResponse {
332    /// The message content from the LLM.
333    pub content: String,
334    /// Token usage from the API response.
335    pub usage: LlmUsage,
336}
337
338/// Extracts token usage from an OpenAI-compatible API response JSON.
339///
340/// Cache reads are read from `usage.prompt_tokens_details.cached_tokens`
341/// (`OpenAI`, `Azure`, `OpenRouter`, Groq, xAI, …), falling back to the
342/// `Anthropic`-compatible `usage.cache_read_input_tokens` and `DeepSeek`
343/// `usage.prompt_cache_hit_tokens` spellings. Cache writes are read from
344/// `usage.prompt_tokens_details.cache_write_tokens` (`OpenRouter`, and
345/// `OpenAI`/`Azure` on GPT-5.6+), falling back to the `Anthropic`-compatible
346/// `usage.cache_creation_input_tokens`. Providers that do not report prompt
347/// caching yield `0`.
348///
349/// Note the provider semantics differ: for `OpenAI`-compatible responses
350/// `cached_tokens`/`cache_write_tokens` are subsets of `prompt_tokens`,
351/// whereas for `Anthropic`/Bedrock-style responses the cache counters are
352/// reported in addition to `inputTokens`.
353#[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/// Fallback extraction for a native Google Gemini `generateContent` response,
376/// which reports usage under `usageMetadata` rather than `usage`:
377/// `promptTokenCount` / `candidatesTokenCount` / `totalTokenCount` and the
378/// prompt-cache read counter `cachedContentTokenCount`. Returns an all-zero
379/// usage when neither shape is present.
380#[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        // 1M input tokens = $0.15, 500K output = $0.30
447        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        // OpenAI-style: cached tokens are a subset of prompt_tokens.
461        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        // All input is cached at 0.1x => $0.10.
464        assert!((cost - 0.10).abs() < 0.001, "got {cost}");
465    }
466
467    #[test]
468    fn test_calculate_llm_cost_cache_additive_bedrock() {
469        // Bedrock/Anthropic-style: cache tokens are additional to input.
470        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        // 1M ordinary input ($1.00) + 1M cached read ($0.10).
474        assert!((cost - 1.10).abs() < 0.001, "got {cost}");
475    }
476
477    #[test]
478    fn test_calculate_llm_cost_cache_pricing_disabled() {
479        // With cache pricing off, every prompt token is billed at input price
480        // even for additive (Bedrock) accounting.
481        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}