Skip to content
Merged
Original file line number Diff line number Diff line change
@@ -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();
Comment thread
nielspardon marked this conversation as resolved.
}

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<FunctionOption> forCall(RexCall call, ScalarFunctionVariant function) {
Binding binding = binding(function);
if (binding == null
|| (call.getOperator() != binding.normal && call.getOperator() != binding.checked)) {
return List.of();
}
List<FunctionOption> 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 =
Comment thread
nielspardon marked this conversation as resolved.
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<String, String> selected = new LinkedHashMap<>();
for (FunctionOption option : expression.options()) {
ScalarFunctionOptionPolicy.requireDeclaredValues(expression, option);
String name = option.getName().toLowerCase(java.util.Locale.ROOT);
List<String> 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;
}
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -50,7 +50,7 @@ public class ScalarFunctionConverter
private final List<ScalarFunctionMapper> mappers;

private final List<ScalarFunctionOptionPolicy> optionPolicies =
List.of(new StringFunctionOptions());
List.of(new StringFunctionOptions(), new IntegerFunctionOptions());

/**
* Creates a converter with the given functions and type factory.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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);
}
}
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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();
}
}
Loading
Loading