diff --git a/isthmus/src/main/java/io/substrait/isthmus/expression/ScalarFunctionConverter.java b/isthmus/src/main/java/io/substrait/isthmus/expression/ScalarFunctionConverter.java index b85db0d66..6193974e0 100644 --- a/isthmus/src/main/java/io/substrait/isthmus/expression/ScalarFunctionConverter.java +++ b/isthmus/src/main/java/io/substrait/isthmus/expression/ScalarFunctionConverter.java @@ -54,7 +54,8 @@ public class ScalarFunctionConverter new StringFunctionOptions(), new IntegerFunctionOptions(), new DecimalFunctionOptions(), - new FloatingPointFunctionOptions()); + new FloatingPointFunctionOptions(), + new UnaryArithmeticOptions()); /** * Creates a converter with the given functions and type factory. diff --git a/isthmus/src/main/java/io/substrait/isthmus/expression/UnaryArithmeticOptions.java b/isthmus/src/main/java/io/substrait/isthmus/expression/UnaryArithmeticOptions.java new file mode 100644 index 000000000..3a297fb19 --- /dev/null +++ b/isthmus/src/main/java/io/substrait/isthmus/expression/UnaryArithmeticOptions.java @@ -0,0 +1,160 @@ +package io.substrait.isthmus.expression; + +import io.substrait.expression.Expression; +import io.substrait.expression.FunctionOption; +import io.substrait.extension.DefaultExtensionCatalog; +import io.substrait.extension.SimpleExtension.ScalarFunctionVariant; +import java.util.List; +import java.util.Locale; +import java.util.Set; +import org.apache.calcite.rex.RexCall; +import org.apache.calcite.sql.SqlOperator; +import org.apache.calcite.sql.fun.SqlStdOperatorTable; +import org.apache.calcite.sql.type.SqlTypeName; + +/** Unary arithmetic option policies. */ +final class UnaryArithmeticOptions implements ScalarFunctionOptionPolicy { + private static final Set NAMES = + Set.of( + "negate", + "abs", + "sqrt", + "exp", + "cos", + "sin", + "tan", + "cosh", + "sinh", + "tanh", + "acos", + "asin", + "atan", + "acosh", + "asinh", + "atanh", + "radians", + "degrees", + "factorial"); + + private static String name(ScalarFunctionVariant function) { + if (!DefaultExtensionCatalog.FUNCTIONS_ARITHMETIC.equals(function.urn()) + || !NAMES.contains(function.name())) { + return null; + } + for (String tag : List.of("i8", "i16", "i32", "i64", "fp32", "fp64")) { + if (function.key().equals(function.name() + ":" + tag)) { + return function.name(); + } + } + return null; + } + + private static boolean integer(ScalarFunctionVariant function) { + return function.key().matches("(?:negate|abs):i(?:8|16|32|64)"); + } + + private static SqlOperator nativeOperator(String name) { + return FunctionMappings.SCALAR_SIGS.stream() + .filter(sig -> sig.name().equals(name)) + .map(FunctionMappings.Sig::operator) + .findFirst() + .orElseThrow(); + } + + @Override + public SqlOperator signatureOperator(RexCall call) { + return checkedIntegerNegation(call) ? SqlStdOperatorTable.UNARY_MINUS : call.getOperator(); + } + + static boolean checkedIntegerNegation(RexCall call) { + SqlTypeName type = call.getType().getSqlTypeName(); + return call.getOperator() == SqlStdOperatorTable.CHECKED_UNARY_MINUS + && call.getOperands().size() == 1 + && call.getOperands().get(0).getType().getSqlTypeName() == type + && SqlTypeName.INT_TYPES.contains(type); + } + + private static FunctionOption option(String name, String value) { + return FunctionOption.builder().name(name).addValues(value).build(); + } + + @Override + public List forCall(RexCall call, ScalarFunctionVariant function) { + String name = name(function); + if (name == null || call.getOperands().size() != 1) { + return List.of(); + } + if (name.equals("negate") && integer(function) && checkedIntegerNegation(call)) { + return List.of(option("overflow", "ERROR")); + } + if (call.getOperator() != nativeOperator(name)) { + return List.of(); + } + if ((name.equals("negate") || name.equals("abs")) && integer(function)) { + return List.of(option("overflow", "SILENT")); + } + // Calcite's FP32 unary conversion path cannot carry a NaN result. + if ((name.equals("acos") || name.equals("asin")) && function.key().endsWith(":fp64")) { + return List.of(option("on_domain_error", "NAN")); + } + // Java's transcendental functions need not be correctly rounded. SQL SQRT is + // represented as POWER(x, 0.5), without a verified explicit option policy. + return List.of(); + } + + @Override + public SqlOperator resolve(Expression.ScalarFunctionInvocation expression, SqlOperator selected) { + String name = name(expression.declaration()); + if (name == null || expression.options().isEmpty()) { + return selected; + } + SqlOperator nativeOperator = nativeOperator(name); + boolean negate = name.equals("negate") && integer(expression.declaration()); + if (selected != nativeOperator + && !(negate && selected == SqlStdOperatorTable.CHECKED_UNARY_MINUS)) { + throw new UnsupportedOperationException( + "No unary arithmetic option policy for Calcite operator " + selected.getName()); + } + String overflow = null; + for (FunctionOption option : expression.options()) { + ScalarFunctionOptionPolicy.requireDeclaredValues(expression, option); + String optionName = option.getName().toLowerCase(Locale.ROOT); + List supported; + if (optionName.equals("overflow") && integer(expression.declaration())) { + supported = negate ? List.of("SILENT", "ERROR") : List.of("SILENT"); + } else if (optionName.equals("on_domain_error") + && (name.equals("acos") || name.equals("asin")) + && expression.declaration().key().endsWith(":fp64")) { + supported = List.of("NAN"); + } else { + supported = List.of(); + } + String value = + option.values().stream() + .map(v -> v.toUpperCase(Locale.ROOT)) + .filter(supported::contains) + .findFirst() + .orElseThrow( + () -> + new UnsupportedOperationException( + "Unsupported unary arithmetic " + + name + + " " + + optionName + + " preferences: " + + option.values())); + if (optionName.equals("overflow")) { + if (overflow != null && !overflow.equals(value)) { + throw new UnsupportedOperationException("Conflicting unary arithmetic option: overflow"); + } + overflow = value; + } + } + if (!negate || overflow == null) { + return selected; + } + return overflow.equals("ERROR") + ? SqlStdOperatorTable.CHECKED_UNARY_MINUS + : SqlStdOperatorTable.UNARY_MINUS; + } +} diff --git a/isthmus/src/test/java/io/substrait/isthmus/SqrtImportTest.java b/isthmus/src/test/java/io/substrait/isthmus/SqrtImportTest.java index 34233d57b..bdfb7a928 100644 --- a/isthmus/src/test/java/io/substrait/isthmus/SqrtImportTest.java +++ b/isthmus/src/test/java/io/substrait/isthmus/SqrtImportTest.java @@ -3,6 +3,7 @@ import static org.junit.jupiter.api.Assertions.assertEquals; import static org.junit.jupiter.api.Assertions.assertInstanceOf; import static org.junit.jupiter.api.Assertions.assertNotEquals; +import static org.junit.jupiter.api.Assertions.assertThrows; import static org.junit.jupiter.api.Assertions.assertTrue; import io.substrait.expression.Expression; @@ -377,12 +378,15 @@ void explicitOptionsAreNotReinterpretedByThePowerExpansion() { .from(expression("fp64", false)) .addOptions(FunctionOption.builder().name("on_domain_error").addValues("ERROR").build()) .build(); - RexCall call = - (RexCall) - expression.accept( - converter(new ScalarFunctionConverter(extensions.scalarFunctions(), typeFactory)), - Context.newContext()); - assertEquals(SqlStdOperatorTable.SQRT, call.getOperator()); + UnsupportedOperationException failure = + assertThrows( + UnsupportedOperationException.class, + () -> + expression.accept( + converter( + new ScalarFunctionConverter(extensions.scalarFunctions(), typeFactory)), + Context.newContext())); + assertTrue(failure.getMessage().contains("sqrt on_domain_error")); } @Test diff --git a/isthmus/src/test/java/io/substrait/isthmus/UnaryArithmeticOptionsTest.java b/isthmus/src/test/java/io/substrait/isthmus/UnaryArithmeticOptionsTest.java new file mode 100644 index 000000000..7f800951d --- /dev/null +++ b/isthmus/src/test/java/io/substrait/isthmus/UnaryArithmeticOptionsTest.java @@ -0,0 +1,436 @@ +package io.substrait.isthmus; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import io.substrait.expression.Expression; +import io.substrait.expression.ExpressionCreator; +import io.substrait.expression.FunctionOption; +import io.substrait.extension.DefaultExtensionCatalog; +import io.substrait.isthmus.SubstraitRelNodeConverter.Context; +import io.substrait.isthmus.expression.CallConverters; +import io.substrait.isthmus.expression.ExpressionRexConverter; +import io.substrait.isthmus.expression.RexExpressionConverter; +import io.substrait.isthmus.expression.ScalarFunctionConverter; +import io.substrait.isthmus.expression.WindowFunctionConverter; +import io.substrait.isthmus.sql.SubstraitCreateStatementParser; +import io.substrait.relation.Project; +import io.substrait.type.Type; +import java.sql.PreparedStatement; +import java.sql.ResultSet; +import java.util.List; +import java.util.Optional; +import java.util.stream.Stream; +import org.apache.calcite.DataContext; +import org.apache.calcite.jdbc.CalciteSchema; +import org.apache.calcite.linq4j.Enumerable; +import org.apache.calcite.linq4j.Linq4j; +import org.apache.calcite.rel.RelNode; +import org.apache.calcite.rel.logical.LogicalProject; +import org.apache.calcite.rel.type.RelDataType; +import org.apache.calcite.rel.type.RelDataTypeFactory; +import org.apache.calcite.rex.RexCall; +import org.apache.calcite.rex.RexNode; +import org.apache.calcite.schema.ScannableTable; +import org.apache.calcite.schema.impl.AbstractTable; +import org.apache.calcite.sql.SqlOperator; +import org.apache.calcite.sql.fun.SqlStdOperatorTable; +import org.apache.calcite.tools.RelBuilder; +import org.apache.calcite.tools.RelRunners; +import org.junit.jupiter.api.Test; + +class UnaryArithmeticOptionsTest extends PlanTestBase { + private final ExpressionRexConverter toRex = + new ExpressionRexConverter( + typeFactory, + new ScalarFunctionConverter(extensions.scalarFunctions(), typeFactory), + new WindowFunctionConverter(extensions.windowFunctions(), typeFactory), + TypeConverter.DEFAULT); + + private double execute(RexNode call, Number a) throws Exception { + RexCall original = (RexCall) call; + CalciteSchema schema = CalciteSchema.createRootSchema(false); + schema.add("inputs", new RuntimeInputs(original, a)); + RelBuilder runtimeBuilder = converterProvider.getRelBuilder(schema); + RelNode input = runtimeBuilder.scan("inputs").build(); + call = + runtimeBuilder + .getRexBuilder() + .makeCall( + original.getType(), + original.getOperator(), + List.of(runtimeBuilder.getRexBuilder().makeInputRef(input, 0))); + RelNode project = LogicalProject.create(input, List.of(), List.of(call), List.of("result")); + try (PreparedStatement statement = RelRunners.run(project); + ResultSet result = statement.executeQuery()) { + if (!result.next()) { + throw new IllegalStateException("No result row"); + } + return result.getDouble(1); + } + } + + private static class RuntimeInputs extends AbstractTable implements ScannableTable { + private final RexCall call; + private final Object[] values; + + private RuntimeInputs(RexCall call, Number a) { + this.call = call; + Object value; + switch (call.getOperands().get(0).getType().getSqlTypeName()) { + case TINYINT: + value = a.byteValue(); + break; + case SMALLINT: + value = a.shortValue(); + break; + case INTEGER: + value = a.intValue(); + break; + case BIGINT: + value = a.longValue(); + break; + case REAL: + value = a.floatValue(); + break; + default: + value = a.doubleValue(); + } + values = new Object[] {value}; + } + + @Override + public RelDataType getRowType(RelDataTypeFactory factory) { + return factory.builder().add("a", call.getOperands().get(0).getType()).build(); + } + + @Override + public Enumerable scan(DataContext context) { + return Linq4j.asEnumerable(new Object[][] {values}); + } + } + + private Type type(String tag) { + switch (tag) { + case "i8": + return R.I8; + case "i16": + return R.I16; + case "i32": + return R.I32; + case "i64": + return R.I64; + case "fp32": + return R.FP32; + default: + return R.FP64; + } + } + + private Expression.ScalarFunctionInvocation invocation( + String name, String tag, List options) { + Expression zero; + switch (tag) { + case "i8": + zero = ExpressionCreator.i8(false, (byte) 0); + break; + case "i16": + zero = ExpressionCreator.i16(false, (short) 0); + break; + case "i32": + zero = ExpressionCreator.i32(false, 0); + break; + case "i64": + zero = ExpressionCreator.i64(false, 0); + break; + case "fp32": + zero = ExpressionCreator.fp32(false, 0); + break; + default: + zero = ExpressionCreator.fp64(false, 0); + } + return sb.scalarFn( + DefaultExtensionCatalog.FUNCTIONS_ARITHMETIC, + name + ":" + tag, + tag.equals("i64") && (name.equals("sqrt") || name.equals("exp")) ? R.FP64 : type(tag), + List.of(zero), + options); + } + + private FunctionOption option(String name, String... values) { + return FunctionOption.builder().name(name).addValues(values).build(); + } + + private Expression.ScalarFunctionInvocation export(RexNode rex) { + ScalarFunctionConverter scalar = + new ScalarFunctionConverter(extensions.scalarFunctions(), typeFactory); + RexExpressionConverter converter = + new RexExpressionConverter( + null, + Stream.concat( + CallConverters.defaults(TypeConverter.DEFAULT).stream(), Stream.of(scalar)) + .toList(), + new WindowFunctionConverter(extensions.windowFunctions(), typeFactory), + TypeConverter.DEFAULT); + return (Expression.ScalarFunctionInvocation) rex.accept(converter); + } + + private void assertOverflow(RexNode rex, Number value) { + Exception failure = assertThrows(Exception.class, () -> execute(rex, value)); + Throwable root = failure; + while (root.getCause() != null) { + root = root.getCause(); + } + assertTrue(root instanceof ArithmeticException, root.toString()); + } + + @Test + void negateSupportsCheckedAndSilentOverflowAtEveryWidth() throws Exception { + for (String tag : List.of("i8", "i16", "i32", "i64")) { + int width = Integer.parseInt(tag.substring(1)); + long min = width == 64 ? Long.MIN_VALUE : -(1L << (width - 1)); + RexNode checked = + invocation("negate", tag, List.of(option("overflow", "ERROR"))) + .accept(toRex, Context.newContext()); + assertOverflow(checked, min); + Expression.ScalarFunctionInvocation exported = export(checked); + assertEquals(List.of(option("overflow", "ERROR")), exported.options()); + assertOverflow(exported.accept(toRex, Context.newContext()), min); + RexNode silent = + invocation("negate", tag, List.of(option("overflow", "SILENT"))) + .accept(toRex, Context.newContext()); + assertEquals((double) min, execute(silent, min)); + assertEquals(List.of(option("overflow", "SILENT")), export(silent).options()); + assertEquals(-2.0, execute(checked, 2)); + } + } + + @Test + void preferencesSelectTheFirstSupportedOverflowMode() { + for (String tag : List.of("i8", "i16", "i32", "i64")) { + for (List preferences : + List.of( + List.of("ERROR", "SILENT"), + List.of("SILENT", "ERROR"), + List.of("SATURATE", "error"))) { + RexCall rex = + (RexCall) + invocation( + "negate", + tag, + List.of( + FunctionOption.builder().name("OVERFLOW").values(preferences).build())) + .accept(toRex, Context.newContext()); + assertEquals( + preferences.get(0).equals("SILENT") + ? SqlStdOperatorTable.UNARY_MINUS + : SqlStdOperatorTable.CHECKED_UNARY_MINUS, + rex.getOperator()); + } + assertThrows( + UnsupportedOperationException.class, + () -> + invocation("negate", tag, List.of(option("overflow", "SATURATE"))) + .accept(toRex, Context.newContext())); + } + } + + @Test + void absOnlySupportsSilentOverflow() throws Exception { + for (String tag : List.of("i8", "i16", "i32", "i64")) { + int width = Integer.parseInt(tag.substring(1)); + long min = width == 64 ? Long.MIN_VALUE : -(1L << (width - 1)); + for (String unsupported : List.of("ERROR", "SATURATE")) { + assertThrows( + UnsupportedOperationException.class, + () -> + invocation("abs", tag, List.of(option("overflow", unsupported))) + .accept(toRex, Context.newContext())); + RexNode rex = + invocation("abs", tag, List.of(option("overflow", unsupported, "SILENT"))) + .accept(toRex, Context.newContext()); + assertEquals((double) min, execute(rex, min)); + assertEquals(2.0, execute(rex, -2)); + assertEquals(List.of(option("overflow", "SILENT")), export(rex).options()); + } + } + } + + @Test + void asinAndAcosHonorNanDomainPreferences() throws Exception { + for (String name : List.of("asin", "acos")) { + for (String tag : List.of("fp64")) { + assertThrows( + UnsupportedOperationException.class, + () -> + invocation(name, tag, List.of(option("on_domain_error", "ERROR"))) + .accept(toRex, Context.newContext())); + RexNode rex = + invocation(name, tag, List.of(option("on_domain_error", "ERROR", "NAN"))) + .accept(toRex, Context.newContext()); + assertTrue(Double.isNaN(execute(rex, 2))); + assertTrue(Double.isNaN(execute(rex, Double.NaN))); + assertEquals(List.of(option("on_domain_error", "NAN")), export(rex).options()); + } + } + } + + @Test + void unprovenRoundingPreferencesAreRejectedForTheWholeUnaryFamily() { + for (String name : + List.of( + "sqrt", "exp", "cos", "sin", "tan", "cosh", "sinh", "tanh", "acos", "asin", "atan", + "acosh", "asinh", "atanh", "radians", "degrees")) { + for (String tag : List.of("fp32", "fp64")) { + for (String rounding : + List.of("TIE_TO_EVEN", "TIE_AWAY_FROM_ZERO", "TRUNCATE", "CEILING", "FLOOR")) { + assertThrows( + UnsupportedOperationException.class, + () -> + invocation(name, tag, List.of(option("rounding", rounding))) + .accept(toRex, Context.newContext()), + name + ":" + tag); + } + } + } + for (String name : List.of("sqrt", "exp")) { + assertThrows( + UnsupportedOperationException.class, + () -> + invocation(name, "i64", List.of(option("rounding", "TIE_TO_EVEN"))) + .accept(toRex, Context.newContext())); + } + } + + @Test + void unimplementedDomainAndFactorialPoliciesAreRejected() { + for (String name : List.of("asin", "acos")) { + for (String domain : List.of("NAN", "ERROR")) { + assertThrows( + UnsupportedOperationException.class, + () -> + invocation(name, "fp32", List.of(option("on_domain_error", domain))) + .accept(toRex, Context.newContext())); + } + } + for (String name : List.of("sqrt", "acosh", "atanh")) { + for (String domain : List.of("NAN", "ERROR")) { + assertThrows( + UnsupportedOperationException.class, + () -> + invocation(name, "fp64", List.of(option("on_domain_error", domain))) + .accept(toRex, Context.newContext())); + } + } + for (String tag : List.of("i32", "i64")) { + for (String overflow : List.of("SILENT", "SATURATE", "ERROR")) { + assertThrows( + UnsupportedOperationException.class, + () -> + invocation("factorial", tag, List.of(option("overflow", overflow))) + .accept(toRex, Context.newContext())); + } + } + } + + @Test + void unknownEmptyAndConflictingPreferencesAreRejected() { + for (List options : + List.of( + List.of(option("overflow")), + List.of(option("unknown", "SILENT")), + List.of(option("overflow", "SILENT"), option("overflow", "ERROR")))) { + assertThrows( + UnsupportedOperationException.class, + () -> invocation("negate", "i32", options).accept(toRex, Context.newContext())); + } + } + + @Test + void callsWithoutOptionsKeepTheExistingOperators() { + for (String name : + List.of( + "negate", "abs", "sqrt", "exp", "asin", "acos", "acosh", "atanh", "sin", "factorial")) { + String tag = + name.equals("factorial") || name.equals("negate") || name.equals("abs") ? "i64" : "fp64"; + RexCall rex = (RexCall) invocation(name, tag, List.of()).accept(toRex, Context.newContext()); + SqlOperator expected = + io.substrait.isthmus.expression.FunctionMappings.SCALAR_SIGS.stream() + .filter(sig -> sig.name().equals(name)) + .findFirst() + .orElseThrow() + .operator(); + assertEquals(name.equals("sqrt") ? SqlStdOperatorTable.POWER : expected, rex.getOperator()); + } + } + + @Test + void customOperatorsDoNotInheritNativeOptions() { + ScalarFunctionConverter custom = + new ScalarFunctionConverter(extensions.scalarFunctions(), typeFactory) { + @Override + public Optional getSqlOperatorFromSubstraitFunc( + String key, Type outputType) { + return Optional.of(SqlStdOperatorTable.UNARY_PLUS); + } + }; + assertEquals( + Optional.of(SqlStdOperatorTable.UNARY_PLUS), + custom.getSqlOperatorFromSubstraitFunc(invocation("negate", "i32", List.of()))); + assertThrows( + UnsupportedOperationException.class, + () -> + custom.getSqlOperatorFromSubstraitFunc( + invocation("negate", "i32", List.of(option("overflow", "SILENT"))))); + } + + @Test + void sqlExportOnlyNamesVerifiedNativePolicies() throws Exception { + String creates = + "CREATE TABLE numbers (i8 TINYINT, i16 SMALLINT, i32 INT, i64 BIGINT, f DOUBLE)"; + String query = + "SELECT -i8, -i16, -i32, -i64, abs(i8), abs(i16), abs(i32), abs(i64), asin(f), acos(f), sqrt(f), sin(f), exp(f) FROM numbers"; + assertFullRoundTrip(query, creates); + Project project = + (Project) + new SqlToSubstrait() + .convert( + query, SubstraitCreateStatementParser.processCreateStatementsToCatalog(creates)) + .getRoots() + .get(0) + .getInput(); + for (Expression expression : project.getExpressions()) { + Expression.ScalarFunctionInvocation function = + (Expression.ScalarFunctionInvocation) expression; + String name = function.declaration().name(); + if (name.equals("negate") || name.equals("abs")) { + assertEquals(List.of(option("overflow", "SILENT")), function.options()); + } else if (name.equals("asin") || name.equals("acos")) { + assertEquals(List.of(option("on_domain_error", "NAN")), function.options()); + } else { + assertTrue(function.options().isEmpty(), name); + } + } + } + + @Test + void undeclaredPreferencesCannotHideBehindASupportedFallback() { + for (String[] values : + List.of(new String[] {"WRAP", "SILENT"}, new String[] {"SILENT", "WRAP"})) { + UnsupportedOperationException failure = + assertThrows( + UnsupportedOperationException.class, + () -> + invocation("negate", "i32", List.of(option("overflow", values))) + .accept(toRex, Context.newContext())); + assertTrue(failure.getMessage().contains("does not declare value WRAP")); + } + assertTrue( + export( + invocation("negate", "i32", List.of(option("overflow", "silent"))) + .accept(toRex, Context.newContext())) + .options() + .contains(option("overflow", "SILENT"))); + } +}