use super::super::select_ast::{ArithmeticOperator, ComparisonOperator};
use crate::{types::Value, Error, Result};
pub(super) fn values_equal(a: &Value, b: &Value) -> bool {
if a == b {
return true;
}
if let (Some(x), Some(y)) = (as_integral_i128(a), as_integral_i128(b)) {
return x == y;
}
if same_numeric_family(a, b) {
if let (Some(x), Some(y)) = (a.as_f64(), b.as_f64()) {
return x == y;
}
}
false
}
pub(super) fn as_integral_i128(v: &Value) -> Option<i128> {
match v {
Value::Integer(i) => Some(*i as i128),
Value::BigInt(i) => Some(*i as i128),
Value::Counter(i) => Some(*i as i128),
Value::TinyInt(i) => Some(*i as i128),
Value::SmallInt(i) => Some(*i as i128),
_ => None,
}
}
pub(in crate::query) fn is_nan_value(v: &Value) -> bool {
match v {
Value::Float(x) => x.is_nan(),
Value::Float32(x) => x.is_nan(),
_ => false,
}
}
pub(super) fn same_numeric_family(a: &Value, b: &Value) -> bool {
a.as_f64().is_some() && b.as_f64().is_some()
}
pub(super) fn compare_values_ordering(a: &Value, b: &Value) -> std::cmp::Ordering {
try_compare_values(a, b).unwrap_or(std::cmp::Ordering::Equal)
}
pub(super) fn try_compare_values(a: &Value, b: &Value) -> Result<std::cmp::Ordering> {
if same_numeric_family(a, b) {
if let (Some(x), Some(y)) = (a.as_f64(), b.as_f64()) {
return Ok(crate::float_cmp::cassandra_double_cmp(x, y));
}
}
if std::mem::discriminant(a) == std::mem::discriminant(b) {
return a.partial_cmp(b).ok_or_else(|| {
Error::query_execution("Cannot compare incompatible types".to_string())
});
}
tracing::debug!(
"Cannot compare values of incompatible types: {:?} vs {:?}",
a.data_type(),
b.data_type()
);
Err(Error::query_execution(
"Cannot compare incompatible types".to_string(),
))
}
pub(super) fn try_compare_values_predicate(
a: &Value,
b: &Value,
) -> Result<Option<std::cmp::Ordering>> {
if is_nan_value(a) || is_nan_value(b) {
return Ok(None);
}
if let (Some(x), Some(y)) = (as_integral_i128(a), as_integral_i128(b)) {
return Ok(Some(x.cmp(&y)));
}
try_compare_values(a, b).map(Some)
}
pub(super) fn compare_values_ordering_predicate(
a: &Value,
b: &Value,
) -> Option<std::cmp::Ordering> {
match try_compare_values_predicate(a, b) {
Ok(ordering) => ordering,
Err(_) => Some(std::cmp::Ordering::Equal),
}
}
pub(super) fn eval_scalar_comparison(
op: &ComparisonOperator,
left: &Value,
right: &Value,
) -> Result<bool> {
use ComparisonOperator::*;
Ok(match op {
Equal => values_equal(left, right),
NotEqual => !values_equal(left, right),
LessThan => try_compare_values_predicate(left, right)?.is_some_and(|o| o.is_lt()),
LessThanOrEqual => try_compare_values_predicate(left, right)?.is_some_and(|o| o.is_le()),
GreaterThan => try_compare_values_predicate(left, right)?.is_some_and(|o| o.is_gt()),
GreaterThanOrEqual => try_compare_values_predicate(left, right)?.is_some_and(|o| o.is_ge()),
other => {
return Err(Error::query_execution(format!(
"operator {:?} is not a scalar comparison",
other
)))
}
})
}
pub(super) fn eval_arithmetic(op: &ArithmeticOperator, left: Value, right: Value) -> Result<Value> {
use ArithmeticOperator::*;
macro_rules! int_op {
($a:expr, $b:expr, $ctor:expr) => {
match op {
Add => Ok($ctor($a + $b)),
Subtract => Ok($ctor($a - $b)),
Multiply => Ok($ctor($a * $b)),
Divide => {
if $b == 0 {
Err(Error::query_execution("Division by zero".to_string()))
} else {
Ok($ctor($a / $b))
}
}
Modulo => {
if $b == 0 {
Err(Error::query_execution("Modulo by zero".to_string()))
} else {
Ok($ctor($a % $b))
}
}
}
};
}
match (left, right) {
(Value::Integer(a), Value::Integer(b)) => int_op!(a, b, Value::Integer),
(Value::BigInt(a), Value::BigInt(b)) => int_op!(a, b, Value::BigInt),
(Value::Float(a), Value::Float(b)) => match op {
Add => Ok(Value::Float(a + b)),
Subtract => Ok(Value::Float(a - b)),
Multiply => Ok(Value::Float(a * b)),
Divide => Ok(Value::Float(a / b)),
Modulo => Ok(Value::Float(a % b)),
},
_ => Err(Error::query_execution(
"Incompatible types for arithmetic".to_string(),
)),
}
}
pub(super) fn const_arithmetic(
op: &ArithmeticOperator,
left: Value,
right: Value,
) -> Result<Value> {
use ArithmeticOperator::*;
if matches!(op, Modulo) {
return match (left, right) {
(Value::Integer(a), Value::Integer(b)) => {
eval_arithmetic(op, Value::Integer(a), Value::Integer(b))
}
(Value::BigInt(a), Value::BigInt(b)) => {
eval_arithmetic(op, Value::BigInt(a), Value::BigInt(b))
}
_ => Err(Error::query_execution(
"Modulo only supported for integers".to_string(),
)),
};
}
let verb = match op {
Add => "add",
Subtract => "subtract",
Multiply => "multiply",
Divide => "divide",
Modulo => unreachable!("handled above"),
};
match (left, right) {
(Value::Integer(a), Value::Integer(b)) => {
eval_arithmetic(op, Value::Integer(a), Value::Integer(b))
}
(Value::BigInt(a), Value::BigInt(b)) => {
eval_arithmetic(op, Value::BigInt(a), Value::BigInt(b))
}
(Value::Float(a), Value::Float(b)) => {
if matches!(op, Divide) && b == 0.0 {
return Err(Error::query_execution("Division by zero".to_string()));
}
eval_arithmetic(op, Value::Float(a), Value::Float(b))
}
_ => Err(Error::query_execution(format!(
"Cannot {} incompatible types",
verb
))),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_value_comparison() {
use std::cmp::Ordering;
assert_eq!(
try_compare_values(&Value::Integer(5), &Value::Integer(3)).unwrap(),
Ordering::Greater
);
assert_eq!(
try_compare_values(&Value::Integer(3), &Value::Integer(5)).unwrap(),
Ordering::Less
);
assert_eq!(
try_compare_values(&Value::Integer(5), &Value::Integer(5)).unwrap(),
Ordering::Equal
);
}
#[test]
fn values_equal_distinguishes_large_i64_across_f64_boundary() {
let two53 = 1_i64 << 53; let plus1 = two53 + 1; assert_eq!(
two53 as f64, plus1 as f64,
"precondition: f64 collapses them"
);
assert!(
!values_equal(&Value::BigInt(two53), &Value::BigInt(plus1)),
"2^53 != 2^53 + 1 as bigint"
);
assert!(
values_equal(&Value::BigInt(plus1), &Value::BigInt(plus1)),
"identical large bigints are equal"
);
assert!(!values_equal(
&Value::BigInt(9_007_199_254_740_992),
&Value::BigInt(9_007_199_254_740_993),
));
assert!(values_equal(&Value::Integer(7), &Value::BigInt(7)));
assert!(values_equal(&Value::TinyInt(7), &Value::SmallInt(7)));
assert!(!values_equal(&Value::Integer(7), &Value::BigInt(8)));
assert!(values_equal(&Value::BigInt(2), &Value::Float(2.0)));
assert!(!values_equal(&Value::BigInt(2), &Value::Float(2.5)));
}
#[test]
fn predicate_ordering_distinguishes_large_i64_across_f64_boundary() {
use std::cmp::Ordering;
let two53 = 1_i64 << 53; let plus1 = two53 + 1;
assert_eq!(
try_compare_values_predicate(&Value::BigInt(plus1), &Value::BigInt(two53)).unwrap(),
Some(Ordering::Greater),
"9007199254740993 > 9007199254740992 must hold exactly"
);
assert_eq!(
compare_values_ordering_predicate(&Value::BigInt(plus1), &Value::BigInt(two53)),
Some(Ordering::Greater)
);
assert_eq!(
try_compare_values_predicate(&Value::BigInt(two53), &Value::BigInt(plus1)).unwrap(),
Some(Ordering::Less)
);
assert_eq!(
try_compare_values_predicate(&Value::BigInt(plus1), &Value::BigInt(plus1)).unwrap(),
Some(Ordering::Equal)
);
}
#[test]
fn nan_predicate_comparison_is_unknown_for_all_relations() {
let nan = Value::Float(f64::NAN);
let bound = Value::Float(1.5);
let cmp = try_compare_values_predicate(&nan, &bound).unwrap();
assert!(cmp.is_none(), "NaN vs 1.5 is UNKNOWN (Gt/Gte)");
assert!(
!cmp.is_some_and(|o| o.is_gt()),
"d > 1.5 with NaN is dropped"
);
assert!(
!cmp.is_some_and(|o| o.is_ge()),
"d >= 1.5 with NaN is dropped"
);
assert!(
!cmp.is_some_and(|o| o.is_lt()),
"d < 1.5 with NaN is dropped"
);
assert!(
!cmp.is_some_and(|o| o.is_le()),
"d <= 1.5 with NaN is dropped"
);
assert!(try_compare_values_predicate(&bound, &nan)
.unwrap()
.is_none());
assert!(try_compare_values_predicate(&nan, &nan).unwrap().is_none());
assert!(
try_compare_values_predicate(&Value::Float32(f32::NAN), &Value::Float32(1.5))
.unwrap()
.is_none()
);
assert!(!values_equal(&nan, &bound), "NaN = 1.5 is false");
assert!(!values_equal(&nan, &nan), "NaN = NaN is false (SQL)");
let ok = try_compare_values_predicate(&Value::Float(2.0), &bound).unwrap();
assert!(ok.is_some_and(|o| o.is_gt()), "2.0 > 1.5 holds");
}
#[test]
fn nan_ordering_total_order_unchanged_by_predicate_fix() {
use std::cmp::Ordering;
assert_eq!(
compare_values_ordering(&Value::Float(f64::NAN), &Value::Float(1.5)),
Ordering::Greater,
"sort order still puts NaN last (unchanged)"
);
assert_eq!(
compare_values_ordering_predicate(&Value::Float(2.0), &Value::Float(1.5)),
Some(Ordering::Greater)
);
assert!(
compare_values_ordering_predicate(&Value::Float(f64::NAN), &Value::Float(1.5))
.is_none()
);
}
#[test]
fn eval_scalar_comparison_dispatches_all_six_operators() {
use ComparisonOperator::*;
let five = Value::Integer(5);
let three = Value::Integer(3);
assert!(!eval_scalar_comparison(&Equal, &five, &three).unwrap());
assert!(eval_scalar_comparison(&Equal, &five, &five).unwrap());
assert!(eval_scalar_comparison(&NotEqual, &five, &three).unwrap());
assert!(!eval_scalar_comparison(&NotEqual, &five, &five).unwrap());
assert!(eval_scalar_comparison(&GreaterThan, &five, &three).unwrap());
assert!(!eval_scalar_comparison(&GreaterThan, &three, &five).unwrap());
assert!(eval_scalar_comparison(&GreaterThanOrEqual, &five, &five).unwrap());
assert!(!eval_scalar_comparison(&GreaterThanOrEqual, &three, &five).unwrap());
assert!(eval_scalar_comparison(&LessThan, &three, &five).unwrap());
assert!(!eval_scalar_comparison(&LessThan, &five, &three).unwrap());
assert!(eval_scalar_comparison(&LessThanOrEqual, &five, &five).unwrap());
assert!(!eval_scalar_comparison(&LessThanOrEqual, &five, &three).unwrap());
assert!(eval_scalar_comparison(&In, &five, &three).is_err());
}
#[test]
fn compare_values_ordering_double_matches_cassandra() {
use std::cmp::Ordering;
let f = Value::Float; assert_eq!(
compare_values_ordering(&f(f64::NAN), &f(f64::INFINITY)),
Ordering::Greater,
"NaN sorts after +Infinity"
);
assert_eq!(
compare_values_ordering(&f(f64::NAN), &f(f64::NAN)),
Ordering::Equal,
"two NaNs compare equal"
);
assert_eq!(
compare_values_ordering(&f(-0.0), &f(0.0)),
Ordering::Less,
"-0.0 < +0.0"
);
assert_eq!(compare_values_ordering(&f(1.0), &f(2.0)), Ordering::Less);
}
#[test]
fn order_by_double_sort_matches_oracle() {
let mut v = vec![
Value::Float(1.0),
Value::Float(f64::NAN),
Value::Float(-0.0),
Value::Float(0.0),
Value::Float(f64::NEG_INFINITY),
Value::Float(f64::INFINITY),
Value::Float(f64::NAN),
];
v.sort_by(compare_values_ordering);
let f = |i: usize| match v[i] {
Value::Float(x) => x,
_ => unreachable!(),
};
assert_eq!(f(0), f64::NEG_INFINITY);
assert!(f(1) == 0.0 && f(1).is_sign_negative(), "index 1 = -0.0");
assert!(f(2) == 0.0 && f(2).is_sign_positive(), "index 2 = +0.0");
assert_eq!(f(3), 1.0);
assert_eq!(f(4), f64::INFINITY);
assert!(f(5).is_nan() && f(6).is_nan(), "NaNs sort last");
}
#[test]
fn order_by_float32_sort_matches_oracle() {
let mut v = [
Value::Float32(f32::NAN),
Value::Float32(0.0),
Value::Float32(-0.0),
Value::Float32(f32::INFINITY),
];
v.sort_by(compare_values_ordering);
let f = |i: usize| match v[i] {
Value::Float32(x) => x,
_ => unreachable!(),
};
assert!(f(0) == 0.0 && f(0).is_sign_negative(), "index 0 = -0.0");
assert!(f(1) == 0.0 && f(1).is_sign_positive(), "index 1 = +0.0");
assert_eq!(f(2), f32::INFINITY);
assert!(f(3).is_nan(), "NaN sorts last");
}
#[test]
fn min_max_double_matches_cassandra() {
let data = [
Value::Float(f64::NAN),
Value::Float(3.0),
Value::Float(-0.0),
Value::Float(0.0),
Value::Float(-2.0),
];
let min = data
.iter()
.min_by(|a, b| compare_values_ordering(a, b))
.unwrap();
assert!(matches!(min, Value::Float(x) if *x == -2.0), "MIN = -2.0");
let max = data
.iter()
.max_by(|a, b| compare_values_ordering(a, b))
.unwrap();
assert!(
matches!(max, Value::Float(x) if x.is_nan()),
"MAX = NaN (sorts last)"
);
let zeros = [Value::Float(0.0), Value::Float(-0.0)];
let zmin = zeros
.iter()
.min_by(|a, b| compare_values_ordering(a, b))
.unwrap();
assert!(
matches!(zmin, Value::Float(x) if *x == 0.0 && x.is_sign_negative()),
"MIN of {{-0.0, +0.0}} = -0.0"
);
}
}