1use std::collections::HashMap;
5use std::sync::Arc;
6
7use parking_lot::Mutex;
8
9use thiserror::Error;
10
11#[derive(Debug, Error)]
12#[error("daily budget exhausted: spent {spent_cents:.2} / {budget_cents:.2} cents")]
13pub struct BudgetExhausted {
14 pub spent_cents: f64,
15 pub budget_cents: f64,
16}
17
18#[derive(Debug, Clone, Default)]
20pub struct ProviderUsage {
21 pub input_tokens: u64,
22 pub cache_read_tokens: u64,
23 pub cache_write_tokens: u64,
24 pub output_tokens: u64,
25 pub cost_cents: f64,
26 pub request_count: u64,
27 pub model: String,
29}
30
31#[derive(Debug, Clone)]
32pub struct ModelPricing {
33 pub prompt_cents_per_1k: f64,
34 pub completion_cents_per_1k: f64,
35 pub cache_read_cents_per_1k: f64,
37 pub cache_write_cents_per_1k: f64,
39}
40
41struct CostState {
42 spent_cents: f64,
43 day: u32,
44 providers: HashMap<String, ProviderUsage>,
45 successful_tasks: u64,
46}
47
48pub struct CostTracker {
49 pricing: HashMap<String, ModelPricing>,
50 state: Arc<Mutex<CostState>>,
51 max_daily_cents: f64,
52 enabled: bool,
53}
54
55fn current_day() -> u32 {
56 use std::time::{SystemTime, UNIX_EPOCH};
57 let secs = SystemTime::now()
58 .duration_since(UNIX_EPOCH)
59 .unwrap_or_default()
60 .as_secs();
61 u32::try_from(secs / 86_400).unwrap_or(0)
63}
64
65fn claude_pricing(prompt: f64, completion: f64) -> ModelPricing {
66 ModelPricing {
67 prompt_cents_per_1k: prompt,
68 completion_cents_per_1k: completion,
69 cache_read_cents_per_1k: prompt * 0.1,
71 cache_write_cents_per_1k: prompt * 1.25,
72 }
73}
74
75fn openai_pricing(prompt: f64, completion: f64) -> ModelPricing {
76 ModelPricing {
77 prompt_cents_per_1k: prompt,
78 completion_cents_per_1k: completion,
79 cache_read_cents_per_1k: prompt * 0.5,
81 cache_write_cents_per_1k: 0.0,
82 }
83}
84
85fn default_pricing() -> HashMap<String, ModelPricing> {
86 let mut m = HashMap::new();
87 m.insert("claude-sonnet-4-20250514".into(), claude_pricing(0.3, 1.5));
89 m.insert("claude-opus-4-20250514".into(), claude_pricing(1.5, 7.5));
90 m.insert("claude-opus-4-1-20250805".into(), claude_pricing(1.5, 7.5));
92 m.insert("claude-haiku-4-5-20251001".into(), claude_pricing(0.1, 0.5));
94 m.insert(
95 "claude-sonnet-4-5-20250929".into(),
96 claude_pricing(0.3, 1.5),
97 );
98 m.insert("claude-opus-4-5-20251101".into(), claude_pricing(0.5, 2.5));
99 m.insert("claude-sonnet-4-6".into(), claude_pricing(0.3, 1.5));
101 m.insert("claude-opus-4-6".into(), claude_pricing(0.5, 2.5));
102 m.insert("claude-sonnet-5".into(), claude_pricing(0.3, 1.5));
104 m.insert("claude-opus-4-8".into(), claude_pricing(0.5, 2.5));
105 m.insert("gpt-4o".into(), openai_pricing(0.25, 1.0));
107 m.insert("gpt-4o-mini".into(), openai_pricing(0.015, 0.06));
108 m.insert("gpt-5".into(), openai_pricing(0.125, 1.0));
110 m.insert("gpt-5-mini".into(), openai_pricing(0.025, 0.2));
112 m
113}
114
115fn reset_if_new_day(state: &mut CostState) {
116 let today = current_day();
117 if state.day != today {
118 state.spent_cents = 0.0;
119 state.day = today;
120 state.providers.clear();
121 state.successful_tasks = 0;
122 }
123}
124
125impl CostTracker {
126 #[must_use]
127 pub fn new(enabled: bool, max_daily_cents: f64) -> Self {
128 Self {
129 pricing: default_pricing(),
130 state: Arc::new(Mutex::new(CostState {
131 spent_cents: 0.0,
132 day: current_day(),
133 providers: HashMap::new(),
134 successful_tasks: 0,
135 })),
136 max_daily_cents,
137 enabled,
138 }
139 }
140
141 #[must_use]
142 pub fn with_pricing(mut self, model: &str, pricing: ModelPricing) -> Self {
143 self.pricing.insert(model.to_owned(), pricing);
144 self
145 }
146
147 #[allow(clippy::too_many_arguments)] pub fn record_usage(
158 &self,
159 provider_name: &str,
160 provider_kind: &str,
161 model: &str,
162 input_tokens: u64,
163 cache_read_tokens: u64,
164 cache_write_tokens: u64,
165 output_tokens: u64,
166 ) {
167 if !self.enabled {
168 return;
169 }
170 let pricing = if let Some(p) = self.pricing.get(model).cloned() {
171 p
172 } else {
173 let is_local = matches!(provider_kind, "ollama" | "candle" | "local");
174 if is_local {
175 tracing::debug!(model, "local model; cost recorded as zero");
176 } else {
177 tracing::warn!(
178 model,
179 "model not found in pricing table; cost recorded as zero"
180 );
181 }
182 ModelPricing {
183 prompt_cents_per_1k: 0.0,
184 completion_cents_per_1k: 0.0,
185 cache_read_cents_per_1k: 0.0,
186 cache_write_cents_per_1k: 0.0,
187 }
188 };
189 #[allow(clippy::cast_precision_loss)]
190 let cost = pricing.prompt_cents_per_1k * (input_tokens as f64) / 1000.0
191 + pricing.completion_cents_per_1k * (output_tokens as f64) / 1000.0
192 + pricing.cache_read_cents_per_1k * (cache_read_tokens as f64) / 1000.0
193 + pricing.cache_write_cents_per_1k * (cache_write_tokens as f64) / 1000.0;
194
195 let mut state = self.state.lock();
196 reset_if_new_day(&mut state);
197 state.spent_cents += cost;
198
199 let entry = state.providers.entry(provider_name.to_owned()).or_default();
200 entry.input_tokens += input_tokens;
201 entry.cache_read_tokens += cache_read_tokens;
202 entry.cache_write_tokens += cache_write_tokens;
203 entry.output_tokens += output_tokens;
204 entry.cost_cents += cost;
205 entry.request_count += 1;
206 model.clone_into(&mut entry.model);
207 }
208
209 pub fn check_budget(&self) -> Result<(), BudgetExhausted> {
213 if !self.enabled {
214 return Ok(());
215 }
216 let mut state = self.state.lock();
217 reset_if_new_day(&mut state);
218 if self.max_daily_cents > 0.0 && state.spent_cents >= self.max_daily_cents {
219 return Err(BudgetExhausted {
220 spent_cents: state.spent_cents,
221 budget_cents: self.max_daily_cents,
222 });
223 }
224 Ok(())
225 }
226
227 #[must_use]
229 pub fn max_daily_cents(&self) -> f64 {
230 self.max_daily_cents
231 }
232
233 #[must_use]
234 pub fn current_spend(&self) -> f64 {
235 let state = self.state.lock();
236 state.spent_cents
237 }
238
239 pub fn record_successful_task(&self) {
243 if !self.enabled {
244 return;
245 }
246 let mut state = self.state.lock();
247 reset_if_new_day(&mut state);
248 state.successful_tasks += 1;
249 }
250
251 #[must_use]
253 pub fn cps(&self) -> Option<f64> {
254 let state = self.state.lock();
255 if state.successful_tasks == 0 {
256 return None;
257 }
258 #[allow(clippy::cast_precision_loss)]
259 Some(state.spent_cents / state.successful_tasks as f64)
260 }
261
262 #[must_use]
264 pub fn successful_tasks(&self) -> u64 {
265 self.state.lock().successful_tasks
266 }
267
268 #[must_use]
270 pub fn provider_breakdown(&self) -> Vec<(String, ProviderUsage)> {
271 let state = self.state.lock();
272 let mut breakdown: Vec<(String, ProviderUsage)> = state
273 .providers
274 .iter()
275 .map(|(k, v)| (k.clone(), v.clone()))
276 .collect();
277 breakdown.sort_by(|a, b| {
278 b.1.cost_cents
279 .partial_cmp(&a.1.cost_cents)
280 .unwrap_or(std::cmp::Ordering::Equal)
281 });
282 breakdown
283 }
284}
285
286#[cfg(test)]
287mod tests {
288 use super::*;
289
290 fn record(tracker: &CostTracker, provider: &str, model: &str, input: u64, output: u64) {
291 tracker.record_usage(provider, "cloud", model, input, 0, 0, output);
292 }
293
294 #[test]
295 fn cost_tracker_records_usage_and_calculates_cost() {
296 let tracker = CostTracker::new(true, 1000.0);
297 record(&tracker, "openai", "gpt-4o", 1000, 1000);
298 let spend = tracker.current_spend();
300 assert!((spend - 1.25).abs() < 0.001);
301 }
302
303 #[test]
304 fn check_budget_passes_when_under_limit() {
305 let tracker = CostTracker::new(true, 100.0);
306 record(&tracker, "openai", "gpt-4o-mini", 100, 100);
307 assert!(tracker.check_budget().is_ok());
308 }
309
310 #[test]
311 fn check_budget_fails_when_over_limit() {
312 let tracker = CostTracker::new(true, 0.01);
313 record(&tracker, "claude", "claude-opus-4-20250514", 10000, 10000);
314 assert!(tracker.check_budget().is_err());
315 }
316
317 #[test]
318 fn daily_reset_clears_spending() {
319 let tracker = CostTracker::new(true, 100.0);
320 record(&tracker, "openai", "gpt-4o", 1000, 1000);
321 assert!(tracker.current_spend() > 0.0);
322 {
324 let mut state = tracker.state.lock();
325 state.day = 0; }
327 assert!(tracker.check_budget().is_ok());
329 assert!((tracker.current_spend() - 0.0).abs() < 0.001);
330 }
331
332 #[test]
333 fn daily_reset_clears_provider_breakdown() {
334 let tracker = CostTracker::new(true, 100.0);
335 record(&tracker, "openai", "gpt-4o", 1000, 1000);
336 assert!(!tracker.provider_breakdown().is_empty());
337 {
339 let mut state = tracker.state.lock();
340 state.day = 0;
341 }
342 assert!(tracker.check_budget().is_ok());
343 assert!(tracker.provider_breakdown().is_empty());
344 }
345
346 #[test]
347 fn ollama_zero_cost() {
348 let tracker = CostTracker::new(true, 100.0);
349 record(&tracker, "ollama", "llama3:8b", 10000, 10000);
350 assert!((tracker.current_spend() - 0.0).abs() < 0.001);
351 }
352
353 #[test]
354 fn ollama_unknown_model_no_warn_no_panic() {
355 let tracker = CostTracker::new(true, 100.0);
357 tracker.record_usage(
358 "local",
359 "ollama",
360 "totally-unknown-ollama-model",
361 5000,
362 0,
363 0,
364 5000,
365 );
366 assert!((tracker.current_spend() - 0.0).abs() < 0.001);
367 }
368
369 #[test]
370 fn cloud_unknown_model_still_records_zero_cost() {
371 let tracker = CostTracker::new(true, 100.0);
373 tracker.record_usage(
374 "openai",
375 "cloud",
376 "totally-unknown-cloud-model",
377 5000,
378 0,
379 0,
380 5000,
381 );
382 assert!((tracker.current_spend() - 0.0).abs() < 0.001);
383 }
384
385 #[test]
386 fn unknown_model_zero_cost() {
387 let tracker = CostTracker::new(true, 100.0);
388 record(&tracker, "unknown", "totally-unknown-model", 5000, 5000);
389 assert!((tracker.current_spend() - 0.0).abs() < 0.001);
390 }
391
392 #[test]
393 fn known_claude_model_has_nonzero_cost() {
394 let tracker = CostTracker::new(true, 1000.0);
395 record(&tracker, "claude", "claude-haiku-4-5-20251001", 1000, 1000);
396 assert!(tracker.current_spend() > 0.0);
397 }
398
399 #[test]
400 fn gpt5_pricing_is_correct() {
401 let tracker = CostTracker::new(true, 1000.0);
402 record(&tracker, "openai", "gpt-5", 1000, 1000);
403 let spend = tracker.current_spend();
405 assert!((spend - 1.125).abs() < 0.001);
406 }
407
408 #[test]
409 fn gpt5_mini_pricing_is_correct() {
410 let tracker = CostTracker::new(true, 1000.0);
411 record(&tracker, "openai", "gpt-5-mini", 1000, 1000);
412 let spend = tracker.current_spend();
414 assert!((spend - 0.225).abs() < 0.001);
415 }
416
417 #[test]
418 fn disabled_tracker_always_passes() {
419 let tracker = CostTracker::new(false, 0.0);
420 record(
421 &tracker,
422 "claude",
423 "claude-opus-4-20250514",
424 1_000_000,
425 1_000_000,
426 );
427 assert!(tracker.check_budget().is_ok());
428 assert!((tracker.current_spend() - 0.0).abs() < 0.001);
429 }
430
431 #[test]
432 fn check_budget_unlimited_when_max_daily_cents_is_zero() {
433 let tracker = CostTracker::new(true, 0.0);
434 record(
435 &tracker,
436 "claude",
437 "claude-opus-4-20250514",
438 100_000,
439 100_000,
440 );
441 assert!(tracker.check_budget().is_ok());
442 }
443
444 #[test]
445 fn per_provider_accumulation() {
446 let tracker = CostTracker::new(true, 1000.0);
447 record(&tracker, "claude", "claude-haiku-4-5-20251001", 1000, 500);
448 record(&tracker, "openai", "gpt-4o", 2000, 1000);
449 record(&tracker, "claude", "claude-haiku-4-5-20251001", 500, 200);
450
451 let breakdown = tracker.provider_breakdown();
452 assert_eq!(breakdown.len(), 2);
453
454 let claude = breakdown.iter().find(|(n, _)| n == "claude").unwrap();
455 assert_eq!(claude.1.request_count, 2);
456 assert_eq!(claude.1.input_tokens, 1500);
457 assert_eq!(claude.1.output_tokens, 700);
458
459 let openai = breakdown.iter().find(|(n, _)| n == "openai").unwrap();
460 assert_eq!(openai.1.request_count, 1);
461 assert_eq!(openai.1.input_tokens, 2000);
462 }
463
464 #[test]
465 fn provider_breakdown_sorted_by_cost_desc() {
466 let tracker = CostTracker::new(true, 1000.0);
467 record(&tracker, "cheap", "gpt-4o-mini", 100, 100);
469 record(&tracker, "expensive", "claude-opus-4-20250514", 10000, 5000);
470
471 let breakdown = tracker.provider_breakdown();
472 assert_eq!(breakdown[0].0, "expensive");
473 }
474
475 #[test]
476 fn cache_tokens_included_in_cost() {
477 let tracker = CostTracker::new(true, 1000.0);
478 tracker.record_usage(
481 "claude",
482 "cloud",
483 "claude-haiku-4-5-20251001",
484 0,
485 1000,
486 0,
487 0,
488 );
489 let spend = tracker.current_spend();
490 assert!(spend > 0.0, "cache read should contribute to cost");
491 }
492
493 #[test]
494 fn cache_write_cost_included_in_total() {
495 let tracker = CostTracker::new(true, 1000.0);
496 tracker.record_usage("claude-provider", "cloud", "claude-opus-4-6", 0, 0, 1000, 0);
500 let cost = tracker.current_spend();
501 assert!((cost - 0.625).abs() < 0.001);
502 }
503
504 #[test]
505 fn claude_5_generation_pricing_matches_retained_4_6_rows() {
506 let sonnet_5 = CostTracker::new(true, 1000.0);
509 sonnet_5.record_usage(
510 "claude-provider",
511 "cloud",
512 "claude-sonnet-5",
513 1000,
514 0,
515 0,
516 1000,
517 );
518 let sonnet_4_6 = CostTracker::new(true, 1000.0);
519 sonnet_4_6.record_usage(
520 "claude-provider",
521 "cloud",
522 "claude-sonnet-4-6",
523 1000,
524 0,
525 0,
526 1000,
527 );
528 assert!((sonnet_5.current_spend() - sonnet_4_6.current_spend()).abs() < f64::EPSILON);
529 assert!(
530 sonnet_5.current_spend() > 0.0,
531 "must not fall back to zero-cost"
532 );
533
534 let opus_4_8 = CostTracker::new(true, 1000.0);
535 opus_4_8.record_usage(
536 "claude-provider",
537 "cloud",
538 "claude-opus-4-8",
539 1000,
540 0,
541 0,
542 1000,
543 );
544 let opus_4_6 = CostTracker::new(true, 1000.0);
545 opus_4_6.record_usage(
546 "claude-provider",
547 "cloud",
548 "claude-opus-4-6",
549 1000,
550 0,
551 0,
552 1000,
553 );
554 assert!((opus_4_8.current_spend() - opus_4_6.current_spend()).abs() < f64::EPSILON);
555 assert!(
556 opus_4_8.current_spend() > 0.0,
557 "must not fall back to zero-cost"
558 );
559 }
560
561 #[test]
562 fn provider_breakdown_empty_when_disabled() {
563 let tracker = CostTracker::new(false, 100.0);
564 tracker.record_usage(
565 "claude",
566 "cloud",
567 "claude-haiku-4-5-20251001",
568 1000,
569 0,
570 0,
571 1000,
572 );
573 assert!(tracker.provider_breakdown().is_empty());
574 }
575
576 #[test]
577 fn cps_none_when_no_tasks() {
578 let tracker = CostTracker::new(true, 100.0);
579 assert!(tracker.cps().is_none());
580 assert_eq!(tracker.successful_tasks(), 0);
581 }
582
583 #[test]
584 fn cps_calculated_correctly() {
585 let tracker = CostTracker::new(true, 100.0);
586 record(&tracker, "openai", "gpt-4o", 1000, 1000);
588 tracker.record_successful_task();
589 tracker.record_successful_task();
590 assert_eq!(tracker.successful_tasks(), 2);
591 let cps = tracker.cps().expect("cps should be Some after tasks");
592 assert!((cps - 0.625).abs() < 0.001, "cps={cps}");
594 }
595
596 #[test]
597 fn cps_resets_on_new_day() {
598 let tracker = CostTracker::new(true, 100.0);
599 record(&tracker, "openai", "gpt-4o", 1000, 1000);
600 tracker.record_successful_task();
601 assert_eq!(tracker.successful_tasks(), 1);
602 {
604 let mut state = tracker.state.lock();
605 state.day = 0;
606 }
607 assert!(tracker.check_budget().is_ok());
609 assert_eq!(tracker.successful_tasks(), 0);
610 assert!(tracker.cps().is_none());
611 }
612}