Add SQRT function

This commit is contained in:
Paul Masurel
2026-08-19 10:47:58 +02:00
parent 918427533f
commit ffbcb199d0
4 changed files with 194 additions and 5 deletions
+2 -5
View File
@@ -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.
+2
View File
@@ -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<Function, DeserializeErro
"NOT" => 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),
+13
View File
@@ -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 => <NotFnCall as FnCall>::call_with_types(args, target_type_set, context),
Function::Or => <OrFnCall as FnCall>::call_with_types(args, target_type_set, context),
Function::Pow => <PowFnCall as FnCall>::call_with_types(args, target_type_set, context),
Function::Sqrt => {
<SqrtFnCall as FnCall>::call_with_types(args, target_type_set, context)
}
Function::RegexpExtract => {
<RegexpExtractFnCall as FnCall>::call_with_types(args, target_type_set, context)
}
@@ -290,6 +297,9 @@ impl Function {
Function::Not => <NotFnCall as FnCall>::infer_types(args, target_type, inferred_types),
Function::Or => <OrFnCall as FnCall>::infer_types(args, target_type, inferred_types),
Function::Pow => <PowFnCall as FnCall>::infer_types(args, target_type, inferred_types),
Function::Sqrt => {
<SqrtFnCall as FnCall>::infer_types(args, target_type, inferred_types)
}
Function::RegexpExtract => {
<RegexpExtractFnCall as FnCall>::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)
}
+177
View File
@@ -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<TypedExpr>,
}
impl FnCall for SqrtFnCall {
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::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<TypedExpr, CompileError> {
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<LoweredValue, CompileError> {
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<SqrtFnCall> 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<f64> {
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
);
}
}