Skip to main content

synapse/
pricing.rs

1//! Static pricing table + cost calculation.
2
3use serde::Deserialize;
4use std::collections::HashMap;
5
6#[derive(Debug, Clone, Copy, Deserialize)]
7pub struct ModelPrice {
8    /// USD per 1M input tokens.
9    pub input: f64,
10    /// USD per 1M output tokens.
11    pub output: f64,
12}
13
14#[derive(Debug, Clone, Default)]
15pub struct PricingTable {
16    prices: HashMap<String, ModelPrice>,
17}
18
19impl PricingTable {
20    pub fn from_toml_str(s: &str) -> anyhow::Result<Self> {
21        let prices: HashMap<String, ModelPrice> = toml::from_str(s)?;
22        Ok(Self { prices })
23    }
24
25    /// Cost in USD for a completed call. Unknown `provider:model` → 0.0
26    /// (self-hosted / unpriced models do not error the request).
27    pub fn cost_usd(
28        &self,
29        provider: &str,
30        model: &str,
31        input_tokens: u64,
32        output_tokens: u64,
33    ) -> f64 {
34        self.prices
35            .get(&format!("{provider}:{model}"))
36            .map(|p| {
37                (input_tokens as f64 * p.input + output_tokens as f64 * p.output) / 1_000_000.0
38            })
39            .unwrap_or(0.0)
40    }
41
42    /// Cost in USD for an embedding call (input tokens only). Unlike `cost_usd`,
43    /// an unknown `provider:model` falls back to `default_input_per_mtok` (USD per
44    /// 1M tokens) so embedding usage is never silently free.
45    pub fn embedding_cost_usd(
46        &self,
47        provider: &str,
48        model: &str,
49        input_tokens: u64,
50        default_input_per_mtok: f64,
51    ) -> f64 {
52        let per_mtok = self
53            .prices
54            .get(&format!("{provider}:{model}"))
55            .map(|p| p.input)
56            .unwrap_or(default_input_per_mtok);
57        input_tokens as f64 * per_mtok / 1_000_000.0
58    }
59}
60
61#[cfg(test)]
62mod tests {
63    use super::*;
64
65    const SAMPLE: &str = r#"
66        ["vertex:gemini-3-pro"]
67        input = 1.25
68        output = 5.0
69    "#;
70
71    #[test]
72    fn computes_cost_from_tokens() {
73        let t = PricingTable::from_toml_str(SAMPLE).unwrap();
74        // 1M input @1.25 + 1M output @5.0 = 6.25
75        let c = t.cost_usd("vertex", "gemini-3-pro", 1_000_000, 1_000_000);
76        assert!((c - 6.25).abs() < 1e-9, "got {c}");
77    }
78
79    #[test]
80    fn unknown_model_is_free_not_error() {
81        let t = PricingTable::from_toml_str(SAMPLE).unwrap();
82        assert_eq!(t.cost_usd("oai_compat", "qwen-local", 1000, 1000), 0.0);
83    }
84
85    #[test]
86    fn embedding_uses_key_when_present_else_default() {
87        let t = PricingTable::from_toml_str(
88            "[\"vertex:text-embedding-004\"]\ninput = 0.025\noutput = 0.0\n",
89        )
90        .unwrap();
91        // priced key: 1M tokens * 0.025 = 0.025
92        assert!(
93            (t.embedding_cost_usd("vertex", "text-embedding-004", 1_000_000, 0.10) - 0.025).abs()
94                < 1e-9
95        );
96        // missing key: falls back to the $0.10/1M default
97        assert!(
98            (t.embedding_cost_usd("openai", "text-embedding-3-large", 1_000_000, 0.10) - 0.10)
99                .abs()
100                < 1e-9
101        );
102    }
103}