use std::cmp::Ordering;
use std::f64::consts::PI;
use std::fmt::{self, Debug, Display, Formatter};
use std::hash;
use std::iter::{Product, Sum};
use std::ops::{self, Add, Div, Mul, Neg, Rem, Sub};
use std::str::FromStr;
use anyhow::{Result, bail, ensure};
use fastnum::D128;
use revision::revisioned;
use rust_decimal::Decimal;
use rust_decimal::prelude::*;
use storekey::{BorrowDecode, Encode};
use surrealdb_types::{SqlFormat, ToSql, fmt_non_finite_f64, write_sql};
use super::IndexFormat;
use crate::err::Error;
use crate::expr::decimal::DecimalLexEncoder;
use crate::fnc::util::math::ToFloat;
use crate::val::{TryAdd, TryDiv, TryFloatDiv, TryMul, TryNeg, TryPow, TryRem, TrySub};
#[derive(Copy, Clone, Encode, BorrowDecode)]
pub(crate) enum NumberKind {
Int,
Float,
Decimal,
}
#[revisioned(revision = 1)]
#[derive(Clone, Copy, Debug)]
#[cfg_attr(feature = "arbitrary", derive(arbitrary::Arbitrary))]
pub(crate) enum Number {
Int(i64),
Float(f64),
Decimal(Decimal),
}
impl Default for Number {
fn default() -> Self {
Self::Int(0)
}
}
macro_rules! from_prim_ints {
($($int: ty),*) => {
$(
impl From<$int> for Number {
fn from(i: $int) -> Self {
Self::Int(i as i64)
}
}
)*
};
}
from_prim_ints!(i8, i16, i32, i64, isize, u8, u16, u32);
impl From<u64> for Number {
fn from(i: u64) -> Self {
match i64::try_from(i) {
Ok(v) => Self::Int(v),
Err(_) => Self::Decimal(Decimal::from(i)),
}
}
}
impl From<usize> for Number {
fn from(i: usize) -> Self {
Self::from(i as u64)
}
}
impl TryFrom<i128> for Number {
type Error = Error;
fn try_from(i: i128) -> Result<Self, Self::Error> {
if let Ok(v) = i64::try_from(i) {
Ok(Self::Int(v))
} else if let Some(v) = Decimal::from_i128(i) {
Ok(Self::Decimal(v))
} else {
Err(Error::TryFrom(i.to_string(), "Number"))
}
}
}
impl TryFrom<u128> for Number {
type Error = Error;
fn try_from(i: u128) -> Result<Self, Self::Error> {
if let Ok(v) = i64::try_from(i) {
Ok(Self::Int(v))
} else if let Some(v) = Decimal::from_u128(i) {
Ok(Self::Decimal(v))
} else {
Err(Error::TryFrom(i.to_string(), "Number"))
}
}
}
impl From<f32> for Number {
fn from(f: f32) -> Self {
Self::Float(f as f64)
}
}
impl From<f64> for Number {
fn from(f: f64) -> Self {
Self::Float(f)
}
}
impl From<Decimal> for Number {
fn from(v: Decimal) -> Self {
Self::Decimal(v)
}
}
impl From<surrealdb_types::Number> for Number {
fn from(v: surrealdb_types::Number) -> Self {
match v {
surrealdb_types::Number::Int(i) => Self::Int(i),
surrealdb_types::Number::Float(f) => Self::Float(f),
surrealdb_types::Number::Decimal(d) => Self::Decimal(d),
}
}
}
impl From<Number> for surrealdb_types::Number {
fn from(v: Number) -> Self {
match v {
Number::Int(i) => Self::Int(i),
Number::Float(f) => Self::Float(f),
Number::Decimal(d) => Self::Decimal(d),
}
}
}
impl FromStr for Number {
type Err = ();
fn from_str(s: &str) -> Result<Self, Self::Err> {
match s.parse::<i64>() {
Ok(v) => Ok(Self::Int(v)),
_ => match s.parse::<f64>() {
Ok(v) => Ok(Self::Float(v)),
_ => Err(()),
},
}
}
}
macro_rules! try_into_prim {
($($int: ty => $to_int: ident),*) => {
$(
impl TryFrom<Number> for $int {
type Error = Error;
fn try_from(value: Number) -> Result<Self, Self::Error> {
match value {
Number::Int(v) => match v.$to_int() {
Some(v) => Ok(v),
None => Err(Error::TryFrom(value.to_sql(), stringify!($int))),
},
Number::Float(v) => match v.$to_int() {
Some(v) => Ok(v),
None => Err(Error::TryFrom(value.to_sql(), stringify!($int))),
},
Number::Decimal(ref v) => match v.$to_int() {
Some(v) => Ok(v),
None => Err(Error::TryFrom(value.to_sql(), stringify!($int))),
},
}
}
}
)*
};
}
try_into_prim!(
i8 => to_i8, i16 => to_i16, i32 => to_i32, i64 => to_i64, i128 => to_i128,
u8 => to_u8, u16 => to_u16, u32 => to_u32, u64 => to_u64, u128 => to_u128,
f32 => to_f32, f64 => to_f64
);
impl TryFrom<Number> for Decimal {
type Error = Error;
fn try_from(value: Number) -> Result<Self, Self::Error> {
match value {
Number::Int(v) => match Decimal::from_i64(v) {
Some(v) => Ok(v),
None => Err(Error::TryFrom(value.to_sql(), "Decimal")),
},
Number::Float(v) => match Decimal::try_from(v) {
Ok(v) => Ok(v),
_ => Err(Error::TryFrom(value.to_sql(), "Decimal")),
},
Number::Decimal(x) => Ok(x),
}
}
}
impl Display for Number {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
match self {
Number::Int(v) => Display::fmt(v, f),
Number::Float(v) => Display::fmt(v, f),
Number::Decimal(v) => Display::fmt(v, f),
}
}
}
impl ToSql for Number {
fn fmt_sql(&self, f: &mut String, sql_fmt: SqlFormat) {
match self {
Number::Int(v) => v.fmt_sql(f, sql_fmt),
Number::Float(v) => {
match fmt_non_finite_f64(*v) {
Some(special) => write_sql!(f, sql_fmt, "{}", special),
None => write_sql!(f, sql_fmt, "{v}f"),
}
}
Number::Decimal(v) => v.fmt_sql(f, sql_fmt),
}
}
}
impl Number {
pub const NAN: Number = Number::Float(f64::NAN);
pub fn is_nan(&self) -> bool {
matches!(self, Number::Float(v) if v.is_nan())
}
pub fn is_int(&self) -> bool {
matches!(self, Number::Int(_))
}
pub fn is_float(&self) -> bool {
matches!(self, Number::Float(_))
}
pub fn is_truthy(&self) -> bool {
match self {
Number::Int(v) => v != &0,
Number::Float(v) => v != &0.0,
Number::Decimal(v) => v != &Decimal::ZERO,
}
}
pub fn is_zero(&self) -> bool {
match self {
Number::Int(v) => v == &0,
Number::Float(v) => v == &0.0,
Number::Decimal(v) => v == &Decimal::ZERO,
}
}
pub fn as_usize(self) -> usize {
match self {
Number::Int(v) => v as usize,
Number::Float(v) => v as usize,
Number::Decimal(v) => v.try_into().unwrap_or_default(),
}
}
pub fn as_int(self) -> i64 {
match self {
Number::Int(v) => v,
Number::Float(v) => v as i64,
Number::Decimal(v) => v.try_into().unwrap_or_default(),
}
}
pub fn as_float(self) -> f64 {
match self {
Number::Int(v) => v as f64,
Number::Float(v) => v,
Number::Decimal(v) => v.try_into().unwrap_or_default(),
}
}
pub fn as_decimal(self) -> Decimal {
match self {
Number::Int(v) => Decimal::from(v),
Number::Float(v) => Decimal::try_from(v).unwrap_or_default(),
Number::Decimal(v) => v,
}
}
pub fn as_int_lossless(self) -> Option<i64> {
match self {
Number::Int(v) => Some(v),
Number::Float(v) => {
if !v.is_finite()
|| v.fract() != 0.0
|| v < i64::MIN as f64
|| v >= -(i64::MIN as f64)
{
return None;
}
Some(v as i64)
}
Number::Decimal(v) => {
if v.fract().is_zero() {
v.try_into().ok()
} else {
None
}
}
}
}
pub fn as_array_index(self) -> Option<usize> {
match self {
Number::Int(v) => usize::try_from(v).ok(),
Number::Float(v) => {
if !v.is_finite() || v < 0.0 || v.fract() != 0.0 {
return None;
}
let idx = v as usize;
(idx as f64 == v).then_some(idx)
}
Number::Decimal(v) => {
if v.fract().is_zero() {
v.to_usize()
} else {
None
}
}
}
}
pub fn to_int(self) -> i64 {
match self {
Number::Int(v) => v,
Number::Float(v) => v as i64,
Number::Decimal(v) => v.to_i64().unwrap_or_default(),
}
}
pub fn to_float(self) -> f64 {
match self {
Number::Int(v) => v as f64,
Number::Float(v) => v,
Number::Decimal(v) => v.try_into().unwrap_or_default(),
}
}
pub fn to_decimal(self) -> Decimal {
match self {
Number::Int(v) => Decimal::from(v),
Number::Float(v) => Decimal::from_f64(v).unwrap_or_default(),
Number::Decimal(v) => v,
}
}
pub(crate) fn as_decimal_buf(&self) -> Vec<u8> {
match self {
Self::Int(v) => {
DecimalLexEncoder::encode(D128::from(*v))
}
Self::Float(v) => {
DecimalLexEncoder::encode(D128::from_f64(*v))
}
Self::Decimal(v) => {
DecimalLexEncoder::encode(DecimalLexEncoder::to_d128(*v))
}
}
}
pub(crate) fn from_decimal_buf(b: &[u8]) -> Result<Self> {
let dec = DecimalLexEncoder::decode(b)?;
if dec.is_finite() {
match DecimalLexEncoder::to_decimal(dec) {
Ok(dec) => Ok(Number::Decimal(dec)),
Err(_) => Ok(Number::Float(dec.to_f64())),
}
} else if dec.is_nan() {
Ok(Number::Float(f64::NAN))
} else if dec.is_infinite() {
if dec.is_negative() {
Ok(Number::Float(f64::NEG_INFINITY))
} else {
Ok(Number::Float(f64::INFINITY))
}
} else {
bail!(Error::Serialization(format!("Invalid decimal value: {dec}")))
}
}
pub(crate) fn from_decimal_buf_kind(b: &[u8], kind: NumberKind) -> Result<Self> {
let dec = DecimalLexEncoder::decode(b)?;
match kind {
NumberKind::Int => {
ensure!(dec.is_finite(), format!("Invalid integer value: {dec}"));
Ok(Number::Int(dec.to_string().parse::<i64>()?))
}
NumberKind::Float => {
if dec.is_nan() {
Ok(Number::Float(f64::NAN))
} else if dec.is_infinite() {
if dec.is_negative() {
Ok(Number::Float(f64::NEG_INFINITY))
} else {
Ok(Number::Float(f64::INFINITY))
}
} else {
let dec = DecimalLexEncoder::to_decimal(dec)?;
dec.to_f64()
.ok_or_else(|| anyhow::Error::msg(format!("Invalid f64 {dec}")))
.map(Number::Float)
}
}
NumberKind::Decimal => {
ensure!(dec.is_finite(), format!("Invalid integer value: {dec}"));
Ok(Number::Decimal(DecimalLexEncoder::to_decimal(dec)?))
}
}
}
pub fn abs(self) -> Self {
match self {
Number::Int(v) => v.abs().into(),
Number::Float(v) => v.abs().into(),
Number::Decimal(v) => v.abs().into(),
}
}
pub fn checked_abs(self) -> Option<Self> {
match self {
Number::Int(v) => v.checked_abs().map(|x| x.into()),
Number::Float(v) => Some(v.abs().into()),
Number::Decimal(v) => Some(v.abs().into()),
}
}
pub fn acos(self) -> Self {
self.to_float().acos().into()
}
pub fn asin(self) -> Self {
self.to_float().asin().into()
}
pub fn atan(self) -> Self {
self.to_float().atan().into()
}
pub fn acot(self) -> Self {
(PI / 2.0 - self.atan().to_float()).into()
}
pub fn ceil(self) -> Self {
match self {
Number::Int(v) => v.into(),
Number::Float(v) => v.ceil().into(),
Number::Decimal(v) => v.ceil().into(),
}
}
pub fn clamp(self, min: Self, max: Self) -> Self {
match (self, min, max) {
(Number::Int(n), Number::Int(min), Number::Int(max)) => n.clamp(min, max).into(),
(Number::Decimal(n), min, max) => n.clamp(min.to_decimal(), max.to_decimal()).into(),
(Number::Float(n), min, max) => n.clamp(min.to_float(), max.to_float()).into(),
(Number::Int(n), min, max) => n.to_float().clamp(min.to_float(), max.to_float()).into(),
}
}
pub fn cos(self) -> Self {
self.to_float().cos().into()
}
pub fn cot(self) -> Self {
(1.0 / self.to_float().tan()).into()
}
pub fn deg2rad(self) -> Self {
self.to_float().to_radians().into()
}
pub fn floor(self) -> Self {
match self {
Number::Int(v) => v.into(),
Number::Float(v) => v.floor().into(),
Number::Decimal(v) => v.floor().into(),
}
}
fn lerp_f64(from: f64, to: f64, factor: f64) -> f64 {
from + factor * (to - from)
}
fn lerp_decimal(from: Decimal, to: Decimal, factor: Decimal) -> Decimal {
from + factor * (to - from)
}
pub fn lerp(self, from: Self, to: Self) -> Self {
match (self, from, to) {
(Number::Decimal(val), from, to) => {
Self::lerp_decimal(from.to_decimal(), to.to_decimal(), val).into()
}
(val, from, to) => {
Self::lerp_f64(from.to_float(), to.to_float(), val.to_float()).into()
}
}
}
fn repeat_f64(t: f64, m: f64) -> f64 {
(t - (t / m).floor() * m).clamp(0.0, m)
}
fn repeat_decimal(t: Decimal, m: Decimal) -> Decimal {
(t - (t / m).floor() * m).clamp(Decimal::ZERO, m)
}
pub fn lerp_angle(self, from: Self, to: Self) -> Self {
match (self, from, to) {
(Number::Decimal(val), from, to) => {
let from = from.to_decimal();
let to = to.to_decimal();
let mut dt = Self::repeat_decimal(to - from, Decimal::from(360));
if dt > Decimal::from(180) {
dt = Decimal::from(360) - dt;
}
Self::lerp_decimal(from, from + dt, val).into()
}
(val, from, to) => {
let val = val.to_float();
let from = from.to_float();
let to = to.to_float();
let mut dt = Self::repeat_f64(to - from, 360.0);
if dt > 180.0 {
dt = 360.0 - dt;
}
Self::lerp_f64(from, from + dt, val).into()
}
}
}
pub fn ln(self) -> Self {
self.to_float().ln().into()
}
pub fn log(self, base: Self) -> Self {
self.to_float().log(base.to_float()).into()
}
pub fn log2(self) -> Self {
self.to_float().log2().into()
}
pub fn log10(self) -> Self {
self.to_float().log10().into()
}
pub fn rad2deg(self) -> Self {
self.to_float().to_degrees().into()
}
pub fn round(self) -> Self {
match self {
Number::Int(v) => v.into(),
Number::Float(v) => v.round().into(),
Number::Decimal(v) => v.round().into(),
}
}
pub fn fixed(self, precision: usize) -> Number {
match self {
Number::Int(v) => v.into(),
Number::Float(v) => format!("{v:.precision$}")
.parse::<f64>()
.expect("formatted f64 always parses back as f64")
.into(),
Number::Decimal(v) => v.round_dp(precision as u32).into(),
}
}
pub fn sign(self) -> Self {
match self {
Number::Int(n) => n.signum().into(),
Number::Float(n) => n.signum().into(),
Number::Decimal(n) => n.signum().into(),
}
}
pub fn sin(self) -> Self {
self.to_float().sin().into()
}
pub fn tan(self) -> Self {
self.to_float().tan().into()
}
pub fn sqrt(self) -> Self {
match self {
Number::Int(v) => (v as f64).sqrt().into(),
Number::Float(v) => v.sqrt().into(),
Number::Decimal(v) => v.sqrt().unwrap_or_default().into(),
}
}
}
impl Eq for Number {}
impl Ord for Number {
fn cmp(&self, other: &Self) -> Ordering {
fn total_cmp_f64(a: f64, b: f64) -> Ordering {
if a == 0.0 && b == 0.0 {
Ordering::Equal
} else {
a.total_cmp(&b)
}
}
macro_rules! greater {
($f:ident) => {
if $f.is_sign_positive() {
Ordering::Greater
} else {
Ordering::Less
}
};
}
match (self, other) {
(Number::Int(v), Number::Int(w)) => v.cmp(w),
(Number::Float(v), Number::Float(w)) => total_cmp_f64(*v, *w),
(Number::Decimal(v), Number::Decimal(w)) => v.cmp(w),
(Number::Int(v), Number::Float(w)) => {
if !w.is_finite() {
return greater!(w).reverse();
}
let l = *v as i128;
let r = *w as i128;
match l.cmp(&r) {
Ordering::Equal => total_cmp_f64(0.0, w.fract()),
ordering => ordering,
}
}
(v @ Number::Float(_), w @ Number::Int(_)) => w.cmp(v).reverse(),
(Number::Int(v), Number::Decimal(w)) => Decimal::from(*v).cmp(w),
(Number::Decimal(v), Number::Int(w)) => v.cmp(&Decimal::from(*w)),
(Number::Float(v), Number::Decimal(w)) => {
macro_rules! compare_fractions {
($l:ident, $r:ident) => {
match ($l == 0.0, $r == Decimal::ZERO) {
(true, true) => {
return Ordering::Equal;
}
(true, false) => {
return greater!($r).reverse();
}
(false, true) => {
return greater!($l);
}
(false, false) => {
continue;
}
}
};
}
if !v.is_finite() {
return greater!(v);
}
let l = *v as i128;
let Ok(r) = i128::try_from(*w) else {
return greater!(w).reverse();
};
match l.cmp(&r) {
Ordering::Equal => {
const SAFE_MULTIPLIER: i64 = 9_007_199_254_740_000;
let mut l = v.fract();
let mut r = w.fract();
for _ in 0..12 {
l *= SAFE_MULTIPLIER as f64;
r *= Decimal::new(SAFE_MULTIPLIER, 0);
match r.to_i64() {
Some(ref right) => match (l as i64).cmp(right) {
Ordering::Equal => {
l = l.fract();
r = r.fract();
compare_fractions!(l, r);
}
ordering => {
return ordering;
}
},
None => {
return greater!(w).reverse();
}
}
}
Ordering::Equal
}
ordering => ordering,
}
}
(v @ Number::Decimal(..), w @ Number::Float(..)) => w.cmp(v).reverse(),
}
}
}
impl hash::Hash for Number {
fn hash<H: hash::Hasher>(&self, state: &mut H) {
self.as_decimal_buf().hash(state);
}
}
impl PartialEq for Number {
fn eq(&self, other: &Self) -> bool {
fn total_eq_f64(a: f64, b: f64) -> bool {
a.to_bits().eq(&b.to_bits()) || (a == 0.0 && b == 0.0)
}
match (self, other) {
(Number::Int(v), Number::Int(w)) => v.eq(w),
(Number::Float(v), Number::Float(w)) => total_eq_f64(*v, *w),
(Number::Decimal(v), Number::Decimal(w)) => v.eq(w),
(v @ Number::Int(_), w @ Number::Float(_)) => v.cmp(w) == Ordering::Equal,
(v @ Number::Float(_), w @ Number::Int(_)) => v.cmp(w) == Ordering::Equal,
(Number::Int(v), Number::Decimal(w)) => Decimal::from(*v).eq(w),
(Number::Decimal(v), Number::Int(w)) => v.eq(&Decimal::from(*w)),
(v @ Number::Float(_), w @ Number::Decimal(_)) => v.cmp(w) == Ordering::Equal,
(v @ Number::Decimal(_), w @ Number::Float(_)) => v.cmp(w) == Ordering::Equal,
}
}
}
impl PartialOrd for Number {
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
Some(self.cmp(other))
}
}
macro_rules! impl_simple_try_op {
($trt:ident, $fn:ident, $unchecked:ident, $checked:ident) => {
impl $trt for Number {
type Output = Self;
fn $fn(self, other: Self) -> Result<Self> {
Ok(match (self, other) {
(Number::Int(v), Number::Int(w)) => Number::Int(
v.$checked(w).ok_or_else(|| Error::$trt(v.to_string(), w.to_string()))?,
),
(Number::Float(v), Number::Float(w)) => Number::Float(v.$unchecked(w)),
(Number::Decimal(v), Number::Decimal(w)) => Number::Decimal(
v.$checked(w).ok_or_else(|| Error::$trt(v.to_string(), w.to_string()))?,
),
(Number::Int(v), Number::Float(w)) => Number::Float((v as f64).$unchecked(w)),
(Number::Float(v), Number::Int(w)) => Number::Float(v.$unchecked(w as f64)),
(v, w) => Number::Decimal(
v.to_decimal()
.$checked(w.to_decimal())
.ok_or_else(|| Error::$trt(v.to_sql(), w.to_sql()))?,
),
})
}
}
};
}
impl_simple_try_op!(TryAdd, try_add, add, checked_add);
impl_simple_try_op!(TrySub, try_sub, sub, checked_sub);
impl_simple_try_op!(TryMul, try_mul, mul, checked_mul);
impl_simple_try_op!(TryDiv, try_div, div, checked_div);
impl_simple_try_op!(TryRem, try_rem, rem, checked_rem);
impl TryPow for Number {
type Output = Self;
fn try_pow(self, power: Self) -> Result<Self> {
Ok(match (self, power) {
(Self::Int(v), Self::Int(p)) => Self::Int(match v {
0 => match p.cmp(&0) {
Ordering::Less => bail!(Error::TryPow(v.to_string(), p.to_string())),
Ordering::Equal => 1,
Ordering::Greater => 0,
},
1 => 1,
-1 => {
if p % 2 == 0 {
1
} else {
-1
}
}
_ => p
.try_into()
.ok()
.and_then(|p| v.checked_pow(p))
.ok_or_else(|| Error::TryPow(v.to_string(), p.to_string()))?,
}),
(Self::Decimal(v), Self::Int(p)) => Self::Decimal(
v.checked_powi(p).ok_or_else(|| Error::TryPow(v.to_string(), p.to_string()))?,
),
(Self::Decimal(v), Self::Float(p)) => Self::Decimal(
v.checked_powf(p).ok_or_else(|| Error::TryPow(v.to_string(), p.to_string()))?,
),
(Self::Decimal(v), Self::Decimal(p)) => Self::Decimal(
v.checked_powd(p).ok_or_else(|| Error::TryPow(v.to_string(), p.to_string()))?,
),
(v, p) => v.as_float().powf(p.as_float()).into(),
})
}
}
impl TryNeg for Number {
type Output = Self;
fn try_neg(self) -> Result<Self::Output> {
Ok(match self {
Self::Int(n) => {
Number::Int(n.checked_neg().ok_or_else(|| Error::TryNeg(n.to_string()))?)
}
Self::Float(n) => Number::Float(-n),
Self::Decimal(n) => Number::Decimal(-n),
})
}
}
impl TryFloatDiv for Number {
type Output = Self;
fn try_float_div(self, other: Self) -> Result<Self> {
Ok(match (self, other) {
(Number::Int(v), Number::Int(w)) => {
let quotient = (v as f64).div(w as f64);
if quotient.fract() != 0.0 {
return Ok(Number::Float(quotient));
}
Number::Int(
v.checked_div(w).ok_or_else(|| Error::TryDiv(v.to_string(), w.to_string()))?,
)
}
(v, w) => v.try_div(w)?,
})
}
}
impl ops::Add for Number {
type Output = Self;
fn add(self, other: Self) -> Self {
match (self, other) {
(Number::Int(v), Number::Int(w)) => Number::Int(v + w),
(Number::Float(v), Number::Float(w)) => Number::Float(v + w),
(Number::Decimal(v), Number::Decimal(w)) => Number::Decimal(v + w),
(Number::Int(v), Number::Float(w)) => Number::Float(v as f64 + w),
(Number::Float(v), Number::Int(w)) => Number::Float(v + w as f64),
(v, w) => Number::from(v.as_decimal() + w.as_decimal()),
}
}
}
impl<'b> ops::Add<&'b Number> for &Number {
type Output = Number;
fn add(self, other: &'b Number) -> Number {
match (self, other) {
(Number::Int(v), Number::Int(w)) => Number::Int(v + w),
(Number::Float(v), Number::Float(w)) => Number::Float(v + w),
(Number::Decimal(v), Number::Decimal(w)) => Number::Decimal(v + w),
(Number::Int(v), Number::Float(w)) => Number::Float(*v as f64 + w),
(Number::Float(v), Number::Int(w)) => Number::Float(v + *w as f64),
(v, w) => Number::from(v.to_decimal() + w.to_decimal()),
}
}
}
impl ops::Sub for Number {
type Output = Self;
fn sub(self, other: Self) -> Self {
match (self, other) {
(Number::Int(v), Number::Int(w)) => Number::Int(v - w),
(Number::Float(v), Number::Float(w)) => Number::Float(v - w),
(Number::Decimal(v), Number::Decimal(w)) => Number::Decimal(v - w),
(Number::Int(v), Number::Float(w)) => Number::Float(v as f64 - w),
(Number::Float(v), Number::Int(w)) => Number::Float(v - w as f64),
(v, w) => Number::from(v.as_decimal() - w.as_decimal()),
}
}
}
impl<'b> ops::Sub<&'b Number> for &Number {
type Output = Number;
fn sub(self, other: &'b Number) -> Number {
match (self, other) {
(Number::Int(v), Number::Int(w)) => Number::Int(v - w),
(Number::Float(v), Number::Float(w)) => Number::Float(v - w),
(Number::Decimal(v), Number::Decimal(w)) => Number::Decimal(v - w),
(Number::Int(v), Number::Float(w)) => Number::Float(*v as f64 - w),
(Number::Float(v), Number::Int(w)) => Number::Float(v - *w as f64),
(v, w) => Number::from(v.to_decimal() - w.to_decimal()),
}
}
}
impl ops::Mul for Number {
type Output = Self;
fn mul(self, other: Self) -> Self {
match (self, other) {
(Number::Int(v), Number::Int(w)) => Number::Int(v * w),
(Number::Float(v), Number::Float(w)) => Number::Float(v * w),
(Number::Decimal(v), Number::Decimal(w)) => Number::Decimal(v * w),
(Number::Int(v), Number::Float(w)) => Number::Float(v as f64 * w),
(Number::Float(v), Number::Int(w)) => Number::Float(v * w as f64),
(v, w) => Number::from(v.as_decimal() * w.as_decimal()),
}
}
}
impl<'b> ops::Mul<&'b Number> for &Number {
type Output = Number;
fn mul(self, other: &'b Number) -> Number {
match (self, other) {
(Number::Int(v), Number::Int(w)) => Number::Int(v * w),
(Number::Float(v), Number::Float(w)) => Number::Float(v * w),
(Number::Decimal(v), Number::Decimal(w)) => Number::Decimal(v * w),
(Number::Int(v), Number::Float(w)) => Number::Float(*v as f64 * w),
(Number::Float(v), Number::Int(w)) => Number::Float(v * *w as f64),
(v, w) => Number::from(v.to_decimal() * w.to_decimal()),
}
}
}
impl ops::Div for Number {
type Output = Self;
fn div(self, other: Self) -> Self {
match (self, other) {
(Number::Int(v), Number::Int(w)) => Number::Int(v / w),
(Number::Float(v), Number::Float(w)) => Number::Float(v / w),
(Number::Decimal(v), Number::Decimal(w)) => Number::Decimal(v / w),
(Number::Int(v), Number::Float(w)) => Number::Float(v as f64 / w),
(Number::Float(v), Number::Int(w)) => Number::Float(v / w as f64),
(v, w) => Number::from(v.as_decimal() / w.as_decimal()),
}
}
}
impl<'b> ops::Div<&'b Number> for &Number {
type Output = Number;
fn div(self, other: &'b Number) -> Number {
match (self, other) {
(Number::Int(v), Number::Int(w)) => Number::Int(v / w),
(Number::Float(v), Number::Float(w)) => Number::Float(v / w),
(Number::Decimal(v), Number::Decimal(w)) => Number::Decimal(v / w),
(Number::Int(v), Number::Float(w)) => Number::Float(*v as f64 / w),
(Number::Float(v), Number::Int(w)) => Number::Float(v / *w as f64),
(v, w) => Number::from(v.to_decimal() / w.to_decimal()),
}
}
}
impl Neg for Number {
type Output = Self;
fn neg(self) -> Self::Output {
match self {
Self::Int(n) => Number::Int(-n),
Self::Float(n) => Number::Float(-n),
Self::Decimal(n) => Number::Decimal(-n),
}
}
}
impl Sum<Self> for Number {
fn sum<I>(iter: I) -> Number
where
I: Iterator<Item = Self>,
{
iter.fold(Number::Int(0), |a, b| a + b)
}
}
impl<'a> Sum<&'a Self> for Number {
fn sum<I>(iter: I) -> Number
where
I: Iterator<Item = &'a Self>,
{
iter.fold(Number::Int(0), |a, b| &a + b)
}
}
impl Product<Self> for Number {
fn product<I>(iter: I) -> Number
where
I: Iterator<Item = Self>,
{
iter.fold(Number::Int(1), |a, b| a * b)
}
}
impl<'a> Product<&'a Self> for Number {
fn product<I>(iter: I) -> Number
where
I: Iterator<Item = &'a Self>,
{
iter.fold(Number::Int(1), |a, b| &a * b)
}
}
pub struct Sorted<T>(pub T);
pub trait Sort {
fn sorted(&mut self) -> Sorted<&Self>
where
Self: Sized;
}
impl Sort for Vec<Number> {
fn sorted(&mut self) -> Sorted<&Vec<Number>> {
self.sort();
Sorted(self)
}
}
impl ToFloat for Number {
fn to_float(&self) -> f64 {
Number::to_float(*self)
}
}
impl Encode<()> for Number {
fn encode<W: std::io::Write>(
&self,
w: &mut storekey::Writer<W>,
) -> std::result::Result<(), storekey::EncodeError> {
let slice = self.as_decimal_buf();
w.write_slice(&slice)?;
let kind = match self {
Number::Int(_) => NumberKind::Int,
Number::Float(_) => NumberKind::Float,
Number::Decimal(_) => NumberKind::Decimal,
};
Encode::<()>::encode(&kind, w)?;
Ok(())
}
}
impl<'de> BorrowDecode<'de, ()> for Number {
fn borrow_decode(
r: &mut storekey::BorrowReader<'de>,
) -> std::result::Result<Self, storekey::DecodeError> {
let slice = r.read_cow()?;
let kind: NumberKind = BorrowDecode::<'de, ()>::borrow_decode(r)?;
Number::from_decimal_buf_kind(slice.as_ref(), kind)
.map_err(|_| storekey::DecodeError::InvalidFormat)
}
}
impl Encode<IndexFormat> for Number {
fn encode<W: std::io::Write>(
&self,
w: &mut storekey::Writer<W>,
) -> std::result::Result<(), storekey::EncodeError> {
let slice = self.as_decimal_buf();
w.write_slice(&slice)
}
}
impl<'de> BorrowDecode<'de, IndexFormat> for Number {
fn borrow_decode(
r: &mut storekey::BorrowReader<'de>,
) -> std::result::Result<Self, storekey::DecodeError> {
let slice = r.read_cow()?;
Number::from_decimal_buf(slice.as_ref()).map_err(|_| storekey::DecodeError::InvalidFormat)
}
}
pub trait DecimalExt {
fn from_str_normalized(s: &str) -> Result<Self, rust_decimal::Error>
where
Self: Sized;
}
impl DecimalExt for Decimal {
fn from_str_normalized(s: &str) -> Result<Decimal, rust_decimal::Error> {
#[allow(clippy::disallowed_methods)]
Ok(Decimal::from_str(s)?.normalize())
}
}
#[cfg(test)]
mod tests {
use std::cmp::Ordering;
use ahash::HashSet;
use rand::Rng;
use rand::seq::SliceRandom;
use rust_decimal::Decimal;
use rust_decimal::prelude::ToPrimitive;
use super::*;
#[test]
fn test_decimal_ext_from_str_normalized() {
let decimal = Decimal::from_str_normalized("0.0").unwrap();
assert_eq!(decimal.to_string(), "0");
assert_eq!(decimal.to_i64(), Some(0));
assert_eq!(decimal.to_f64(), Some(0.0));
let decimal = Decimal::from_str_normalized("123.456").unwrap();
assert_eq!(decimal.to_string(), "123.456");
assert_eq!(decimal.to_i64(), Some(123));
assert_eq!(decimal.to_f64(), Some(123.456));
let decimal =
Decimal::from_str_normalized("13.5719384719384719385639856394139476937756394756")
.unwrap();
assert_eq!(decimal.to_string(), "13.571938471938471938563985639");
assert_eq!(decimal.to_i64(), Some(13));
assert_eq!(decimal.to_f64(), Some(13.571_938_471_938_472));
}
#[test]
fn test_try_float_div() {
let (sum_one, count_one) = (Number::Int(5), Number::Int(2));
assert_eq!(sum_one.try_float_div(count_one).unwrap(), Number::Float(2.5));
let (sum_two, count_two) = (Number::Int(10), Number::Int(5));
assert_eq!(sum_two.try_float_div(count_two).unwrap(), Number::Int(2));
let (sum_three, count_three) = (Number::Float(6.3), Number::Int(3));
assert_eq!(sum_three.try_float_div(count_three).unwrap(), Number::Float(2.1));
}
#[test]
fn ord_test() {
let a = Number::Float(-f64::NAN);
let b = Number::Float(-f64::INFINITY);
let c = Number::Float(1f64);
let d = Number::Decimal(
Decimal::from_str_normalized("1.0000000000000000000000000002").unwrap(),
);
let e = Number::Decimal(Decimal::from_str_normalized("1.1").unwrap());
let f = Number::Float(1.1f64);
let g = Number::Float(1.5f64);
let h = Number::Decimal(Decimal::from_str_normalized("1.5").unwrap());
let i = Number::Float(f64::INFINITY);
let j = Number::Float(f64::NAN);
let original = vec![a, b, c, d, e, f, g, h, i, j];
let mut copy = original.clone();
let mut rng = rand::rng();
copy.shuffle(&mut rng);
copy.sort();
assert_eq!(original, copy);
}
#[test]
fn ord_fuzz() {
fn random_number() -> Number {
let mut rng = rand::rng();
match rng.random_range(0..3) {
0 => Number::Int(rng.random()),
1 => Number::Float(f64::from_bits(rng.random())),
_ => Number::Decimal(Number::Float(f64::from_bits(rng.random())).as_decimal()),
}
}
fn random_permutation(number: Number) -> Number {
let mut rng = rand::rng();
let value = match rng.random_range(0..4) {
0 => number + Number::from(rng.random::<f64>()),
1 if !matches!(number, Number::Int(i64::MIN)) => number * Number::from(-1),
2 => Number::Float(number.as_float().next_down()),
_ => number,
};
match rng.random_range(0..3) {
0 => Number::Int(value.as_int()),
1 => Number::Float(value.as_float()),
_ => Number::Decimal(value.as_decimal()),
}
}
fn assert_partial_ord(x: Number, y: Number) {
assert_eq!(x == y, x.partial_cmp(&y) == Some(Ordering::Equal), "{x:?} {y:?}");
assert_eq!(x.partial_cmp(&y), Some(x.cmp(&y)), "{x:?} {y:?}");
}
fn assert_consistent(a: Number, b: Number, c: Number) {
assert_partial_ord(a, b);
assert_partial_ord(b, c);
assert_partial_ord(c, a);
if a == b && b == c {
assert_eq!(a, c, "{a:?} {b:?} {c:?}");
}
if a != b && b == c {
assert_ne!(a, c, "{a:?} {b:?} {c:?}");
}
if a < b && b < c {
assert!(a < c, "{a:?} {b:?} {c:?}");
}
if a > b && b > c {
assert!(a > c, "{a:?} {b:?} {c:?}");
}
assert_eq!(a == b, b == a, "{a:?} {b:?}");
assert_eq!(a < b, b > a, "{a:?} {b:?}");
}
for _ in 0..100000 {
let base = random_number();
let a = random_permutation(base);
let b = random_permutation(a);
let c = random_permutation(b);
assert_consistent(a, b, c);
}
}
#[test]
fn serialised_ord_test() {
let ordering = [
Number::from(f64::NEG_INFINITY),
Number::from(f64::MIN),
Number::Int(i64::MIN),
Number::from(-1000),
Number::from(-100),
Number::from(-10),
Number::from(-1.5),
Number::from(-1),
Number::from(0),
Number::from(1),
Number::from(1.5),
Number::from(2),
Number::from(10),
Number::from(100),
Number::from(1000),
Number::from(i64::MAX),
Number::from(f64::MAX),
Number::from(f64::INFINITY),
Number::from(f64::NAN),
];
for window in ordering.windows(2) {
let n1 = &window[0];
let n2 = &window[1];
assert!(n1 < n2, "{n1:?} < {n2:?} (before serialization)");
let b1 = n1.as_decimal_buf();
let b2 = n2.as_decimal_buf();
assert!(b1 < b2, "{n1:?} < {n2:?} (after serialization) - {b1:?} < {b2:?}");
let r1 = Number::from_decimal_buf(&b1).unwrap();
let r2 = Number::from_decimal_buf(&b2).unwrap();
assert!(r1.eq(n1), "{r1:?} = {n1:?} (after deserialization)");
assert!(r2.eq(n2), "{r2:?} = {n2:?} (after deserialization)");
}
}
#[test]
fn serialised_test() {
let check = |numbers: &[Number]| {
let mut buffers = HashSet::default();
for n1 in numbers {
let b = n1.as_decimal_buf();
let n2 = Number::from_decimal_buf(&b).unwrap();
buffers.insert(b);
assert!(n1.eq(&n2), "{n1:?} = {n2:?} (after deserialization)");
}
assert_eq!(buffers.len(), 1, "{numbers:?}");
};
check(&[Number::Int(0), Number::Float(0.0), Number::Decimal(Decimal::ZERO)]);
check(&[Number::Int(1), Number::Float(1.0), Number::Decimal(Decimal::ONE)]);
check(&[Number::Int(-1), Number::Float(-1.0), Number::Decimal(Decimal::NEGATIVE_ONE)]);
check(&[Number::Float(1.5), Number::Decimal(Decimal::from_str_normalized("1.5").unwrap())]);
}
#[test]
fn as_int_lossless_accepts_valid() {
assert_eq!(Number::Int(0).as_int_lossless(), Some(0));
assert_eq!(Number::Int(i64::MIN).as_int_lossless(), Some(i64::MIN));
assert_eq!(Number::Int(i64::MAX).as_int_lossless(), Some(i64::MAX));
assert_eq!(Number::Float(0.0).as_int_lossless(), Some(0));
assert_eq!(Number::Float(-0.0).as_int_lossless(), Some(0));
assert_eq!(Number::Float(7.0).as_int_lossless(), Some(7));
assert_eq!(Number::Float(i64::MIN as f64).as_int_lossless(), Some(i64::MIN));
assert_eq!(Number::Decimal(Decimal::ZERO).as_int_lossless(), Some(0));
assert_eq!(Number::Decimal(Decimal::ONE).as_int_lossless(), Some(1));
assert_eq!(Number::Decimal(Decimal::NEGATIVE_ONE).as_int_lossless(), Some(-1));
}
#[test]
fn as_int_lossless_rejects_invalid() {
assert_eq!(Number::Float(1.5).as_int_lossless(), None);
assert_eq!(Number::Float(-1.5).as_int_lossless(), None);
assert_eq!(Number::Float(f64::NAN).as_int_lossless(), None);
assert_eq!(Number::Float(f64::INFINITY).as_int_lossless(), None);
assert_eq!(Number::Float(f64::NEG_INFINITY).as_int_lossless(), None);
assert_eq!(Number::Float(i64::MAX as f64).as_int_lossless(), None);
assert_eq!(Number::Float(1e30).as_int_lossless(), None);
assert_eq!(Number::Float(-1e30).as_int_lossless(), None);
assert_eq!(
Number::Decimal(Decimal::from_str_normalized("1.5").unwrap()).as_int_lossless(),
None,
);
}
#[test]
fn as_array_index_accepts_valid() {
assert_eq!(Number::Int(0).as_array_index(), Some(0));
assert_eq!(Number::Int(7).as_array_index(), Some(7));
assert_eq!(Number::Float(0.0).as_array_index(), Some(0));
assert_eq!(Number::Float(-0.0).as_array_index(), Some(0));
assert_eq!(Number::Float(3.0).as_array_index(), Some(3));
assert_eq!(Number::Decimal(Decimal::ZERO).as_array_index(), Some(0));
assert_eq!(Number::Decimal(Decimal::ONE).as_array_index(), Some(1));
}
#[test]
fn as_array_index_rejects_invalid() {
assert_eq!(Number::Int(-1).as_array_index(), None);
assert_eq!(Number::Int(i64::MIN).as_array_index(), None);
assert_eq!(Number::Float(-1.0).as_array_index(), None);
assert_eq!(Number::Float(1.5).as_array_index(), None);
assert_eq!(Number::Float(f64::NAN).as_array_index(), None);
assert_eq!(Number::Float(f64::INFINITY).as_array_index(), None);
assert_eq!(Number::Float(f64::NEG_INFINITY).as_array_index(), None);
assert_eq!(Number::Float(1e30).as_array_index(), None);
assert_eq!(Number::Decimal(Decimal::NEGATIVE_ONE).as_array_index(), None);
assert_eq!(
Number::Decimal(Decimal::from_str_normalized("1.5").unwrap()).as_array_index(),
None,
);
}
#[test]
fn fixed_int_is_unchanged_regardless_of_precision() {
assert_eq!(Number::Int(101).fixed(0), Number::Int(101));
assert_eq!(Number::Int(101).fixed(2), Number::Int(101));
assert_eq!(Number::Int(101).fixed(319), Number::Int(101));
assert_eq!(Number::Int(-7).fixed(5), Number::Int(-7));
}
#[test]
fn fixed_float_rounds_to_requested_precision() {
assert_eq!(Number::Float(101.5).fixed(2), Number::Float(101.5));
assert_eq!(Number::Float(101.1111111).fixed(2), Number::Float(101.11));
assert_eq!(
Number::Float(2.2250738585072014e-308).fixed(319),
Number::Float(2.22507385851e-308),
);
}
#[test]
fn fixed_float_preserves_non_finite() {
let Number::Float(nan) = Number::Float(f64::NAN).fixed(2) else {
panic!("expected Number::Float(NaN)");
};
assert!(nan.is_nan());
assert_eq!(Number::Float(f64::INFINITY).fixed(2), Number::Float(f64::INFINITY));
assert_eq!(Number::Float(f64::NEG_INFINITY).fixed(2), Number::Float(f64::NEG_INFINITY));
}
#[test]
fn fixed_decimal_uses_round_dp() {
let d = Decimal::from_str_normalized("1.2345").unwrap();
assert_eq!(
Number::Decimal(d).fixed(2),
Number::Decimal(Decimal::from_str_normalized("1.23").unwrap()),
);
}
}