Add MAX function

This commit is contained in:
Paul Masurel
2026-08-18 15:37:06 +02:00
parent 4764a9cdc3
commit c58d3ee657
4 changed files with 217 additions and 2 deletions
+2 -2
View File
@@ -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
+2
View File
@@ -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<Function, DeserializeErro
"IS_NULL" => 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),
+204
View File
@@ -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<InferredTypeSet, TypeError> {
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<TypedExpr, CompileError> {
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::<Result<Vec<_>, _>>()?;
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<LoweredValue, CompileError> {
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<MaxFnCall> 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
);
}
}
+9
View File
@@ -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 => {
<LowerFnCall as FnCall>::call_with_types(args, target_type_set, context)
}
Function::Max => <MaxFnCall as FnCall>::call_with_types(args, target_type_set, context),
Function::Min => <MinFnCall as FnCall>::call_with_types(args, target_type_set, context),
Function::Multiply => {
<MultiplyFnCall as FnCall>::call_with_types(args, target_type_set, context)
@@ -211,6 +216,7 @@ impl Function {
Function::Lower => {
<LowerFnCall as FnCall>::infer_types(args, target_type, inferred_types)
}
Function::Max => <MaxFnCall as FnCall>::infer_types(args, target_type, inferred_types),
Function::Min => <MinFnCall as FnCall>::infer_types(args, target_type, inferred_types),
Function::Multiply => {
<MultiplyFnCall as FnCall>::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),