use crate::{Error, Result, Value};
use rust_decimal::Decimal;
use std::cmp::Ordering;
pub const NAN_TEXT: &str = "NaN";
pub const POS_INF_TEXT: &str = "Infinity";
pub const NEG_INF_TEXT: &str = "-Infinity";
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Special {
NaN,
PosInf,
NegInf,
}
impl Special {
pub fn to_text(self) -> &'static str {
match self {
Special::NaN => NAN_TEXT,
Special::PosInf => POS_INF_TEXT,
Special::NegInf => NEG_INF_TEXT,
}
}
pub fn to_value(self) -> Value {
Value::Numeric(self.to_text().to_string())
}
pub fn to_f64(self) -> f64 {
match self {
Special::NaN => f64::NAN,
Special::PosInf => f64::INFINITY,
Special::NegInf => f64::NEG_INFINITY,
}
}
fn from_neg(neg: bool) -> Special {
if neg {
Special::NegInf
} else {
Special::PosInf
}
}
fn is_neg(self) -> bool {
matches!(self, Special::NegInf)
}
pub fn abs(self) -> Special {
match self {
Special::NaN => Special::NaN,
_ => Special::PosInf,
}
}
fn rank(self) -> i8 {
match self {
Special::NegInf => 0,
Special::PosInf => 2,
Special::NaN => 3,
}
}
}
pub fn parse_special(s: &str) -> Option<Special> {
let t = s.trim();
if t.eq_ignore_ascii_case("nan") {
return Some(Special::NaN);
}
let (neg, rest) = if let Some(r) = t.strip_prefix('+') {
(false, r)
} else if let Some(r) = t.strip_prefix('-') {
(true, r)
} else {
(false, t)
};
if rest.eq_ignore_ascii_case("inf") || rest.eq_ignore_ascii_case("infinity") {
return Some(Special::from_neg(neg));
}
None
}
#[inline]
pub fn special_of(s: &str) -> Option<Special> {
match s {
"NaN" => Some(Special::NaN),
"Infinity" => Some(Special::PosInf),
"-Infinity" => Some(Special::NegInf),
_ => parse_special(s),
}
}
#[inline]
pub fn is_special(s: &str) -> bool {
special_of(s).is_some()
}
pub fn parse_numeric_text(s: &str) -> Option<String> {
if let Some(sp) = parse_special(s) {
return Some(sp.to_text().to_string());
}
s.trim().parse::<Decimal>().ok().map(|d| format!("{}", d))
}
pub fn float_to_numeric_text(f: f64) -> String {
if f.is_nan() {
NAN_TEXT.to_string()
} else if f.is_infinite() {
if f.is_sign_positive() {
POS_INF_TEXT.to_string()
} else {
NEG_INF_TEXT.to_string()
}
} else {
format!("{f}")
}
}
pub fn cmp_sort(a: &str, b: &str) -> Ordering {
match (special_of(a), special_of(b)) {
(None, None) => cmp_finite_numeric(a, b).unwrap_or_else(|| a.cmp(b)),
(sa, sb) => rank_of(sa).cmp(&rank_of(sb)),
}
}
fn cmp_finite_numeric(a: &str, b: &str) -> Option<Ordering> {
if let (Ok(ad), Ok(bd)) = (a.parse::<Decimal>(), b.parse::<Decimal>()) {
return Some(ad.cmp(&bd));
}
let af: f64 = a.parse().ok()?;
let bf: f64 = b.parse().ok()?;
af.partial_cmp(&bf)
}
fn rank_of(s: Option<Special>) -> i8 {
match s {
None => 1, Some(sp) => sp.rank(),
}
}
fn numeric_rank(v: &Value) -> Option<i8> {
Some(match v {
Value::Numeric(s) => rank_of(special_of(s)),
Value::Float4(f) => float_rank(f64::from(*f)),
Value::Float8(f) => float_rank(*f),
Value::Int2(_) | Value::Int4(_) | Value::Int8(_) => 1,
_ => return None,
})
}
fn float_rank(f: f64) -> i8 {
if f.is_nan() {
3
} else if f.is_infinite() {
if f.is_sign_positive() {
2
} else {
0
}
} else {
1
}
}
#[inline]
fn is_special_numeric(v: &Value) -> bool {
matches!(v, Value::Numeric(s) if is_special(s))
}
pub fn special_operator_cmp(left: &Value, right: &Value) -> Option<Ordering> {
if !is_special_numeric(left) && !is_special_numeric(right) {
return None;
}
let lr = numeric_rank(left)?;
let rr = numeric_rank(right)?;
Some(lr.cmp(&rr))
}
#[derive(Debug, Clone, Copy)]
pub enum ArithOp {
Add,
Sub,
Mul,
Div,
}
#[derive(Debug, Clone, Copy)]
enum Operand {
Special(Special),
Finite { neg: bool, zero: bool },
}
fn arith_operand(v: &Value) -> Option<Operand> {
Some(match v {
Value::Numeric(s) => match special_of(s) {
Some(sp) => Operand::Special(sp),
None => {
let d = s.parse::<Decimal>().ok()?;
let zero = d == Decimal::from(0);
Operand::Finite {
neg: d.is_sign_negative() && !zero,
zero,
}
}
},
Value::Int2(i) => Operand::Finite {
neg: *i < 0,
zero: *i == 0,
},
Value::Int4(i) => Operand::Finite {
neg: *i < 0,
zero: *i == 0,
},
Value::Int8(i) => Operand::Finite {
neg: *i < 0,
zero: *i == 0,
},
_ => return None,
})
}
pub fn special_arith(op: ArithOp, left: &Value, right: &Value) -> Result<Option<Value>> {
if !is_special_numeric(left) && !is_special_numeric(right) {
return Ok(None);
}
if matches!(left, Value::Float4(_) | Value::Float8(_)) || matches!(right, Value::Float4(_) | Value::Float8(_)) {
return Ok(None);
}
let (l, r) = match (arith_operand(left), arith_operand(right)) {
(Some(l), Some(r)) => (l, r),
_ => return Ok(None),
};
compute_special(op, l, r).map(Some)
}
fn compute_special(op: ArithOp, l: Operand, r: Operand) -> Result<Value> {
use Operand::{Finite, Special as Sp};
if matches!(l, Sp(Special::NaN)) || matches!(r, Sp(Special::NaN)) {
return Ok(Special::NaN.to_value());
}
let out: Special = match op {
ArithOp::Add => match (l, r) {
(Sp(a), Sp(b)) => {
if a == b {
a } else {
Special::NaN }
}
(Sp(a), Finite { .. }) | (Finite { .. }, Sp(a)) => a,
(Finite { .. }, Finite { .. }) => Special::NaN,
},
ArithOp::Sub => match (l, r) {
(Sp(a), Sp(b)) => {
if a == b {
Special::NaN } else {
a }
}
(Sp(a), Finite { .. }) => a, (Finite { .. }, Sp(b)) => Special::from_neg(!b.is_neg()), (Finite { .. }, Finite { .. }) => Special::NaN,
},
ArithOp::Mul => match (l, r) {
(Sp(a), Sp(b)) => Special::from_neg(a.is_neg() ^ b.is_neg()),
(Sp(a), Finite { zero, neg }) | (Finite { zero, neg }, Sp(a)) => {
if zero {
Special::NaN } else {
Special::from_neg(a.is_neg() ^ neg)
}
}
(Finite { .. }, Finite { .. }) => Special::NaN,
},
ArithOp::Div => match (l, r) {
(Sp(_), Sp(_)) => Special::NaN, (Sp(a), Finite { zero, neg }) => {
if zero {
return Err(Error::query_execution("division by zero"));
}
Special::from_neg(a.is_neg() ^ neg)
}
(Finite { .. }, Sp(_)) => return Ok(Value::Numeric("0".to_string())), (Finite { .. }, Finite { .. }) => Special::NaN,
},
};
Ok(out.to_value())
}
pub fn fold_sum_specials<'a, I>(values: I) -> Option<Special>
where
I: IntoIterator<Item = &'a Value>,
{
let mut acc: Option<Special> = None;
for v in values {
if let Value::Numeric(s) = v {
if let Some(sp) = special_of(s) {
acc = Some(combine_sum(acc, sp));
}
}
}
acc
}
pub fn combine_sum(acc: Option<Special>, incoming: Special) -> Special {
match acc {
None => incoming,
Some(cur) => {
if cur == Special::NaN || incoming == Special::NaN {
Special::NaN
} else if cur == incoming {
cur } else {
Special::NaN }
}
}
}
pub fn cmp_numeric_pushdown(a: &str, b: &str) -> Option<Ordering> {
match (special_of(a), special_of(b)) {
(None, None) => cmp_finite_numeric(a, b),
(sa, sb) => Some(rank_of(sa).cmp(&rank_of(sb))),
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used)]
mod tests {
use super::*;
#[test]
fn parses_special_case_and_sign_insensitively() {
assert_eq!(parse_numeric_text("NaN").as_deref(), Some("NaN"));
assert_eq!(parse_numeric_text("nan").as_deref(), Some("NaN"));
assert_eq!(parse_numeric_text(" NAN ").as_deref(), Some("NaN"));
assert_eq!(parse_numeric_text("Infinity").as_deref(), Some("Infinity"));
assert_eq!(parse_numeric_text("inf").as_deref(), Some("Infinity"));
assert_eq!(parse_numeric_text("+infinity").as_deref(), Some("Infinity"));
assert_eq!(parse_numeric_text("-Infinity").as_deref(), Some("-Infinity"));
assert_eq!(parse_numeric_text("-inf").as_deref(), Some("-Infinity"));
assert_eq!(parse_numeric_text("6").as_deref(), Some("6"));
assert_eq!(parse_numeric_text(" 6.50 ").as_deref(), Some("6.50"));
assert_eq!(parse_numeric_text("NotANumber"), None);
assert_eq!(parse_numeric_text("infi"), None);
}
#[test]
fn nan_sorts_greatest_and_equals_itself() {
assert_eq!(cmp_sort("NaN", "Infinity"), Ordering::Greater);
assert_eq!(cmp_sort("NaN", "5"), Ordering::Greater);
assert_eq!(cmp_sort("NaN", "NaN"), Ordering::Equal);
assert_eq!(cmp_sort("Infinity", "5"), Ordering::Greater);
assert_eq!(cmp_sort("-Infinity", "-999999"), Ordering::Less);
assert_eq!(cmp_sort("-Infinity", "-Infinity"), Ordering::Equal);
}
#[test]
fn cmp_sort_orders_finite_numerically_not_lexicographically() {
assert_eq!(cmp_sort("2", "10"), Ordering::Less);
assert_eq!(cmp_sort("10", "2"), Ordering::Greater);
assert_eq!(cmp_sort("9", "10"), Ordering::Less);
assert_eq!(cmp_sort("25", "100"), Ordering::Less);
assert_eq!(cmp_sort("100", "9"), Ordering::Greater);
let mut v = vec!["2", "10", "9", "100", "25"];
v.sort_by(|a, b| cmp_sort(a, b));
assert_eq!(v, vec!["2", "9", "10", "25", "100"]);
assert_eq!(cmp_sort("1.0", "1.00"), Ordering::Equal);
assert_eq!(cmp_sort("2.5", "2.50"), Ordering::Equal);
assert_eq!(cmp_sort("-10", "-2"), Ordering::Less);
assert_eq!(cmp_sort("0.1", "0.09"), Ordering::Greater);
assert_eq!(cmp_sort("100", "Infinity"), Ordering::Less);
assert_eq!(cmp_sort("100", "-Infinity"), Ordering::Greater);
assert_eq!(cmp_sort("100", "NaN"), Ordering::Less);
}
#[test]
fn cmp_numeric_pushdown_orders_finite_numerically() {
assert_eq!(cmp_numeric_pushdown("2", "10"), Some(Ordering::Less));
assert_eq!(cmp_numeric_pushdown("100", "25"), Some(Ordering::Greater));
assert_eq!(cmp_numeric_pushdown("1.0", "1.00"), Some(Ordering::Equal));
assert_eq!(cmp_numeric_pushdown("-Infinity", "5"), Some(Ordering::Less));
assert_eq!(cmp_numeric_pushdown("NaN", "Infinity"), Some(Ordering::Greater));
}
#[test]
fn operator_cmp_matches_pg() {
assert_eq!(
special_operator_cmp(&Value::Numeric("NaN".into()), &Value::Numeric("NaN".into())),
Some(Ordering::Equal)
);
assert_eq!(
special_operator_cmp(&Value::Numeric("NaN".into()), &Value::Int4(5)),
Some(Ordering::Greater)
);
assert_eq!(
special_operator_cmp(&Value::Numeric("-Infinity".into()), &Value::Int4(5)),
Some(Ordering::Less)
);
assert_eq!(
special_operator_cmp(&Value::Numeric("NaN".into()), &Value::Float8(f64::NAN)),
Some(Ordering::Equal)
);
assert_eq!(
special_operator_cmp(&Value::Numeric("1".into()), &Value::Numeric("2".into())),
None
);
}
#[test]
fn arithmetic_follows_pg_infinity_rules() {
let n = |s: &str| Value::Numeric(s.into());
let go = |op, a: Value, b: Value| special_arith(op, &a, &b).unwrap();
assert_eq!(go(ArithOp::Add, n("NaN"), n("5")), Some(n("NaN")));
assert_eq!(go(ArithOp::Add, n("Infinity"), n("Infinity")), Some(n("Infinity")));
assert_eq!(go(ArithOp::Add, n("Infinity"), n("-Infinity")), Some(n("NaN")));
assert_eq!(go(ArithOp::Sub, n("Infinity"), n("Infinity")), Some(n("NaN")));
assert_eq!(go(ArithOp::Sub, n("Infinity"), Value::Int4(3)), Some(n("Infinity")));
assert_eq!(go(ArithOp::Mul, n("Infinity"), Value::Int4(-2)), Some(n("-Infinity")));
assert_eq!(go(ArithOp::Mul, n("Infinity"), Value::Int4(0)), Some(n("NaN")));
assert_eq!(go(ArithOp::Div, n("Infinity"), n("Infinity")), Some(n("NaN")));
assert_eq!(go(ArithOp::Div, Value::Int4(5), n("Infinity")), Some(n("0")));
assert_eq!(go(ArithOp::Div, n("NaN"), Value::Int4(0)), Some(n("NaN")));
assert!(special_arith(ArithOp::Div, &n("Infinity"), &Value::Int4(0)).is_err());
assert_eq!(go(ArithOp::Add, n("1"), n("2")), None);
assert_eq!(go(ArithOp::Add, n("NaN"), Value::Float8(1.0)), None);
}
#[test]
fn sum_fold_combines_specials() {
let vals = vec![
Value::Numeric("1".into()),
Value::Numeric("NaN".into()),
Value::Numeric("2".into()),
];
assert_eq!(fold_sum_specials(&vals), Some(Special::NaN));
let inf = vec![Value::Int4(1), Value::Numeric("Infinity".into())];
assert_eq!(fold_sum_specials(&inf), Some(Special::PosInf));
let mixed = vec![Value::Numeric("Infinity".into()), Value::Numeric("-Infinity".into())];
assert_eq!(fold_sum_specials(&mixed), Some(Special::NaN));
let finite = vec![Value::Int4(1), Value::Numeric("2".into())];
assert_eq!(fold_sum_specials(&finite), None);
}
}