Skip to main content

zeph_core/
cost.rs

1// SPDX-FileCopyrightText: 2026 Andrei G <bug-ops>
2// SPDX-License-Identifier: MIT OR Apache-2.0
3
4use 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/// Per-provider usage and cost breakdown for the current session/day.
19#[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    /// Last model seen for this provider (informational only — may change per-call).
28    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    /// Cache read (cache hit) price. Claude: 10% of prompt; `OpenAI`: 50%; others: 0%.
36    pub cache_read_cents_per_1k: f64,
37    /// Cache write (cache creation) price. Claude: 125% of prompt; others: 0%.
38    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    // UTC day number (days since epoch)
62    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        // Claude: cache read = 10% of prompt, cache write = 125% of prompt
70        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        // OpenAI: cache read = 50% of prompt, no cache write charge
80        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    // Claude 4 (sonnet-4 / opus-4 base releases)
88    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    // Claude 4.1 Opus ($15/$75 per 1M tokens)
91    m.insert("claude-opus-4-1-20250805".into(), claude_pricing(1.5, 7.5));
92    // Claude 4.5 family
93    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    // Claude 4.6 family
100    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    // Claude 5 / Opus 4.8 family
103    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    // OpenAI
106    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    // GPT-5 family ($1.25/$10 per 1M tokens)
109    m.insert("gpt-5".into(), openai_pricing(0.125, 1.0));
110    // GPT-5 mini ($0.25/$2 per 1M tokens)
111    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    /// Record token usage for a single LLM call, attributed to `provider_name`.
148    ///
149    /// `provider_kind` must be the value returned by `AnyProvider::provider_kind_str()`:
150    /// `"ollama"` or `"candle"` for local providers, `"cloud"` for API providers.
151    /// Local providers always have zero cost by design; the missing-pricing WARN is
152    /// suppressed for them to avoid log floods on every Ollama call.
153    ///
154    /// Cache token counts are optional (pass 0 when not available). Cost is computed
155    /// using model-specific pricing including cache read/write rates.
156    #[allow(clippy::too_many_arguments)] // function with many required inputs; a *Params struct would be more verbose without simplifying the call site
157    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    /// # Errors
210    ///
211    /// Returns `BudgetExhausted` when daily spend exceeds the configured limit.
212    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    /// Returns the configured daily budget in cents. Zero means unlimited.
228    #[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    /// Increment the successful-task counter.
240    ///
241    /// Call after each turn that completes without error and produces a usable agent response.
242    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    /// Returns cost-per-successful-task in cents, or `None` if no tasks recorded yet.
252    #[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    /// Returns total number of successful tasks recorded today.
263    #[must_use]
264    pub fn successful_tasks(&self) -> u64 {
265        self.state.lock().successful_tasks
266    }
267
268    /// Returns per-provider breakdown sorted by cost descending.
269    #[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        // 0.25 + 1.0 = 1.25
299        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        // Simulate day change
323        {
324            let mut state = tracker.state.lock();
325            state.day = 0; // force a past day
326        }
327        // check_budget should reset
328        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        // Simulate day change
338        {
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        // Local providers should silently record zero cost for unknown models.
356        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        // Cloud providers record zero cost for unknown models (WARN emitted separately).
372        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        // 0.125 + 1.0 = 1.125
404        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        // 0.025 + 0.2 = 0.225
413        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        // gpt-4o: cheap; claude-opus: expensive
468        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        // claude-haiku prompt=0.1, cache_read=0.01 per 1k
479        // 1000 cache_read tokens = 0.01 cents; 0 input/output for isolation
480        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        // Claude pricing: cache_write = 125% of prompt price
497        // claude-opus-4-6: prompt = 0.5 cents/1k
498        // 1000 cache_write tokens = (0.5 * 1.25 * 1000) / 1000 = 0.625 cents
499        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        // The new claude-sonnet-5/claude-opus-4-8 rows must price identically to the
507        // retained claude-sonnet-4-6/claude-opus-4-6 rows (#5889: same rates, new IDs).
508        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        // 0.25 (input) + 1.0 (output) = 1.25 cents
587        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        // 1.25 / 2 = 0.625
593        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        // Force day change
603        {
604            let mut state = tracker.state.lock();
605            state.day = 0;
606        }
607        // Any state-touching call triggers reset
608        assert!(tracker.check_budget().is_ok());
609        assert_eq!(tracker.successful_tasks(), 0);
610        assert!(tracker.cps().is_none());
611    }
612}