1use crate::model_resolver::ModelResolver;
4use crate::provider::Usage as ProviderUsage;
5use vtcode_config::models::ModelPricing;
6
7pub fn provider_reports_exclusive_input(provider: &str) -> bool {
16 matches!(provider.trim().to_ascii_lowercase().as_str(), "anthropic" | "minimax")
17}
18
19pub fn normalized_turn_usage(provider: &str, usage: &ProviderUsage) -> vtcode_exec_events::Usage {
24 let totals = usage.billable_totals();
25 let cached = u64::from(totals.cache_read_tokens);
26 let creation = u64::from(totals.cache_creation_tokens);
27 let mut input = u64::from(totals.prompt_tokens);
28 if provider_reports_exclusive_input(provider) {
29 input = input.saturating_add(cached).saturating_add(creation);
30 }
31 let output = u64::from(totals.completion_tokens);
32
33 vtcode_exec_events::Usage {
34 input_tokens: input,
35 cached_input_tokens: cached,
36 cache_creation_tokens: creation,
37 output_tokens: output,
38 }
39}
40
41#[derive(Debug, Clone, Copy, PartialEq)]
43pub struct SessionCostEstimate {
44 pub raw_usd: f64,
48 pub effective_usd: f64,
51}
52
53#[derive(Debug, Clone)]
56pub struct SessionCostAccumulator {
57 total: Option<SessionCostEstimate>,
58}
59
60impl Default for SessionCostAccumulator {
61 fn default() -> Self {
62 Self {
63 total: Some(SessionCostEstimate { raw_usd: 0.0, effective_usd: 0.0 }),
64 }
65 }
66}
67
68impl SessionCostAccumulator {
69 pub fn record(&mut self, estimate: Option<SessionCostEstimate>) -> Option<SessionCostEstimate> {
70 self.total = self.total.zip(estimate).and_then(|(total, turn)| {
71 let raw_usd = total.raw_usd + turn.raw_usd;
72 let effective_usd = total.effective_usd + turn.effective_usd;
73 (raw_usd.is_finite() && effective_usd.is_finite()).then_some(SessionCostEstimate { raw_usd, effective_usd })
74 });
75 self.total
76 }
77
78 pub fn total(&self) -> Option<SessionCostEstimate> {
79 self.total
80 }
81}
82
83pub fn estimate_session_costs(
87 provider: &str,
88 model: &str,
89 usage: &vtcode_exec_events::Usage,
90) -> Option<SessionCostEstimate> {
91 let resolved = ModelResolver::resolve(Some(provider), model, &[], None)?;
92 let pricing = resolved.pricing()?;
93 estimate_session_costs_with_pricing(pricing, usage)
94}
95
96pub fn require_budget_pricing(provider: &str, model: &str, max_budget_usd: Option<f64>) -> anyhow::Result<()> {
99 if let Some(maximum) = max_budget_usd {
100 anyhow::ensure!(maximum.is_finite() && maximum >= 0.0, "Session USD budget must be finite and non-negative");
101 anyhow::ensure!(
102 estimate_session_costs(provider, model, &vtcode_exec_events::Usage::default()).is_some(),
103 "Cannot enforce session USD budget for `{provider}/{model}`: complete valid pricing metadata is unavailable"
104 );
105 }
106 Ok(())
107}
108
109pub fn estimate_session_costs_with_pricing(
116 pricing: ModelPricing,
117 usage: &vtcode_exec_events::Usage,
118) -> Option<SessionCostEstimate> {
119 let input_rate = pricing.input?;
120 let output_rate = pricing.output?;
121 if [pricing.input, pricing.output, pricing.cache_read, pricing.cache_write]
122 .into_iter()
123 .flatten()
124 .any(|rate| !rate.is_finite() || rate < 0.0)
125 {
126 return None;
127 }
128
129 let input_tokens = usage.input_tokens as f64;
130 let output_tokens = usage.output_tokens as f64;
131 let cached_tokens = usage.cached_input_tokens as f64;
132 let creation_tokens = usage.cache_creation_tokens as f64;
133
134 let raw_usd = input_tokens * input_rate + output_tokens * output_rate;
135
136 let read_rate = pricing.cache_read.unwrap_or(input_rate * 0.10);
141 let write_rate = pricing.cache_write.unwrap_or(input_rate * DEFAULT_CACHE_WRITE_MULTIPLIER);
142
143 let uncached_tokens = usage
144 .input_tokens
145 .saturating_sub(usage.cached_input_tokens)
146 .saturating_sub(usage.cache_creation_tokens) as f64;
147
148 let effective_usd = uncached_tokens * input_rate
149 + cached_tokens * read_rate
150 + creation_tokens * write_rate
151 + output_tokens * output_rate;
152
153 (raw_usd.is_finite() && effective_usd.is_finite()).then_some(SessionCostEstimate { raw_usd, effective_usd })
154}
155
156pub const DEFAULT_CACHE_WRITE_MULTIPLIER: f64 = 1.25;
158
159pub const EXTENDED_TTL_CACHE_WRITE_MULTIPLIER: f64 = 2.0;
161
162#[must_use]
165pub fn cache_write_rate(input_rate: f64, configured: Option<f64>, extended_ttl: bool) -> f64 {
166 if let Some(rate) = configured {
167 return rate;
168 }
169 let multiplier = if extended_ttl {
170 EXTENDED_TTL_CACHE_WRITE_MULTIPLIER
171 } else {
172 DEFAULT_CACHE_WRITE_MULTIPLIER
173 };
174 input_rate * multiplier
175}
176
177#[must_use]
180pub fn prompt_tokens_for_rate_limit(usage: &vtcode_exec_events::Usage) -> u64 {
181 let uncached = usage
182 .input_tokens
183 .saturating_sub(usage.cached_input_tokens)
184 .saturating_sub(usage.cache_creation_tokens);
185 uncached
186 .saturating_add(usage.cached_input_tokens)
187 .saturating_add(usage.cache_creation_tokens)
188}
189
190#[cfg(test)]
191mod tests {
192 use super::*;
193
194 #[test]
195 fn cost_normalization_matches_across_three_provider_families() {
196 let pricing = ModelPricing {
197 input: Some(0.01),
198 output: Some(0.02),
199 cache_read: Some(0.001),
200 cache_write: Some(0.0125),
201 };
202 for provider in ["openai", "anthropic", "gemini"] {
203 let usage = ProviderUsage {
204 prompt_tokens: if provider == "anthropic" { 150 } else { 1000 },
205 completion_tokens: 100,
206 total_tokens: 1100,
207 cached_prompt_tokens: None,
208 cache_read_tokens: Some(800),
209 cache_creation_tokens: Some(50),
210 iterations: None,
211 };
212 let normalized = normalized_turn_usage(provider, &usage);
213 let cost = estimate_session_costs_with_pricing(pricing, &normalized).expect("priced");
214 assert!((cost.raw_usd - 12.0).abs() < 1e-12, "{provider}");
215 assert!((cost.effective_usd - 4.925).abs() < 1e-12, "{provider}");
216 }
217 }
218
219 #[test]
220 fn anthropic_compaction_iterations_are_included_in_normalized_usage() {
221 let usage = ProviderUsage {
222 prompt_tokens: 10,
224 completion_tokens: 1,
225 total_tokens: 11,
226 cached_prompt_tokens: None,
227 cache_read_tokens: None,
228 cache_creation_tokens: None,
229 iterations: Some(vec![
230 serde_json::json!({
231 "type": "compaction",
232 "input_tokens": 50,
233 "output_tokens": 5,
234 }),
235 serde_json::json!({
236 "type": "message",
237 "input_tokens": 10,
238 "output_tokens": 2,
239 }),
240 ]),
241 };
242
243 let normalized = normalized_turn_usage("anthropic", &usage);
244 assert_eq!(normalized.input_tokens, 60);
245 assert_eq!(normalized.output_tokens, 7);
246 }
247
248 #[test]
249 fn astra_pricing_is_resolved_per_route_without_inference() {
250 let usage = vtcode_exec_events::Usage {
251 input_tokens: 1000,
252 output_tokens: 100,
253 cached_input_tokens: 800,
254 cache_creation_tokens: 0,
255 };
256 for (provider, model, priced) in [
257 ("openai", "gpt-6-astra", true),
258 ("openrouter", "openai/gpt-6-astra", true),
259 ("merge-gateway", "openai/gpt-6-astra", false),
260 ] {
261 assert_eq!(estimate_session_costs(provider, model, &usage).is_some(), priced, "{provider}");
262 assert_eq!(require_budget_pricing(provider, model, Some(1.0)).is_ok(), priced, "{provider}");
263 }
264 }
265
266 #[test]
267 fn switching_to_a_cheaper_model_never_reprices_previous_spend() {
268 let usage = vtcode_exec_events::Usage { input_tokens: 100, ..Default::default() };
269 let mut session = SessionCostAccumulator::default();
270 for (rate, expected) in [(0.10, 10.0), (0.001, 10.1), (0.10, 20.1)] {
271 let pricing = ModelPricing {
272 input: Some(rate),
273 output: Some(rate),
274 cache_read: None,
275 cache_write: None,
276 };
277 let total = session
278 .record(estimate_session_costs_with_pricing(pricing, &usage))
279 .expect("priced");
280 assert!((total.raw_usd - expected).abs() < 1e-12);
281 assert!((total.effective_usd - expected).abs() < 1e-12);
282 }
283 }
284
285 #[test]
286 fn an_unpriced_turn_keeps_session_cost_unknown() {
287 let mut session = SessionCostAccumulator::default();
288 assert!(session.record(None).is_none());
289 assert!(
290 session
291 .record(Some(SessionCostEstimate { raw_usd: 1.0, effective_usd: 1.0 }))
292 .is_none()
293 );
294 }
295
296 #[test]
297 fn missing_pricing_requires_removing_the_budget_explicitly() {
298 assert!(require_budget_pricing("openai", "unknown-dynamic-model", Some(1.0)).is_err());
299 assert!(require_budget_pricing("openai", "unknown-dynamic-model", None).is_ok());
300 }
301
302 #[test]
303 fn invalid_pricing_cannot_bypass_budget_enforcement() {
304 for invalid in [f64::NAN, f64::INFINITY, -1.0] {
305 let pricing = ModelPricing {
306 input: Some(invalid),
307 output: Some(0.01),
308 cache_read: None,
309 cache_write: None,
310 };
311 assert!(estimate_session_costs_with_pricing(pricing, &vtcode_exec_events::Usage::default()).is_none());
312 }
313 }
314
315 #[test]
316 fn overflowing_estimates_are_treated_as_unpriced() {
317 let pricing = ModelPricing {
318 input: Some(f64::MAX),
319 output: Some(f64::MAX),
320 cache_read: None,
321 cache_write: None,
322 };
323 let usage = vtcode_exec_events::Usage {
324 input_tokens: u64::MAX,
325 output_tokens: u64::MAX,
326 ..Default::default()
327 };
328 assert!(estimate_session_costs_with_pricing(pricing, &usage).is_none());
329 }
330
331 #[test]
332 fn cache_write_rate_uses_extended_ttl_multiplier() {
333 #[allow(clippy::float_cmp, reason = "exact constants under test")]
334 {
335 assert_eq!(cache_write_rate(1.0, None, false), DEFAULT_CACHE_WRITE_MULTIPLIER);
336 assert_eq!(cache_write_rate(1.0, None, true), EXTENDED_TTL_CACHE_WRITE_MULTIPLIER);
337 assert_eq!(cache_write_rate(1.0, Some(0.5), true), 0.5);
339 }
340 }
341
342 #[test]
343 fn prompt_tokens_for_rate_limit_count_cache_traffic() {
344 let usage = vtcode_exec_events::Usage {
345 input_tokens: 1000,
346 cached_input_tokens: 400,
347 cache_creation_tokens: 200,
348 output_tokens: 50,
349 };
350 assert_eq!(prompt_tokens_for_rate_limit(&usage), 1000);
352 }
353}