1#[cfg(test)]
13use vtcode_config::models::ModelPricing;
14
15use crate::llm::provider::ToolDefinition;
16#[cfg(test)]
17use crate::llm::provider::Usage as ProviderUsage;
18
19pub fn estimate_tool_definition_tokens(tools: &[ToolDefinition]) -> u64 {
28 let mut buf = Vec::new();
32 let total_bytes: u64 = tools
33 .iter()
34 .map(|tool| {
35 buf.clear();
36 serde_json::to_writer(&mut buf, tool).map(|_| buf.len() as u64).unwrap_or(0)
37 })
38 .sum();
39 total_bytes.div_ceil(4)
40}
41
42pub use vtcode_llm::usage_cost::{
43 SessionCostAccumulator, SessionCostEstimate, estimate_decisions_cost, estimate_session_costs,
44 estimate_session_costs_with_pricing, normalized_turn_usage, provider_reports_exclusive_input,
45 require_budget_pricing,
46};
47
48#[cfg(test)]
49mod tests {
50 use super::*;
51 use serde_json::json;
52
53 fn approx_eq(a: f64, b: f64) {
54 assert!((a - b).abs() < 1e-12, "expected {a} to approx-equal {b}");
55 }
56
57 #[test]
58 fn estimate_tool_definition_tokens_is_zero_for_empty_slice() {
59 assert_eq!(estimate_tool_definition_tokens(&[]), 0);
60 }
61
62 #[test]
63 fn estimate_tool_definition_tokens_matches_serialized_byte_length() {
64 let tool = ToolDefinition::function(
65 "read_file".to_string(),
66 "Read the contents of a file from the workspace.".to_string(),
67 json!({
68 "type": "object",
69 "properties": {
70 "path": { "type": "string" }
71 },
72 "required": ["path"],
73 }),
74 );
75
76 let expected_bytes = serde_json::to_string(&tool).expect("tool serializes").len() as u64;
77 let expected_tokens = expected_bytes.div_ceil(4);
78
79 assert_eq!(estimate_tool_definition_tokens(&[tool]), expected_tokens);
80 }
81
82 #[test]
83 fn normalized_turn_usage_adds_cache_tokens_for_anthropic() {
84 let usage = ProviderUsage {
85 prompt_tokens: 100,
86 completion_tokens: 20,
87 total_tokens: 120,
88 cached_prompt_tokens: None,
89 cache_creation_tokens: Some(50),
90 cache_read_tokens: Some(400),
91 iterations: None,
92 };
93
94 let normalized = normalized_turn_usage("anthropic", &usage);
95 assert_eq!(normalized.input_tokens, 550);
96 assert_eq!(normalized.cached_input_tokens, 400);
97 assert_eq!(normalized.cache_creation_tokens, 50);
98 assert_eq!(normalized.output_tokens, 20);
99 }
100
101 #[test]
102 fn normalized_turn_usage_treats_minimax_like_anthropic() {
103 let usage = ProviderUsage {
104 prompt_tokens: 100,
105 completion_tokens: 20,
106 total_tokens: 120,
107 cached_prompt_tokens: None,
108 cache_creation_tokens: Some(50),
109 cache_read_tokens: Some(400),
110 iterations: None,
111 };
112
113 let normalized = normalized_turn_usage("minimax", &usage);
114 assert_eq!(normalized.input_tokens, 550);
115 assert_eq!(normalized.cached_input_tokens, 400);
116 assert_eq!(normalized.cache_creation_tokens, 50);
117 assert_eq!(normalized.output_tokens, 20);
118 }
119
120 #[test]
121 fn normalized_turn_usage_keeps_openai_prompt_tokens_as_total() {
122 let usage = ProviderUsage {
123 prompt_tokens: 500,
124 completion_tokens: 30,
125 total_tokens: 530,
126 cached_prompt_tokens: Some(400),
127 cache_creation_tokens: None,
128 cache_read_tokens: None,
129 iterations: None,
130 };
131
132 let normalized = normalized_turn_usage("openai", &usage);
133 assert_eq!(normalized.input_tokens, 500);
134 assert_eq!(normalized.cached_input_tokens, 400);
135 assert_eq!(normalized.cache_creation_tokens, 0);
136 }
137
138 #[test]
139 fn provider_reports_exclusive_input_is_case_insensitive() {
140 assert!(provider_reports_exclusive_input("Anthropic"));
141 assert!(provider_reports_exclusive_input("ANTHROPIC"));
142 assert!(!provider_reports_exclusive_input("OpenAI"));
143 assert!(!provider_reports_exclusive_input("openai"));
144 }
145
146 #[test]
147 fn estimate_session_costs_with_pricing_discounts_cache_reads() {
148 let pricing = ModelPricing {
149 input: Some(0.01),
150 output: Some(0.02),
151 cache_read: Some(0.001),
152 cache_write: Some(0.0125),
153 };
154 let usage = vtcode_exec_events::Usage {
155 input_tokens: 1_000,
156 cached_input_tokens: 800,
157 cache_creation_tokens: 0,
158 output_tokens: 100,
159 };
160
161 let estimate = estimate_session_costs_with_pricing(pricing, &usage).expect("estimate");
162
163 approx_eq(estimate.raw_usd, 1_000.0 * 0.01 + 100.0 * 0.02);
165 approx_eq(estimate.effective_usd, 200.0 * 0.01 + 800.0 * 0.001 + 100.0 * 0.02);
167 assert!(estimate.effective_usd < estimate.raw_usd);
168 }
169
170 #[test]
171 fn estimate_session_costs_with_pricing_matches_raw_when_no_cache_activity() {
172 let pricing = ModelPricing {
173 input: Some(0.01),
174 output: Some(0.02),
175 cache_read: Some(0.001),
176 cache_write: Some(0.0125),
177 };
178 let usage = vtcode_exec_events::Usage {
179 input_tokens: 1_000,
180 cached_input_tokens: 0,
181 cache_creation_tokens: 0,
182 output_tokens: 100,
183 };
184
185 let estimate = estimate_session_costs_with_pricing(pricing, &usage).expect("estimate");
186 approx_eq(estimate.raw_usd, estimate.effective_usd);
187 }
188
189 #[test]
190 fn estimate_session_costs_with_pricing_uses_heuristic_fallback_rates() {
191 let pricing = ModelPricing {
192 input: Some(0.01),
193 output: Some(0.02),
194 cache_read: None,
195 cache_write: None,
196 };
197 let usage = vtcode_exec_events::Usage {
198 input_tokens: 1_000,
199 cached_input_tokens: 50,
200 cache_creation_tokens: 500,
201 output_tokens: 50,
202 };
203
204 let estimate = estimate_session_costs_with_pricing(pricing, &usage).expect("estimate");
205
206 let read_rate = 0.01 * 0.10;
207 let write_rate = 0.01 * 1.25;
208 let uncached = 1_000.0 - 50.0 - 500.0;
209 let expected_effective = uncached * 0.01 + 50.0 * read_rate + 500.0 * write_rate + 50.0 * 0.02;
210 approx_eq(estimate.effective_usd, expected_effective);
211 approx_eq(estimate.raw_usd, 1_000.0 * 0.01 + 50.0 * 0.02);
212 assert!(estimate.effective_usd > estimate.raw_usd);
216 }
217
218 #[test]
219 fn estimate_session_costs_with_pricing_returns_none_without_full_pricing() {
220 let missing_input = ModelPricing {
221 input: None,
222 output: Some(0.02),
223 cache_read: None,
224 cache_write: None,
225 };
226 let missing_output = ModelPricing {
227 input: Some(0.01),
228 output: None,
229 cache_read: None,
230 cache_write: None,
231 };
232 let usage = vtcode_exec_events::Usage::default();
233
234 assert!(estimate_session_costs_with_pricing(missing_input, &usage).is_none());
235 assert!(estimate_session_costs_with_pricing(missing_output, &usage).is_none());
236 }
237
238 #[test]
239 fn session_budget_tracks_spend_and_thresholds() {
240 let mut budget = SessionBudget::new(Some(1.0));
241 assert_eq!(budget.status(), BudgetStatus::Ok);
242 assert_eq!(budget.record(0.5), BudgetStatus::Ok);
244 assert_eq!(budget.record(0.3), BudgetStatus::Warning { spent: 0.8, max: 1.0 });
246 assert_eq!(budget.record(0.3), BudgetStatus::Exceeded { spent: 1.1, max: 1.0 });
248 assert!((budget.spent_usd() - 1.1).abs() < 1e-9);
249 assert!((budget.remaining_usd().unwrap() - 0.0).abs() < 1e-9);
250 }
251
252 #[test]
253 fn session_budget_unlimited_is_always_ok() {
254 let mut budget = SessionBudget::new(None);
255 assert_eq!(budget.record(1000.0), BudgetStatus::Ok);
256 assert_eq!(budget.remaining_usd(), None);
257 }
258
259 #[test]
260 fn budget_status_classify_matches_harness_semantics() {
261 assert_eq!(BudgetStatus::classify(999.0, None, 0.75), BudgetStatus::Ok);
263 assert_eq!(BudgetStatus::classify(0.5, Some(1.0), 0.75), BudgetStatus::Ok);
265 assert_eq!(BudgetStatus::classify(0.8, Some(1.0), 0.75), BudgetStatus::Warning { spent: 0.8, max: 1.0 });
267 assert!(!BudgetStatus::classify(1.0, Some(1.0), 0.75).is_exceeded());
269 assert!(BudgetStatus::classify(1.01, Some(1.0), 0.75).is_exceeded());
271 assert_eq!(BudgetStatus::classify(0.6, Some(1.0), 0.5), BudgetStatus::Warning { spent: 0.6, max: 1.0 });
273 }
274}
275
276pub const DEFAULT_BUDGET_WARNING_RATIO: f64 = 0.75;
281
282#[derive(Debug, Clone, Copy, PartialEq)]
289pub enum BudgetStatus {
290 Ok,
292 Warning {
294 spent: f64,
296 max: f64,
298 },
299 Exceeded {
301 spent: f64,
303 max: f64,
305 },
306}
307
308impl BudgetStatus {
309 #[must_use]
318 pub fn classify(spent_usd: f64, max_usd: Option<f64>, warning_threshold: f64) -> Self {
319 let Some(max) = max_usd else {
320 return BudgetStatus::Ok;
321 };
322 if spent_usd > max {
323 BudgetStatus::Exceeded { spent: spent_usd, max }
324 } else if spent_usd >= warning_threshold * max {
325 BudgetStatus::Warning { spent: spent_usd, max }
326 } else {
327 BudgetStatus::Ok
328 }
329 }
330
331 #[must_use]
333 pub fn is_exceeded(&self) -> bool {
334 matches!(self, BudgetStatus::Exceeded { .. })
335 }
336}
337
338#[derive(Debug, Clone)]
347pub struct SessionBudget {
348 max_usd: Option<f64>,
349 warning_threshold: f64,
350 spent_usd: f64,
351}
352
353impl SessionBudget {
354 #[must_use]
357 pub fn new(max_usd: Option<f64>) -> Self {
358 Self::with_warning_threshold(max_usd, DEFAULT_BUDGET_WARNING_RATIO)
359 }
360
361 #[must_use]
364 pub fn with_warning_threshold(max_usd: Option<f64>, warning_threshold: f64) -> Self {
365 Self { max_usd, warning_threshold, spent_usd: 0.0 }
366 }
367
368 pub fn record(&mut self, raw_usd: f64) -> BudgetStatus {
370 self.spent_usd += raw_usd.max(0.0);
371 self.status()
372 }
373
374 #[must_use]
376 pub fn status(&self) -> BudgetStatus {
377 BudgetStatus::classify(self.spent_usd, self.max_usd, self.warning_threshold)
378 }
379
380 #[must_use]
382 pub fn spent_usd(&self) -> f64 {
383 self.spent_usd
384 }
385
386 #[must_use]
388 pub fn remaining_usd(&self) -> Option<f64> {
389 self.max_usd.map(|m| (m - self.spent_usd).max(0.0))
390 }
391}