use std::cmp::Ordering;
use std::fmt;
use std::ops::{Add, Div, Mul, Neg, Sub};
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub struct Rational {
pub num: i64,
pub den: i64,
}
impl Rational {
pub const fn new(num: i64, den: i64) -> Self {
Self { num, den }
}
pub const fn zero() -> Self {
Self { num: 0, den: 1 }
}
pub fn is_zero(&self) -> bool {
self.num == 0
}
pub fn as_f64(&self) -> f64 {
self.num as f64 / self.den as f64
}
pub fn reduced(self) -> Self {
reduce_i128(self.num as i128, self.den as i128)
}
pub fn invert(self) -> Self {
Self {
num: self.den,
den: self.num,
}
}
pub fn cmp_value(&self, other: &Self) -> Ordering {
let (an, ad) = sign_normalized(self.num, self.den);
let (bn, bd) = sign_normalized(other.num, other.den);
if ad != 0 && bd != 0 {
return (an * bd).cmp(&(bn * ad));
}
fn rank(n: i128, d: i128) -> i128 {
if d == 0 {
n.signum() * 2
} else {
n.signum()
}
}
rank(an, ad).cmp(&rank(bn, bd))
}
pub fn equals_value(&self, other: &Self) -> bool {
self.cmp_value(other) == Ordering::Equal
}
pub fn signum(&self) -> i64 {
let (n, _) = sign_normalized(self.num, self.den);
n.signum() as i64
}
pub fn abs(self) -> Self {
let (n, d) = sign_normalized(self.num, self.den);
let n = n.abs();
if n <= i64::MAX as i128 && d <= i64::MAX as i128 {
return Self {
num: n as i64,
den: d as i64,
};
}
reduce_i128(n, d)
}
pub fn checked_add(self, rhs: Self) -> Option<Self> {
let a = self.num as i128 * rhs.den as i128;
let b = rhs.num as i128 * self.den as i128;
let den = self.den as i128 * rhs.den as i128;
reduce_exact_i128(a.checked_add(b)?, den)
}
pub fn checked_sub(self, rhs: Self) -> Option<Self> {
let a = self.num as i128 * rhs.den as i128;
let b = rhs.num as i128 * self.den as i128;
let den = self.den as i128 * rhs.den as i128;
reduce_exact_i128(a.checked_sub(b)?, den)
}
pub fn checked_mul(self, rhs: Self) -> Option<Self> {
let num = self.num as i128 * rhs.num as i128;
let den = self.den as i128 * rhs.den as i128;
reduce_exact_i128(num, den)
}
pub fn checked_div(self, rhs: Self) -> Option<Self> {
let num = self.num as i128 * rhs.den as i128;
let den = self.den as i128 * rhs.num as i128;
reduce_exact_i128(num, den)
}
}
#[inline]
const fn sign_normalized(num: i64, den: i64) -> (i128, i128) {
let (n, d) = (num as i128, den as i128);
if d < 0 {
(-n, -d)
} else {
(n, d)
}
}
fn reduce_exact_i128(mut num: i128, mut den: i128) -> Option<Rational> {
if den < 0 {
num = -num;
den = -den;
}
let g = gcd_i128(num.unsigned_abs(), den.unsigned_abs()) as i128;
if g > 1 {
num /= g;
den /= g;
}
Some(Rational {
num: i64::try_from(num).ok()?,
den: i64::try_from(den).ok()?,
})
}
fn reduce_i128(mut num: i128, mut den: i128) -> Rational {
if den < 0 {
num = -num;
den = -den;
}
let g = gcd_i128(num.unsigned_abs(), den.unsigned_abs()) as i128;
if g > 1 {
num /= g;
den /= g;
}
if let (Ok(n), Ok(d)) = (i64::try_from(num), i64::try_from(den)) {
return Rational { num: n, den: d };
}
approx_narrow(num, den)
}
fn approx_narrow(num: i128, den: i128) -> Rational {
fn sat(v: i128) -> i64 {
if v > i64::MAX as i128 {
i64::MAX
} else if v < i64::MIN as i128 {
i64::MIN
} else {
v as i64
}
}
if den <= i64::MAX as i128 {
return Rational {
num: sat(num),
den: den as i64,
};
}
if num.unsigned_abs() <= i64::MAX as u128 {
let n = (num * i64::MAX as i128 + num.signum() * (den / 2)) / den;
return Rational {
num: n as i64,
den: i64::MAX,
};
}
let dbits = 128 - den.leading_zeros();
let k = dbits - 63;
let half = 1u128 << (k - 1);
let n_abs = (num.unsigned_abs() + half) >> k;
let d = ((den as u128 + half) >> k).max(1);
let n = if num < 0 {
-(n_abs as i128)
} else {
n_abs as i128
};
Rational {
num: sat(n),
den: sat(d as i128),
}
}
impl Add for Rational {
type Output = Rational;
fn add(self, rhs: Self) -> Self {
let a = self.num as i128 * rhs.den as i128;
let b = rhs.num as i128 * self.den as i128;
let den = self.den as i128 * rhs.den as i128;
reduce_i128(a.saturating_add(b), den)
}
}
impl Sub for Rational {
type Output = Rational;
fn sub(self, rhs: Self) -> Self {
let a = self.num as i128 * rhs.den as i128;
let b = rhs.num as i128 * self.den as i128;
let den = self.den as i128 * rhs.den as i128;
reduce_i128(a.saturating_sub(b), den)
}
}
impl Mul for Rational {
type Output = Rational;
fn mul(self, rhs: Self) -> Self {
let num = self.num as i128 * rhs.num as i128;
let den = self.den as i128 * rhs.den as i128;
reduce_i128(num, den)
}
}
impl Div for Rational {
type Output = Rational;
fn div(self, rhs: Self) -> Self {
let num = self.num as i128 * rhs.den as i128;
let den = self.den as i128 * rhs.num as i128;
reduce_i128(num, den)
}
}
impl Neg for Rational {
type Output = Rational;
fn neg(self) -> Self {
if self.num != i64::MIN {
Self {
num: -self.num,
den: self.den,
}
} else if self.den == 0 {
Self {
num: i64::MAX,
den: 0,
}
} else if self.den != i64::MIN {
Self {
num: self.num,
den: -self.den,
}
} else {
Self { num: -1, den: 1 }
}
}
}
impl fmt::Display for Rational {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}/{}", self.num, self.den)
}
}
fn gcd_i128(mut a: u128, mut b: u128) -> u128 {
while b != 0 {
let t = b;
b = a % b;
a = t;
}
a.max(1)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn reduce() {
assert_eq!(Rational::new(10, 20).reduced(), Rational::new(1, 2));
assert_eq!(Rational::new(-6, 9).reduced(), Rational::new(-2, 3));
assert_eq!(Rational::new(6, -9).reduced(), Rational::new(-2, 3));
}
#[test]
fn invert() {
assert_eq!(Rational::new(1, 2).invert(), Rational::new(2, 1));
}
#[test]
fn cmp_value_orders_by_number_not_fields() {
assert!(Rational::new(1, 2).equals_value(&Rational::new(2, 4)));
assert_ne!(Rational::new(1, 2), Rational::new(2, 4));
assert_eq!(
Rational::new(1, 3).cmp_value(&Rational::new(1, 2)),
std::cmp::Ordering::Less
);
assert_eq!(
Rational::new(30_000, 1001).cmp_value(&Rational::new(30, 1)),
std::cmp::Ordering::Less
);
assert!(Rational::new(1, -2).equals_value(&Rational::new(-1, 2)));
assert_eq!(
Rational::new(-1, 2).cmp_value(&Rational::new(1, 2)),
std::cmp::Ordering::Less
);
}
#[test]
fn cmp_value_zero_denominator_is_total() {
let pos_inf = Rational::new(1, 0);
let neg_inf = Rational::new(-1, 0);
let finite = Rational::new(1_000_000, 1);
assert_eq!(pos_inf.cmp_value(&finite), std::cmp::Ordering::Greater);
assert_eq!(finite.cmp_value(&pos_inf), std::cmp::Ordering::Less);
assert_eq!(neg_inf.cmp_value(&finite), std::cmp::Ordering::Less);
assert_eq!(pos_inf.cmp_value(&neg_inf), std::cmp::Ordering::Greater);
assert!(pos_inf.equals_value(&Rational::new(7, 0)));
}
#[test]
fn signum_and_abs() {
assert_eq!(Rational::new(3, 4).signum(), 1);
assert_eq!(Rational::new(-3, 4).signum(), -1);
assert_eq!(Rational::new(3, -4).signum(), -1);
assert_eq!(Rational::new(0, 4).signum(), 0);
assert_eq!(Rational::new(-3, 4).abs(), Rational::new(3, 4));
assert_eq!(Rational::new(3, -4).abs(), Rational::new(3, 4));
}
#[test]
fn arithmetic_reduces() {
assert_eq!(
Rational::new(1, 2) + Rational::new(1, 3),
Rational::new(5, 6)
);
assert_eq!(
Rational::new(1, 2) - Rational::new(1, 3),
Rational::new(1, 6)
);
assert_eq!(
Rational::new(2, 4) * Rational::new(3, 9),
Rational::new(1, 6)
);
assert_eq!(
Rational::new(1, 2) / Rational::new(3, 4),
Rational::new(2, 3)
);
assert_eq!(-Rational::new(1, 2), Rational::new(-1, 2));
}
#[test]
fn reduced_handles_i64_min_terms() {
assert_eq!(
Rational::new(i64::MIN, i64::MIN).reduced(),
Rational::new(1, 1)
);
assert_eq!(
Rational::new(i64::MIN, -2).reduced(),
Rational::new(1 << 62, 1)
);
assert_eq!(
Rational::new(i64::MIN, 2).reduced(),
Rational::new(-(1 << 62), 1)
);
assert_eq!(
Rational::new(-2, i64::MIN).reduced(),
Rational::new(1, 1 << 62)
);
assert_eq!(
Rational::new(i64::MIN, -3).reduced(),
Rational::new(i64::MAX, 3)
);
}
#[test]
fn neg_handles_i64_min_numerator() {
let r = -Rational::new(i64::MIN, 5);
assert_eq!(r, Rational::new(i64::MIN, -5));
assert!((-r).equals_value(&Rational::new(i64::MIN, 5)));
assert_eq!(-Rational::new(i64::MIN, i64::MIN), Rational::new(-1, 1));
let inf = -Rational::new(i64::MIN, 0);
assert_eq!(inf, Rational::new(i64::MAX, 0));
}
#[test]
fn abs_signum_cmp_handle_i64_min_terms() {
assert_eq!(Rational::new(i64::MIN, 2).abs(), Rational::new(1 << 62, 1));
assert_eq!(Rational::new(i64::MIN, 3).abs(), Rational::new(i64::MAX, 3));
assert_eq!(Rational::new(3, i64::MIN).signum(), -1);
assert_eq!(
Rational::new(3, i64::MIN).cmp_value(&Rational::zero()),
std::cmp::Ordering::Less
);
assert!(Rational::new(i64::MIN, i64::MIN).equals_value(&Rational::new(1, 1)));
}
#[test]
fn checked_ops_exact_or_none() {
assert_eq!(
Rational::new(1, 2).checked_add(Rational::new(1, 3)),
Some(Rational::new(5, 6))
);
assert_eq!(
Rational::new(1, 2).checked_sub(Rational::new(1, 3)),
Some(Rational::new(1, 6))
);
assert_eq!(
Rational::new(2, 4).checked_mul(Rational::new(3, 9)),
Some(Rational::new(1, 6))
);
assert_eq!(
Rational::new(1, 2).checked_div(Rational::new(3, 4)),
Some(Rational::new(2, 3))
);
let max = Rational::new(i64::MAX, 1);
assert_eq!(max.checked_add(Rational::new(1, 1)), None);
assert_eq!(
Rational::new(1 << 32, 1).checked_mul(Rational::new(1 << 32, 1)),
None
);
assert_eq!(
Rational::new(1, 1 << 32).checked_mul(Rational::new(1, 1 << 32)),
None
);
assert_eq!(
Rational::new(1, 2).checked_div(Rational::zero()),
Some(Rational::new(1, 0))
);
}
#[test]
fn operators_approximate_instead_of_wrapping() {
let r = Rational::new(i64::MAX, 1) + Rational::new(1, 1);
assert_eq!(r, Rational::new(i64::MAX, 1));
let r = Rational::new(i64::MIN, 1) - Rational::new(1, 1);
assert_eq!(r, Rational::new(i64::MIN, 1));
let tiny = Rational::new(1, i64::MAX) * Rational::new(1, 2);
assert_eq!(tiny.den, i64::MAX);
assert!(tiny.num == 0 || tiny.num == 1);
let tiny_neg = Rational::new(-3, i64::MAX) * Rational::new(1, 2);
assert_eq!(tiny_neg.den, i64::MAX);
assert!(tiny_neg.num <= 0 && tiny_neg.num >= -2);
let r = Rational::new(i64::MAX, 3) * Rational::new(3, i64::MAX);
assert_eq!(r, Rational::new(1, 1));
}
#[test]
fn arithmetic_uses_128bit_intermediates() {
let big = Rational::new(i64::MAX / 2, 3);
let r = big * Rational::new(3, 1);
assert_eq!(r, Rational::new(i64::MAX / 2, 1));
let a = Rational::new(1_000_000_000, 1);
let sum = a + a; assert_eq!(sum, Rational::new(2_000_000_000, 1));
}
}