1use tracing::{info, warn};
24
25pub trait PricingSource: Send + Sync {
42 fn price_per_1m(&self, model: &str) -> Option<(f64, f64)>;
46
47 fn model_pricing(&self, model: &str) -> Option<ModelPricing> {
66 self.price_per_1m(model)
67 .map(|(i, o)| ModelPricing::without_cache(i, o))
68 }
69}
70
71#[derive(Debug, Clone, Copy, PartialEq)]
86pub struct ModelPricing {
87 pub input_per_1m: f64,
89 pub output_per_1m: f64,
91 pub cache_read_per_1m: f64,
93 pub cache_write_per_1m: f64,
95}
96
97impl ModelPricing {
98 #[must_use]
114 pub fn without_cache(input_per_1m: f64, output_per_1m: f64) -> Self {
115 Self {
116 input_per_1m,
117 output_per_1m,
118 cache_read_per_1m: input_per_1m,
119 cache_write_per_1m: input_per_1m,
120 }
121 }
122}
123
124#[derive(Debug, Clone, Copy, PartialEq)]
140pub struct CostBreakdown {
141 pub prompt_usd: f64,
143 pub cache_read_usd: f64,
145 pub cache_write_usd: f64,
147 pub completion_usd: f64,
149 pub total_usd: f64,
152}
153
154fn round6(v: f64) -> f64 {
156 (v * 1_000_000.0).round() / 1_000_000.0
157}
158
159const SONNET_FALLBACK: (f64, f64) = (3.0, 15.0);
162
163impl CostBreakdown {
164 pub fn compute(
179 source: &dyn PricingSource,
180 model: &str,
181 input_tokens: u64,
182 output_tokens: u64,
183 ) -> Self {
184 Self::compute_with_cache(source, model, input_tokens, 0, 0, output_tokens)
185 }
186
187 pub fn compute_with_cache(
211 source: &dyn PricingSource,
212 model: &str,
213 input_tokens: u64,
214 cache_read_tokens: u64,
215 cache_write_tokens: u64,
216 output_tokens: u64,
217 ) -> Self {
218 let rates = source.model_pricing(model).unwrap_or_else(|| {
219 warn!(
220 model,
221 fallback_input = SONNET_FALLBACK.0,
222 fallback_output = SONNET_FALLBACK.1,
223 "unknown model, falling back to Claude Sonnet pricing"
224 );
225 ModelPricing::without_cache(SONNET_FALLBACK.0, SONNET_FALLBACK.1)
226 });
227
228 let prompt_usd = round6(input_tokens as f64 / 1_000_000.0 * rates.input_per_1m);
229 let cache_read_usd =
230 round6(cache_read_tokens as f64 / 1_000_000.0 * rates.cache_read_per_1m);
231 let cache_write_usd =
232 round6(cache_write_tokens as f64 / 1_000_000.0 * rates.cache_write_per_1m);
233 let completion_usd = round6(output_tokens as f64 / 1_000_000.0 * rates.output_per_1m);
234 let total_usd = round6(prompt_usd + cache_read_usd + cache_write_usd + completion_usd);
235
236 Self {
237 prompt_usd,
238 cache_read_usd,
239 cache_write_usd,
240 completion_usd,
241 total_usd,
242 }
243 }
244}
245
246pub struct StaticPricing {
264 entries: Vec<(&'static str, f64, f64)>,
265}
266
267impl StaticPricing {
268 #[must_use]
270 pub fn new() -> Self {
271 let mut entries = vec![
272 ("claude-fable-5", 10.0, 50.0),
274 ("claude-fable-5-1", 10.0, 50.0),
275 ("claude-mythos-5", 10.0, 50.0),
276 ("claude-mythos-5-1", 10.0, 50.0),
277 ("claude-opus-5-5", 4.0, 20.0),
278 ("claude-opus-5", 5.0, 25.0),
279 ("claude-sonnet-5", 2.0, 10.0),
280 ("claude-opus-4-8", 5.0, 25.0),
281 ("claude-opus-4-7", 5.0, 25.0),
282 ("claude-sonnet-4-6", 3.0, 15.0),
283 ("claude-opus-4-6", 5.0, 25.0),
284 ("claude-sonnet-4-5", 3.0, 15.0),
285 ("claude-haiku-4-5", 1.0, 5.0),
286 ("sonnet", 3.0, 15.0),
288 ("opus", 5.0, 25.0),
289 ("haiku", 1.0, 5.0),
290 ("gpt-5.5", 5.0, 30.0),
292 ("gpt-5.4-mini", 0.75, 4.5),
293 ("gpt-5.4-nano", 0.20, 1.25),
294 ("gpt-5.4", 2.5, 15.0),
295 ("gpt-4.1-mini", 0.40, 1.60),
296 ("gpt-4.1-nano", 0.10, 0.40),
297 ("gpt-4.1", 2.0, 8.0),
298 ("gpt-4o-mini", 0.15, 0.60),
299 ("gpt-4o", 2.5, 10.0),
300 ("mistral-medium-3.5", 1.5, 7.5),
302 ("mistral-large", 0.50, 1.50),
303 ("mistral-small", 0.10, 0.30),
304 ("mistral-medium", 1.0, 3.0),
305 ("codestral", 0.30, 0.90),
306 ("gemini-3.5-flash", 0.15, 0.60),
308 ("gemini-3.1-flash-lite", 0.05, 0.20),
309 ("gemini-2.5-pro", 1.25, 10.0),
310 ("gemini-2.5-flash", 0.15, 0.60),
311 ("gemini-2.5-flash-lite", 0.05, 0.20),
312 ("nemotron-nano-9b", 0.04, 0.16),
314 ("nemotron-super-49b", 0.10, 0.40),
315 ("nemotron-ultra-253b", 0.90, 0.90),
316 ];
317 entries.sort_by_key(|e| std::cmp::Reverse(e.0.len()));
319 Self { entries }
320 }
321}
322
323impl Default for StaticPricing {
324 fn default() -> Self {
325 Self::new()
326 }
327}
328
329fn cache_multipliers(key: &str) -> Option<(f64, f64)> {
334 if key.starts_with("claude-") || matches!(key, "sonnet" | "opus" | "haiku") {
335 Some((0.1, 1.25))
336 } else if key.starts_with("gpt-5") {
337 Some((0.1, 1.0))
338 } else if key.starts_with("gpt-4.1") {
339 Some((0.25, 1.0))
340 } else if key.starts_with("gpt-4o") {
341 Some((0.5, 1.0))
342 } else if key.starts_with("gemini-") {
343 Some((0.25, 1.0))
344 } else {
345 None
346 }
347}
348
349impl StaticPricing {
350 fn find_entry(&self, model: &str) -> Option<&(&'static str, f64, f64)> {
352 self.entries.iter().find(|(key, _, _)| model.contains(key))
353 }
354}
355
356impl PricingSource for StaticPricing {
357 fn price_per_1m(&self, model: &str) -> Option<(f64, f64)> {
358 self.find_entry(model).map(|(_, inp, out)| (*inp, *out))
359 }
360
361 fn model_pricing(&self, model: &str) -> Option<ModelPricing> {
362 let (key, input, output) = *self.find_entry(model)?;
363 Some(match cache_multipliers(key) {
364 Some((read_mult, write_mult)) => ModelPricing {
365 input_per_1m: input,
366 output_per_1m: output,
367 cache_read_per_1m: input * read_mult,
368 cache_write_per_1m: input * write_mult,
369 },
370 None => ModelPricing::without_cache(input, output),
371 })
372 }
373}
374
375pub fn spawn_log(step_name: &str, model: &str, breakdown: CostBreakdown) {
390 let step = step_name.to_string();
391 let model = model.to_string();
392 tokio::spawn(async move {
393 info!(
394 step = %step,
395 model = %model,
396 prompt_usd = breakdown.prompt_usd,
397 cache_read_usd = breakdown.cache_read_usd,
398 cache_write_usd = breakdown.cache_write_usd,
399 completion_usd = breakdown.completion_usd,
400 total_usd = breakdown.total_usd,
401 "agent step cost"
402 );
403 });
404}
405
406#[cfg(test)]
407mod tests {
408 use super::*;
409
410 #[test]
411 fn known_model_returns_correct_price() {
412 let pricing = StaticPricing::new();
413 let price = pricing.price_per_1m("claude-opus-5");
414 assert_eq!(price, Some((5.0, 25.0)));
415 }
416
417 #[test]
418 fn sonnet_5_and_opus_5_5_have_current_pricing() {
419 let pricing = StaticPricing::new();
420 assert_eq!(pricing.price_per_1m("claude-sonnet-5"), Some((2.0, 10.0)));
421 assert_eq!(pricing.price_per_1m("claude-opus-5-5"), Some((4.0, 20.0)));
422 assert_eq!(
423 pricing.price_per_1m("claude-mythos-5-1"),
424 Some((10.0, 50.0))
425 );
426 }
427
428 #[test]
429 fn substring_resolution_most_specific_wins() {
430 let pricing = StaticPricing::new();
431 let price = pricing.price_per_1m("claude-sonnet-4-6[1m]");
434 assert_eq!(price, Some((3.0, 15.0)));
435
436 let price = pricing.price_per_1m("gpt-4.1-mini");
439 assert_eq!(price, Some((0.40, 1.60)));
440 }
441
442 #[test]
443 fn unknown_model_falls_back_to_sonnet() {
444 let pricing = StaticPricing::new();
445 assert!(pricing.price_per_1m("totally-unknown-model-xyz").is_none());
447
448 let bd =
450 CostBreakdown::compute(&pricing, "totally-unknown-model-xyz", 1_000_000, 1_000_000);
451 assert_eq!(bd.prompt_usd, 3.0);
453 assert_eq!(bd.completion_usd, 15.0);
454 assert_eq!(bd.total_usd, 18.0);
455 }
456
457 #[test]
458 fn cost_breakdown_rounds_to_six_decimals() {
459 let pricing = StaticPricing::new();
460 let bd = CostBreakdown::compute(&pricing, "claude-opus-5", 7, 3);
463 assert_eq!(bd.prompt_usd, 0.000035);
464 assert_eq!(bd.completion_usd, 0.000075);
465 assert_eq!(bd.total_usd, 0.00011);
466
467 let bd = CostBreakdown::compute(&pricing, "claude-sonnet-4-6", 1, 1);
470 assert_eq!(bd.prompt_usd, 0.000003);
471 assert_eq!(bd.completion_usd, 0.000015);
472 assert_eq!(bd.total_usd, 0.000018);
473
474 let bd = CostBreakdown::compute(&pricing, "gemini-2.5-pro", 3, 0);
476 assert_eq!(bd.prompt_usd, 0.000004);
477 }
478
479 #[test]
480 fn zero_tokens_returns_zero_cost() {
481 let pricing = StaticPricing::new();
482 let bd = CostBreakdown::compute(&pricing, "claude-opus-5", 0, 0);
483 assert_eq!(bd.prompt_usd, 0.0);
484 assert_eq!(bd.completion_usd, 0.0);
485 assert_eq!(bd.total_usd, 0.0);
486 }
487
488 #[tokio::test]
489 async fn spawn_log_does_not_block() {
490 use std::time::Duration;
491
492 tokio::time::timeout(Duration::from_secs(5), async {
493 let pricing = StaticPricing::new();
494 let bd = CostBreakdown::compute(&pricing, "claude-opus-5", 1000, 500);
495 spawn_log("test-step", "claude-opus-5", bd);
496 tokio::task::yield_now().await;
497 })
498 .await
499 .expect("spawn_log timed out");
500 }
501
502 #[test]
503 fn all_known_models_have_prices() {
504 let pricing = StaticPricing::new();
505 let models = [
506 "claude-fable-5",
507 "claude-fable-5-1",
508 "claude-mythos-5",
509 "claude-mythos-5-1",
510 "claude-opus-5-5",
511 "claude-opus-5",
512 "claude-sonnet-5",
513 "claude-opus-4-8",
514 "claude-opus-4-7",
515 "claude-sonnet-4-6",
516 "claude-opus-4-6",
517 "claude-sonnet-4-5",
518 "claude-haiku-4-5",
519 "sonnet",
520 "opus",
521 "haiku",
522 "gpt-5.5",
523 "gpt-5.4",
524 "gpt-4.1",
525 "gpt-4o",
526 "gpt-4o-mini",
527 "mistral-large",
528 "mistral-small",
529 "codestral",
530 "gemini-2.5-pro",
531 "gemini-2.5-flash",
532 "nemotron-nano-9b",
533 "nemotron-super-49b",
534 "nemotron-ultra-253b",
535 ];
536 for model in models {
537 assert!(
538 pricing.price_per_1m(model).is_some(),
539 "missing price for {model}"
540 );
541 }
542 }
543
544 #[test]
545 fn entries_sorted_by_length_descending() {
546 let pricing = StaticPricing::new();
547 for window in pricing.entries.windows(2) {
548 assert!(
549 window[0].0.len() >= window[1].0.len(),
550 "entries not sorted: {:?} before {:?}",
551 window[0].0,
552 window[1].0
553 );
554 }
555 }
556
557 #[test]
558 fn default_matches_new() {
559 let a = StaticPricing::new();
560 let b = StaticPricing::default();
561 assert_eq!(a.entries.len(), b.entries.len());
562 for (ea, eb) in a.entries.iter().zip(b.entries.iter()) {
563 assert_eq!(ea.0, eb.0);
564 assert_eq!(ea.1, eb.1);
565 assert_eq!(ea.2, eb.2);
566 }
567 }
568
569 #[test]
570 fn cost_breakdown_with_large_token_counts() {
571 let pricing = StaticPricing::new();
572 let bd = CostBreakdown::compute(&pricing, "claude-opus-5", 1_000_000, 500_000);
575 assert_eq!(bd.prompt_usd, 5.0);
576 assert_eq!(bd.completion_usd, 12.5);
577 assert_eq!(bd.total_usd, 17.5);
578 }
579
580 fn assert_close(actual: f64, expected: f64) {
581 assert!(
582 (actual - expected).abs() < 1e-9,
583 "expected {expected}, got {actual}"
584 );
585 }
586
587 #[test]
588 fn pricing_cache_read_rate_is_ten_percent_for_claude() {
589 let pricing = StaticPricing::new();
590 let p = pricing.model_pricing("claude-sonnet-4-6").unwrap();
591 assert_close(p.input_per_1m, 3.0);
592 assert_close(p.output_per_1m, 15.0);
593 assert_close(p.cache_read_per_1m, 0.3);
594 assert_close(p.cache_write_per_1m, 3.75);
595
596 let alias = pricing.model_pricing("sonnet").unwrap();
597 assert_close(alias.cache_read_per_1m, 0.3);
598 }
599
600 #[test]
601 fn pricing_unknown_cache_rate_defaults_to_input() {
602 let pricing = StaticPricing::new();
603 let p = pricing.model_pricing("mistral-large").unwrap();
604 assert_close(p.cache_read_per_1m, 0.50);
605 assert_close(p.cache_write_per_1m, 0.50);
606 }
607
608 #[test]
609 fn pricing_cache_read_breakdown_total_includes_four_parts() {
610 let pricing = StaticPricing::new();
611 let bd = CostBreakdown::compute_with_cache(
612 &pricing,
613 "claude-opus-5",
614 1_000_000,
615 1_000_000,
616 1_000_000,
617 1_000_000,
618 );
619 assert_eq!(bd.prompt_usd, 5.0);
620 assert_eq!(bd.cache_read_usd, 0.5);
621 assert_eq!(bd.cache_write_usd, 6.25);
622 assert_eq!(bd.completion_usd, 25.0);
623 assert_eq!(bd.total_usd, 36.75);
624 }
625
626 #[test]
627 fn pricing_compute_without_cache_matches_legacy() {
628 let pricing = StaticPricing::new();
629 let bd = CostBreakdown::compute(&pricing, "claude-opus-5", 1_000_000, 500_000);
630 assert_eq!(bd.cache_read_usd, 0.0);
631 assert_eq!(bd.cache_write_usd, 0.0);
632 assert_eq!(
633 bd,
634 CostBreakdown::compute_with_cache(&pricing, "claude-opus-5", 1_000_000, 0, 0, 500_000)
635 );
636 }
637
638 #[test]
639 fn pricing_unknown_model_cache_fallback() {
640 let pricing = StaticPricing::new();
641 assert!(pricing.model_pricing("totally-unknown-model-xyz").is_none());
642 let bd = CostBreakdown::compute_with_cache(
643 &pricing,
644 "totally-unknown-model-xyz",
645 0,
646 1_000_000,
647 1_000_000,
648 0,
649 );
650 assert_eq!(bd.cache_read_usd, 3.0);
652 assert_eq!(bd.cache_write_usd, 3.0);
653 assert_eq!(bd.total_usd, 6.0);
654 }
655
656 #[test]
657 fn pricing_openai_cache_read_rate() {
658 let pricing = StaticPricing::new();
659 let p = pricing.model_pricing("gpt-4.1").unwrap();
660 assert_close(p.cache_read_per_1m, 0.5);
661 assert_close(p.cache_write_per_1m, 2.0);
662
663 let gpt5 = pricing.model_pricing("gpt-5.4").unwrap();
664 assert_close(gpt5.cache_read_per_1m, 0.25);
665
666 let gpt4o = pricing.model_pricing("gpt-4o").unwrap();
667 assert_close(gpt4o.cache_read_per_1m, 1.25);
668
669 let gemini = pricing.model_pricing("gemini-2.5-pro").unwrap();
670 assert_close(gemini.cache_read_per_1m, 0.3125);
671 }
672
673 #[test]
674 fn pricing_default_model_pricing_impl_without_cache() {
675 struct FlatPricing;
676
677 impl PricingSource for FlatPricing {
678 fn price_per_1m(&self, model: &str) -> Option<(f64, f64)> {
679 (model == "flat").then_some((2.0, 8.0))
680 }
681 }
682
683 let p = FlatPricing.model_pricing("flat").unwrap();
684 assert_eq!(p, ModelPricing::without_cache(2.0, 8.0));
685 assert_eq!(p.cache_read_per_1m, 2.0);
686 assert_eq!(p.cache_write_per_1m, 2.0);
687 assert!(FlatPricing.model_pricing("other").is_none());
688 }
689
690 #[test]
691 fn sonnet_fallback_never_zero() {
692 let pricing = StaticPricing::new();
693 let bd = CostBreakdown::compute(&pricing, "unknown-model", 1, 0);
695 assert!(bd.total_usd > 0.0);
696 }
697}