diff --git a/isthmus/src/main/java/io/substrait/isthmus/expression/IntegerFunctionOptions.java b/isthmus/src/main/java/io/substrait/isthmus/expression/IntegerFunctionOptions.java new file mode 100644 index 000000000..7af462dcd --- /dev/null +++ b/isthmus/src/main/java/io/substrait/isthmus/expression/IntegerFunctionOptions.java @@ -0,0 +1,182 @@ +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.ArrayList; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; +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; + +/** Signed integer arithmetic options supported by Calcite. */ +final class IntegerFunctionOptions implements ScalarFunctionOptionPolicy { + @Override + public SqlOperator signatureOperator(RexCall call) { + return integerCall(call) ? unchecked(call.getOperator()) : call.getOperator(); + } + + static SqlOperator unchecked(SqlOperator operator) { + if (operator == SqlStdOperatorTable.CHECKED_PLUS) { + return SqlStdOperatorTable.PLUS; + } + if (operator == SqlStdOperatorTable.CHECKED_MINUS) { + return SqlStdOperatorTable.MINUS; + } + if (operator == SqlStdOperatorTable.CHECKED_MULTIPLY) { + return SqlStdOperatorTable.MULTIPLY; + } + if (operator == SqlStdOperatorTable.CHECKED_DIVIDE) { + return SqlStdOperatorTable.DIVIDE; + } + return operator; + } + + static boolean integerCall(RexCall call) { + return SqlTypeName.INT_TYPES.contains(call.getType().getSqlTypeName()) + && call.getOperands().stream() + .allMatch(arg -> SqlTypeName.INT_TYPES.contains(arg.getType().getSqlTypeName())); + } + + private static Binding binding(ScalarFunctionVariant function) { + if (!DefaultExtensionCatalog.FUNCTIONS_ARITHMETIC.equals(function.urn())) { + return null; + } + String name = function.name(); + for (String type : List.of("i8", "i16", "i32", "i64")) { + if (!function.key().equals(name + ":" + type + "_" + type)) { + continue; + } + boolean narrow = type.equals("i8") || type.equals("i16"); + switch (name) { + case "add": + return new Binding( + SqlStdOperatorTable.PLUS, SqlStdOperatorTable.CHECKED_PLUS, narrow, false, false); + case "subtract": + return new Binding( + SqlStdOperatorTable.MINUS, SqlStdOperatorTable.CHECKED_MINUS, narrow, false, false); + case "multiply": + return new Binding( + SqlStdOperatorTable.MULTIPLY, + SqlStdOperatorTable.CHECKED_MULTIPLY, + narrow, + false, + false); + case "divide": + return new Binding( + SqlStdOperatorTable.DIVIDE, SqlStdOperatorTable.CHECKED_DIVIDE, narrow, true, false); + case "modulus": + return new Binding(SqlStdOperatorTable.MOD, SqlStdOperatorTable.MOD, false, false, true); + default: + return null; + } + } + return null; + } + + @Override + public List forCall(RexCall call, ScalarFunctionVariant function) { + Binding binding = binding(function); + if (binding == null + || (call.getOperator() != binding.normal && call.getOperator() != binding.checked)) { + return List.of(); + } + List options = new ArrayList<>(); + // Export follows the actual Rex operator. Conformance-specific preparation may + // replace a plain operator with a checked equivalent before execution. + // Plain narrow arithmetic range-checks required results but wraps nullable ones. + // Planner nullability inference can change that choice, so do not promise a policy. + if (!binding.narrow || call.getOperator() == binding.checked) { + String overflow = + call.getOperator() == binding.checked && !binding.modulus ? "ERROR" : "SILENT"; + options.add(option("overflow", overflow)); + } + if (binding.divide) { + options.add(option("on_domain_error", "ERROR")); + options.add(option("on_division_by_zero", "ERROR")); + } + if (binding.modulus) { + options.add(option("division_type", "TRUNCATE")); + options.add(option("on_domain_error", "ERROR")); + } + return options; + } + + private static FunctionOption option(String name, String value) { + return FunctionOption.builder().name(name).addValues(value).build(); + } + + @Override + public SqlOperator resolve(Expression.ScalarFunctionInvocation expression, SqlOperator operator) { + Binding binding = binding(expression.declaration()); + if (binding == null) { + return operator; + } + if (operator != binding.normal && operator != binding.checked) { + if (!expression.options().isEmpty()) { + throw new UnsupportedOperationException( + "No integer option policy for Calcite operator " + operator.getName()); + } + return operator; + } + Map selected = new LinkedHashMap<>(); + for (FunctionOption option : expression.options()) { + ScalarFunctionOptionPolicy.requireDeclaredValues(expression, option); + String name = option.getName().toLowerCase(java.util.Locale.ROOT); + List supported; + if (name.equals("overflow")) { + supported = binding.narrow ? List.of("ERROR") : List.of("SILENT", "ERROR"); + } else if ((binding.divide + && (name.equals("on_domain_error") || name.equals("on_division_by_zero"))) + || (binding.modulus && name.equals("on_domain_error"))) { + supported = List.of("ERROR"); + } else if (binding.modulus && name.equals("division_type")) { + supported = List.of("TRUNCATE"); + } else { + throw new UnsupportedOperationException("Unsupported integer arithmetic option: " + name); + } + String value = + option.values().stream() + .map(v -> v.toUpperCase(java.util.Locale.ROOT)) + .filter(supported::contains) + .findFirst() + .orElseThrow( + () -> + new UnsupportedOperationException( + "Unsupported integer arithmetic " + + name + + " preferences: " + + option.values())); + String previous = selected.putIfAbsent(name, value); + if (previous != null && !previous.equals(value)) { + throw new UnsupportedOperationException("Conflicting integer arithmetic option: " + name); + } + } + String overflow = selected.get("overflow"); + if (overflow == null) { + return operator; + } + return overflow.equals("ERROR") ? binding.checked : binding.normal; + } + + private static final class Binding { + private final SqlOperator normal; + private final SqlOperator checked; + private final boolean narrow; + private final boolean divide; + private final boolean modulus; + + private Binding( + SqlOperator normal, SqlOperator checked, boolean narrow, boolean divide, boolean modulus) { + this.normal = normal; + this.checked = checked; + this.narrow = narrow; + this.divide = divide; + this.modulus = modulus; + } + } +} 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 0781c4ac3..38c4ece2e 100644 --- a/isthmus/src/main/java/io/substrait/isthmus/expression/ScalarFunctionConverter.java +++ b/isthmus/src/main/java/io/substrait/isthmus/expression/ScalarFunctionConverter.java @@ -50,7 +50,7 @@ public class ScalarFunctionConverter private final List mappers; private final List optionPolicies = - List.of(new StringFunctionOptions()); + List.of(new StringFunctionOptions(), new IntegerFunctionOptions()); /** * Creates a converter with the given functions and type factory. diff --git a/isthmus/src/main/java/io/substrait/isthmus/expression/ScalarFunctionOptionPolicy.java b/isthmus/src/main/java/io/substrait/isthmus/expression/ScalarFunctionOptionPolicy.java index ec1ebcbd4..aabb112a7 100644 --- a/isthmus/src/main/java/io/substrait/isthmus/expression/ScalarFunctionOptionPolicy.java +++ b/isthmus/src/main/java/io/substrait/isthmus/expression/ScalarFunctionOptionPolicy.java @@ -17,4 +17,24 @@ interface ScalarFunctionOptionPolicy { default SqlOperator signatureOperator(RexCall call) { return call.getOperator(); } + + /** Rejects undeclared values while comparing option names and values without regard to case. */ + static void requireDeclaredValues( + Expression.ScalarFunctionInvocation expression, FunctionOption option) { + for (String value : option.values()) { + boolean declared = + expression.declaration().options().entrySet().stream() + .filter(entry -> entry.getKey().equalsIgnoreCase(option.getName())) + .flatMap(entry -> entry.getValue().getValues().stream()) + .anyMatch(value::equalsIgnoreCase); + if (!declared) { + throw new UnsupportedOperationException( + expression.declaration().name() + + " option " + + option.getName() + + " does not declare value " + + value); + } + } + } } diff --git a/isthmus/src/main/java/io/substrait/isthmus/expression/StringFunctionOptions.java b/isthmus/src/main/java/io/substrait/isthmus/expression/StringFunctionOptions.java index 0d0ac3c0a..6e6d8a478 100644 --- a/isthmus/src/main/java/io/substrait/isthmus/expression/StringFunctionOptions.java +++ b/isthmus/src/main/java/io/substrait/isthmus/expression/StringFunctionOptions.java @@ -71,21 +71,7 @@ public SqlOperator resolve(Expression.ScalarFunctionInvocation expression, SqlOp throw new UnsupportedOperationException( "Unsupported " + expression.declaration().name() + " option: " + option.getName()); } - for (String value : option.values()) { - boolean declared = - expression.declaration().options().entrySet().stream() - .filter(entry -> entry.getKey().equalsIgnoreCase(option.getName())) - .flatMap(entry -> entry.getValue().getValues().stream()) - .anyMatch(value::equalsIgnoreCase); - if (!declared) { - throw new UnsupportedOperationException( - expression.declaration().name() - + " option " - + option.getName() - + " does not declare value " - + value); - } - } + ScalarFunctionOptionPolicy.requireDeclaredValues(expression, option); if (binding.value() == null) { // The spec lists initcap's charsets without defining ASCII word boundaries. throw new UnsupportedOperationException( diff --git a/isthmus/src/test/java/io/substrait/isthmus/DdlRoundtripTest.java b/isthmus/src/test/java/io/substrait/isthmus/DdlRoundtripTest.java index 97ae80be1..86335917e 100644 --- a/isthmus/src/test/java/io/substrait/isthmus/DdlRoundtripTest.java +++ b/isthmus/src/test/java/io/substrait/isthmus/DdlRoundtripTest.java @@ -314,8 +314,8 @@ private Rel computedColumns() { .input(scan) .remap(Rel.Remap.offset(2, 2)) .addExpressions( - sb.add(sb.fieldReference(scan, 0), sb.i32(1)), - sb.add(sb.fieldReference(scan, 0), sb.i32(2))) + withSilentOverflow(sb.add(sb.fieldReference(scan, 0), sb.i32(1))), + withSilentOverflow(sb.add(sb.fieldReference(scan, 0), sb.i32(2)))) .build(); } } diff --git a/isthmus/src/test/java/io/substrait/isthmus/IntegerArithmeticOptionsTest.java b/isthmus/src/test/java/io/substrait/isthmus/IntegerArithmeticOptionsTest.java new file mode 100644 index 000000000..86e6ad561 --- /dev/null +++ b/isthmus/src/test/java/io/substrait/isthmus/IntegerArithmeticOptionsTest.java @@ -0,0 +1,479 @@ +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.math.BigDecimal; +import java.sql.PreparedStatement; +import java.sql.ResultSet; +import java.util.List; +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.RexBuilder; +import org.apache.calcite.rex.RexCall; +import org.apache.calcite.rex.RexLiteral; +import org.apache.calcite.rex.RexNode; +import org.apache.calcite.schema.ScannableTable; +import org.apache.calcite.schema.impl.AbstractTable; +import org.apache.calcite.sql.fun.SqlStdOperatorTable; +import org.apache.calcite.sql.validate.SqlConformanceEnum; +import org.apache.calcite.tools.RelBuilder; +import org.apache.calcite.tools.RelRunners; +import org.junit.jupiter.api.Test; + +class IntegerArithmeticOptionsTest extends PlanTestBase { + private final ExpressionRexConverter toRex = + new ExpressionRexConverter( + typeFactory, + new ScalarFunctionConverter(extensions.scalarFunctions(), typeFactory), + new WindowFunctionConverter(extensions.windowFunctions(), typeFactory), + TypeConverter.DEFAULT); + + private String execute(RexNode call) { + RexCall original = (RexCall) call; + CalciteSchema schema = CalciteSchema.createRootSchema(false); + schema.add("inputs", new RuntimeInputs(original)); + 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), + runtimeBuilder.getRexBuilder().makeInputRef(input, 1))); + 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 "VALUE=" + result.getObject(1); + } catch (Exception | LinkageError failure) { + Throwable root = failure; + while (root.getCause() != null) { + root = root.getCause(); + } + if (root instanceof NullPointerException) { + throw new IllegalStateException("Probe execution failed", failure); + } + return "ERROR=" + root.getClass().getSimpleName() + ": " + root.getMessage(); + } + } + + private static class RuntimeInputs extends AbstractTable implements ScannableTable { + private final RexCall call; + + private RuntimeInputs(RexCall call) { + this.call = call; + } + + @Override + public RelDataType getRowType(RelDataTypeFactory factory) { + return factory + .builder() + .add( + "a", + factory.createTypeWithNullability( + call.getOperands().get(0).getType(), call.getType().isNullable())) + .add( + "b", + factory.createTypeWithNullability( + call.getOperands().get(1).getType(), call.getType().isNullable())) + .build(); + } + + private Object value(RexNode expression) { + RexLiteral literal = (RexLiteral) expression; + switch (literal.getType().getSqlTypeName()) { + case TINYINT: + return literal.getValueAs(Integer.class).byteValue(); + case SMALLINT: + return literal.getValueAs(Integer.class).shortValue(); + case INTEGER: + return literal.getValueAs(Integer.class); + case BIGINT: + return literal.getValueAs(Long.class); + default: + return literal.getValueAs(BigDecimal.class); + } + } + + @Override + public Enumerable scan(DataContext context) { + return Linq4j.asEnumerable( + new Object[][] {{value(call.getOperands().get(0)), value(call.getOperands().get(1))}}); + } + } + + private Expression integer(int width, long value) { + switch (width) { + case 8: + return ExpressionCreator.i8(false, (byte) value); + case 16: + return ExpressionCreator.i16(false, (short) value); + case 32: + return ExpressionCreator.i32(false, (int) value); + default: + return ExpressionCreator.i64(false, value); + } + } + + private Type integerType(int width) { + return width == 8 ? R.I8 : width == 16 ? R.I16 : width == 32 ? R.I32 : R.I64; + } + + private Expression.ScalarFunctionInvocation invocation( + String name, int width, long a, long b, List options) { + return sb.scalarFn( + DefaultExtensionCatalog.FUNCTIONS_ARITHMETIC, + name + ":i" + width + "_i" + width, + integerType(width), + List.of(integer(width, a), integer(width, b)), + 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); + } + + @Test + void checkedOverflowIsExecutedAndSurvivesExportForEveryWidth() { + for (int width : List.of(8, 16, 32, 64)) { + long max = width == 64 ? Long.MAX_VALUE : (1L << (width - 1)) - 1; + long min = width == 64 ? Long.MIN_VALUE : -(1L << (width - 1)); + String[] names = {"add", "subtract", "multiply", "divide"}; + long[] left = {max, min, max, min}; + long[] right = {1, 1, 2, -1}; + long[] wrapped = {min, max, -2, min}; + for (int i = 0; i < names.length; i++) { + Expression.ScalarFunctionInvocation expression = + invocation(names[i], width, left[i], right[i], List.of(option("overflow", "ERROR"))); + RexNode rex = expression.accept(toRex, Context.newContext()); + assertTrue(execute(rex).startsWith("ERROR=ArithmeticException"), width + " " + names[i]); + assertTrue(export(rex).options().contains(option("overflow", "ERROR"))); + Expression.ScalarFunctionInvocation silent = + invocation(names[i], width, left[i], right[i], List.of(option("overflow", "SILENT"))); + if (width <= 16) { + assertThrows( + UnsupportedOperationException.class, + () -> silent.accept(toRex, Context.newContext())); + } else { + assertEquals("VALUE=" + wrapped[i], execute(silent.accept(toRex, Context.newContext()))); + } + } + } + } + + @Test + void narrowOverflowDependsOnResultNullabilityUnlessTheOperatorIsChecked() { + for (int width : List.of(8, 16)) { + long max = (1L << (width - 1)) - 1; + long min = -(1L << (width - 1)); + String[] names = {"add", "subtract", "multiply", "divide"}; + long[] left = {max, min, max, min}; + long[] right = {1, 1, 2, -1}; + long[] wrapped = {min, max, -2, min}; + for (int i = 0; i < names.length; i++) { + for (boolean nullable : List.of(false, true)) { + Type type = + width == 8 ? Type.withNullability(nullable).I8 : Type.withNullability(nullable).I16; + Expression.ScalarFunctionInvocation plain = + Expression.ScalarFunctionInvocation.builder() + .from(invocation(names[i], width, left[i], right[i], List.of())) + .outputType(type) + .build(); + RexNode call = plain.accept(toRex, Context.newContext()); + String result = execute(call); + if (nullable) { + assertEquals("VALUE=" + wrapped[i], result, names[i]); + } else { + assertTrue(result.startsWith("ERROR=ArithmeticException"), result); + } + assertTrue( + export(call).options().stream().noneMatch(o -> o.getName().equals("overflow"))); + Expression.ScalarFunctionInvocation checked = + Expression.ScalarFunctionInvocation.builder() + .from(plain) + .addOptions(option("overflow", "ERROR")) + .build(); + RexNode checkedCall = checked.accept(toRex, Context.newContext()); + assertTrue(execute(checkedCall).startsWith("ERROR=ArithmeticException"), names[i]); + assertTrue(export(checkedCall).options().contains(option("overflow", "ERROR"))); + } + } + } + } + + @Test + void preferenceOrderChoosesTheFirstSupportedBehavior() { + for (List preferences : + List.of( + List.of("ERROR", "SILENT"), List.of("SILENT", "ERROR"), List.of("SATURATE", "ERROR"))) { + RexNode rex = + invocation( + "add", + 32, + Integer.MAX_VALUE, + 1, + List.of(FunctionOption.builder().name("OVERFLOW").values(preferences).build())) + .accept(toRex, Context.newContext()); + String result = execute(rex); + if (preferences.get(0).equals("SILENT")) { + assertEquals("VALUE=" + Integer.MIN_VALUE, result); + } else { + assertTrue(result.startsWith("ERROR=ArithmeticException")); + } + } + RexNode rex = + invocation("add", 32, 1, 2, List.of(option("overflow", "error"))) + .accept(toRex, Context.newContext()); + assertEquals("VALUE=3", execute(rex)); + assertThrows( + UnsupportedOperationException.class, + () -> + invocation("add", 32, 1, 2, List.of(option("overflow", "SATURATE"))) + .accept(toRex, Context.newContext())); + } + + @Test + void zeroDivisionAndModulusRejectNullAndAcceptErrorFallback() { + for (String name : List.of("divide", "modulus")) { + String optionName = name.equals("divide") ? "on_division_by_zero" : "on_domain_error"; + assertThrows( + UnsupportedOperationException.class, + () -> + invocation(name, 32, 1, 0, List.of(option(optionName, "NULL"))) + .accept(toRex, Context.newContext())); + RexNode rex = + invocation(name, 32, 1, 0, List.of(option(optionName, "NULL", "ERROR"))) + .accept(toRex, Context.newContext()); + assertTrue(execute(rex).startsWith("ERROR=ArithmeticException")); + } + } + + @Test + void modulusRejectsFloorAndKeepsTruncation() { + assertThrows( + UnsupportedOperationException.class, + () -> + invocation("modulus", 32, -5, 3, List.of(option("division_type", "FLOOR"))) + .accept(toRex, Context.newContext())); + RexNode rex = + invocation("modulus", 32, -5, 3, List.of(option("division_type", "FLOOR", "TRUNCATE"))) + .accept(toRex, Context.newContext()); + assertEquals("VALUE=-2", execute(rex)); + assertTrue(export(rex).options().contains(option("division_type", "TRUNCATE"))); + for (int width : List.of(8, 16, 32, 64)) { + long min = width == 64 ? Long.MIN_VALUE : -(1L << (width - 1)); + RexNode remainder = + invocation("modulus", width, min, -1, List.of(option("overflow", "ERROR"))) + .accept(toRex, Context.newContext()); + assertEquals("VALUE=0", execute(remainder)); + } + } + + @Test + void omittedOptionsKeepTheExistingOperator() { + RexCall rex = + (RexCall) + invocation("add", 32, Integer.MAX_VALUE, 1, List.of()) + .accept(toRex, Context.newContext()); + assertEquals(SqlStdOperatorTable.PLUS, rex.getOperator()); + assertEquals("VALUE=" + Integer.MIN_VALUE, execute(rex)); + } + + @Test + void unknownEmptyAndConflictingOptionsAreRejected() { + for (List options : + List.of( + List.of(option("unknown", "ERROR")), + List.of(option("overflow")), + List.of(option("overflow", "ERROR"), option("overflow", "SILENT")))) { + assertThrows( + UnsupportedOperationException.class, + () -> invocation("add", 32, 1, 2, options).accept(toRex, Context.newContext())); + } + } + + @Test + void sqlExportNamesTheActualNativeBehavior() throws Exception { + for (int width : List.of(8, 16, 32, 64)) { + String sqlType = + width == 8 ? "TINYINT" : width == 16 ? "SMALLINT" : width == 32 ? "INT" : "BIGINT"; + Project project = + (Project) + new SqlToSubstrait() + .convert( + "SELECT a+b, a-b, a*b, a/b, mod(a,b) FROM numbers", + SubstraitCreateStatementParser.processCreateStatementsToCatalog( + "CREATE TABLE numbers (a " + sqlType + ", b " + sqlType + ")")) + .getRoots() + .get(0) + .getInput(); + for (Expression expression : project.getExpressions()) { + Expression.ScalarFunctionInvocation function = + (Expression.ScalarFunctionInvocation) expression; + if (width <= 16 && !function.declaration().name().equals("modulus")) { + assertTrue(function.options().stream().noneMatch(o -> o.getName().equals("overflow"))); + } else { + assertTrue(function.options().contains(option("overflow", "SILENT"))); + } + if (function.declaration().name().equals("divide")) { + assertTrue(function.options().contains(option("on_division_by_zero", "ERROR"))); + } + } + } + } + + @Test + void mixedIntegerWidthsExportTheBoundVariantsOptions() throws Exception { + for (String[] types : + List.of( + new String[] {"INTEGER", "SMALLINT", "i32"}, + new String[] {"BIGINT", "INTEGER", "i64"})) { + String query = "SELECT a+b, a-b, a*b, a/b, mod(a,b) FROM numbers"; + String creates = "CREATE TABLE numbers (a " + types[0] + ", b " + types[1] + ")"; + 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; + assertTrue( + function.declaration().key().endsWith(":" + types[2] + "_" + types[2]), + function.declaration().key()); + assertTrue( + function.options().contains(option("overflow", "SILENT")), + function.declaration().key() + " " + function.options()); + if (function.declaration().name().equals("divide")) { + assertTrue(function.options().contains(option("on_domain_error", "ERROR"))); + assertTrue(function.options().contains(option("on_division_by_zero", "ERROR"))); + } + if (function.declaration().name().equals("modulus")) { + assertTrue(function.options().contains(option("division_type", "TRUNCATE"))); + assertTrue(function.options().contains(option("on_domain_error", "ERROR"))); + } + } + assertFullRoundTrip(query, creates); + } + } + + @Test + void undeclaredPreferencesCannotHideBehindASupportedFallback() { + for (String[] values : + List.of(new String[] {"WRAP", "SILENT"}, new String[] {"SILENT", "WRAP"})) { + UnsupportedOperationException failure = + assertThrows( + UnsupportedOperationException.class, + () -> + invocation("add", 32, 1, 2, List.of(option("overflow", values))) + .accept(toRex, Context.newContext())); + assertTrue(failure.getMessage().contains("does not declare value WRAP")); + } + assertTrue( + export( + invocation("add", 32, 1, 2, List.of(option("overflow", "silent"))) + .accept(toRex, Context.newContext())) + .options() + .contains(option("overflow", "SILENT"))); + } + + @Test + void mixedWidthCheckedCallsExportTheBoundVariantsOptions() { + RexBuilder rexBuilder = new RexBuilder(typeFactory); + RexNode left = + rexBuilder.makeExactLiteral( + BigDecimal.ONE, TypeConverter.DEFAULT.toCalcite(typeFactory, R.I32)); + RexNode right = + rexBuilder.makeExactLiteral( + BigDecimal.valueOf(2), TypeConverter.DEFAULT.toCalcite(typeFactory, R.I64)); + for (org.apache.calcite.sql.SqlOperator operator : + List.of( + SqlStdOperatorTable.CHECKED_PLUS, + SqlStdOperatorTable.CHECKED_MINUS, + SqlStdOperatorTable.CHECKED_MULTIPLY, + SqlStdOperatorTable.CHECKED_DIVIDE)) { + RexNode call = + rexBuilder.makeCall( + TypeConverter.DEFAULT.toCalcite(typeFactory, R.I64), operator, List.of(left, right)); + Expression.ScalarFunctionInvocation exported = export(call); + assertTrue(exported.declaration().key().endsWith(":i64_i64")); + assertTrue(exported.options().contains(option("overflow", "ERROR"))); + RexCall imported = (RexCall) exported.accept(toRex, Context.newContext()); + assertEquals(operator, imported.getOperator()); + } + } + + @Test + void parserConformanceDoesNotChangeTheRexOperatorExportPolicy() throws Exception { + for (SqlConformanceEnum conformance : + List.of( + SqlConformanceEnum.LENIENT, + SqlConformanceEnum.BIG_QUERY, + SqlConformanceEnum.SQL_SERVER_2008, + SqlConformanceEnum.MYSQL_5)) { + ConverterProvider provider = + ConverterProvider.builder() + .sqlParserConfig( + ConverterProvider.DEFAULT_SQL_PARSER_CONFIG.withConformance(conformance)) + .build(); + Project project = + (Project) + new SqlToSubstrait(provider) + .convert( + "SELECT a+b FROM numbers", + SubstraitCreateStatementParser.processCreateStatementsToCatalog( + "CREATE TABLE numbers (a BIGINT, b BIGINT)")) + .getRoots() + .get(0) + .getInput(); + Expression.ScalarFunctionInvocation function = + (Expression.ScalarFunctionInvocation) project.getExpressions().get(0); + assertEquals(List.of(option("overflow", "SILENT")), function.options(), conformance.name()); + } + } +} diff --git a/isthmus/src/test/java/io/substrait/isthmus/LambdaExpressionTest.java b/isthmus/src/test/java/io/substrait/isthmus/LambdaExpressionTest.java index 476f8a8df..cc4e174af 100644 --- a/isthmus/src/test/java/io/substrait/isthmus/LambdaExpressionTest.java +++ b/isthmus/src/test/java/io/substrait/isthmus/LambdaExpressionTest.java @@ -81,7 +81,9 @@ void lambdaWithArithmeticBody() { Expression.Lambda lambda = lb.lambda( List.of(R.I64), - params -> sb.scalarFn(ARITH, "add:i64_i64", R.I64, params.ref(0), params.ref(0))); + params -> + withSilentOverflow( + sb.scalarFn(ARITH, "add:i64_i64", R.I64, params.ref(0), params.ref(0)))); List exprs = new ArrayList<>(); exprs.add(lambda); diff --git a/isthmus/src/test/java/io/substrait/isthmus/NestedExpressionsTest.java b/isthmus/src/test/java/io/substrait/isthmus/NestedExpressionsTest.java index 5144c9875..8c083cbe9 100644 --- a/isthmus/src/test/java/io/substrait/isthmus/NestedExpressionsTest.java +++ b/isthmus/src/test/java/io/substrait/isthmus/NestedExpressionsTest.java @@ -29,8 +29,10 @@ class NestedExpressionsTest extends PlanTestBase { Expression literalExpression = Expression.BoolLiteral.builder().value(true).build(); - Expression.ScalarFunctionInvocation nonLiteralExpression = sb.add(sb.i32(7), sb.i32(42)); - Expression.ScalarFunctionInvocation nonLiteralExpression2 = sb.add(sb.i32(3), sb.i32(4)); + Expression.ScalarFunctionInvocation nonLiteralExpression = + withSilentOverflow(sb.add(sb.i32(7), sb.i32(42))); + Expression.ScalarFunctionInvocation nonLiteralExpression2 = + withSilentOverflow(sb.add(sb.i32(3), sb.i32(4))); final List tableType = List.of(R.I32, R.FP32, N.STRING, N.BOOLEAN, N.STRING); final Rel commonTable = diff --git a/isthmus/src/test/java/io/substrait/isthmus/PlanTestBase.java b/isthmus/src/test/java/io/substrait/isthmus/PlanTestBase.java index c4be9290a..d992ef80b 100644 --- a/isthmus/src/test/java/io/substrait/isthmus/PlanTestBase.java +++ b/isthmus/src/test/java/io/substrait/isthmus/PlanTestBase.java @@ -8,6 +8,7 @@ import com.google.common.io.Resources; import io.substrait.dsl.SubstraitBuilder; import io.substrait.expression.Expression; +import io.substrait.expression.FunctionOption; import io.substrait.extension.ExtensionCollector; import io.substrait.extension.SimpleExtension; import io.substrait.isthmus.sql.SubstraitCreateStatementParser; @@ -88,6 +89,14 @@ protected PlanTestBase(ConverterProvider converterProvider) { this.substraitToCalcite = new SubstraitToCalcite(converterProvider); } + protected static Expression.ScalarFunctionInvocation withSilentOverflow( + Expression.ScalarFunctionInvocation expression) { + return Expression.ScalarFunctionInvocation.builder() + .from(expression) + .addOptions(FunctionOption.builder().name("overflow").addValues("SILENT").build()) + .build(); + } + public static String asString(String resource) throws IOException { return Resources.toString(Resources.getResource(resource), StandardCharsets.UTF_8); } diff --git a/isthmus/src/test/java/io/substrait/isthmus/ProjectTest.java b/isthmus/src/test/java/io/substrait/isthmus/ProjectTest.java index b1f1c636b..c0cdb5161 100644 --- a/isthmus/src/test/java/io/substrait/isthmus/ProjectTest.java +++ b/isthmus/src/test/java/io/substrait/isthmus/ProjectTest.java @@ -10,7 +10,10 @@ class ProjectTest extends PlanTestBase { @Test void avoidProjectRemapOnEmptyInput() { Rel projection = - Project.builder().input(emptyTable).addExpressions(sb.add(sb.i32(1), sb.i32(2))).build(); + Project.builder() + .input(emptyTable) + .addExpressions(withSilentOverflow(sb.add(sb.i32(1), sb.i32(2)))) + .build(); assertFullRoundTrip(projection); } } diff --git a/isthmus/src/test/java/io/substrait/isthmus/VirtualTableScanTest.java b/isthmus/src/test/java/io/substrait/isthmus/VirtualTableScanTest.java index 84551bf38..4ae064c42 100644 --- a/isthmus/src/test/java/io/substrait/isthmus/VirtualTableScanTest.java +++ b/isthmus/src/test/java/io/substrait/isthmus/VirtualTableScanTest.java @@ -59,7 +59,7 @@ void expressionContainingVirtualTable() { virtualTable( schema, List.of(sb.i32(2), sb.add(sb.fp64(4.4), sb.fp64(4.5))), - List.of(sb.multiply(sb.i32(6), sb.i32(2)), sb.fp64(8.8))); + List.of(withSilentOverflow(sb.multiply(sb.i32(6), sb.i32(2))), sb.fp64(8.8))); // Check the specific Calcite encoding RelNode relNode = substraitToCalcite.convert(virtualTableScan); @@ -242,7 +242,7 @@ void structColumnWithAComputedField() { schema, List.of( ExpressionCreator.nestedStruct( - false, sb.multiply(sb.i32(6), sb.i32(2)), sb.fp64(2.0)))); + false, withSilentOverflow(sb.multiply(sb.i32(6), sb.i32(2))), sb.fp64(2.0)))); RelNode relNode = substraitToCalcite.convert(virtualTableScan); assertEquals( @@ -414,7 +414,8 @@ void aComputedFieldInsideANullableStructConvertsBack() { schema, List.of( ExpressionCreator.nestedStruct( - true, List.of(sb.multiply(sb.i32(6), sb.i32(2)), sb.fp64(2.0))))); + true, + List.of(withSilentOverflow(sb.multiply(sb.i32(6), sb.i32(2))), sb.fp64(2.0))))); RelNode relNode = substraitToCalcite.convert(virtualTableScan); assertEquals( diff --git a/isthmus/src/test/java/io/substrait/isthmus/VirtualTableTest.java b/isthmus/src/test/java/io/substrait/isthmus/VirtualTableTest.java index 2d9c5d9a7..5263543ca 100644 --- a/isthmus/src/test/java/io/substrait/isthmus/VirtualTableTest.java +++ b/isthmus/src/test/java/io/substrait/isthmus/VirtualTableTest.java @@ -46,7 +46,7 @@ private VirtualTableScan computedRows() { return virtualTable( schema, List.of(sb.i32(2), sb.add(sb.fp64(4.4), sb.fp64(4.5))), - List.of(sb.multiply(sb.i32(6), sb.i32(2)), sb.fp64(8.8))); + List.of(withSilentOverflow(sb.multiply(sb.i32(6), sb.i32(2))), sb.fp64(8.8))); } @Test @@ -74,7 +74,9 @@ void theRuleExpandsASingleRowIntoAProjection() { RelNode expanded = plan( substraitToCalcite.convert( - virtualTable(schema, List.of(sb.multiply(sb.i32(6), sb.i32(2)), sb.fp64(8.8)))), + virtualTable( + schema, + List.of(withSilentOverflow(sb.multiply(sb.i32(6), sb.i32(2))), sb.fp64(8.8)))), VirtualTableExpansionRule.instance()); assertEquals( @@ -318,7 +320,7 @@ void sqlGenerationReachesATableInsideASubquery() { VirtualTableScan oneRow = virtualTable( NamedStruct.of(List.of("col1"), R.struct(R.I32)), - List.of(sb.multiply(sb.i32(6), sb.i32(2)))); + List.of(withSilentOverflow(sb.multiply(sb.i32(6), sb.i32(2))))); Rel root = sb.project( input -> List.of(sb.scalarSubquery(oneRow, R.I32)), diff --git a/isthmus/src/test/resources/lambdas/lambda-with-function.json b/isthmus/src/test/resources/lambdas/lambda-with-function.json index 6252a2f63..3a5dbe4d8 100644 --- a/isthmus/src/test/resources/lambdas/lambda-with-function.json +++ b/isthmus/src/test/resources/lambdas/lambda-with-function.json @@ -118,6 +118,7 @@ "body": { "scalarFunction": { "functionReference": 1, + "options": [{"name": "overflow", "preference": ["SILENT"]}], "outputType": { "i32": { "nullability": "NULLABILITY_REQUIRED"