use core::{cmp::Ordering, fmt::Display};
use thiserror::Error;
use crate::common::Fraction;
impl Display for Fraction {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
write!(f, "{}/{}", self.numerator, self.denominator)
}
}
#[derive(Debug, Error, PartialEq, Eq, Clone)]
#[non_exhaustive]
pub enum FractionError {
#[error("Denominator cannot be zero")]
ZeroDenominator,
#[error("Fraction arithmetic operation resulted in an overflow")]
Overflow,
#[error("Fraction arithmetic operation resulted in an undefined state")]
Undefined,
}
impl Fraction {
#[must_use]
#[inline]
pub const fn gcd(a: i64, b: i64) -> u64 {
let mut ua = a.unsigned_abs();
let mut ub = b.unsigned_abs();
while ub != 0 {
let temp = ub;
ub = ua % ub;
ua = temp;
}
ua
}
#[inline]
pub fn lcm(a: i64, b: i64) -> Result<i128, FractionError> {
if a == 0 || b == 0 {
return Err(FractionError::ZeroDenominator);
}
let common_divisor = i128::from(Self::gcd(a, b));
let val_a = i128::from(a);
let val_b = i128::from(b);
let term1 = val_a
.checked_div(common_divisor)
.ok_or(FractionError::Overflow)?;
term1
.checked_mul(val_b)
.ok_or(FractionError::Overflow)
}
#[inline]
pub const fn new(numerator: i64, denominator: i64) -> Result<Self, FractionError> {
if denominator == 0 {
return Err(FractionError::ZeroDenominator);
}
if denominator == i64::MIN {
return Err(FractionError::Overflow);
}
if denominator < 0 && numerator == i64::MIN {
return Err(FractionError::Overflow);
}
let (mut num, mut den) = (numerator, denominator);
if den < 0 {
num = -num;
den = -den;
}
let common_divisor = Self::gcd(num, den);
let cd = common_divisor.cast_signed();
Ok(Self {
numerator: num / cd,
denominator: den / cd,
})
}
#[inline]
pub const fn reduce(&mut self) {
if self.denominator == 0 {
return;
}
if self.denominator == i64::MIN {
return;
}
if self.denominator < 0 && self.numerator == i64::MIN {
return;
}
if self.denominator < 0 {
self.numerator = -self.numerator;
self.denominator = -self.denominator;
}
let common_divisor = Self::gcd(self.numerator, self.denominator);
let cd = common_divisor.cast_signed();
self.numerator /= cd;
self.denominator /= cd;
}
#[must_use]
#[inline]
pub const fn reduced(mut self) -> Self {
self.reduce();
self
}
#[inline]
pub fn checked_add(self, other: Self) -> Result<Self, FractionError> {
let common_denominator_i128 = Self::lcm(self.denominator, other.denominator)?;
let factor_self = common_denominator_i128
.checked_div(i128::from(self.denominator))
.ok_or(FractionError::Overflow)?;
let factor_other = common_denominator_i128
.checked_div(i128::from(other.denominator))
.ok_or(FractionError::Overflow)?;
let new_numerator_left = i128::from(self.numerator)
.checked_mul(factor_self)
.ok_or(FractionError::Overflow)?;
let new_numerator_right = i128::from(other.numerator)
.checked_mul(factor_other)
.ok_or(FractionError::Overflow)?;
let new_numerator = new_numerator_left
.checked_add(new_numerator_right)
.ok_or(FractionError::Overflow)?;
let num_i64 = i64::try_from(new_numerator).map_err(|_| FractionError::Overflow)?;
let den_i64 =
i64::try_from(common_denominator_i128).map_err(|_| FractionError::Overflow)?;
Self::new(num_i64, den_i64)
}
#[inline]
pub fn checked_sub(self, other: Self) -> Result<Self, FractionError> {
let common_denominator_i128 = Self::lcm(self.denominator, other.denominator)?;
let factor_self = common_denominator_i128
.checked_div(i128::from(self.denominator))
.ok_or(FractionError::Overflow)?;
let factor_other = common_denominator_i128
.checked_div(i128::from(other.denominator))
.ok_or(FractionError::Overflow)?;
let new_numerator_left = i128::from(self.numerator)
.checked_mul(factor_self)
.ok_or(FractionError::Overflow)?;
let new_numerator_right = i128::from(other.numerator)
.checked_mul(factor_other)
.ok_or(FractionError::Overflow)?;
let new_numerator = new_numerator_left
.checked_sub(new_numerator_right)
.ok_or(FractionError::Overflow)?;
let num_i64 = i64::try_from(new_numerator).map_err(|_| FractionError::Overflow)?;
let den_i64 =
i64::try_from(common_denominator_i128).map_err(|_| FractionError::Overflow)?;
Self::new(num_i64, den_i64)
}
#[inline]
pub fn checked_mul(self, other: Self) -> Result<Self, FractionError> {
let new_numerator = i128::from(self.numerator)
.checked_mul(i128::from(other.numerator))
.ok_or(FractionError::Overflow)?;
let new_denominator = i128::from(self.denominator)
.checked_mul(i128::from(other.denominator))
.ok_or(FractionError::Overflow)?;
let num_i64 = i64::try_from(new_numerator).map_err(|_| FractionError::Overflow)?;
let den_i64 = i64::try_from(new_denominator).map_err(|_| FractionError::Overflow)?;
Self::new(num_i64, den_i64)
}
#[inline]
pub fn checked_div(self, other: Self) -> Result<Self, FractionError> {
if other.numerator == 0 {
return Err(FractionError::Undefined);
}
let new_numerator = i128::from(self.numerator)
.checked_mul(i128::from(other.denominator))
.ok_or(FractionError::Overflow)?;
let new_denominator = i128::from(self.denominator)
.checked_mul(i128::from(other.numerator))
.ok_or(FractionError::Overflow)?;
let num_i64 = i64::try_from(new_numerator).map_err(|_| FractionError::Overflow)?;
let den_i64 = i64::try_from(new_denominator).map_err(|_| FractionError::Overflow)?;
Self::new(num_i64, den_i64)
}
#[must_use]
#[inline]
pub fn to_f64_unchecked(self) -> f64 {
self.try_into().unwrap()
}
}
impl PartialOrd for Fraction {
#[inline]
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
if self.denominator <= 0 || other.denominator <= 0 {
return None;
}
let self_val = i128::from(self.numerator) * i128::from(other.denominator);
let other_val = i128::from(other.numerator) * i128::from(self.denominator);
Some(self_val.cmp(&other_val))
}
}
impl TryFrom<Fraction> for f64 {
type Error = FractionError;
#[inline]
fn try_from(fraction: Fraction) -> Result<Self, Self::Error> {
if fraction.denominator == 0 {
return Err(FractionError::ZeroDenominator);
}
let num_f64 = fraction.numerator as Self;
let den_f64 = fraction.denominator as Self;
Ok(num_f64 / den_f64)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn frac(n: i64, d: i64) -> Result<Fraction, FractionError> {
Fraction::new(n, d)
}
#[test]
fn test_creation_and_reduction() {
let f = frac(2, 4).unwrap();
assert_eq!(f.numerator, 1);
assert_eq!(f.denominator, 2);
let f = frac(1, -2).unwrap();
assert_eq!(f.numerator, -1);
assert_eq!(f.denominator, 2);
let f = frac(-2, -4).unwrap();
assert_eq!(f.numerator, 1);
assert_eq!(f.denominator, 2);
let f = frac(0, 5).unwrap();
assert_eq!(f.numerator, 0);
assert_eq!(f.denominator, 1); }
#[test]
fn test_creation_edge_cases() {
assert_eq!(frac(1, 0), Err(FractionError::ZeroDenominator));
let f = frac(i64::MIN, 1).unwrap();
assert_eq!(f.numerator, i64::MIN);
assert_eq!(f.denominator, 1);
assert_eq!(frac(1, i64::MIN), Err(FractionError::Overflow));
assert_eq!(frac(i64::MIN, -1), Err(FractionError::Overflow));
}
#[test]
fn test_arithmetic_add() {
let f1 = frac(1, 2).unwrap();
let f2 = frac(1, 3).unwrap();
let res = f1.checked_add(f2).unwrap();
assert_eq!(res.numerator, 5);
assert_eq!(res.denominator, 6);
let f1 = frac(1, 2).unwrap();
let f2 = frac(1, -2).unwrap();
let res = f1.checked_add(f2).unwrap();
assert_eq!(res.numerator, 0);
assert_eq!(res.denominator, 1);
}
#[test]
fn test_arithmetic_sub() {
let f1 = frac(1, 2).unwrap();
let f2 = frac(1, 3).unwrap();
let res = f1.checked_sub(f2).unwrap();
assert_eq!(res.numerator, 1);
assert_eq!(res.denominator, 6);
}
#[test]
fn test_arithmetic_mul() {
let f1 = frac(2, 3).unwrap();
let f2 = frac(3, 4).unwrap();
let res = f1.checked_mul(f2).unwrap();
assert_eq!(res.numerator, 1);
assert_eq!(res.denominator, 2);
}
#[test]
fn test_arithmetic_div() {
let f1 = frac(1, 2).unwrap();
let f2 = frac(1, 2).unwrap();
let res = f1.checked_div(f2).unwrap();
assert_eq!(res.numerator, 1);
assert_eq!(res.denominator, 1);
let f1 = frac(1, 2).unwrap();
let f2 = frac(0, 1).unwrap();
assert_eq!(f1.checked_div(f2), Err(FractionError::Undefined));
}
#[test]
fn test_ordering() {
let f1 = frac(1, 2).unwrap();
let f2 = frac(1, 3).unwrap();
let f3 = frac(2, 4).unwrap();
assert!(f1 > f2); assert!(f2 < f1);
assert_eq!(f1.partial_cmp(&f3), Some(Ordering::Equal));
let neg = frac(-1, 2).unwrap();
assert!(neg < f1);
}
#[test]
fn test_f64_conversion() {
let f = frac(1, 2).unwrap();
let val: f64 = f.try_into().unwrap();
assert!((val - 0.5).abs() < f64::EPSILON);
let bad_frac = Fraction {
numerator: 1,
denominator: 0,
};
assert_eq!(f64::try_from(bad_frac), Err(FractionError::ZeroDenominator));
}
#[test]
fn test_overflow_checks() {
let f1 = frac(i64::MAX - 1, 1).unwrap();
let f2 = frac(2, 1).unwrap();
assert_eq!(f1.checked_add(f2), Err(FractionError::Overflow));
}
#[test]
fn test_gcd_edge_cases() {
assert_eq!(Fraction::gcd(10, 5), 5);
assert_eq!(Fraction::gcd(-10, 5), 5);
assert_eq!(Fraction::gcd(i64::MIN, 0), 9_223_372_036_854_775_808);
}
#[test]
fn test_manual_corruption_resilience() {
let mut f = Fraction {
numerator: 1,
denominator: i64::MIN,
};
f.reduce();
assert_eq!(f.denominator, i64::MIN);
let mut f2 = Fraction {
numerator: i64::MIN,
denominator: -1,
};
f2.reduce();
assert_eq!(f2.denominator, -1);
}
#[test]
fn test_reduce_works_normally() {
let mut f = Fraction {
numerator: 2,
denominator: 4,
};
f.reduce();
assert_eq!(f.numerator, 1);
assert_eq!(f.denominator, 2);
let mut f2 = Fraction {
numerator: 2,
denominator: -4,
};
f2.reduce();
assert_eq!(f2.numerator, -1);
assert_eq!(f2.denominator, 2);
}
}