use std::{
error::Error,
fmt::{self, Display, Formatter},
iter::{Product, Sum},
ops::{Add, AddAssign, Mul, MulAssign},
str::FromStr
};
use num_bigint::BigUint;
use num_traits::{CheckedSub, ToPrimitive};
#[cfg(feature = "serde")]
use serde::{
Deserialize, Deserializer, Serialize, Serializer,
de::{self, Visitor}
};
#[derive(Debug, Clone, Default, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub struct Weight(Repr);
#[derive(Debug, Clone, PartialEq, Eq, Hash, PartialOrd, Ord)]
enum Repr
{
Small(u128),
Big(BigUint)
}
impl Default for Repr
{
fn default() -> Self { Repr::Small(0) }
}
impl Weight
{
pub const ZERO: Self = Self(Repr::Small(0));
pub const ONE: Self = Self(Repr::Small(1));
fn from_big(big: BigUint) -> Self
{
match u128::try_from(&big)
{
Ok(small) => Self(Repr::Small(small)),
Err(_) => Self(Repr::Big(big))
}
}
pub(crate) fn to_biguint(&self) -> BigUint
{
match &self.0
{
Repr::Small(n) => BigUint::from(*n),
Repr::Big(n) => n.clone()
}
}
pub(crate) fn write_u32_digits(&self, digits: &mut [u32])
{
debug_assert!(
digits.len() as u64 * 32 >= self.bits(),
"{} bits exceed {} digits",
self.bits(),
digits.len()
);
match &self.0
{
Repr::Small(n) =>
{
for (i, digit) in digits.iter_mut().take(4).enumerate()
{
*digit = (n >> (32 * i)) as u32;
}
},
Repr::Big(n) =>
{
for (digit, d) in digits.iter_mut().zip(n.iter_u32_digits())
{
*digit = d;
}
}
}
}
pub(crate) fn to_u32_digits(&self) -> Vec<u32>
{
match &self.0
{
Repr::Small(n) =>
{
let len = (self.bits() as usize).div_ceil(32);
(0..len).map(|i| (n >> (32 * i)) as u32).collect()
},
Repr::Big(n) => n.to_u32_digits()
}
}
pub(crate) fn from_u32_digits(digits: &[u32]) -> Self
{
let len = digits.iter().rposition(|d| *d != 0).map_or(0, |i| i + 1);
if len <= 4
{
let n = digits[..len]
.iter()
.rev()
.fold(0u128, |n, d| (n << 32) | *d as u128);
Self(Repr::Small(n))
}
else
{
Self(Repr::Big(BigUint::from_slice(&digits[..len])))
}
}
#[inline]
pub fn is_zero(&self) -> bool { matches!(self.0, Repr::Small(0)) }
pub fn bits(&self) -> u64
{
match &self.0
{
Repr::Small(n) => (u128::BITS - n.leading_zeros()) as u64,
Repr::Big(n) => n.bits()
}
}
pub fn to_f64(&self) -> f64
{
match &self.0
{
Repr::Small(n) => *n as f64,
Repr::Big(n) => n.to_f64().unwrap_or(f64::INFINITY)
}
}
pub(crate) fn ratio_to_f64(&self, denominator: &Weight) -> f64
{
assert!(!denominator.is_zero(), "division by zero weight");
if self.is_zero()
{
return 0.0
}
if self.bits() <= f64::MANTISSA_DIGITS as u64
&& denominator.bits() <= f64::MANTISSA_DIGITS as u64
{
return self.to_f64() / denominator.to_f64()
}
let e = self.bits() as i64 - denominator.bits() as i64;
if e >= 1025
{
return f64::INFINITY
}
if e <= -1076
{
return 0.0
}
let s = (56 - e).min(1076);
let numerator = self.to_biguint();
let denominator = denominator.to_biguint();
let (numerator, denominator) = if s >= 0
{
(numerator << s as u64, denominator)
}
else
{
(numerator, denominator << -s as u64)
};
let quotient = &numerator / &denominator;
let remainder = numerator % denominator;
let scaled = quotient.to_u64().expect("scaled quotient fits in u64")
| u64::from(remainder != BigUint::ZERO);
let leading = (u64::BITS - scaled.leading_zeros()) as i64 - 1 - s;
let unit = (leading - 52).max(-1074);
let drop = (unit + s) as u32;
let half = 1u64 << (drop - 1);
let mut rounded = scaled >> drop;
let rest = scaled & ((1u64 << drop) - 1);
if rest > half || (rest == half && rounded & 1 == 1)
{
rounded += 1;
}
if unit + (u64::BITS - rounded.leading_zeros()) as i64 > 1024
{
return f64::INFINITY
}
rounded as f64 * pow2(unit)
}
}
fn pow2(exponent: i64) -> f64
{
debug_assert!((-1074..=1023).contains(&exponent));
if exponent >= -1022
{
f64::from_bits(((exponent + 1023) as u64) << 52)
}
else
{
f64::from_bits(1 << (exponent + 1074))
}
}
macro_rules! weight_from_unsigned {
($($t:ty),*) => {
$(
impl From<$t> for Weight
{
fn from(n: $t) -> Self { Self(Repr::Small(n as u128)) }
}
)*
};
}
weight_from_unsigned!(u8, u16, u32, u64, u128, usize);
#[cfg(test)]
impl Weight
{
pub(crate) fn from_biguint(big: BigUint) -> Self { Self::from_big(big) }
pub(crate) fn is_small(&self) -> bool { matches!(self.0, Repr::Small(_)) }
}
impl Weight
{
pub fn pow(&self, exponent: u32) -> Self
{
match &self.0
{
Repr::Small(base) => match base.checked_pow(exponent)
{
Some(power) => Self(Repr::Small(power)),
None => Self(Repr::Big(BigUint::from(*base).pow(exponent)))
},
Repr::Big(_) if exponent == 0 => Self::ONE,
Repr::Big(base) => Self(Repr::Big(base.pow(exponent)))
}
}
pub fn checked_sub(&self, rhs: &Self) -> Option<Self>
{
match (&self.0, &rhs.0)
{
(Repr::Small(a), Repr::Small(b)) =>
{
u128::checked_sub(*a, *b).map(|d| Self(Repr::Small(d)))
},
(Repr::Small(_), Repr::Big(_)) => None,
(Repr::Big(a), Repr::Small(b)) => Some(Self::from_big(a - *b)),
(Repr::Big(a), Repr::Big(b)) => a.checked_sub(b).map(Self::from_big)
}
}
}
impl Weight
{
pub(crate) fn gcd(&self, other: &Self) -> Self
{
match (&self.0, &other.0)
{
(Repr::Small(a), Repr::Small(b)) =>
{
let (mut a, mut b) = (*a, *b);
while b != 0
{
(a, b) = (b, a % b);
}
Self(Repr::Small(a))
},
_ =>
{
let (mut a, mut b) = (self.to_biguint(), other.to_biguint());
while b != BigUint::ZERO
{
let r = &a % &b;
(a, b) = (b, r);
}
Self::from_big(a)
}
}
}
pub(crate) fn lcm(&self, other: &Self) -> Self
{
if self.is_zero() || other.is_zero()
{
return Self::ZERO
}
self.exact_div(&self.gcd(other)) * other
}
pub(crate) fn exact_div(&self, divisor: &Self) -> Self
{
match (&self.0, &divisor.0)
{
(Repr::Small(a), Repr::Small(b)) =>
{
debug_assert!(a % b == 0, "{b} does not divide {a}");
Self(Repr::Small(a / b))
},
(Repr::Small(a), Repr::Big(b)) =>
{
debug_assert!(*a == 0, "{b} does not divide {a}");
Self::ZERO
},
_ =>
{
let (a, b) = (self.to_biguint(), divisor.to_biguint());
debug_assert!(
&a % &b == BigUint::ZERO,
"{b} does not divide {a}"
);
Self::from_big(a / b)
}
}
}
}
impl AddAssign<&Weight> for Weight
{
fn add_assign(&mut self, rhs: &Weight)
{
match (&mut self.0, &rhs.0)
{
(Repr::Small(a), Repr::Small(b)) => match a.checked_add(*b)
{
Some(sum) => *a = sum,
None => self.0 = Repr::Big(BigUint::from(*a) + *b)
},
(Repr::Small(a), Repr::Big(b)) => self.0 = Repr::Big(b + *a),
(Repr::Big(a), Repr::Small(b)) => *a += *b,
(Repr::Big(a), Repr::Big(b)) => *a += b
}
}
}
impl AddAssign for Weight
{
fn add_assign(&mut self, rhs: Weight)
{
match (&self.0, rhs.0)
{
(Repr::Small(a), Repr::Big(mut b)) =>
{
b += *a;
self.0 = Repr::Big(b);
},
(_, rhs) => *self += &Weight(rhs)
}
}
}
impl MulAssign<&Weight> for Weight
{
fn mul_assign(&mut self, rhs: &Weight)
{
match (&mut self.0, &rhs.0)
{
(Repr::Small(a), Repr::Small(b)) => match a.checked_mul(*b)
{
Some(product) => *a = product,
None => self.0 = Repr::Big(BigUint::from(*a) * *b)
},
(Repr::Small(0), Repr::Big(_)) =>
{},
(Repr::Big(_), Repr::Small(0)) => self.0 = Repr::Small(0),
(Repr::Small(a), Repr::Big(b)) => self.0 = Repr::Big(b * *a),
(Repr::Big(a), Repr::Small(b)) => *a *= *b,
(Repr::Big(a), Repr::Big(b)) => *a *= b
}
}
}
impl MulAssign for Weight
{
fn mul_assign(&mut self, rhs: Weight)
{
match (&self.0, rhs.0)
{
(Repr::Small(a), Repr::Big(mut b)) if *a != 0 =>
{
b *= *a;
self.0 = Repr::Big(b);
},
(_, rhs) => *self *= &Weight(rhs)
}
}
}
macro_rules! weight_binop {
($op:ident, $method:ident, $assign:ident) => {
impl $op for Weight
{
type Output = Weight;
fn $method(mut self, rhs: Weight) -> Weight
{
self.$assign(rhs);
self
}
}
impl $op<&Weight> for Weight
{
type Output = Weight;
fn $method(mut self, rhs: &Weight) -> Weight
{
self.$assign(rhs);
self
}
}
impl $op<Weight> for &Weight
{
type Output = Weight;
fn $method(self, mut rhs: Weight) -> Weight
{
rhs.$assign(self);
rhs
}
}
};
}
weight_binop!(Add, add, add_assign);
weight_binop!(Mul, mul, mul_assign);
impl Add<&Weight> for &Weight
{
type Output = Weight;
fn add(self, rhs: &Weight) -> Weight
{
let mut sum = self.clone();
sum += rhs;
sum
}
}
impl Mul<&Weight> for &Weight
{
type Output = Weight;
fn mul(self, rhs: &Weight) -> Weight
{
match (&self.0, &rhs.0)
{
(Repr::Small(a), Repr::Small(b)) => match a.checked_mul(*b)
{
Some(product) => Weight(Repr::Small(product)),
None => Weight(Repr::Big(BigUint::from(*a) * *b))
},
(Repr::Small(0), Repr::Big(_)) | (Repr::Big(_), Repr::Small(0)) =>
{
Weight::ZERO
},
(Repr::Small(a), Repr::Big(b)) => Weight(Repr::Big(b * *a)),
(Repr::Big(a), Repr::Small(b)) => Weight(Repr::Big(a * *b)),
(Repr::Big(a), Repr::Big(b)) => Weight(Repr::Big(a * b))
}
}
}
impl Sum for Weight
{
fn sum<I: Iterator<Item = Weight>>(iter: I) -> Self
{
iter.fold(Weight::ZERO, |sum, weight| sum + weight)
}
}
impl<'a> Sum<&'a Weight> for Weight
{
fn sum<I: Iterator<Item = &'a Weight>>(iter: I) -> Self
{
iter.fold(Weight::ZERO, |sum, weight| sum + weight)
}
}
impl Product for Weight
{
fn product<I: Iterator<Item = Weight>>(iter: I) -> Self
{
iter.fold(Weight::ONE, |product, weight| product * weight)
}
}
impl<'a> Product<&'a Weight> for Weight
{
fn product<I: Iterator<Item = &'a Weight>>(iter: I) -> Self
{
iter.fold(Weight::ONE, |product, weight| product * weight)
}
}
impl Display for Weight
{
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result
{
match &self.0
{
Repr::Small(n) => Display::fmt(n, f),
Repr::Big(n) => Display::fmt(n, f)
}
}
}
impl FromStr for Weight
{
type Err = ParseWeightError;
fn from_str(s: &str) -> Result<Self, Self::Err>
{
let digits = s.strip_prefix('+').unwrap_or(s);
if digits.is_empty()
{
return Err(ParseWeightError::Empty)
}
if !digits.bytes().all(|b| b.is_ascii_digit())
{
return Err(ParseWeightError::InvalidDigit)
}
match digits.parse::<u128>()
{
Ok(n) => Ok(Self(Repr::Small(n))),
Err(_) => Ok(Self(Repr::Big(
digits.parse().expect("decimal digits parse as big")
)))
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ParseWeightError
{
Empty,
InvalidDigit
}
impl Display for ParseWeightError
{
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result
{
match self
{
Self::Empty => write!(f, "cannot parse weight from empty string"),
Self::InvalidDigit => write!(f, "invalid digit found in string")
}
}
}
impl Error for ParseWeightError {}
#[cfg(feature = "serde")]
impl Serialize for Weight
{
fn serialize<S: Serializer>(&self, serializer: S)
-> Result<S::Ok, S::Error>
{
serializer.collect_str(self)
}
}
#[cfg(feature = "serde")]
impl<'de> Deserialize<'de> for Weight
{
fn deserialize<D: Deserializer<'de>>(
deserializer: D
) -> Result<Self, D::Error>
{
deserializer.deserialize_str(WeightVisitor)
}
}
#[cfg(feature = "serde")]
struct WeightVisitor;
#[cfg(feature = "serde")]
impl Visitor<'_> for WeightVisitor
{
type Value = Weight;
fn expecting(&self, f: &mut Formatter<'_>) -> fmt::Result
{
f.write_str("a nonnegative integer as a decimal string")
}
fn visit_str<E: de::Error>(self, v: &str) -> Result<Weight, E>
{
v.parse()
.map_err(|e| E::custom(format_args!("invalid weight {v:?}: {e}")))
}
}