Skip to content
Merged
Original file line number Diff line number Diff line change
Expand Up @@ -592,7 +592,7 @@ public RexNode visit(Expression.ScalarFunctionInvocation expr, Context context)
throws RuntimeException {
SqlOperator operator =
scalarFunctionConverter
.getSqlOperatorFromSubstraitFunc(expr.declaration().key(), expr.outputType())
.getSqlOperatorFromSubstraitFunc(expr)
.orElseThrow(
() ->
new IllegalArgumentException(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
import io.substrait.expression.Expression;
import io.substrait.expression.ExpressionCreator;
import io.substrait.expression.FunctionArg;
import io.substrait.expression.FunctionOption;
import io.substrait.extension.DefaultExtensionCatalog;
import io.substrait.extension.SimpleExtension;
import io.substrait.isthmus.CallConverter;
Expand Down Expand Up @@ -48,6 +49,9 @@ public class ScalarFunctionConverter
*/
private final List<ScalarFunctionMapper> mappers;

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

/**
* Creates a converter with the given functions and type factory.
*
Expand Down Expand Up @@ -144,6 +148,14 @@ public Stream<RexNode> getOperands() {
private Optional<Expression> defaultConvert(
RexCall call, Function<RexNode, Expression> topLevelConverter) {
FunctionFinder finder = signatures.get(call.op);
if (finder == null) {
for (ScalarFunctionOptionPolicy policy : optionPolicies) {
finder = signatures.get(policy.signatureOperator(call));
if (finder != null) {
break;
}
}
}
WrappedScalarCall wrapped = new WrappedScalarCall(call);

return attemptMatch(finder, wrapped, topLevelConverter);
Expand Down Expand Up @@ -185,7 +197,12 @@ protected Expression generateBinding(
List<? extends FunctionArg> arguments,
Type outputType) {
if (!DefaultExtensionCatalog.FUNCTIONS_DATETIME.equals(function.getAnchor().urn())) {
return ExpressionCreator.scalarFunction(function, outputType, arguments);
return Expression.ScalarFunctionInvocation.builder()
.declaration(function)
.outputType(outputType)
.addAllArguments(arguments)
.options(options(call.delegate, function))
.build();
}
// The datetime extension declares its results by parameter, where Calcite keeps an operand's
// own type: add(date, interval_day<P>) is a precision_timestamp<P> there and a DATE here. The
Expand Down Expand Up @@ -378,6 +395,49 @@ public List<FunctionArg> getExpressionArguments(Expression.ScalarFunctionInvocat
return getMappedExpressionArguments(expression).orElseGet(expression::arguments);
}

/**
* Resolves an invocation through the existing operator mapping and its option policy.
*
* @param expression the Substrait scalar invocation
* @return the selected operator, or empty when no mapping exists
*/
public Optional<SqlOperator> getSqlOperatorFromSubstraitFunc(
Expression.ScalarFunctionInvocation expression) {
return getSqlOperatorFromSubstraitFunc(expression.declaration().key(), expression.outputType())
.map(operator -> resolveOptions(expression, operator));
}

/**
* Selects an operator that honors the requested options. Custom converters can override this
* together with {@link #options} to supply their own semantics in both directions.
*
* @param expression the Substrait invocation
* @param operator the selected Calcite operator
* @return the operator implementing a supported preference
* @throws UnsupportedOperationException if no requested preference is supported
*/
protected SqlOperator resolveOptions(
Expression.ScalarFunctionInvocation expression, SqlOperator operator) {
for (ScalarFunctionOptionPolicy policy : optionPolicies) {
operator = policy.resolve(expression, operator);
}
return operator;
}

/**
* Returns the option preferences implemented by the matched Calcite call.
*
* @param call the Calcite call
* @param function the bound Substrait variant
* @return the preferences to export
*/
protected List<FunctionOption> options(
RexCall call, SimpleExtension.ScalarFunctionVariant function) {
return optionPolicies.stream()
.flatMap(policy -> policy.forCall(call, function).stream())
.collect(Collectors.toList());
}

/**
* Builds a Calcite call, applying a reverse mapping when the native operator is not executable.
*
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,20 @@
package io.substrait.isthmus.expression;

import io.substrait.expression.Expression;
import io.substrait.expression.FunctionOption;
import io.substrait.extension.SimpleExtension.ScalarFunctionVariant;
import java.util.List;
import org.apache.calcite.rex.RexCall;
import org.apache.calcite.sql.SqlOperator;

/** Option semantics for one family of scalar functions, in both conversion directions. */
interface ScalarFunctionOptionPolicy {
SqlOperator resolve(Expression.ScalarFunctionInvocation expression, SqlOperator operator);

List<FunctionOption> forCall(RexCall call, ScalarFunctionVariant function);

/** Returns the operator whose signature also binds this call, such as an unchecked equivalent. */
default SqlOperator signatureOperator(RexCall call) {
return call.getOperator();
}
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,140 @@
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.Map;
import java.util.Set;
import org.apache.calcite.rex.RexCall;
import org.apache.calcite.sql.SqlOperator;
import org.apache.calcite.sql.fun.SqlLibraryOperators;
import org.apache.calcite.sql.fun.SqlStdOperatorTable;

/** The spec v0.103.0 string options supported by the corresponding Calcite operators. */
final class StringFunctionOptions implements ScalarFunctionOptionPolicy {
private static final Map<String, Binding> BINDINGS =
Map.of(
"concat",
new Binding(
"null_handling",
"ACCEPT_NULLS",
Set.of(SqlStdOperatorTable.CONCAT, SqlLibraryOperators.CONCAT_FUNCTION)),
"like",
new Binding("case_sensitivity", "CASE_SENSITIVE", Set.of(SqlStdOperatorTable.LIKE)),
"replace",
new Binding(
"case_sensitivity", "CASE_SENSITIVE", Set.of(SqlStdOperatorTable.REPLACE)),
"starts_with",
new Binding(
"case_sensitivity", "CASE_SENSITIVE", Set.of(SqlLibraryOperators.STARTS_WITH)),
"ends_with",
new Binding(
"case_sensitivity", "CASE_SENSITIVE", Set.of(SqlLibraryOperators.ENDS_WITH)),
"strpos",
new Binding(
"case_sensitivity", "CASE_SENSITIVE", Set.of(SqlStdOperatorTable.POSITION)),
"substring",
new Binding(
"negative_start", "LEFT_OF_BEGINNING", Set.of(SqlStdOperatorTable.SUBSTRING)),
"lower", new Binding("char_set", "UTF8", Set.of(SqlStdOperatorTable.LOWER)),
"upper", new Binding("char_set", "UTF8", Set.of(SqlStdOperatorTable.UPPER)),
"initcap", new Binding("char_set", null, Set.of(SqlStdOperatorTable.INITCAP)));

@Override
public List<FunctionOption> forCall(RexCall call, ScalarFunctionVariant function) {
Binding binding = binding(function);
if (binding == null
|| binding.value() == null
|| !binding.operators().contains(call.getOperator())) {
return List.of();
}
// An omitted option lets a consumer choose any supported behavior. Emit the one
// the producer's Calcite operator actually implements instead of a YAML default.
return List.of(
FunctionOption.builder().name(binding.name()).addValues(binding.value()).build());
}

@Override
public SqlOperator resolve(Expression.ScalarFunctionInvocation expression, SqlOperator operator) {
Binding binding = binding(expression.declaration());
if (binding == null) {
return operator;
}
if (!expression.options().isEmpty() && !binding.operators().contains(operator)) {
throw new UnsupportedOperationException(
"No string option policy for Calcite operator " + operator.getName());
}
for (FunctionOption option : expression.options()) {
if (!binding.name().equalsIgnoreCase(option.getName())) {
throw new UnsupportedOperationException(
"Unsupported " + expression.declaration().name() + " option: " + option.getName());
}
Comment thread
nielspardon marked this conversation as resolved.
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);
}
}
if (binding.value() == null) {
// The spec lists initcap's charsets without defining ASCII word boundaries.
throw new UnsupportedOperationException(
"No established Calcite initcap charset option policy");
}
// Preferences name acceptable behaviors, in order. These operators implement one
// behavior each, so skip other preferences and reject when none is supported.
if (option.values().stream().noneMatch(binding.value()::equalsIgnoreCase)) {
throw new UnsupportedOperationException(
"Calcite "
+ expression.declaration().name()
+ " requires "
+ binding.name()
+ " "
+ binding.value()
+ "; preferences: "
+ option.values());
}
}
return operator;
}

private static Binding binding(ScalarFunctionVariant function) {
return DefaultExtensionCatalog.FUNCTIONS_STRING.equals(function.urn())
? BINDINGS.get(function.name())
: null;
}

private static final class Binding {
private final String name;
private final String value;
private final Set<SqlOperator> operators;

private Binding(String name, String value, Set<SqlOperator> operators) {
this.name = name;
this.value = value;
this.operators = operators;
}

private String name() {
return name;
}

private String value() {
return value;
}

private Set<SqlOperator> operators() {
return operators;
}
}
}
Loading
Loading