nanocodex_oai_api/pricing/
estimate.rs1use serde::{Deserialize, Serialize};
2
3use super::UsdAmount;
4use crate::Usage;
5
6const STANDARD: TokenRates = TokenRates {
9 input: 5_000,
10 cached_input: 500,
11 cache_write_input: 6_250,
12 output: 30_000,
13};
14const PRIORITY: TokenRates = TokenRates {
15 input: 10_000,
16 cached_input: 1_000,
17 cache_write_input: 12_500,
18 output: 60_000,
19};
20
21#[derive(Clone, Copy)]
22struct TokenRates {
23 input: u64,
24 cached_input: u64,
25 cache_write_input: u64,
26 output: u64,
27}
28
29impl TokenRates {
30 const fn for_service_tier(service_tier: ServiceTier) -> Self {
31 match service_tier {
32 ServiceTier::Standard => STANDARD,
33 ServiceTier::Priority => PRIORITY,
34 }
35 }
36}
37
38#[derive(Clone, Copy, Debug, Default, Deserialize, Eq, PartialEq, Serialize)]
40#[serde(rename_all = "snake_case")]
41pub enum ServiceTier {
42 #[default]
44 Standard,
45 Priority,
47}
48
49impl ServiceTier {
50 #[must_use]
52 pub const fn as_str(self) -> &'static str {
53 match self {
54 Self::Standard => "standard",
55 Self::Priority => "priority",
56 }
57 }
58}
59
60#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
66pub struct EstimatedUsdCost {
67 #[serde(rename = "usd")]
68 amount: UsdAmount,
69 #[serde(rename = "input_usd")]
70 input: UsdAmount,
71 #[serde(rename = "cached_input_usd")]
72 cached_input: UsdAmount,
73 #[serde(rename = "cache_write_input_usd")]
74 cache_write_input: UsdAmount,
75 #[serde(rename = "output_usd")]
76 output: UsdAmount,
77 #[serde(default)]
78 service_tier: ServiceTier,
79}
80
81impl EstimatedUsdCost {
82 #[must_use]
84 pub const fn amount(&self) -> UsdAmount {
85 self.amount
86 }
87
88 #[must_use]
90 pub const fn input(&self) -> UsdAmount {
91 self.input
92 }
93
94 #[must_use]
96 pub const fn cached_input(&self) -> UsdAmount {
97 self.cached_input
98 }
99
100 #[must_use]
102 pub const fn cache_write_input(&self) -> UsdAmount {
103 self.cache_write_input
104 }
105
106 #[must_use]
108 pub const fn output(&self) -> UsdAmount {
109 self.output
110 }
111
112 #[must_use]
114 pub const fn service_tier(&self) -> ServiceTier {
115 self.service_tier
116 }
117}
118
119#[must_use]
146pub fn estimate(usage: &Usage, service_tier: ServiceTier) -> EstimatedUsdCost {
147 let cached_input_tokens = usage
148 .input_tokens_details
149 .as_ref()
150 .map_or(0, |details| details.cached_tokens);
151 let cache_write_input_tokens = usage
152 .input_tokens_details
153 .as_ref()
154 .map_or(0, |details| details.cache_write_tokens);
155 estimate_tokens(
156 usage.input_tokens,
157 cached_input_tokens,
158 cache_write_input_tokens,
159 usage.output_tokens,
160 service_tier,
161 )
162}
163
164pub(crate) fn estimate_tokens(
165 input_tokens: u64,
166 cached_input_tokens: u64,
167 cache_write_input_tokens: u64,
168 output_tokens: u64,
169 service_tier: ServiceTier,
170) -> EstimatedUsdCost {
171 let rates = TokenRates::for_service_tier(service_tier);
172 let cached_input_tokens = cached_input_tokens.min(input_tokens);
173 let remaining_input = input_tokens.saturating_sub(cached_input_tokens);
174 let cache_write_input_tokens = cache_write_input_tokens.min(remaining_input);
175 let ordinary_input_tokens = remaining_input.saturating_sub(cache_write_input_tokens);
176
177 let input = UsdAmount::saturating_mul(ordinary_input_tokens, rates.input);
178 let cached_input = UsdAmount::saturating_mul(cached_input_tokens, rates.cached_input);
179 let cache_write_input =
180 UsdAmount::saturating_mul(cache_write_input_tokens, rates.cache_write_input);
181 let output = UsdAmount::saturating_mul(output_tokens, rates.output);
182 let amount = input
183 .saturating_add(cached_input)
184 .saturating_add(cache_write_input)
185 .saturating_add(output);
186
187 EstimatedUsdCost {
188 amount,
189 input,
190 cached_input,
191 cache_write_input,
192 output,
193 service_tier,
194 }
195}
196
197#[cfg(test)]
198mod tests {
199 use super::{ServiceTier, estimate, estimate_tokens};
200 use crate::responses::{InputTokenDetails, OutputTokenDetails, Usage};
201
202 #[test]
203 fn standard_rates_price_each_input_class_once() {
204 let estimate = estimate(
205 &Usage {
206 input_tokens: 1_000_000,
207 input_tokens_details: Some(InputTokenDetails {
208 cached_tokens: 250_000,
209 cache_write_tokens: 100_000,
210 }),
211 output_tokens: 200_000,
212 output_tokens_details: Some(OutputTokenDetails {
213 reasoning_tokens: 150_000,
214 }),
215 total_tokens: 1_200_000,
216 },
217 ServiceTier::Standard,
218 );
219
220 assert_eq!(estimate.input().decimal(), "3.25");
221 assert_eq!(estimate.cached_input().decimal(), "0.125");
222 assert_eq!(estimate.cache_write_input().decimal(), "0.625");
223 assert_eq!(estimate.output().decimal(), "6");
224 assert_eq!(estimate.amount().decimal(), "10");
225 }
226
227 #[test]
228 fn priority_rates_are_selected_by_fast_mode() {
229 let standard = estimate_tokens(1_000_000, 0, 0, 1_000_000, ServiceTier::Standard);
230 let priority = estimate_tokens(1_000_000, 0, 0, 1_000_000, ServiceTier::Priority);
231
232 assert_eq!(standard.amount().decimal(), "35");
233 assert_eq!(priority.amount().decimal(), "70");
234 assert_eq!(priority.service_tier(), ServiceTier::Priority);
235 assert_eq!(priority.service_tier().as_str(), "priority");
236 }
237
238 #[test]
239 fn malformed_detail_counts_do_not_double_charge_input() {
240 let estimate = estimate_tokens(10, 8, 8, 0, ServiceTier::Standard);
241
242 assert_eq!(estimate.input().nano_usd(), 0);
243 assert_eq!(estimate.cached_input().nano_usd(), 4_000);
244 assert_eq!(estimate.cache_write_input().nano_usd(), 12_500);
245 }
246}