diff --git a/jitexpr/docs/FUNCTIONS.md b/jitexpr/docs/FUNCTIONS.md index 60065290c..63aa6e3d9 100644 --- a/jitexpr/docs/FUNCTIONS.md +++ b/jitexpr/docs/FUNCTIONS.md @@ -26,7 +26,7 @@ excluded from this pass. | 4 | `ADD` | done | | 5 | `SUBTRACT` | done | | 6 | `MULTIPLY` | done | -| 7 | `DIVIDE` | pending | +| 7 | `DIVIDE` | done | | 8 | `EQ` | done | | 9 | `GT` | pending | | 10 | `LT` | pending | @@ -76,4 +76,4 @@ excluded from this pass. | 79 | `SUBSTRING_COUNT` | pending | | 80 | `REGEXP_LIKE` | pending | -Progress: **11 / 48 in-scope** functions implemented; **7** functions are out-of-scope. +Progress: **12 / 48 in-scope** functions implemented; **7** functions are out-of-scope. diff --git a/jitexpr/src/ast/serialize.rs b/jitexpr/src/ast/serialize.rs index e2a457f7e..cad056593 100644 --- a/jitexpr/src/ast/serialize.rs +++ b/jitexpr/src/ast/serialize.rs @@ -120,6 +120,7 @@ fn function_name(function: Function) -> &'static str { match function { Function::And => "AND", Function::Add => "ADD", + Function::Divide => "DIVIDE", Function::Eq => "EQ", Function::IsNull => "IS_NULL", Function::IsNotNull => "IS_NOT_NULL", @@ -136,6 +137,7 @@ fn parse_function(name: &str, offset: usize) -> Result Ok(Function::And), "ADD" => Ok(Function::Add), + "DIVIDE" => Ok(Function::Divide), "EQ" => Ok(Function::Eq), "IS_NULL" => Ok(Function::IsNull), "IS_NOT_NULL" => Ok(Function::IsNotNull), diff --git a/jitexpr/src/functions/divide.rs b/jitexpr/src/functions/divide.rs new file mode 100644 index 000000000..1734b1553 --- /dev/null +++ b/jitexpr/src/functions/divide.rs @@ -0,0 +1,210 @@ +//! `DIVIDE` performs floating-point division. +//! +//! It accepts exactly two numeric arguments and always coerces both to `f64`, even when both are +//! integers. This avoids the surprising integer-division behavior explicitly called out by the +//! dd-go type checker (`BinaryExpression_DIVIDE` in `expression_type_checker.go`). jitexpr's +//! scalar API accepts its native numeric types; the reader's additional string-parsing coercion is +//! outside the current expression type model. +//! +//! The result is absent if either operand is absent or if the divisor is positive or negative +//! zero. Division by zero therefore yields NULL rather than infinity or NaN. Otherwise IEEE-754 +//! behavior applies, including propagation of NaN and infinities. This matches the zero guard and +//! null propagation in dd-go's `arithmeticDIVIDE` kernels (`vector_generated.go`). + +use std::collections::HashMap; + +use cranelift::prelude::{FloatCC, FunctionBuilder, InstBuilder, types}; + +use crate::ast::{Function, InferredTypeSet, TypeError, UntypedExpr}; +use crate::compile::{ + CompileError, CompileFnBuilder, LoweredValue, LoweringContext, TypedExpr, TypedExprAst, +}; +use crate::functions::{FnCall, FnCallEnum}; +use crate::types::VarType; + +#[derive(Clone, Debug, PartialEq)] +pub(crate) struct DivideFnCall { + pub(crate) args: Box<[TypedExpr]>, +} + +impl FnCall for DivideFnCall { + fn infer_types<'a>( + args: &'a [UntypedExpr], + target_type: InferredTypeSet, + inferred_types: &mut HashMap<&'a str, InferredTypeSet>, + ) -> Result { + if target_type.intersect(InferredTypeSet::F64).is_none() { + return Err(TypeError::WrongFunctionReturnType { + function: Function::Divide, + expected: target_type, + got: InferredTypeSet::F64, + }); + } + if args.len() != 2 { + return Err(TypeError::InvalidNumberOfArguments { + function: Function::Divide, + expected: 2, + got: args.len(), + }); + } + + for arg in args { + crate::ast::infer_types_aux(arg, InferredTypeSet::NUMERICAL, inferred_types)?; + } + Ok(InferredTypeSet::F64) + } + + fn call_with_types( + args: &[UntypedExpr], + target_type_set: InferredTypeSet, + context: &mut CompileFnBuilder<'_, '_>, + ) -> Result { + assert_eq!(args.len(), 2, "expected 2 args for DIVIDE"); + debug_assert!(target_type_set.contains(VarType::F64)); + + let typed_args = args + .iter() + .map(|arg| context.apply_types(arg, InferredTypeSet::F64)) + .collect::, _>>()?; + if typed_args + .iter() + .any(|typed_arg| typed_arg.return_type == VarType::None) + { + return Ok(TypedExpr::none()); + } + + Ok(TypedExpr { + return_type: VarType::F64, + ast: TypedExprAst::from_call(DivideFnCall { + args: typed_args.into_boxed_slice(), + }), + }) + } + + fn args_mut(&mut self) -> &mut [TypedExpr] { + &mut self.args + } + + fn emit_cranelift_ir( + &self, + return_type: VarType, + context: &mut LoweringContext<'_>, + builder: &mut FunctionBuilder<'_>, + ) -> Result { + if return_type != VarType::F64 { + return Err(CompileError::UnsupportedFunctionType { + function: Function::Divide, + return_type, + }); + } + + let dividend = context.compile_expr(&self.args[0], builder)?; + let divisor = context.compile_expr(&self.args[1], builder)?; + let value = builder.ins().fdiv(dividend.value, divisor.value); + let zero = builder.ins().f64const(0.0); + let divisor_is_zero = builder.ins().fcmp(FloatCC::Equal, divisor.value, zero); + let divisor_is_nonzero = builder.ins().bxor_imm_u(divisor_is_zero, 1); + let both_present = builder.ins().band(dividend.is_present, divisor.is_present); + let is_present = builder.ins().band(both_present, divisor_is_nonzero); + Ok(LoweredValue { + value, + is_present, + string_len: builder.ins().iconst(types::I64, 0), + }) + } +} + +impl From for FnCallEnum { + fn from(call: DivideFnCall) -> Self { + FnCallEnum::Divide(call) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::ast::{deserialize, infer_types}; + use crate::compile::compile; + use crate::types::VariableValue; + + #[test] + fn test_infer_types_requires_two_numeric_arguments_and_returns_float() { + let expression = deserialize("(DIVIDE left right)").unwrap(); + let inferred_types = infer_types(&expression).unwrap(); + assert_eq!( + inferred_types.get("left"), + Some(&InferredTypeSet::NUMERICAL) + ); + assert_eq!( + inferred_types.get("right"), + Some(&InferredTypeSet::NUMERICAL) + ); + + for expression in ["(DIVIDE 1i64)", "(DIVIDE 1i64 2i64 3i64)"] { + let expression = deserialize(expression).unwrap(); + assert!(matches!( + infer_types(&expression), + Err(TypeError::InvalidNumberOfArguments { + function: Function::Divide, + expected: 2, + .. + }) + )); + } + } + + #[test] + fn test_integer_inputs_use_floating_point_division() { + let expression = deserialize("(DIVIDE 5i64 2i64)").unwrap(); + let mut compiled = compile(&expression, &HashMap::new()).unwrap(); + + assert_eq!(compiled.result_type(), VarType::F64); + // SAFETY: The expression has no inputs and returns f64. + assert_eq!(unsafe { compiled.call(&[]).as_f64() }, Some(2.5)); + } + + #[test] + fn test_positive_and_negative_zero_divisors_return_none() { + for expression in ["(DIVIDE 1f64 0f64)", "(DIVIDE 1f64 -0f64)"] { + let expression = deserialize(expression).unwrap(); + let mut compiled = compile(&expression, &HashMap::new()).unwrap(); + // SAFETY: The expression has no inputs and returns nullable f64. + assert_eq!(unsafe { compiled.call(&[]).as_f64() }, None); + } + } + + #[test] + fn test_nan_divisor_remains_present() { + let expression = deserialize("(DIVIDE 1f64 nanf64)").unwrap(); + let mut compiled = compile(&expression, &HashMap::new()).unwrap(); + + // SAFETY: The expression has no inputs and returns f64. + assert!(unsafe { compiled.call(&[]).as_f64() }.unwrap().is_nan()); + } + + #[test] + fn test_runtime_null_propagation() { + let expression = deserialize("(DIVIDE left right)").unwrap(); + let variable_types = HashMap::from([("left", VarType::I64), ("right", VarType::U64)]); + let mut compiled = compile(&expression, &variable_types).unwrap(); + + // SAFETY: The input and output types match the compiled signature. + assert_eq!( + unsafe { + compiled + .call(&[VariableValue::some(9i64), VariableValue::some(2u64)]) + .as_f64() + }, + Some(4.5) + ); + // SAFETY: The input and output types match the compiled signature. + assert_eq!( + unsafe { + compiled + .call(&[VariableValue::none(), VariableValue::some(2u64)]) + .as_f64() + }, + None + ); + } +} diff --git a/jitexpr/src/functions/mod.rs b/jitexpr/src/functions/mod.rs index 9a761137f..d239066f6 100644 --- a/jitexpr/src/functions/mod.rs +++ b/jitexpr/src/functions/mod.rs @@ -1,5 +1,6 @@ mod add; mod and; +mod divide; mod eq; mod is_not_null; mod is_null; @@ -17,6 +18,7 @@ use cranelift::frontend::FunctionBuilder; pub(crate) use self::add::AddFnCall; pub(crate) use self::and::AndFnCall; +pub(crate) use self::divide::DivideFnCall; pub(crate) use self::eq::EqFnCall; pub(crate) use self::is_not_null::IsNotNullFnCall; pub(crate) use self::is_null::IsNullFnCall; @@ -40,6 +42,8 @@ pub enum Function { And, /// Adds zero or more numerical expressions. Add, + /// Divides two numeric arguments using floating-point arithmetic. + Divide, /// Compares two expressions for value equality. Eq, /// Tests whether an expression produced a present value. @@ -70,6 +74,9 @@ impl Function { match self { Function::And => ::call_with_types(args, target_type_set, context), Function::Add => ::call_with_types(args, target_type_set, context), + Function::Divide => { + ::call_with_types(args, target_type_set, context) + } Function::Eq => ::call_with_types(args, target_type_set, context), Function::IsNotNull => { ::call_with_types(args, target_type_set, context) @@ -103,6 +110,9 @@ impl Function { match self { Function::And => ::infer_types(args, target_type, inferred_types), Function::Add => ::infer_types(args, target_type, inferred_types), + Function::Divide => { + ::infer_types(args, target_type, inferred_types) + } Function::Eq => ::infer_types(args, target_type, inferred_types), Function::IsNotNull => { ::infer_types(args, target_type, inferred_types) @@ -139,6 +149,7 @@ impl Function { pub(crate) enum FnCallEnum { And(AndFnCall), Add(AddFnCall), + Divide(DivideFnCall), Eq(EqFnCall), IsNull(IsNullFnCall), IsNotNull(IsNotNullFnCall), @@ -155,6 +166,7 @@ impl FnCallEnum { match self { FnCallEnum::And(call) => call.args_mut(), FnCallEnum::Add(call) => call.args_mut(), + FnCallEnum::Divide(call) => call.args_mut(), FnCallEnum::Eq(call) => call.args_mut(), FnCallEnum::IsNull(call) => call.args_mut(), FnCallEnum::IsNotNull(call) => call.args_mut(), @@ -177,6 +189,7 @@ impl FnCallEnum { match self { FnCallEnum::And(call) => call.emit_cranelift_ir(return_type, context, builder), FnCallEnum::Add(call) => call.emit_cranelift_ir(return_type, context, builder), + FnCallEnum::Divide(call) => call.emit_cranelift_ir(return_type, context, builder), FnCallEnum::Eq(call) => call.emit_cranelift_ir(return_type, context, builder), FnCallEnum::IsNull(call) => call.emit_cranelift_ir(return_type, context, builder), FnCallEnum::IsNotNull(call) => call.emit_cranelift_ir(return_type, context, builder),