mirror of
https://github.com/quickwit-oss/tantivy.git
synced 2026-10-06 11:52:40 +00:00
Add MAX function
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -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),
|
||||
|
||||
Reference in New Issue
Block a user