use std::collections::BTreeMap;
use std::sync::Mutex;
use serde::de::{Deserialize, Deserializer};
use serde::ser::{Serialize, Serializer};
#[derive(Debug, Default)]
pub struct CostReport {
counters: Mutex<BTreeMap<String, f64>>,
}
impl CostReport {
pub fn new() -> Self {
Self::default()
}
pub fn record(&self, metric: impl Into<String>, amount: f64) {
let mut counters = self.counters.lock().unwrap();
*counters.entry(metric.into()).or_insert(0.0) += amount;
}
pub fn get(&self, metric: &str) -> f64 {
self.counters
.lock()
.unwrap()
.get(metric)
.copied()
.unwrap_or(0.0)
}
pub fn merge(&self, other: &CostReport) {
let snapshot = other.counters.lock().unwrap().clone();
let mut counters = self.counters.lock().unwrap();
for (metric, amount) in snapshot {
*counters.entry(metric).or_insert(0.0) += amount;
}
}
pub fn entries(&self) -> Vec<(String, f64)> {
self.counters
.lock()
.unwrap()
.iter()
.map(|(k, v)| (k.clone(), *v))
.collect()
}
pub fn is_empty(&self) -> bool {
self.counters.lock().unwrap().is_empty()
}
}
impl Clone for CostReport {
fn clone(&self) -> Self {
Self {
counters: Mutex::new(self.counters.lock().unwrap().clone()),
}
}
}
impl Serialize for CostReport {
fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
self.counters.lock().unwrap().serialize(serializer)
}
}
impl<'de> Deserialize<'de> for CostReport {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
let counters = BTreeMap::<String, f64>::deserialize(deserializer)?;
Ok(Self {
counters: Mutex::new(counters),
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn record_accumulates_per_metric() {
let report = CostReport::new();
report.record("tokens", 100.0);
report.record("tokens", 50.0);
report.record("calls", 1.0);
assert_eq!(report.get("tokens"), 150.0);
assert_eq!(report.get("calls"), 1.0);
}
#[test]
fn get_returns_zero_for_unknown_metric() {
let report = CostReport::new();
assert_eq!(report.get("missing"), 0.0);
}
#[test]
fn merge_sums_overlapping_and_disjoint_metrics() {
let a = CostReport::new();
a.record("tokens", 100.0);
a.record("usd", 1.0);
let b = CostReport::new();
b.record("tokens", 25.0);
b.record("calls", 3.0);
a.merge(&b);
assert_eq!(a.get("tokens"), 125.0);
assert_eq!(a.get("usd"), 1.0);
assert_eq!(a.get("calls"), 3.0);
}
#[test]
fn entries_are_name_ordered() {
let report = CostReport::new();
report.record("zeta", 1.0);
report.record("alpha", 2.0);
assert_eq!(
report.entries(),
vec![("alpha".to_string(), 2.0), ("zeta".to_string(), 1.0)],
);
}
#[test]
fn serde_round_trips_through_rmp() {
let report = CostReport::new();
report.record("tokens", 42.0);
report.record("usd", 0.5);
let bytes = rmp_serde::to_vec_named(&report).unwrap();
let restored: CostReport = rmp_serde::from_slice(&bytes).unwrap();
assert_eq!(restored.entries(), report.entries());
}
#[test]
fn clone_is_an_independent_snapshot() {
let original = CostReport::new();
original.record("tokens", 10.0);
let snapshot = original.clone();
original.record("tokens", 5.0);
assert_eq!(snapshot.get("tokens"), 10.0);
assert_eq!(original.get("tokens"), 15.0);
}
}