use std::{cmp::Ordering, collections::BTreeMap};
use num_bigint::{BigInt, BigUint, Sign};
use proptest::{collection::vec, prelude::*};
use super::weight::value;
use crate::{
Distribution, EmptyDistributionError, Probability, ProbabilityError,
Rational, Weight, ZeroDenominatorError
};
fn w(n: u128) -> Weight { Weight::from(n) }
fn p(numerator: u128, denominator: u128) -> Probability
{
Probability::new(w(numerator), w(denominator)).unwrap()
}
fn r(negative: bool, numerator: u128, denominator: u128) -> Rational
{
Rational::new(negative, w(numerator), w(denominator)).unwrap()
}
fn distribution(weights: &[(i32, u128)]) -> Distribution
{
Distribution::from_weights(weights.iter().map(|&(v, n)| (v, w(n)))).unwrap()
}
fn one_d3_d3() -> Distribution
{
let weights = [9, 12, 16, 12, 12, 10, 6, 3, 1];
distribution(&(1..).zip(weights).collect::<Vec<_>>())
}
fn pow2(exponent: u64) -> BigUint { BigUint::from(1u8) << exponent }
type Ratio = (BigInt, BigUint);
fn exact(x: f64) -> Ratio
{
assert!(x.is_finite());
let bits = x.to_bits();
let biased = ((bits >> 52) & 0x7ff) as i64;
let fraction = bits & ((1 << 52) - 1);
let (significand, exponent) = match biased
{
0 => (fraction, -1074),
_ => (fraction | (1 << 52), biased - 1075)
};
let sign = if x.is_sign_negative()
{
Sign::Minus
}
else
{
Sign::Plus
};
let significand = BigInt::from_biguint(sign, BigUint::from(significand));
match exponent
{
e if e >= 0 => (significand << e as u64, BigUint::from(1u8)),
e => (significand, pow2(-e as u64))
}
}
fn compare(a: &Ratio, b: &Ratio) -> Ordering
{
(&a.0 * BigInt::from(b.1.clone())).cmp(&(&b.0 * BigInt::from(a.1.clone())))
}
fn midpoint(a: &Ratio, b: &Ratio) -> Ratio
{
let numerator =
&a.0 * BigInt::from(b.1.clone()) + &b.0 * BigInt::from(a.1.clone());
(numerator, &a.1 * &b.1 * 2u8)
}
fn check_rounding(answer: f64, q: &Ratio) -> Result<(), TestCaseError>
{
prop_assert!(!answer.is_nan(), "{:?} rounded to NaN", q);
prop_assert!(
answer == 0.0
|| answer.is_sign_negative() == (q.0.sign() == Sign::Minus),
"{} has the wrong sign for {:?}",
answer,
q
);
let q = (BigInt::from(q.0.magnitude().clone()), q.1.clone());
let answer = answer.abs();
let even = answer.to_bits() & 1 == 0;
let within = |bound: &Ratio, side: Ordering| match compare(&q, bound)
{
Ordering::Equal => even,
ordering => ordering == side
};
let overflow = (BigInt::from(pow2(1024) - pow2(970)), BigUint::from(1u8));
if answer.is_infinite()
{
prop_assert!(
compare(&q, &overflow) != Ordering::Less,
"{:?} rounded to infinity",
q
);
return Ok(())
}
let value = exact(answer);
if answer > 0.0
{
let below = midpoint(&exact(answer.next_down()), &value);
prop_assert!(
within(&below, Ordering::Greater),
"{:?} rounded up to {}",
q,
answer
);
}
let above = match answer == f64::MAX
{
true => overflow,
false => midpoint(&value, &exact(answer.next_up()))
};
prop_assert!(
within(&above, Ordering::Less),
"{:?} rounded down to {}",
q,
answer
);
Ok(())
}
fn ratio(numerator: &BigUint, denominator: &BigUint) -> Ratio
{
(BigInt::from(numerator.clone()), denominator.clone())
}
fn signed(negative: bool, numerator: &BigUint, denominator: &BigUint) -> Ratio
{
let sign = if negative { Sign::Minus } else { Sign::Plus };
(
BigInt::from_biguint(sign, numerator.clone()),
denominator.clone()
)
}
fn reference(weights: &[(i32, BigUint)]) -> BTreeMap<i32, BigUint>
{
let mut merged = BTreeMap::<i32, BigUint>::new();
for (outcome, weight) in weights
{
*merged.entry(*outcome).or_default() += weight;
}
merged.retain(|_, weight| *weight != BigUint::ZERO);
merged
}
fn from_big(
weights: &[(i32, BigUint)]
) -> Result<Distribution, EmptyDistributionError>
{
Distribution::from_weights(
weights
.iter()
.map(|(v, n)| (*v, Weight::from_biguint(n.clone())))
)
}
fn outcome() -> impl Strategy<Value = i32>
{
prop_oneof![-8i32..=8, Just(i32::MIN), Just(i32::MAX), any::<i32>()]
}
fn weights() -> impl Strategy<Value = Vec<(i32, BigUint)>>
{
vec((outcome(), value()), 0..=12)
}
fn nonempty_weights() -> impl Strategy<Value = Vec<(i32, BigUint)>>
{
weights().prop_filter("some weight is nonzero", |weights| {
!reference(weights).is_empty()
})
}
fn wide() -> impl Strategy<Value = BigUint>
{
prop_oneof![
value(),
(value(), 0u64..=1200).prop_map(|(v, shift)| v << shift),
(1u64..=8, 1070u64..=1080).prop_map(|(v, shift)| pow2(shift) + v),
(1u64..=8, 1070u64..=1080).prop_map(|(v, shift)| pow2(shift) - v)
]
}
fn probability() -> impl Strategy<Value = (BigUint, BigUint)>
{
(wide(), wide()).prop_map(|(a, b)| match a.cmp(&b)
{
_ if b == BigUint::ZERO && a == BigUint::ZERO =>
{
(BigUint::ZERO, BigUint::from(1u8))
},
Ordering::Greater => (b, a),
_ => (a, b)
})
}
fn rational() -> impl Strategy<Value = (bool, BigUint, BigUint)>
{
(
any::<bool>(),
wide(),
wide().prop_filter("nonzero", |d| *d != BigUint::ZERO)
)
}
#[test]
fn test_probability_construction()
{
assert_eq!(
Probability::new(w(0), w(0)),
Err(ProbabilityError::ZeroDenominator)
);
assert_eq!(
Probability::new(w(3), w(2)),
Err(ProbabilityError::ExceedsOne)
);
let half = p(2, 4);
assert_eq!(half.numerator(), &w(2));
assert_eq!(half.denominator(), &w(4));
assert_eq!(Probability::ZERO.to_string(), "0/1");
assert_eq!(Probability::ONE.to_string(), "1/1");
}
#[test]
fn test_probability_comparison()
{
assert_eq!(p(2, 4), p(1, 2));
assert_eq!(p(2, 4).to_string(), "2/4");
assert_eq!(p(0, 7), Probability::ZERO);
assert_eq!(p(7, 7), Probability::ONE);
assert!(p(1, 3) < p(1, 2));
assert!(p(2, 3) > p(3, 5));
assert_eq!(p(1, 3).cmp(&p(2, 6)), Ordering::Equal);
}
#[test]
fn test_rational_construction()
{
assert_eq!(Rational::new(false, w(1), w(0)), Err(ZeroDenominatorError));
assert_eq!(Rational::new(true, w(0), w(0)), Err(ZeroDenominatorError));
let loss = r(true, 21, 6);
assert!(loss.is_negative());
assert_eq!(loss.numerator(), &w(21));
assert_eq!(loss.denominator(), &w(6));
assert_eq!(loss.to_string(), "-21/6");
assert_eq!(r(false, 21, 6).to_string(), "21/6");
let zero = r(true, 0, 5);
assert!(!zero.is_negative());
assert_eq!(zero.to_string(), "0/5");
assert_eq!(Rational::ZERO.to_string(), "0/1");
assert_eq!(Rational::ONE.to_string(), "1/1");
}
#[test]
fn test_rational_comparison()
{
assert_eq!(r(true, 21, 6), r(true, 7, 2));
assert_ne!(r(true, 7, 2), r(false, 7, 2));
assert_eq!(r(true, 0, 5), Rational::ZERO);
assert_eq!(r(false, 4, 4), Rational::ONE);
assert!(r(true, 1, 2) < Rational::ZERO);
assert!(Rational::ZERO < r(false, 1, 1000));
assert!(r(true, 9, 1) < r(false, 1, 9));
assert!(r(false, 3, 7) < r(false, 4, 7));
assert!(r(true, 3, 7) > r(true, 4, 7));
assert!(r(true, 1, 2) < r(true, 1, 3));
assert_eq!(r(true, 1, 3).cmp(&r(true, 2, 6)), Ordering::Equal);
}
#[test]
fn test_rational_to_f64()
{
assert_eq!(r(true, 21, 6).to_f64(), -3.5);
assert_eq!(r(false, 21, 6).to_f64(), 3.5);
assert_eq!(r(true, 0, 6).to_f64().to_bits(), 0f64.to_bits());
let big = Weight::from_biguint;
let tiny = Rational::new(true, w(1), big(pow2(5000))).unwrap();
assert_eq!(tiny.to_f64().to_bits(), (-0f64).to_bits());
let vast = Rational::new(true, big(pow2(5000)), w(1)).unwrap();
assert_eq!(vast.to_f64(), f64::NEG_INFINITY);
}
#[test]
fn test_ratio_to_f64_edges()
{
let big = Weight::from_biguint;
let ratio = |n: BigUint, d: BigUint| big(n).ratio_to_f64(&big(d));
let one = || BigUint::from(1u8);
let min_subnormal = f64::from_bits(1);
assert_eq!(ratio(one(), pow2(1074)), min_subnormal);
assert_eq!(ratio(one(), pow2(1075)), 0.0);
assert_eq!(ratio(one() + 0u8, pow2(1075) - 1u8), min_subnormal);
assert_eq!(ratio(BigUint::from(3u8), pow2(1076)), min_subnormal);
assert_eq!(ratio(one(), pow2(5000)), 0.0);
assert_eq!(ratio(one(), pow2(1022)), f64::MIN_POSITIVE);
assert_eq!(ratio(pow2(1023), one()), 2f64.powi(1023));
assert_eq!(ratio(pow2(1024) - pow2(970), one()), f64::INFINITY);
assert_eq!(ratio(pow2(1024) - pow2(970) - 1u8, one()), f64::MAX);
assert_eq!(ratio(pow2(3000), pow2(1976)), f64::INFINITY);
let total = BigUint::from(6u8).pow(1000);
assert_eq!(ratio(&total - 1u8, total.clone()), 1.0);
assert_eq!(ratio(total.clone(), &total * 2u8), 0.5);
assert_eq!(ratio(total.clone() * 3u8, total), 3.0);
assert_eq!(w(0).ratio_to_f64(&w(5)), 0.0);
}
#[test]
#[should_panic(expected = "division by zero weight")]
fn test_ratio_to_f64_by_zero() { w(1).ratio_to_f64(&w(0)); }
#[test]
fn test_point()
{
let point = Distribution::point(-4);
assert_eq!(point.len(), 1);
assert_eq!(point.get(-4), &w(1));
assert_eq!(point.get(4), &w(0));
assert_eq!(point.total(), &w(1));
assert_eq!((point.min(), point.max()), (-4, -4));
assert_eq!(point.probability(-4), Probability::ONE);
assert_eq!(point.cdf(-5), Probability::ZERO);
assert_eq!(point.cdf(-4), Probability::ONE);
assert_eq!(point.quantile(&Probability::ZERO), -4);
assert_eq!(point.quantile(&Probability::ONE), -4);
assert_eq!(point.mean().to_string(), "-4/1");
assert_eq!(point.to_f64(), BTreeMap::from([(-4, 1.0)]));
let least = Distribution::point(i32::MIN).mean();
assert_eq!(least.to_string(), "-2147483648/1");
assert_eq!(least.to_f64(), i32::MIN as f64);
let greatest = Distribution::point(i32::MAX).mean();
assert_eq!(greatest.to_string(), "2147483647/1");
assert_eq!(greatest.to_f64(), i32::MAX as f64);
}
#[test]
fn test_from_weights()
{
let merged = distribution(&[(3, 2), (1, 0), (3, 5), (2, 1), (2, 0)]);
assert_eq!(merged.len(), 2);
assert_eq!(merged.iter().collect::<Vec<_>>(), [(2, &w(1)), (3, &w(7))]);
assert_eq!(merged.total(), &w(8));
assert_eq!(merged.get(1), &w(0));
assert_eq!(Distribution::from_weights([]), Err(EmptyDistributionError));
assert_eq!(
Distribution::from_weights([(1, w(0)), (2, w(0))]),
Err(EmptyDistributionError)
);
}
#[test]
fn test_unreduced_equality()
{
let coin = distribution(&[(0, 1), (1, 1)]);
let doubled = distribution(&[(0, 2), (1, 2)]);
assert_ne!(coin, doubled);
assert_eq!(coin.probability(1), doubled.probability(1));
assert_eq!(doubled.probability(1).to_string(), "2/4");
assert_eq!(coin, distribution(&[(1, 1), (0, 1)]));
}
#[test]
fn test_one_d3_d3()
{
let d = one_d3_d3();
assert_eq!(d.total(), &w(81));
assert_eq!((d.min(), d.max(), d.len()), (1, 9, 9));
let cumulative = [0, 9, 21, 37, 49, 61, 71, 77, 80, 81, 81];
for (outcome, expected) in (0..).zip(cumulative)
{
assert_eq!(d.cdf(outcome).to_string(), format!("{expected}/81"));
}
assert_eq!(d.probability(3).to_string(), "16/81");
assert_eq!(d.probability(10).to_string(), "0/81");
assert_eq!(d.quantile(&Probability::ZERO), 1);
assert_eq!(d.quantile(&p(9, 81)), 1);
assert_eq!(d.quantile(&p(10, 81)), 2);
assert_eq!(d.quantile(&p(1, 2)), 4);
assert_eq!(d.quantile(&p(80, 81)), 8);
assert_eq!(d.quantile(&Probability::ONE), 9);
assert_eq!(d.mean().to_string(), "324/81");
assert_eq!(d.mean(), r(false, 4, 1));
assert_eq!(d.mean().to_f64(), 4.0);
assert_eq!(d.to_f64()[&3], 16.0 / 81.0);
}
#[test]
fn test_mean_with_negative_outcomes()
{
let cases = [
(distribution(&[(-3, 1), (1, 1)]), "-2/2", -1.0f64),
(distribution(&[(-3, 1), (3, 1)]), "0/2", 0.0),
(distribution(&[(-1, 1), (0, 2)]), "-1/3", -1.0 / 3.0),
(distribution(&[(i32::MIN, 1), (i32::MAX, 1)]), "-1/2", -0.5)
];
for (d, exact, approximate) in cases
{
assert_eq!(d.mean().to_string(), exact);
assert_eq!(d.mean().to_f64().to_bits(), approximate.to_bits());
}
assert!(!distribution(&[(-3, 1), (3, 1)]).mean().is_negative());
}
#[test]
fn test_to_f64_with_huge_total()
{
let total = BigUint::from(6u8).pow(1000);
let huge = from_big(&[(0, &total - 1u8), (1, BigUint::from(1u8))]).unwrap();
assert_eq!(huge.total().to_f64(), f64::INFINITY);
assert_eq!(huge.to_f64(), BTreeMap::from([(0, 1.0), (1, 0.0)]));
assert_eq!(huge.probability(0).to_f64(), 1.0);
assert_eq!(huge.mean().numerator(), &w(1));
assert_eq!(huge.mean().to_f64(), 0.0);
let even = from_big(&[(2, total.clone()), (4, total)]).unwrap();
assert_eq!(even.to_f64(), BTreeMap::from([(2, 0.5), (4, 0.5)]));
assert_eq!(even.mean(), r(false, 3, 1));
assert_eq!(even.mean().to_f64(), 3.0);
}
#[test]
fn test_iteration()
{
let d = distribution(&[(5, 1), (-2, 3), (0, 2)]);
let forward: Vec<_> = d.iter().collect();
assert_eq!(forward, [(-2, &w(3)), (0, &w(2)), (5, &w(1))]);
let backward: Vec<_> = d.iter().rev().collect();
assert_eq!(backward, [(5, &w(1)), (0, &w(2)), (-2, &w(3))]);
assert_eq!(d.iter().len(), 3);
let borrowed: Vec<_> = (&d).into_iter().map(|(v, _)| v).collect();
assert_eq!(borrowed, [-2, 0, 5]);
let owned: Vec<_> = d.into_iter().collect();
assert_eq!(owned, [(-2, w(3)), (0, w(2)), (5, w(1))]);
}
#[test]
fn test_display()
{
let d = distribution(&[(2, 1), (-1, 3)]);
assert_eq!(d.to_string(), "-1: 3\n2: 1\n");
}
#[cfg(feature = "serde")]
#[test]
fn test_serde()
{
let d = distribution(&[(2, 1), (-1, 3)]);
let json = serde_json::to_string(&d).unwrap();
assert_eq!(json, r#"{"weights":{"-1":"3","2":"1"},"total":"4"}"#);
assert_eq!(serde_json::from_str::<Distribution>(&json).unwrap(), d);
for (invalid, reason) in [
(r#"{"weights":{},"total":"0"}"#, "no outcomes"),
(
r#"{"weights":{"1":"0","2":"1"},"total":"1"}"#,
"outcome 1 as zero"
),
(r#"{"weights":{"1":"2"},"total":"3"}"#, "not the sum"),
(r#"{"weights":{"1":"1"}}"#, "missing field"),
(
r#"{"weights":{"1":"1"},"total":"1","x":1}"#,
"unknown field"
),
(r#"{"weights":{"1":1},"total":"1"}"#, "decimal string")
]
{
let error = serde_json::from_str::<Distribution>(invalid).unwrap_err();
assert!(error.to_string().contains(reason), "{invalid}: {error}");
}
}
proptest! {
#[test]
fn test_from_weights_agrees(weights in weights(), probe in outcome())
{
let expected = reference(&weights);
let d = match from_big(&weights)
{
Ok(d) => d,
Err(EmptyDistributionError) =>
{
prop_assert!(expected.is_empty());
return Ok(())
}
};
let actual: BTreeMap<_, _> =
d.iter().map(|(v, n)| (v, n.to_biguint())).collect();
prop_assert_eq!(&actual, &expected);
prop_assert!(d.iter().all(|(_, n)| !n.is_zero()));
prop_assert_eq!(d.total().to_biguint(), expected.values().sum::<BigUint>());
prop_assert_eq!(d.len(), expected.len());
prop_assert_eq!(d.iter().len(), expected.len());
prop_assert_eq!(d.min(), *expected.keys().next().unwrap());
prop_assert_eq!(d.max(), *expected.keys().next_back().unwrap());
prop_assert_eq!(
d.get(probe).to_biguint(),
expected.get(&probe).cloned().unwrap_or_default()
);
}
#[test]
fn test_probability_and_cdf_agree(
weights in nonempty_weights(),
probe in outcome()
)
{
let expected = reference(&weights);
let total: BigUint = expected.values().sum();
let d = from_big(&weights).unwrap();
let outcomes = expected.keys().copied().chain([probe]);
for outcome in outcomes
{
let probability = d.probability(outcome);
prop_assert_eq!(probability.denominator().to_biguint(), total.clone());
prop_assert_eq!(
probability.numerator().to_biguint(),
expected.get(&outcome).cloned().unwrap_or_default()
);
let cdf = d.cdf(outcome);
prop_assert_eq!(cdf.denominator().to_biguint(), total.clone());
prop_assert_eq!(
cdf.numerator().to_biguint(),
expected.range(..=outcome).map(|(_, n)| n).sum::<BigUint>()
);
}
prop_assert_eq!(d.cdf(d.max()), Probability::ONE);
prop_assert_eq!(d.cdf(i32::MAX), Probability::ONE);
if d.min() > i32::MIN
{
prop_assert_eq!(d.cdf(d.min() - 1), Probability::ZERO);
}
}
#[test]
fn test_quantile_agrees(
weights in nonempty_weights(),
(numerator, denominator) in probability()
)
{
let expected = reference(&weights);
let total: BigUint = expected.values().sum();
let d = from_big(&weights).unwrap();
let p = Probability::new(
Weight::from_biguint(numerator.clone()),
Weight::from_biguint(denominator.clone())
)
.unwrap();
let reaches =
|cumulative: &BigUint| cumulative * &denominator >= &numerator * &total;
let mut cumulative = BigUint::ZERO;
let least = expected
.iter()
.find(|(_, n)| {
cumulative += *n;
reaches(&cumulative)
})
.map(|(v, _)| *v);
prop_assert_eq!(Some(d.quantile(&p)), least);
}
#[test]
fn test_probability_ordering_agrees(
(a, b) in probability(),
(c, e) in probability()
)
{
let x = Probability::new(
Weight::from_biguint(a.clone()),
Weight::from_biguint(b.clone())
)
.unwrap();
let y = Probability::new(
Weight::from_biguint(c.clone()),
Weight::from_biguint(e.clone())
)
.unwrap();
let expected = compare(&ratio(&a, &b), &ratio(&c, &e));
prop_assert_eq!(x.cmp(&y), expected);
prop_assert_eq!(x == y, expected == Ordering::Equal);
prop_assert_eq!(x.partial_cmp(&y), Some(expected));
}
#[test]
fn test_ratio_to_f64_rounds_correctly(
numerator in wide(),
denominator in wide().prop_filter("nonzero", |d| *d != BigUint::ZERO)
)
{
let answer = Weight::from_biguint(numerator.clone())
.ratio_to_f64(&Weight::from_biguint(denominator.clone()));
check_rounding(answer, &ratio(&numerator, &denominator))?;
}
#[test]
fn test_probability_to_f64_rounds_correctly(
(numerator, denominator) in probability()
)
{
let p = Probability::new(
Weight::from_biguint(numerator.clone()),
Weight::from_biguint(denominator.clone())
)
.unwrap();
let answer = p.to_f64();
prop_assert!((0.0..=1.0).contains(&answer));
check_rounding(answer, &ratio(&numerator, &denominator))?;
}
#[test]
fn test_rational_ordering_agrees(
(a, b, c) in rational(),
(d, e, f) in rational()
)
{
let x = Rational::new(
a,
Weight::from_biguint(b.clone()),
Weight::from_biguint(c.clone())
)
.unwrap();
let y = Rational::new(
d,
Weight::from_biguint(e.clone()),
Weight::from_biguint(f.clone())
)
.unwrap();
let expected = compare(&signed(a, &b, &c), &signed(d, &e, &f));
prop_assert_eq!(x.cmp(&y), expected);
prop_assert_eq!(x == y, expected == Ordering::Equal);
prop_assert_eq!(x.partial_cmp(&y), Some(expected));
prop_assert_eq!(x.is_negative(), a && b != BigUint::ZERO);
}
#[test]
fn test_rational_to_f64_rounds_correctly(
(negative, numerator, denominator) in rational()
)
{
let x = Rational::new(
negative,
Weight::from_biguint(numerator.clone()),
Weight::from_biguint(denominator.clone())
)
.unwrap();
check_rounding(x.to_f64(), &signed(negative, &numerator, &denominator))?;
}
#[test]
fn test_mean_and_to_f64_round_correctly(weights in nonempty_weights())
{
let expected = reference(&weights);
let total: BigUint = expected.values().sum();
let d = from_big(&weights).unwrap();
let sum: BigInt = expected
.iter()
.map(|(v, n)| BigInt::from(*v) * BigInt::from(n.clone()))
.sum();
let mean = d.mean();
prop_assert_eq!(mean.is_negative(), sum.sign() == Sign::Minus);
prop_assert_eq!(&mean.numerator().to_biguint(), sum.magnitude());
prop_assert_eq!(mean.denominator().to_biguint(), total.clone());
check_rounding(mean.to_f64(), &(sum, total.clone()))?;
let view = d.to_f64();
prop_assert!(view.keys().eq(expected.keys()));
for (outcome, n) in &expected
{
prop_assert_eq!(view[outcome], d.probability(*outcome).to_f64());
check_rounding(view[outcome], &ratio(n, &total))?;
}
}
#[test]
fn test_display_agrees(weights in nonempty_weights())
{
let expected: String = reference(&weights)
.iter()
.map(|(v, n)| format!("{v}: {n}\n"))
.collect();
prop_assert_eq!(from_big(&weights).unwrap().to_string(), expected);
}
#[cfg(feature = "serde")]
#[test]
fn test_serde_round_trips(weights in nonempty_weights())
{
let d = from_big(&weights).unwrap();
let json = serde_json::to_string(&d).unwrap();
prop_assert_eq!(serde_json::from_str::<Distribution>(&json).unwrap(), d);
}
}