use std::fmt::{Display, Formatter}; use Num::{Float, Int}; #[derive(Clone, Copy, Debug, PartialEq)] pub enum Num { Int(i128), Float(f64), } impl Num { fn as_float(self) -> f64 { match self { Int(i) => i as f64, Float(f) => f, } } fn powi(self, p: i128) -> Num { if p == 0 { return Int(1); } if p == 1 { return self; } if p % 2 == 0 { return self.mul(self).powi(p / 2); } self.mul(self).powi(p / 2).mul(self) } pub fn pow(self, other: Num) -> Num { match other { Int(i) => { if i >= 0 { return self.powi(i); } if let Ok(i_32) = i32::try_from(i) { return Float(self.as_float().powi(i_32)); } Float(1.).div(self).powi(-i) } Float(f) => Float(self.as_float().powf(f)), } } pub fn mul(self, other: Num) -> Num { if let (Int(i1), Int(i2)) = (self, other) { if let Some(p) = i1.checked_mul(i2) { return Int(p); } } Float(self.as_float() * other.as_float()) } pub fn div(self, other: Num) -> Num { Float(self.as_float() / other.as_float()) } pub fn int_div(self, other: Num) -> Num { if let (Int(i1), Int(i2)) = (self, other) { if let Some(q) = i1.checked_div(i2) { return Int(q); } } Int((self.as_float() / other.as_float()) as i128) } pub fn modulo(self, other: Num) -> Num { if let (Int(i1), Int(i2)) = (self, other) { if let Some(r) = i1.checked_rem(i2) { if i2 > 0 && r < 0 || i2 < 0 && r > 0 { return Int(r + i2); } return Int(r); } } let n1 = self.as_float(); let n2 = other.as_float(); let r = n1 % n2; if n2 > 0. && r < 0. || n2 < 0. && r > 0. { return Float(r + n2); } Float(r) } pub fn add(self, other: Num) -> Num { if let (Int(i1), Int(i2)) = (self, other) { if let Some(s) = i1.checked_add(i2) { return Int(s); } } Float(self.as_float() + other.as_float()) } pub fn sub(self, other: Num) -> Num { if let (Int(i1), Int(i2)) = (self, other) { if let Some(d) = i1.checked_sub(i2) { return Int(d); } } Float(self.as_float() - other.as_float()) } pub fn sin(self) -> Num { Float(self.as_float().sin()) } pub fn cos(self) -> Num { Float(self.as_float().cos()) } pub fn tan(self) -> Num { Float(self.as_float().tan()) } pub fn asin(self) -> Num { Float(self.as_float().asin()) } pub fn acos(self) -> Num { Float(self.as_float().acos()) } pub fn atan(self) -> Num { Float(self.as_float().atan()) } pub fn sqrt(self) -> Num { match self { Int(i) if i >= 0 => { if i <= 1 { return Int(i); } let mut x0 = i / 2; let mut x1 = (x0 + i / x0) / 2; while x1 < x0 { x0 = x1; x1 = (x0 + i / x0) / 2; } if x0 * x0 == i { return Int(x0); } } _ => (), } Float(self.as_float().sqrt()) } pub fn log(self) -> Num { Float(self.as_float().ln()) } pub fn log10(self) -> Num { Float(self.as_float().log10()) } pub fn log2(self) -> Num { Float(self.as_float().log2()) } pub fn floor(self) -> Num { match self { Int(i) => Int(i), Float(f) => Int(f.floor() as i128), } } pub fn ceil(self) -> Num { match self { Int(i) => Int(i), Float(f) => Int(f.ceil() as i128), } } pub fn round(self) -> Num { match self { Int(i) => Int(i), Float(f) => Int(f.round_ties_even() as i128), } } pub fn abs(self) -> Num { match self { Int(i) => Int(i.abs()), Float(f) => Float(f.abs()), } } } impl Display for Num { fn fmt(&self, f: &mut Formatter<'_>) -> Result<(), std::fmt::Error> { match self { Int(n) => write!(f, "{n}"), Float(n) => { if !(-1e30..=1e30).contains(n) { return write!(f, "{n:.7e}"); } let mut buf = format!("{n:.7}"); while let Some(b'0') = buf.as_bytes().get(buf.len() - 1) { buf.truncate(buf.len() - 1); } write!(f, "{buf}") } } } } #[cfg(test)] mod tests { use super::*; #[test] fn pow_positive_int() { assert_eq!(Int(2).pow(Int(16)), Int(65536)); } #[test] fn pow_negative_int() { assert_eq!(Int(2).pow(Int(-3)), Float(0.125)); } #[test] fn pow_float() { assert_eq!(Int(4).pow(Float(0.5)), Float(2.)); } #[test] fn pow_overflow() { assert_eq!(Int(2).pow(Int(1 << 126)), Float(f64::INFINITY)); } #[test] fn pow_underflow() { assert_eq!(Int(2).pow(Int(-(1 << 126))), Float(0.)); } #[test] fn modulo_pos() { assert_eq!(Int(5).modulo(Int(3)), Int(2)); } #[test] fn modulo_pos_neg() { assert_eq!(Int(5).modulo(Int(-3)), Int(-1)); } #[test] fn modulo_neg_pos() { assert_eq!(Int(-5).modulo(Int(3)), Int(1)); } #[test] fn modulo_neg() { assert_eq!(Int(-5).modulo(Int(-3)), Int(-2)); } #[test] fn modulo_float_pos() { assert_eq!(Float(5.).modulo(Int(3)), Float(2.)); } #[test] fn modulo_float_pos_neg() { assert_eq!(Float(5.).modulo(Int(-3)), Float(-1.)); } #[test] fn modulo_float_neg_pos() { assert_eq!(Float(-5.).modulo(Int(3)), Float(1.)); } #[test] fn modulo_float_neg() { assert_eq!(Float(-5.).modulo(Int(-3)), Float(-2.)); } #[test] fn sqrt() { for n in 0..65536 { assert_eq!( Int(n * n).sqrt(), Int(n), "Int({}).sqrt() is not equal to {}", n * n, n ); } } #[test] fn sqrt_big() { assert_eq!(Int(1 << 126).sqrt(), Int(1 << 63)); } }