From ffbcb199d0a6517019a8afbc3a7954cefc83355d Mon Sep 17 00:00:00 2001 From: Paul Masurel Date: Wed, 19 Aug 2026 10:47:58 +0200 Subject: [PATCH] Add SQRT function --- jitexpr/docs/FUNCTIONS.md | 7 +- jitexpr/src/ast/serialize.rs | 2 + jitexpr/src/functions/mod.rs | 13 +++ jitexpr/src/functions/sqrt.rs | 177 ++++++++++++++++++++++++++++++++++ 4 files changed, 194 insertions(+), 5 deletions(-) create mode 100644 jitexpr/src/functions/sqrt.rs diff --git a/jitexpr/docs/FUNCTIONS.md b/jitexpr/docs/FUNCTIONS.md index 6124ec5bc..7c6a85d99 100644 --- a/jitexpr/docs/FUNCTIONS.md +++ b/jitexpr/docs/FUNCTIONS.md @@ -48,7 +48,7 @@ excluded from this pass, or is complex enough to warrant a separate implementati | 28 | `FLOOR` | out-of-scope | | 29 | `CEIL` | out-of-scope | | 34 | `POW` | done | -| 35 | `SQRT` | out-of-scope | +| 35 | `SQRT` | done | | 38 | `MIN` | done | | 39 | `MAX` | done | | 40 | `LEFT` | done | @@ -76,7 +76,7 @@ excluded from this pass, or is complex enough to warrant a separate implementati | 79 | `SUBSTRING_COUNT` | done | | 80 | `REGEXP_LIKE` | done | -Progress: **34 / 34 in-scope** functions implemented; **21** functions are out-of-scope. +Progress: **35 / 35 in-scope** functions implemented; **20** functions are out-of-scope. ## Deferred implementation notes @@ -99,9 +99,6 @@ Progress: **34 / 34 in-scope** functions implemented; **21** functions are out-o architecture-dependent. It is deferred pending an explicit production-parity policy. - `CEIL` shares `FLOOR`'s unwritten integer-output defect and architecture-dependent exceptional float-to-`int64` conversions, so it is deferred under the same parity policy. -- `SQRT` has a contradictory production type contract: the dd-go type checker returns the selected - input type (including integer), while the registry declares a `float64` output and both integer - and float kernels write `float64`. It is deferred until one of those contracts is chosen. - `COALESCE` uses dd-go's distinct n-ary common-type algorithm, including an ordinal fallback that can choose string where the binary `IF` unifier chooses numeric. Implementing it correctly needs a broader coercion-policy change rather than only a new call node. diff --git a/jitexpr/src/ast/serialize.rs b/jitexpr/src/ast/serialize.rs index b571c3daa..df50dcc5f 100644 --- a/jitexpr/src/ast/serialize.rs +++ b/jitexpr/src/ast/serialize.rs @@ -141,6 +141,7 @@ fn function_name(function: Function) -> &'static str { Function::Not => "NOT", Function::Or => "OR", Function::Pow => "POW", + Function::Sqrt => "SQRT", Function::RegexpExtract => "REGEXP_EXTRACT", Function::RegexpLike => "REGEXP_LIKE", Function::Right => "RIGHT", @@ -180,6 +181,7 @@ fn parse_function(name: &str, offset: usize) -> Result Ok(Function::Not), "OR" => Ok(Function::Or), "POW" => Ok(Function::Pow), + "SQRT" => Ok(Function::Sqrt), "REGEXP_EXTRACT" => Ok(Function::RegexpExtract), "REGEXP_LIKE" => Ok(Function::RegexpLike), "RIGHT" => Ok(Function::Right), diff --git a/jitexpr/src/functions/mod.rs b/jitexpr/src/functions/mod.rs index 4bbade4be..4550f2c11 100644 --- a/jitexpr/src/functions/mod.rs +++ b/jitexpr/src/functions/mod.rs @@ -28,6 +28,7 @@ mod regexp_like; mod right; mod split_after; mod split_before; +mod sqrt; mod substring; mod substring_count; mod subtract; @@ -70,6 +71,7 @@ pub(crate) use self::regexp_like::RegexpLikeFnCall; pub(crate) use self::right::RightFnCall; pub(crate) use self::split_after::SplitAfterFnCall; pub(crate) use self::split_before::SplitBeforeFnCall; +pub(crate) use self::sqrt::SqrtFnCall; pub(crate) use self::substring::SubstringFnCall; pub(crate) use self::substring_count::SubstringCountFnCall; pub(crate) use self::subtract::SubtractFnCall; @@ -129,6 +131,8 @@ pub enum Function { Or, /// Raises a numeric base to a numeric exponent and returns a float. Pow, + /// Returns the floating-point square root of a number, or null for a NaN result. + Sqrt, /// Extracts a capture group from a string using a constant regular expression. RegexpExtract, /// Tests whether a constant regular expression matches a string. @@ -204,6 +208,9 @@ impl Function { Function::Not => ::call_with_types(args, target_type_set, context), Function::Or => ::call_with_types(args, target_type_set, context), Function::Pow => ::call_with_types(args, target_type_set, context), + Function::Sqrt => { + ::call_with_types(args, target_type_set, context) + } Function::RegexpExtract => { ::call_with_types(args, target_type_set, context) } @@ -290,6 +297,9 @@ impl Function { Function::Not => ::infer_types(args, target_type, inferred_types), Function::Or => ::infer_types(args, target_type, inferred_types), Function::Pow => ::infer_types(args, target_type, inferred_types), + Function::Sqrt => { + ::infer_types(args, target_type, inferred_types) + } Function::RegexpExtract => { ::infer_types(args, target_type, inferred_types) } @@ -359,6 +369,7 @@ pub(crate) enum FnCallEnum { Not(NotFnCall), Or(OrFnCall), Pow(PowFnCall), + Sqrt(SqrtFnCall), RegexpExtract(RegexpExtractFnCall), RegexpLike(RegexpLikeFnCall), Right(RightFnCall), @@ -398,6 +409,7 @@ impl FnCallEnum { FnCallEnum::Not(call) => call.args_mut(), FnCallEnum::Or(call) => call.args_mut(), FnCallEnum::Pow(call) => call.args_mut(), + FnCallEnum::Sqrt(call) => call.args_mut(), FnCallEnum::RegexpExtract(call) => call.args_mut(), FnCallEnum::RegexpLike(call) => call.args_mut(), FnCallEnum::Right(call) => call.args_mut(), @@ -443,6 +455,7 @@ impl FnCallEnum { FnCallEnum::Not(call) => call.emit_cranelift_ir(return_type, context, builder), FnCallEnum::Or(call) => call.emit_cranelift_ir(return_type, context, builder), FnCallEnum::Pow(call) => call.emit_cranelift_ir(return_type, context, builder), + FnCallEnum::Sqrt(call) => call.emit_cranelift_ir(return_type, context, builder), FnCallEnum::RegexpExtract(call) => { call.emit_cranelift_ir(return_type, context, builder) } diff --git a/jitexpr/src/functions/sqrt.rs b/jitexpr/src/functions/sqrt.rs new file mode 100644 index 000000000..99236b2fa --- /dev/null +++ b/jitexpr/src/functions/sqrt.rs @@ -0,0 +1,177 @@ +//! `SQRT` computes the square root of one numeric argument. +//! +//! The argument is coerced to `f64`, and the result is always `f64`. Null input returns null. A +//! NaN result also returns null, covering NaN input and negative values other than negative zero. +//! Positive infinity and signed zero retain their IEEE-754 behavior. + +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 SqrtFnCall { + arg: Box, +} + +impl FnCall for SqrtFnCall { + 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::Sqrt, + expected: target_type, + got: InferredTypeSet::F64, + }); + } + if args.len() != 1 { + return Err(TypeError::InvalidNumberOfArguments { + function: Function::Sqrt, + expected: 1, + got: args.len(), + }); + } + crate::ast::infer_types_aux(&args[0], 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(), 1, "expected 1 arg for SQRT"); + debug_assert!(target_type_set.contains(VarType::F64)); + let arg = context.apply_types(&args[0], InferredTypeSet::F64)?; + if arg.return_type == VarType::None { + return Ok(TypedExpr::none()); + } + Ok(TypedExpr { + return_type: VarType::F64, + ast: TypedExprAst::from_call(SqrtFnCall { arg: Box::new(arg) }), + }) + } + + fn args_mut(&mut self) -> &mut [TypedExpr] { + std::slice::from_mut(&mut self.arg) + } + + 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::Sqrt, + return_type, + }); + } + let arg = context.compile_expr(&self.arg, builder)?; + let value = builder.ins().sqrt(arg.value); + let is_nan = builder.ins().fcmp(FloatCC::Unordered, value, value); + let is_not_nan = builder.ins().bxor_imm_u(is_nan, 1); + let is_present = builder.ins().band(arg.is_present, is_not_nan); + Ok(LoweredValue { + value, + is_present, + string_len: builder.ins().iconst(types::I64, 0), + }) + } +} + +impl From for FnCallEnum { + fn from(call: SqrtFnCall) -> Self { + FnCallEnum::Sqrt(call) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::ast::{deserialize, infer_types}; + use crate::compile::compile; + use crate::types::VariableValue; + + fn eval(expression: &str) -> Option { + let expression = deserialize(expression).unwrap(); + let mut compiled = compile(&expression, &HashMap::new()).unwrap(); + // SAFETY: The expression has no inputs and returns a nullable f64 value. + unsafe { compiled.call(&[]).as_f64() } + } + + #[test] + fn test_signature_and_output_type() { + let expression = deserialize("(SQRT value)").unwrap(); + let inferred = infer_types(&expression).unwrap(); + assert_eq!(inferred.get("value"), Some(&InferredTypeSet::NUMERICAL)); + + for expression in ["(SQRT)", "(SQRT 1i64 2i64)"] { + let expression = deserialize(expression).unwrap(); + assert!(matches!( + infer_types(&expression), + Err(TypeError::InvalidNumberOfArguments { + function: Function::Sqrt, + expected: 1, + .. + }) + )); + } + + let expression = deserialize("(SQRT 9i64)").unwrap(); + assert_eq!( + compile(&expression, &HashMap::new()).unwrap().result_type(), + VarType::F64 + ); + } + + #[test] + fn test_numeric_inputs_are_converted_to_float() { + assert_eq!(eval("(SQRT 9i64)"), Some(3.0)); + assert_eq!(eval("(SQRT 2.25f64)"), Some(1.5)); + assert_eq!(eval("(SQRT 2u64)"), Some(2.0f64.sqrt())); + } + + #[test] + fn test_null_nan_and_ieee_edges() { + assert_eq!(eval("(SQRT none)"), None); + assert_eq!(eval("(SQRT -1i64)"), None); + assert_eq!(eval("(SQRT nanf64)"), None); + assert_eq!(eval("(SQRT inff64)"), Some(f64::INFINITY)); + + let negative_zero = eval("(SQRT -0f64)").unwrap(); + assert_eq!(negative_zero.to_bits(), (-0.0f64).to_bits()); + } + + #[test] + fn test_runtime_null_and_negative_input() { + let expression = deserialize("(SQRT value)").unwrap(); + let mut compiled = compile(&expression, &HashMap::from([("value", VarType::I64)])).unwrap(); + + // SAFETY: The compiled expression expects one nullable i64 argument. + assert_eq!( + unsafe { compiled.call(&[VariableValue::some(16i64)]).as_f64() }, + Some(4.0) + ); + // SAFETY: The compiled expression expects one nullable i64 argument. + assert_eq!( + unsafe { compiled.call(&[VariableValue::some(-16i64)]).as_f64() }, + None + ); + // SAFETY: The compiled expression expects one nullable i64 argument. + assert_eq!( + unsafe { compiled.call(&[VariableValue::none()]).as_f64() }, + None + ); + } +}