From 764b207eea4c7d3e7c62d26f8c775f61d9b92b8d Mon Sep 17 00:00:00 2001 From: Liam Date: Fri, 31 Jul 2026 12:26:01 -0400 Subject: [PATCH 1/2] Prevent invalid number arithmetic from panicking --- src/builtins.rs | 27 ++++++++++--- src/evaluator.rs | 36 ++++++----------- src/lib.rs | 2 +- src/number.rs | 100 ++++++++++++++++++++++++++++++++++++++++------- 4 files changed, 119 insertions(+), 46 deletions(-) diff --git a/src/builtins.rs b/src/builtins.rs index 2357222..53280a2 100644 --- a/src/builtins.rs +++ b/src/builtins.rs @@ -249,7 +249,8 @@ fn acot<'a>(payload: &BuiltinFunctionPayload<'a>) -> Result, Error> { ) .0, ) - .div(&Number::from(2_i64), payload.config); + .div(&Number::from(2_i64), payload.config) + .unwrap(); Ok(Value::Number( pi_div_2.sub(&argument.atan(payload.config), payload.config), @@ -266,7 +267,9 @@ fn acsc<'a>(payload: &BuiltinFunctionPayload<'a>) -> Result, Error> { )); } - let reciprocal = Number::from(1_i64).div(&argument, payload.config); + let reciprocal = Number::from(1_i64) + .div(&argument, payload.config) + .map_err(|error| Error::new(payload.span, error.to_string()))?; Ok(Value::Number(reciprocal.asin(payload.config))) } @@ -299,7 +302,9 @@ fn asec<'a>(payload: &BuiltinFunctionPayload<'a>) -> Result, Error> { )); } - let reciprocal = Number::from(1_i64).div(&argument, payload.config); + let reciprocal = Number::from(1_i64) + .div(&argument, payload.config) + .map_err(|error| Error::new(payload.span, error.to_string()))?; Ok(Value::Number(reciprocal.acos(payload.config))) } @@ -349,6 +354,7 @@ fn constant_phi(config: Config) -> Number { Number::from(1_i64) .add(&Number::from(5_i64).sqrt(config), config) .div(&Number::from(2_i64), config) + .unwrap() } fn constant_pi(config: Config) -> Number { @@ -394,7 +400,10 @@ fn cot<'a>(payload: &BuiltinFunctionPayload<'a>) -> Result, Error> { )); } - Ok(Value::Number(Number::from(1_i64).div(&tan, payload.config))) + Number::from(1_i64) + .div(&tan, payload.config) + .map(Value::Number) + .map_err(|error| Error::new(payload.span, error.to_string())) } fn csc<'a>(payload: &BuiltinFunctionPayload<'a>) -> Result, Error> { @@ -409,7 +418,10 @@ fn csc<'a>(payload: &BuiltinFunctionPayload<'a>) -> Result, Error> { )); } - Ok(Value::Number(Number::from(1_i64).div(&sin, payload.config))) + Number::from(1_i64) + .div(&sin, payload.config) + .map(Value::Number) + .map_err(|error| Error::new(payload.span, error.to_string())) } fn e<'a>(payload: &BuiltinFunctionPayload<'a>) -> Result, Error> { @@ -760,7 +772,10 @@ fn sec<'a>(payload: &BuiltinFunctionPayload<'a>) -> Result, Error> { return Err(Error::new(payload.span, "Cannot compute sec of π/2 + nπ")); } - Ok(Value::Number(Number::from(1_i64).div(&cos, payload.config))) + Number::from(1_i64) + .div(&cos, payload.config) + .map(Value::Number) + .map_err(|error| Error::new(payload.span, error.to_string())) } fn sin<'a>(payload: &BuiltinFunctionPayload<'a>) -> Result, Error> { diff --git a/src/evaluator.rs b/src/evaluator.rs index ac7c0b1..410c00b 100644 --- a/src/evaluator.rs +++ b/src/evaluator.rs @@ -166,13 +166,10 @@ impl<'a> Evaluator<'a> { let (lhs_num, rhs_num) = (lhs_val.number(lhs.1)?, rhs_val.number(rhs.1)?); - if rhs_num.is_zero() { - return Err(Error::new(rhs.1, "Division by zero")); - } - - Ok(Value::Number( - lhs_num.div(&rhs_num, self.environment.config), - )) + lhs_num + .div(&rhs_num, self.environment.config) + .map(Value::Number) + .map_err(|error| Error::new(rhs.1, error.to_string())) } Expression::BinaryOp(BinaryOp::Equal, lhs, rhs) => Ok(Value::Boolean( self.evaluate_expression(lhs)? == self.evaluate_expression(rhs)?, @@ -241,13 +238,10 @@ impl<'a> Evaluator<'a> { let (lhs_num, rhs_num) = (lhs_val.number(lhs.1)?, rhs_val.number(rhs.1)?); - if rhs_num.is_zero() { - return Err(Error::new(rhs.1, "Modulo by zero")); - } - - Ok(Value::Number( - lhs_num.rem(&rhs_num, self.environment.config), - )) + lhs_num + .rem(&rhs_num, self.environment.config) + .map(Value::Number) + .map_err(|error| Error::new(rhs.1, error.to_string())) } Expression::BinaryOp(BinaryOp::Multiply, lhs, rhs) => Ok(Value::Number( self.evaluate_expression(lhs)?.number(lhs.1)?.mul( @@ -267,16 +261,10 @@ impl<'a> Evaluator<'a> { let (lhs_num, rhs_num) = (lhs_val.number(lhs.1)?, rhs_val.number(rhs.1)?); - if lhs_num.is_zero() && rhs_num.is_negative() { - return Err(Error::new( - rhs.1, - "Zero cannot be raised to a negative power", - )); - } - - Ok(Value::Number( - lhs_num.pow(&rhs_num, self.environment.config), - )) + lhs_num + .pow(&rhs_num, self.environment.config) + .map(Value::Number) + .map_err(|error| Error::new(rhs.1, error.to_string())) } Expression::BinaryOp(BinaryOp::Subtract, lhs, rhs) => Ok(Value::Number( self.evaluate_expression(lhs)?.number(lhs.1)?.sub( diff --git a/src/lib.rs b/src/lib.rs index dcc20db..ec9ea26 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -37,7 +37,7 @@ pub use crate::{ error::Error, evaluator::Evaluator, function::Function, - number::{Number, ParseDecimalError}, + number::{ArithmeticError, Number, ParseDecimalError}, parser::parse, rounding_mode::RoundingMode, value::Value, diff --git a/src/number.rs b/src/number.rs index 213910d..5a98c41 100644 --- a/src/number.rs +++ b/src/number.rs @@ -3,6 +3,13 @@ use super::*; #[derive(Clone, Copy, Debug, Eq, PartialEq)] pub struct ParseDecimalError; +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum ArithmeticError { + DivisionByZero, + ModuloByZero, + ZeroToNegativePower, +} + #[derive(Clone, Debug)] pub enum Number { Approx(Float), @@ -118,19 +125,27 @@ impl Number { } } - #[must_use] - pub fn div(&self, rhs: &Self, config: Config) -> Self { - if let (Self::Exact(lhs), Self::Exact(rhs)) = (self, rhs) { - Self::Exact((lhs / rhs).complete()) + /// # Errors + /// + /// Returns [`ArithmeticError::DivisionByZero`] if `rhs` is zero. + pub fn div( + &self, + rhs: &Self, + config: Config, + ) -> std::result::Result { + if rhs.is_zero() { + Err(ArithmeticError::DivisionByZero) + } else if let (Self::Exact(lhs), Self::Exact(rhs)) = (self, rhs) { + Ok(Self::Exact((lhs / rhs).complete())) } else { - Self::Approx( + Ok(Self::Approx( Float::with_val_round( config.precision(), &self.to_float(config) / &rhs.to_float(config), config.rounding_mode, ) .0, - ) + )) } } @@ -209,25 +224,46 @@ impl Number { } } - #[must_use] - pub fn pow(&self, rhs: &Self, config: Config) -> Self { + /// # Errors + /// + /// Returns [`ArithmeticError::ZeroToNegativePower`] if `self` is zero and + /// `rhs` is negative. + pub fn pow( + &self, + rhs: &Self, + config: Config, + ) -> std::result::Result { + if self.is_zero() && rhs.is_negative() { + return Err(ArithmeticError::ZeroToNegativePower); + } + match (self, rhs) { (Self::Exact(lhs), Self::Exact(exponent)) => { if exponent.is_integer() && let Some(exponent) = exponent.numer().to_i32() { - return Self::Exact(lhs.clone().pow(exponent)); + return Ok(Self::Exact(lhs.clone().pow(exponent))); } - self.approx_pow(rhs, config) + Ok(self.approx_pow(rhs, config)) } - _ => self.approx_pow(rhs, config), + _ => Ok(self.approx_pow(rhs, config)), } } - #[must_use] - pub fn rem(&self, rhs: &Self, config: Config) -> Self { - self.sub(&self.div(rhs, config).floor().mul(rhs, config), config) + /// # Errors + /// + /// Returns [`ArithmeticError::ModuloByZero`] if `rhs` is zero. + pub fn rem( + &self, + rhs: &Self, + config: Config, + ) -> std::result::Result { + if rhs.is_zero() { + return Err(ArithmeticError::ModuloByZero); + } + + Ok(self.sub(&self.div(rhs, config)?.floor().mul(rhs, config), config)) } #[must_use] @@ -360,6 +396,18 @@ impl Display for Number { } } +impl Display for ArithmeticError { + fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result { + f.write_str(match self { + Self::DivisionByZero => "Division by zero", + Self::ModuloByZero => "Modulo by zero", + Self::ZeroToNegativePower => "Zero cannot be raised to a negative power", + }) + } +} + +impl std::error::Error for ArithmeticError {} + impl From for Number { fn from(value: bool) -> Self { Self::from(i64::from(value)) @@ -452,7 +500,8 @@ mod tests { fn display_approx_configured_digits() { let number = Number::from(2_i64) .to_approx(Config::default()) - .div(&Number::from(5_555_222_222_222_i64), Config::default()); + .div(&Number::from(5_555_222_222_222_i64), Config::default()) + .unwrap(); let config = Config { digits: NonZeroUsize::new(4).unwrap(), @@ -513,6 +562,7 @@ mod tests { Number::from(2_i64) .to_approx(Config::default()) .div(&Number::from(5_555_222_222_222_i64), Config::default()) + .unwrap() .to_string(), "3.600216012960922e-13" ); @@ -556,6 +606,26 @@ mod tests { assert!(greater > approx); } + #[test] + fn undefined_exact_arithmetic_returns_error() { + let zero = Number::from(0_i64); + + assert_eq!( + Number::from(1_i64).div(&zero, Config::default()), + Err(ArithmeticError::DivisionByZero) + ); + + assert_eq!( + Number::from(1_i64).rem(&zero, Config::default()), + Err(ArithmeticError::ModuloByZero) + ); + + assert_eq!( + zero.pow(&Number::from(-1_i64), Config::default()), + Err(ArithmeticError::ZeroToNegativePower) + ); + } + #[test] fn zero_precision_uses_minimum() { let config = Config { From d8b0cfc23b9967e4a7d4ad966f2f1fef7d53b8d8 Mon Sep 17 00:00:00 2001 From: Liam Date: Fri, 31 Jul 2026 12:49:34 -0400 Subject: [PATCH 2/2] Move arithmetic error into module --- src/arithmetic_error.rs | 20 ++++++++++++++++++++ src/lib.rs | 4 +++- src/number.rs | 19 ------------------- 3 files changed, 23 insertions(+), 20 deletions(-) create mode 100644 src/arithmetic_error.rs diff --git a/src/arithmetic_error.rs b/src/arithmetic_error.rs new file mode 100644 index 0000000..51cdf9b --- /dev/null +++ b/src/arithmetic_error.rs @@ -0,0 +1,20 @@ +use super::*; + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum ArithmeticError { + DivisionByZero, + ModuloByZero, + ZeroToNegativePower, +} + +impl Display for ArithmeticError { + fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result { + f.write_str(match self { + Self::DivisionByZero => "Division by zero", + Self::ModuloByZero => "Modulo by zero", + Self::ZeroToNegativePower => "Zero cannot be raised to a negative power", + }) + } +} + +impl std::error::Error for ArithmeticError {} diff --git a/src/lib.rs b/src/lib.rs index ec9ea26..6497098 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -27,6 +27,7 @@ use { }; pub use crate::{ + arithmetic_error::ArithmeticError, builtin::Builtin, builtin_arity::BuiltinArity, builtin_function::BuiltinFunction, @@ -37,7 +38,7 @@ pub use crate::{ error::Error, evaluator::Evaluator, function::Function, - number::{ArithmeticError, Number, ParseDecimalError}, + number::{Number, ParseDecimalError}, parser::parse, rounding_mode::RoundingMode, value::Value, @@ -48,6 +49,7 @@ pub type Spanned = (T, Span); type Result = std::result::Result; +mod arithmetic_error; pub mod ast; mod builtin; mod builtin_arity; diff --git a/src/number.rs b/src/number.rs index 5a98c41..b13ae93 100644 --- a/src/number.rs +++ b/src/number.rs @@ -3,13 +3,6 @@ use super::*; #[derive(Clone, Copy, Debug, Eq, PartialEq)] pub struct ParseDecimalError; -#[derive(Clone, Copy, Debug, Eq, PartialEq)] -pub enum ArithmeticError { - DivisionByZero, - ModuloByZero, - ZeroToNegativePower, -} - #[derive(Clone, Debug)] pub enum Number { Approx(Float), @@ -396,18 +389,6 @@ impl Display for Number { } } -impl Display for ArithmeticError { - fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result { - f.write_str(match self { - Self::DivisionByZero => "Division by zero", - Self::ModuloByZero => "Modulo by zero", - Self::ZeroToNegativePower => "Zero cannot be raised to a negative power", - }) - } -} - -impl std::error::Error for ArithmeticError {} - impl From for Number { fn from(value: bool) -> Self { Self::from(i64::from(value))