Skip to content

Cache Savings Calculator — Rust source

See what prompt caching saves — uncached vs cached cost over N requests, with the write-premium break-even point.

This is the Rust implementation — the same logic the interactive tool runs, in a shareable, citable form.

//! cache_savings — uncached vs prompt-cached LLM cost comparison.
//!
//! Language: Rust (edition 2021, standard library only)
//! Source:   CosmoDev polyglot showcase port of the Cache Savings Calculator
//!           tool, ported from src/lib/cacheSavings.ts (the canonical
//!           TypeScript implementation).
//! Tool:     https://dev.cosmolabs.org/tools/cache-savings-calculator
//! License:  display source — part of CosmoDev's polyglot tool pages.
//!
//! Design goals:
//!   - Pure + deterministic; never panics (plain f64 math, no unwraps).
//!   - Functionally equivalent to the TS reference: same inputs -> same outputs.
//!   - Self-contained: std only (no crates.io dependencies).
//!
//! The TS original takes a full AiModel record but reads only its four pricing
//! rates, so this port narrows the parameter to exactly those fields. Any
//! `None` rate makes every output `None` — the caller renders an explanatory
//! empty state instead of partial math. All rates are per-1M-token USD,
//! mirroring the cost conventions of llmCost.ts.
//!
//! Numeric mapping: TS has one `number` type, so tokens and hits stay `f64`
//! here (fractional `hits` clamps up to 1.0 exactly like `Math.max(1, hits)`);
//! `breakEvenHits` is a whole hit count, so it lands in `u64`.

/// The four per-1M-token USD pricing rates `cache_math` reads from the TS
/// AiModel record.
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct ModelRates {
    /// Uncached prompt (input) rate, USD per 1M tokens.
    pub input_per_m: Option<f64>,
    /// Completion (output) rate, USD per 1M tokens.
    pub output_per_m: Option<f64>,
    /// Cached prompt read rate, USD per 1M tokens.
    pub cache_read_per_m: Option<f64>,
    /// Cache write premium rate, USD per 1M tokens.
    pub cache_write_per_m: Option<f64>,
}

/// Request shape (TS `CacheInput`).
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct CacheInput {
    /// Prompt (input) tokens per request.
    pub prompt_tokens: f64,
    /// Completion (output) tokens per request.
    pub output_tokens: f64,
    /// Requests reusing the cached prompt. Values < 1 are treated as 1.
    pub hits: f64,
}

/// Result shape. Any missing rate nulls every field.
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct CacheMath {
    /// `hits × (prompt·in$/M + output·out$/M) / 1e6`.
    pub uncached: Option<f64>,
    /// `(prompt·write$/M + hits × (prompt·read$/M + output·out$/M)) / 1e6` —
    /// one cache write, `hits` cache reads, output billed every request.
    pub cached: Option<f64>,
    /// uncached − cached (negative when caching costs more).
    pub savings: Option<f64>,
    /// savings / uncached × 100; 0 when uncached is 0.
    pub savings_pct: Option<f64>,
    /// `ceil(write$/M / read$/M)` when read$/M > 0 — cache hits needed for
    /// cumulative READ spend to equal ONE write premium; `None` otherwise.
    pub break_even_hits: Option<u64>,
}

impl CacheMath {
    /// The all-`None` result used when any pricing rate is missing.
    fn nulled() -> Self {
        CacheMath {
            uncached: None,
            cached: None,
            savings: None,
            savings_pct: None,
            break_even_hits: None,
        }
    }
}

/// Compare uncached vs prompt-cached cost for one model. Any missing rate
/// (input, output, cache read, cache write) nulls every field — the caller
/// renders an explanatory empty state instead of partial math.
pub fn cache_math(model: ModelRates, input: CacheInput) -> CacheMath {
    let (Some(ipm), Some(opm), Some(cr), Some(cw)) = (
        model.input_per_m,
        model.output_per_m,
        model.cache_read_per_m,
        model.cache_write_per_m,
    ) else {
        return CacheMath::nulled();
    };

    let hits = input.hits.max(1.0);
    let in_t = input.prompt_tokens;
    let out_t = input.output_tokens;

    // One cache write, `hits` cache reads; output tokens are billed on every request.
    let uncached = (hits * (in_t * ipm + out_t * opm)) / 1_000_000.0;
    let cached = (in_t * cw + hits * (in_t * cr + out_t * opm)) / 1_000_000.0;
    let savings = uncached - cached;
    let savings_pct = if uncached == 0.0 { 0.0 } else { (savings / uncached) * 100.0 };
    let break_even_hits = if cr > 0.0 { Some((cw / cr).ceil() as u64) } else { None };

    CacheMath {
        uncached: Some(uncached),
        cached: Some(cached),
        savings: Some(savings),
        savings_pct: Some(savings_pct),
        break_even_hits,
    }
}

#[cfg(test)]
mod tests {
    use super::*;

    /// Fixture model F from cacheSavings.test.ts: inputPerM 10, outputPerM 50,
    /// cacheReadPerM 1, cacheWritePerM 12.5. These vectors are the lock-step
    /// contract shared with the TS lib (and the future Go twin). Override a
    /// rate with `ModelRates { field: …, ..base() }`.
    fn base() -> ModelRates {
        ModelRates {
            input_per_m: Some(10.0),
            output_per_m: Some(50.0),
            cache_read_per_m: Some(1.0),
            cache_write_per_m: Some(12.5),
        }
    }

    fn input(prompt: f64, output: f64, hits: f64) -> CacheInput {
        CacheInput { prompt_tokens: prompt, output_tokens: output, hits }
    }

    fn close(a: f64, b: f64) -> bool {
        (a - b).abs() < 1e-9
    }

    #[test]
    fn spec_vector_5_hits() {
        let r = cache_math(base(), input(10_000.0, 1_000.0, 5.0));
        let u = r.uncached.unwrap();
        let c = r.cached.unwrap();
        let s = r.savings.unwrap();
        let p = r.savings_pct.unwrap();
        assert!(close(u, 0.75), "uncached {u}");
        assert!(close(c, 0.425), "cached {c}");
        assert!(close(s, 0.325), "savings {s}");
        assert!(close(p, 43.333_333_333_3), "savingsPct {p}");
        assert_eq!(r.break_even_hits, Some(13)); // ceil(12.5 / 1)
    }

    #[test]
    fn honest_negative_saving_at_one_hit() {
        let r = cache_math(base(), input(10_000.0, 1_000.0, 1.0));
        assert!(close(r.uncached.unwrap(), 0.15));
        assert!(close(r.cached.unwrap(), 0.185));
        assert!(close(r.savings.unwrap(), -0.035));
        assert!(close(r.savings_pct.unwrap(), -23.333_333_333_3));
        assert_eq!(r.break_even_hits, Some(13));
    }

    #[test]
    fn unpriced_rates_null_everything() {
        for missing in [
            ModelRates { input_per_m: None, ..base() },
            ModelRates { output_per_m: None, ..base() },
            ModelRates { cache_read_per_m: None, ..base() },
            ModelRates { cache_write_per_m: None, ..base() },
        ] {
            assert_eq!(cache_math(missing, input(10_000.0, 1_000.0, 5.0)), CacheMath::nulled());
        }
    }

    #[test]
    fn hits_below_one_counts_as_one() {
        assert_eq!(
            cache_math(base(), input(10_000.0, 1_000.0, 0.0)),
            cache_math(base(), input(10_000.0, 1_000.0, 1.0))
        );
    }

    #[test]
    fn zero_tokens_no_division_error() {
        let r = cache_math(base(), input(0.0, 0.0, 5.0));
        assert_eq!(r.uncached, Some(0.0));
        assert_eq!(r.cached, Some(0.0));
        assert_eq!(r.savings, Some(0.0));
        assert_eq!(r.savings_pct, Some(0.0));
        assert_eq!(r.break_even_hits, Some(13));
    }

    #[test]
    fn zero_read_rate_nulls_break_even_but_keeps_costs() {
        let r = cache_math(
            ModelRates { cache_read_per_m: Some(0.0), ..base() },
            input(10_000.0, 1_000.0, 5.0),
        );
        assert_eq!(r.break_even_hits, None);
        assert!(close(r.uncached.unwrap(), 0.75));
        assert!(close(r.cached.unwrap(), 0.375)); // 10k×$12.5 + 5×(0 + 1k×$50)
    }

    #[test]
    fn breaks_even_exactly_at_integer_boundary() {
        let r = cache_math(
            ModelRates { cache_write_per_m: Some(4.0), cache_read_per_m: Some(2.0), ..base() },
            input(1_000.0, 0.0, 3.0),
        );
        assert_eq!(r.break_even_hits, Some(2)); // ceil(4/2), no rounding up
    }
}

Also available in 13 other languages

Every CosmoDev tool ships its pure logic in TypeScript (web) and Go (CLI), with authored implementations in a dozen-plus languages — the same contract, ported. Compare all languages side by side →