use std::iter::Sum;
use std::ops::{Add, AddAssign};
use serde::{Deserialize, Serialize};
use super::Usage;
#[derive(Clone, Copy, Debug, Default, PartialEq, Serialize, Deserialize)]
pub struct CacheRates {
pub input: f64,
pub cached_read: f64,
pub cache_write: f64,
pub storage_per_hour: f64,
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Serialize, Deserialize)]
pub struct CacheCost {
pub uncached_input: u64,
pub cache_reads: u64,
pub cache_writes: u64,
pub storage_token_hours: f64,
}
impl CacheCost {
pub fn from_usage(usage: &Usage) -> Self {
let input = usage.input_tokens.unwrap_or(0);
let cache_reads = usage.cached_input_tokens.unwrap_or(0);
let cache_writes = usage.cache_creation_input_tokens.unwrap_or(0);
Self {
uncached_input: input.saturating_sub(cache_reads + cache_writes),
cache_reads,
cache_writes,
storage_token_hours: 0.0,
}
}
pub fn prompt_tokens(&self) -> u64 {
self.uncached_input + self.cache_reads + self.cache_writes
}
pub fn usd(&self, rates: &CacheRates) -> f64 {
(self.uncached_input as f64 * rates.input
+ self.cache_reads as f64 * rates.cached_read
+ self.cache_writes as f64 * rates.cache_write
+ self.storage_token_hours * rates.storage_per_hour)
/ 1e6
}
pub fn uncached_usd(&self, rates: &CacheRates) -> f64 {
self.prompt_tokens() as f64 * rates.input / 1e6
}
pub fn saving(&self, rates: &CacheRates) -> f64 {
let uncached = self.uncached_usd(rates);
if uncached > 0.0 {
1.0 - self.usd(rates) / uncached
} else {
0.0
}
}
}
impl Add for CacheCost {
type Output = Self;
fn add(mut self, other: Self) -> Self {
self += other;
self
}
}
impl AddAssign for CacheCost {
fn add_assign(&mut self, other: Self) {
self.uncached_input += other.uncached_input;
self.cache_reads += other.cache_reads;
self.cache_writes += other.cache_writes;
self.storage_token_hours += other.storage_token_hours;
}
}
impl Sum for CacheCost {
fn sum<I: Iterator<Item = Self>>(costs: I) -> Self {
costs.fold(Self::default(), Add::add)
}
}
#[cfg(test)]
mod tests;