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    /// Accumulated cost in USD.
20    pub cost: f64,
21    /// Model names observed on this endpoint, sorted and deduplicated.
22    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/// Aggregated usage across all endpoints for a test or scenario run.
32#[derive(Debug, Default, Clone)]
33pub struct UsageSnapshot {
34    /// Per-endpoint usage.
35    pub endpoints: HashMap<String, EndpointUsage>,
36    /// Total cost across all endpoints.
37    pub total_cost: f64,
38    /// Total calls across all endpoints.
39    pub total_calls: u64,
40    /// Total tokens across all endpoints.
41    pub total_tokens: u64,
42    /// Total input (prompt) tokens across all endpoints.
43    pub total_input_tokens: u64,
44    /// Total output (completion) tokens across all endpoints.
45    pub total_output_tokens: u64,
46    /// Total input tokens served from provider prompt caches.
47    pub total_cached_input_tokens: u64,
48    /// Model names observed across all endpoints, sorted and deduplicated.
49    pub models: Vec<String>,
50}
51
52impl UsageSnapshot {
53    /// Creates a snapshot from per-endpoint usage data.
54    #[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
81/// Thread-safe usage tracker for the test runner.
82pub struct UsageTracker {
83    inner: Mutex<UsageInner>,
84}
85
86struct UsageInner {
87    /// Per-endpoint usage for the current test.
88    per_endpoint: HashMap<String, EndpointUsage>,
89    /// Aggregated usage across all completed tests.
90    global: UsageSnapshot,
91    /// Per-test snapshots keyed by test name.
92    per_test: Vec<(String, UsageSnapshot)>,
93}
94
95impl UsageTracker {
96    /// Creates a new empty usage tracker.
97    #[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    /// Records a completed LLM call, adding usage, cost, and the answering
109    /// model to the endpoint's accumulator.
110    ///
111    /// # Panics
112    ///
113    /// Panics if the mutex is poisoned.
114    #[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    /// Records a flat-cost call (MCP tool, agent task).
139    ///
140    /// # Panics
141    ///
142    /// Panics if the mutex is poisoned.
143    #[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    /// Reads current usage without locking for the full snapshot.
155    ///
156    /// # Panics
157    ///
158    /// Panics if the mutex is poisoned.
159    #[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    /// Reads the global aggregated snapshot.
166    ///
167    /// # Panics
168    ///
169    /// Panics if the mutex is poisoned.
170    #[must_use]
171    pub fn global_snapshot(&self) -> UsageSnapshot {
172        let inner = self.inner.lock().unwrap();
173        inner.global.clone()
174    }
175
176    /// Reads per-test snapshots.
177    ///
178    /// # Panics
179    ///
180    /// Panics if the mutex is poisoned.
181    #[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    /// Resets the per-test accumulator. Call at the start of each test.
188    ///
189    /// # Panics
190    ///
191    /// Panics if the mutex is poisoned.
192    pub fn reset_per_test(&self) {
193        let mut inner = self.inner.lock().unwrap();
194        inner.per_endpoint.clear();
195    }
196
197    /// Commits the current test's usage to the global accumulator and stores
198    /// it as a per-test snapshot.
199    ///
200    /// # Panics
201    ///
202    /// Panics if the mutex is poisoned.
203    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        // Merge into global
207        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/// Calculates the cost of an LLM call based on token pricing.
242#[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/// Usage info extracted from an LLM API response.
255#[derive(Debug, Default, Clone, Copy)]
256pub struct LlmUsage {
257    /// Number of prompt / input tokens.
258    pub prompt_tokens: u64,
259    /// Number of completion / output tokens.
260    pub completion_tokens: u64,
261    /// Total tokens used.
262    pub total_tokens: u64,
263    /// Input tokens served from the provider's prompt cache (cache hit).
264    pub cached_input_tokens: u64,
265}
266
267/// Result of an LLM chat call including usage data.
268#[derive(Debug, Clone)]
269pub struct LlmResponse {
270    /// The message content from the LLM.
271    pub content: String,
272    /// Token usage from the API response.
273    pub usage: LlmUsage,
274}
275
276/// Extracts token usage from an OpenAI-compatible API response JSON.
277///
278/// Cached prompt tokens are read from the `OpenAI`/`OpenRouter`
279/// `usage.prompt_tokens_details.cached_tokens` field, falling back to the
280/// `Anthropic`-compatible `usage.cache_read_input_tokens` and `DeepSeek`
281/// `usage.prompt_cache_hit_tokens` spellings. Providers that do not report
282/// prompt caching yield `0`.
283#[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        // 1M input tokens = $0.15, 500K output = $0.30
347        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}