use std::collections::BTreeMap;
use std::fmt;
#[derive(Clone, Copy, Debug, Default, PartialEq)]
struct Entry {
total: f64,
scale: f64,
}
#[derive(Clone, Debug, Default, PartialEq)]
pub struct Ledger(BTreeMap<&'static str, Entry>);
pub mod quantity {
pub const ENERGY: &str = "energy";
pub const MOMENTUM: &str = "momentum";
pub const MASS: &str = "mass";
pub const CHARGE: &str = "charge";
pub const PHOTONS: &str = "photons";
}
impl Ledger {
pub fn new() -> Ledger {
Ledger(BTreeMap::new())
}
pub fn with(mut self, quantity: &'static str, si_total: f64) -> Ledger {
self.add(quantity, si_total);
self
}
pub fn add(&mut self, quantity: &'static str, si_total: f64) {
let entry = self.0.entry(quantity).or_default();
entry.total += si_total;
entry.scale = entry.scale.max(si_total.abs());
}
pub fn get(&self, quantity: &str) -> Option<f64> {
self.0.get(quantity).map(|e| e.total)
}
pub fn scale_of(&self, quantity: &str) -> Option<f64> {
self.0.get(quantity).map(|e| e.scale)
}
pub fn is_empty(&self) -> bool {
self.0.is_empty()
}
pub fn quantities(&self) -> impl Iterator<Item = (&'static str, f64)> + '_ {
self.0.iter().map(|(k, e)| (*k, e.total))
}
pub fn merged(mut self, other: &Ledger) -> Ledger {
for (name, entry) in other.0.iter() {
let mine = self.0.entry(name).or_default();
mine.total += entry.total;
mine.scale = mine.scale.max(entry.scale);
}
self
}
}
#[derive(Clone, Debug, PartialEq)]
pub struct Violation {
pub quantity: String,
pub site: String,
pub before: f64,
pub after: f64,
pub scale: f64,
pub tolerance: f64,
}
impl Violation {
pub fn at(site: impl Into<String>, quantity: impl Into<String>, detail: f64) -> Violation {
Violation {
quantity: quantity.into(),
site: site.into(),
before: detail,
after: detail,
scale: detail.abs(),
tolerance: 0.0,
}
}
pub fn error(&self) -> f64 {
(self.after - self.before).abs()
}
pub fn relative_error(&self) -> f64 {
let scale = if self.scale > 0.0 {
self.scale
} else {
self.before.abs().max(self.after.abs())
};
if scale == 0.0 {
0.0
} else {
self.error() / scale
}
}
}
impl fmt::Display for Violation {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
if self.tolerance == 0.0 && self.before == self.after {
return write!(f, "at {}: {} ({})", self.site, self.quantity, self.before);
}
if self.before == self.after {
return write!(
f,
"{} is not conserved at {}: {}",
self.quantity, self.site, self.before
);
}
let verb = if self.after > self.before {
"created"
} else {
"destroyed"
};
write!(
f,
"{} {} at {}: {:.6e} became {:.6e}, a relative change of {:.3e} against a \
tolerance of {:.3e}",
self.quantity,
verb,
self.site,
self.before,
self.after,
self.relative_error(),
self.tolerance
)
}
}
impl std::error::Error for Violation {}
pub fn audit(site: &str, before: &Ledger, after: &Ledger, rel_tol: f64) -> Result<(), Violation> {
let mut names: Vec<&'static str> = before.0.keys().copied().collect();
for name in after.0.keys() {
if !before.0.contains_key(name) {
names.push(name);
}
}
names.sort_unstable();
for name in names {
let b = before.get(name).unwrap_or(0.0);
let a = after.get(name).unwrap_or(0.0);
let scale = b
.abs()
.max(a.abs())
.max(before.scale_of(name).unwrap_or(0.0))
.max(after.scale_of(name).unwrap_or(0.0));
if scale < 1e-300 {
continue;
}
if (a - b).abs() / scale > rel_tol {
return Err(Violation {
quantity: name.to_string(),
site: site.to_string(),
before: b,
after: a,
scale,
tolerance: rel_tol,
});
}
}
Ok(())
}
pub trait Conserves {
fn ledger(&self) -> Ledger;
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_ledger_that_did_not_move_passes() {
let before = Ledger::new()
.with(quantity::ENERGY, 3.7)
.with(quantity::MASS, 2.0);
let after = Ledger::new()
.with(quantity::ENERGY, 3.7 + 4e-16)
.with(quantity::MASS, 2.0);
assert!(audit("test", &before, &after, 1e-12).is_ok());
}
#[test]
fn a_leak_is_named_and_sited() {
let before = Ledger::new().with(quantity::ENERGY, 1.0);
let after = Ledger::new().with(quantity::ENERGY, 0.6);
let err = audit("thermal", &before, &after, 1e-9).expect_err("40% is not arithmetic");
assert_eq!(err.quantity, "energy");
assert_eq!(err.site, "thermal");
assert!((err.relative_error() - 0.4).abs() < 1e-12);
let text = err.to_string();
assert!(text.contains("destroyed"), "{text}");
assert!(text.contains("thermal"), "{text}");
}
#[test]
fn creating_something_reads_differently_from_losing_it() {
let before = Ledger::new().with(quantity::PHOTONS, 1e6);
let after = Ledger::new().with(quantity::PHOTONS, 1.5e6);
let err = audit("optics", &before, &after, 1e-6).unwrap_err();
assert!(err.to_string().contains("created"), "{err}");
}
#[test]
fn a_quantity_absent_before_is_still_audited() {
let before = Ledger::new().with(quantity::ENERGY, 1.0);
let after = Ledger::new()
.with(quantity::ENERGY, 1.0)
.with(quantity::MOMENTUM, 5.0);
let err = audit("contact", &before, &after, 1e-9).expect_err("momentum from nowhere");
assert_eq!(err.quantity, "momentum");
assert_eq!(err.before, 0.0);
}
#[test]
fn the_audit_order_is_fixed() {
let before = Ledger::new()
.with(quantity::MOMENTUM, 1.0)
.with(quantity::CHARGE, 1.0)
.with(quantity::ENERGY, 1.0);
let after = Ledger::new()
.with(quantity::MOMENTUM, 2.0)
.with(quantity::CHARGE, 2.0)
.with(quantity::ENERGY, 2.0);
for _ in 0..8 {
let err = audit("s", &before, &after, 1e-9).unwrap_err();
assert_eq!(err.quantity, "charge");
}
}
#[test]
fn ledgers_merge_by_summing() {
let a = Ledger::new().with(quantity::ENERGY, 1.5);
let b = Ledger::new()
.with(quantity::ENERGY, 2.5)
.with(quantity::MASS, 1.0);
let total = a.merged(&b);
assert_eq!(total.get(quantity::ENERGY), Some(4.0));
assert_eq!(total.get(quantity::MASS), Some(1.0));
}
#[test]
fn nothing_compared_to_nothing_is_fine() {
let z = Ledger::new().with(quantity::ENERGY, 0.0);
assert!(audit("s", &z, &z, 0.0).is_ok());
}
}