diff --git a/jitexpr/docs/FUNCTIONS.md b/jitexpr/docs/FUNCTIONS.md index 511274b6a..d2a1d05ae 100644 --- a/jitexpr/docs/FUNCTIONS.md +++ b/jitexpr/docs/FUNCTIONS.md @@ -50,7 +50,7 @@ excluded from this pass, or is complex enough to warrant a separate implementati | 34 | `POW` | done | | 35 | `SQRT` | out-of-scope | | 38 | `MIN` | done | -| 39 | `MAX` | pending | +| 39 | `MAX` | done | | 40 | `LEFT` | pending | | 41 | `RIGHT` | pending | | 42 | `SUBSTRING` | pending | @@ -76,7 +76,7 @@ excluded from this pass, or is complex enough to warrant a separate implementati | 79 | `SUBSTRING_COUNT` | pending | | 80 | `REGEXP_LIKE` | pending | -Progress: **24 / 41 in-scope** functions implemented; **14** functions are out-of-scope. +Progress: **25 / 41 in-scope** functions implemented; **14** functions are out-of-scope. ## Deferred implementation notes diff --git a/jitexpr/src/ast/serialize.rs b/jitexpr/src/ast/serialize.rs index 06127a307..30ffce093 100644 --- a/jitexpr/src/ast/serialize.rs +++ b/jitexpr/src/ast/serialize.rs @@ -132,6 +132,7 @@ fn function_name(function: Function) -> &'static str { Function::IsNull => "IS_NULL", Function::IsNotNull => "IS_NOT_NULL", Function::Lower => "LOWER", + Function::Max => "MAX", Function::Min => "MIN", Function::Multiply => "MULTIPLY", Function::Neq => "NEQ", @@ -161,6 +162,7 @@ fn parse_function(name: &str, offset: usize) -> Result Ok(Function::IsNull), "IS_NOT_NULL" => Ok(Function::IsNotNull), "LOWER" => Ok(Function::Lower), + "MAX" => Ok(Function::Max), "MIN" => Ok(Function::Min), "MULTIPLY" => Ok(Function::Multiply), "NEQ" => Ok(Function::Neq), diff --git a/jitexpr/src/functions/max.rs b/jitexpr/src/functions/max.rs new file mode 100644 index 000000000..ea5c5109a --- /dev/null +++ b/jitexpr/src/functions/max.rs @@ -0,0 +1,204 @@ +//! `MAX` returns the greatest scalar value among one or more numeric arguments. +//! +//! This is dd-go's `MAX_EXPR`, not an aggregate. All arguments are coerced to one numeric type and +//! the result has that type. Any null argument makes the result null. The float implementation +//! starts at negative infinity and replaces it only on `>`, so NaNs are ignored; if every argument +//! is NaN, the result is negative infinity. Production also scans multivalued arguments; arrays are +//! outside jitexpr's scalar model. + +use std::collections::HashMap; + +use cranelift::frontend::FunctionBuilder; +use cranelift::prelude::{FloatCC, InstBuilder, IntCC, types}; + +use super::add::{is_numerical, select_return_type, with_float_fallback}; +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 MaxFnCall { + args: Box<[TypedExpr]>, +} + +impl FnCall for MaxFnCall { + fn infer_types<'a>( + args: &'a [UntypedExpr], + target_type: InferredTypeSet, + inferred_types: &mut HashMap<&'a str, InferredTypeSet>, + ) -> Result { + if target_type.intersect(InferredTypeSet::NUMERICAL).is_none() { + return Err(TypeError::WrongFunctionReturnType { + function: Function::Max, + expected: target_type, + got: InferredTypeSet::NUMERICAL, + }); + } + if args.is_empty() { + return Err(TypeError::InvalidNumberOfArguments { + function: Function::Max, + expected: 1, + got: 0, + }); + } + let mut return_types = InferredTypeSet::NUMERICAL; + for arg in args { + return_types = return_types.intersect(crate::ast::infer_types_aux( + arg, + InferredTypeSet::NUMERICAL, + inferred_types, + )?); + } + let result = with_float_fallback(return_types).intersect(target_type); + if result.is_none() { + return Err(TypeError::WrongFunctionReturnType { + function: Function::Max, + expected: target_type, + got: return_types, + }); + } + Ok(result) + } + + fn call_with_types( + args: &[UntypedExpr], + target_type_set: InferredTypeSet, + context: &mut CompileFnBuilder<'_, '_>, + ) -> Result { + assert!(!args.is_empty(), "expected at least 1 arg for MAX"); + let mut return_types = InferredTypeSet::NUMERICAL.intersect(target_type_set); + for arg in args { + return_types = return_types.intersect(crate::ast::infer_type_with_variable_types( + arg, + InferredTypeSet::NUMERICAL, + context.variable_types(), + )?); + } + let return_type = select_return_type(with_float_fallback(return_types)); + let args = args + .iter() + .map(|arg| context.apply_types(arg, InferredTypeSet::singleton(return_type))) + .collect::, _>>()?; + if args.iter().any(|arg| !is_numerical(arg.return_type)) { + return Ok(TypedExpr::none()); + } + Ok(TypedExpr { + return_type, + ast: TypedExprAst::from_call(MaxFnCall { + args: 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 { + let mut value = match return_type { + VarType::I64 => builder.ins().iconst(types::I64, i64::MIN), + VarType::U64 => builder.ins().iconst(types::I64, 0), + VarType::F64 => builder.ins().f64const(f64::NEG_INFINITY), + _ => { + return Err(CompileError::UnsupportedFunctionType { + function: Function::Max, + return_type, + }); + } + }; + let mut is_present = builder.ins().iconst(types::I8, 1); + for arg in &self.args { + let arg = context.compile_expr(arg, builder)?; + let is_greater = match return_type { + VarType::I64 => builder + .ins() + .icmp(IntCC::SignedGreaterThan, arg.value, value), + VarType::U64 => builder + .ins() + .icmp(IntCC::UnsignedGreaterThan, arg.value, value), + VarType::F64 => builder.ins().fcmp(FloatCC::GreaterThan, arg.value, value), + _ => unreachable!(), + }; + value = builder.ins().select(is_greater, arg.value, value); + is_present = builder.ins().band(is_present, arg.is_present); + } + Ok(LoweredValue { + value, + is_present, + string_len: builder.ins().iconst(types::I64, 0), + }) + } +} + +impl From for FnCallEnum { + fn from(call: MaxFnCall) -> Self { + FnCallEnum::Max(call) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::ast::{deserialize, infer_types}; + use crate::compile::compile; + use crate::types::VariableValue; + + #[test] + fn test_signature_and_variadic_values() { + let expression = deserialize("(MAX a b c)").unwrap(); + let inferred = infer_types(&expression).unwrap(); + for name in ["a", "b", "c"] { + assert_eq!(inferred.get(name), Some(&InferredTypeSet::NUMERICAL)); + } + let empty = deserialize("(MAX)").unwrap(); + assert!(matches!( + infer_types(&empty), + Err(TypeError::InvalidNumberOfArguments { + function: Function::Max, + expected: 1, + got: 0 + }) + )); + let expression = deserialize("(MAX 7i64 -3i64 12i64)").unwrap(); + let mut compiled = compile(&expression, &HashMap::new()).unwrap(); + assert_eq!(unsafe { compiled.call(&[]).as_i64() }, Some(12)); + let expression = deserialize("(MAX 7u64 13u64)").unwrap(); + let mut compiled = compile(&expression, &HashMap::new()).unwrap(); + assert_eq!(unsafe { compiled.call(&[]).as_u64() }, Some(13)); + } + + #[test] + fn test_float_nan_and_null_behavior() { + let expression = deserialize("(MAX nanf64 3f64 -2f64)").unwrap(); + let mut compiled = compile(&expression, &HashMap::new()).unwrap(); + assert_eq!(unsafe { compiled.call(&[]).as_f64() }, Some(3.0)); + let expression = deserialize("(MAX nanf64)").unwrap(); + let mut compiled = compile(&expression, &HashMap::new()).unwrap(); + assert_eq!( + unsafe { compiled.call(&[]).as_f64() }, + Some(f64::NEG_INFINITY) + ); + let expression = deserialize("(MAX left right)").unwrap(); + let mut compiled = compile( + &expression, + &HashMap::from([("left", VarType::F64), ("right", VarType::F64)]), + ) + .unwrap(); + assert_eq!( + unsafe { + compiled + .call(&[VariableValue::some(1.0f64), VariableValue::none()]) + .as_f64() + }, + None + ); + } +} diff --git a/jitexpr/src/functions/mod.rs b/jitexpr/src/functions/mod.rs index 83ca09575..d1faaf2b8 100644 --- a/jitexpr/src/functions/mod.rs +++ b/jitexpr/src/functions/mod.rs @@ -13,6 +13,7 @@ mod is_null; mod lower; mod lt; mod lt_eq; +mod max; mod min; mod multiply; mod native_function; @@ -43,6 +44,7 @@ pub(crate) use self::is_null::IsNullFnCall; pub(crate) use self::lower::LowerFnCall; pub(crate) use self::lt::LtFnCall; pub(crate) use self::lt_eq::LtEqFnCall; +pub(crate) use self::max::MaxFnCall; pub(crate) use self::min::MinFnCall; pub(crate) use self::multiply::MultiplyFnCall; pub(crate) use self::native_function::{ @@ -91,6 +93,8 @@ pub enum Function { IsNull, /// Constructs the Unicode-lowercase form of a string. Lower, + /// Returns the greatest of one or more scalar numbers. + Max, /// Returns the least of one or more scalar numbers. Min, /// Multiplies two numeric arguments. @@ -151,6 +155,7 @@ impl Function { Function::Lower => { ::call_with_types(args, target_type_set, context) } + Function::Max => ::call_with_types(args, target_type_set, context), Function::Min => ::call_with_types(args, target_type_set, context), Function::Multiply => { ::call_with_types(args, target_type_set, context) @@ -211,6 +216,7 @@ impl Function { Function::Lower => { ::infer_types(args, target_type, inferred_types) } + Function::Max => ::infer_types(args, target_type, inferred_types), Function::Min => ::infer_types(args, target_type, inferred_types), Function::Multiply => { ::infer_types(args, target_type, inferred_types) @@ -258,6 +264,7 @@ pub(crate) enum FnCallEnum { IsNull(IsNullFnCall), IsNotNull(IsNotNullFnCall), Lower(LowerFnCall), + Max(MaxFnCall), Min(MinFnCall), Multiply(MultiplyFnCall), Neq(NeqFnCall), @@ -287,6 +294,7 @@ impl FnCallEnum { FnCallEnum::IsNull(call) => call.args_mut(), FnCallEnum::IsNotNull(call) => call.args_mut(), FnCallEnum::Lower(call) => call.args_mut(), + FnCallEnum::Max(call) => call.args_mut(), FnCallEnum::Min(call) => call.args_mut(), FnCallEnum::Multiply(call) => call.args_mut(), FnCallEnum::Neq(call) => call.args_mut(), @@ -322,6 +330,7 @@ impl FnCallEnum { FnCallEnum::IsNull(call) => call.emit_cranelift_ir(return_type, context, builder), FnCallEnum::IsNotNull(call) => call.emit_cranelift_ir(return_type, context, builder), FnCallEnum::Lower(call) => call.emit_cranelift_ir(return_type, context, builder), + FnCallEnum::Max(call) => call.emit_cranelift_ir(return_type, context, builder), FnCallEnum::Min(call) => call.emit_cranelift_ir(return_type, context, builder), FnCallEnum::Multiply(call) => call.emit_cranelift_ir(return_type, context, builder), FnCallEnum::Neq(call) => call.emit_cranelift_ir(return_type, context, builder),