Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,8 @@
import io.substrait.isthmus.Utils;
import io.substrait.isthmus.expression.FunctionMappings.Sig;
import io.substrait.type.Type;
import io.substrait.type.TypeCreator;
import io.substrait.type.TypeExpressionEvaluator;
import io.substrait.util.Util;
import java.util.ArrayList;
import java.util.Collection;
Expand All @@ -47,6 +49,7 @@
import org.apache.calcite.rex.RexNode;
import org.apache.calcite.sql.SqlKind;
import org.apache.calcite.sql.SqlOperator;
import org.apache.calcite.sql.type.SqlTypeUtil;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;

Expand Down Expand Up @@ -528,6 +531,23 @@ private Optional<F> signatureMatch(List<Type> inputTypes, Type outputType) {
return Optional.empty();
}

private List<F> signatureMatches(List<Type> inputTypes, Type outputType) {
List<F> matches = new ArrayList<>();
for (F function : functions) {
List<SimpleExtension.Argument> args = function.requiredArguments();
// A variant without arguments takes no operands; signatureMatch never reaches one here
// because it stops at the first match, and the check below would index it at -1.
if (args.isEmpty() && !inputTypes.isEmpty()) {
continue;
}
if (isReturnTypeMatch(outputType, function.returnType())
&& inputTypesMatchDefinedArguments(inputTypes, args)) {
matches.add(function);
}
}
return matches;
}

private boolean isReturnTypeMatch(final Type outputType, final TypeExpression funcReturnType) {
if (funcReturnType instanceof ParameterizedType) {
return isMatch(outputType, (ParameterizedType) funcReturnType);
Expand Down Expand Up @@ -817,9 +837,28 @@ private Optional<T> matchByLeastRestrictive(
Type type = typeConverter.toSubstrait(leastRestrictive);
Optional<F> out = singularInputType.orElseThrow().tryMatch(type, outputType);

List<RelDataType> calciteTypes =
call.getOperands().map(RexNode::getType).collect(Collectors.toList());
return out.map(
declaration -> {
List<Expression> coercedArgs = coerceArguments(operands, type);
List<Expression> coercedArgs =
Streams.zip(
operands.stream(),
calciteTypes.stream(),
(operand, calciteType) -> coerceToward(operand, calciteType, type))
.collect(Collectors.toList());
Optional<Type> common = losslessCommonType(operands, calciteTypes, type);
if (!binds(declaration, coercedArgs, outputType) && common.isPresent()) {
// A parameter the operands share, like the any1 of gte(any1, any1), only binds once
// they have one type, so they all take a type that holds each of them exactly.
List<Expression> exact =
operands.stream()
.map(operand -> castUnlessEqual(operand, common.get()))
.collect(Collectors.toList());
if (binds(declaration, exact, outputType)) {
coercedArgs = exact;
}
}
declaration.validateOutputType(coercedArgs, outputType);
return generateBinding(call, out.get(), coercedArgs, outputType);
});
Expand All @@ -842,16 +881,32 @@ private Optional<T> matchCoerced(C call, Type outputType, List<Expression> expre
.collect(Collectors.toList());

// See if all the input types can be made to match the function
Optional<F> matchFunction = signatureMatch(operandTypes, outputType);
if (matchFunction.isEmpty()) {
List<F> candidates = signatureMatches(operandTypes, outputType);
if (candidates.isEmpty()) {
return Optional.empty();
}

List<Expression> coercedArgs =
Streams.zip(
expressions.stream(), operandTypes.stream(), FunctionConverter::coerceArgument)
.collect(Collectors.toList());
return Optional.of(generateBinding(call, matchFunction.get(), coercedArgs, outputType));
// The signature match takes any string for any other, so a candidate may not bind the
// operands as they are: a char(n) operand of a varchar declaration becomes a varchar(n), and
// one of a string declaration a string. The first candidate that binds is taken, and the
// first one as before when none does.
List<Expression> varchars =
coercedArgs.stream()
.map(FunctionConverter::fixedCharAsVarChar)
.collect(Collectors.toList());
for (F candidate : candidates) {
for (List<Expression> arguments :
List.of(coercedArgs, varchars, asDeclaredStrings(candidate, coercedArgs))) {
if (binds(candidate, arguments, outputType)) {
return Optional.of(generateBinding(call, candidate, arguments, outputType));
}
}
}
return Optional.of(generateBinding(call, candidates.get(0), coercedArgs, outputType));
}

/**
Expand Down Expand Up @@ -895,14 +950,159 @@ public interface GenericCall {
}

/**
* Coerces arguments to the target type when mismatched (ignores nullability/parameters).
*
* @param arguments input expressions
* @param targetType target Substrait type
* @return list of coerced expressions (casts applied as needed)
* Returns whether the declaration binds the given arguments: each of its parameters takes one
* value from the argument types, each argument it declares concretely has exactly that type, and
* a string result is not shorter than the call's own. concat:vchar declares one operand's length
* as its result, so {@code CHAR(2) || CHAR(2)} does not bind it: the four characters would come
* out as a varchar(2).
*/
private static List<Expression> coerceArguments(List<Expression> arguments, Type targetType) {
return arguments.stream().map(a -> coerceArgument(a, targetType)).collect(Collectors.toList());
private static boolean binds(
SimpleExtension.Function declaration, List<Expression> arguments, Type outputType) {
List<Type> types = arguments.stream().map(Expression::getType).collect(Collectors.toList());
try {
TypeExpressionEvaluator.checkBindings(declaration.args(), declaration.variadic(), types);
} catch (UnsupportedOperationException e) {
return false;
}
List<ParameterizedType> declared = declaredArgumentTypes(declaration);
for (int index = 0; index < types.size() && !declared.isEmpty(); index++) {
ParameterizedType argument = declared.get(Math.min(index, declared.size() - 1));
if (argument instanceof Type
&& !((Type) argument).equalsIgnoringNullability(types.get(index))) {
return false;
}
}
Optional<Integer> callLength = stringLength(outputType);
if (callLength.isEmpty()) {
return true;
}
try {
return stringLength(declaration.resolveType(types))
.map(length -> length >= callLength.get())
.orElse(true);
} catch (UnsupportedOperationException e) {
return true;
}
}

private static List<ParameterizedType> declaredArgumentTypes(
SimpleExtension.Function declaration) {
List<ParameterizedType> types = new ArrayList<>();
for (SimpleExtension.Argument argument : declaration.args()) {
if (argument instanceof SimpleExtension.ValueArgument) {
types.add(((SimpleExtension.ValueArgument) argument).value());
} else if (argument instanceof SimpleExtension.TypeArgument) {
types.add(((SimpleExtension.TypeArgument) argument).type());
}
}
return types;
}

private static Optional<Integer> stringLength(Type type) {
if (type instanceof Type.VarChar) {
return Optional.of(((Type.VarChar) type).length());
}
if (type instanceof Type.FixedChar) {
return Optional.of(((Type.FixedChar) type).length());
}
return Optional.empty();
}

/**
* Coerces an operand toward the least restrictive type of a call's operands. An integer operand
* of a decimal call becomes the decimal that holds every value of its type, as Calcite types it
* for decimal arithmetic, rather than the least restrictive decimal, whose scale it does not
* have: {@code decimal(7,2) * INTEGER} multiplies by a {@code decimal(10,0)}. The digits are
* Calcite's, so they follow the type system in use.
*/
private static Expression coerceToward(Expression operand, RelDataType calciteType, Type target) {
if (target instanceof Type.Decimal && SqlTypeUtil.isIntType(calciteType)) {
return ExpressionCreator.cast(
TypeCreator.of(operand.getType().nullable()).decimal(calciteType.getPrecision(), 0),
operand,
Expression.FailureBehavior.THROW_EXCEPTION);
}
return coerceArgument(operand, target);
}

/**
* Returns a type every operand converts to without losing a value. Calcite's least restrictive
* type is that, except for decimals: past precision 38 it gives up scale, so a {@code
* decimal(38,0)} and a {@code decimal(3,2)} meet at {@code decimal(38,0)} and the second becomes
* 0. For decimals the type is built from the most integer digits and the largest scale instead,
* and there is none when that needs more than 38 digits.
*/
private static Optional<Type> losslessCommonType(
List<Expression> operands, List<RelDataType> calciteTypes, Type leastRestrictive) {
if (!(leastRestrictive instanceof Type.Decimal)) {
return Optional.of(leastRestrictive);
}
int integerDigits = 0;
int scale = 0;
for (int index = 0; index < operands.size(); index++) {
Type type = operands.get(index).getType();
if (type instanceof Type.Decimal) {
Type.Decimal decimal = (Type.Decimal) type;
integerDigits = Math.max(integerDigits, decimal.precision() - decimal.scale());
scale = Math.max(scale, decimal.scale());
} else if (SqlTypeUtil.isIntType(calciteTypes.get(index))) {
integerDigits = Math.max(integerDigits, calciteTypes.get(index).getPrecision());
} else {
return Optional.empty();
}
}
if (integerDigits + scale > 38) {
return Optional.empty();
}
return Optional.of(
TypeCreator.of(leastRestrictive.nullable()).decimal(integerDigits + scale, scale));
}

private static Expression castUnlessEqual(Expression operand, Type target) {
Type type = operand.getType();
if (type.equalsIgnoringNullability(target)) {
return operand;
}
return ExpressionCreator.cast(
target.withNullable(type.nullable()), operand, Expression.FailureBehavior.THROW_EXCEPTION);
}

/**
* Casts each char or varchar operand to string where the declaration takes a string there, as
* like:str_str and replace:str_str_str do.
*/
private static List<Expression> asDeclaredStrings(
SimpleExtension.Function declaration, List<Expression> operands) {
List<ParameterizedType> declared = declaredArgumentTypes(declaration);
List<Expression> arguments = new ArrayList<>(operands.size());
for (int index = 0; index < operands.size(); index++) {
Expression operand = operands.get(index);
Type type = operand.getType();
ParameterizedType argument =
declared.isEmpty() ? null : declared.get(Math.min(index, declared.size() - 1));
if (argument instanceof Type.Str
&& (type instanceof Type.VarChar || type instanceof Type.FixedChar)) {
arguments.add(
ExpressionCreator.cast(
TypeCreator.of(type.nullable()).STRING,
operand,
Expression.FailureBehavior.THROW_EXCEPTION));
} else {
arguments.add(operand);
}
}
return arguments;
}

private static Expression fixedCharAsVarChar(Expression operand) {
if (!(operand.getType() instanceof Type.FixedChar)) {
return operand;
}
Type.FixedChar fixedChar = (Type.FixedChar) operand.getType();
return ExpressionCreator.cast(
TypeCreator.of(fixedChar.nullable()).varChar(fixedChar.length()),
operand,
Expression.FailureBehavior.THROW_EXCEPTION);
}

/**
Expand Down
Loading
Loading