From 04f928e46a39395a67e27d09ef3c36ec959ecaba Mon Sep 17 00:00:00 2001 From: Aleksandr Efimov Date: Wed, 30 Sep 2026 16:44:00 +0300 Subject: [PATCH 1/3] fix(isthmus): cast operands until the function's declaration binds them An integer operand of decimal arithmetic was cast to the least restrictive decimal, whose scale it does not have, so decimal(7,2) * INT multiplied by a decimal(12,2). It now becomes the decimal that holds its type, decimal(10,0) for an INT. Where the operands share a parameter, like the any1 of gte(any1, any1), and do not bind it, they take the least restrictive type exactly; before, decimals of two precisions were left as they were. A char(n) operand of a varchar declaration becomes a varchar(n). --- .../isthmus/expression/FunctionConverter.java | 100 ++++++++++++++++-- .../isthmus/OperandCoercionTest.java | 80 ++++++++++++++ 2 files changed, 172 insertions(+), 8 deletions(-) create mode 100644 isthmus/src/test/java/io/substrait/isthmus/OperandCoercionTest.java diff --git a/isthmus/src/main/java/io/substrait/isthmus/expression/FunctionConverter.java b/isthmus/src/main/java/io/substrait/isthmus/expression/FunctionConverter.java index 913f4adc3..41c64d062 100644 --- a/isthmus/src/main/java/io/substrait/isthmus/expression/FunctionConverter.java +++ b/isthmus/src/main/java/io/substrait/isthmus/expression/FunctionConverter.java @@ -23,6 +23,7 @@ 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.util.Util; import java.util.ArrayList; import java.util.Collection; @@ -819,7 +820,21 @@ private Optional matchByLeastRestrictive( return out.map( declaration -> { - List coercedArgs = coerceArguments(operands, type); + List coercedArgs = + operands.stream() + .map(operand -> coerceToward(operand, type)) + .collect(Collectors.toList()); + if (!binds(declaration, coercedArgs)) { + // A parameter the operands share, like the any1 of gte(any1, any1), only binds once + // they have one type, so they all take the least restrictive one exactly. + List exact = + operands.stream() + .map(operand -> castUnlessEqual(operand, type)) + .collect(Collectors.toList()); + if (binds(declaration, exact)) { + coercedArgs = exact; + } + } declaration.validateOutputType(coercedArgs, outputType); return generateBinding(call, out.get(), coercedArgs, outputType); }); @@ -851,6 +866,17 @@ private Optional matchCoerced(C call, Type outputType, List expre Streams.zip( expressions.stream(), operandTypes.stream(), FunctionConverter::coerceArgument) .collect(Collectors.toList()); + if (!binds(matchFunction.get(), coercedArgs)) { + // The signature match takes a fixed-length string for a varchar one, and the declaration + // then does not bind it: a char(n) operand of a varchar declaration becomes a varchar(n). + List varchars = + coercedArgs.stream() + .map(FunctionConverter::fixedCharAsVarChar) + .collect(Collectors.toList()); + if (binds(matchFunction.get(), varchars)) { + coercedArgs = varchars; + } + } return Optional.of(generateBinding(call, matchFunction.get(), coercedArgs, outputType)); } @@ -895,14 +921,72 @@ 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, that is, whether its parameters take + * one value each from the argument types. + */ + private static boolean binds(SimpleExtension.Function declaration, List arguments) { + try { + declaration.resolveType( + arguments.stream().map(Expression::getType).collect(Collectors.toList())); + return true; + } catch (UnsupportedOperationException e) { + return false; + } + } + + /** + * 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)}. */ - private static List coerceArguments(List arguments, Type targetType) { - return arguments.stream().map(a -> coerceArgument(a, targetType)).collect(Collectors.toList()); + private static Expression coerceToward(Expression operand, Type target) { + if (target instanceof Type.Decimal) { + Optional digits = integerDigits(operand.getType()); + if (digits.isPresent()) { + return ExpressionCreator.cast( + TypeCreator.of(operand.getType().nullable()).decimal(digits.get(), 0), + operand, + Expression.FailureBehavior.THROW_EXCEPTION); + } + } + return coerceArgument(operand, target); + } + + private static Optional integerDigits(Type type) { + if (type instanceof Type.I8) { + return Optional.of(3); + } + if (type instanceof Type.I16) { + return Optional.of(5); + } + if (type instanceof Type.I32) { + return Optional.of(10); + } + if (type instanceof Type.I64) { + return Optional.of(19); + } + return Optional.empty(); + } + + private static Expression castUnlessEqual(Expression operand, Type target) { + Type type = operand.getType(); + if (type.withNullable(target.nullable()).equals(target)) { + return operand; + } + return ExpressionCreator.cast( + target.withNullable(type.nullable()), operand, Expression.FailureBehavior.THROW_EXCEPTION); + } + + 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); } /** diff --git a/isthmus/src/test/java/io/substrait/isthmus/OperandCoercionTest.java b/isthmus/src/test/java/io/substrait/isthmus/OperandCoercionTest.java new file mode 100644 index 000000000..cc42eae82 --- /dev/null +++ b/isthmus/src/test/java/io/substrait/isthmus/OperandCoercionTest.java @@ -0,0 +1,80 @@ +package io.substrait.isthmus; + +import static org.junit.jupiter.api.Assertions.assertDoesNotThrow; +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertInstanceOf; + +import io.substrait.expression.Expression; +import io.substrait.isthmus.sql.SubstraitCreateStatementParser; +import io.substrait.plan.Plan; +import io.substrait.relation.Project; +import io.substrait.type.Type; +import java.util.List; +import java.util.stream.Collectors; +import org.junit.jupiter.api.Test; + +/** + * The operands of a converted call are cast just far enough for the function's declaration to bind + * them: each of its parameters has to take one value from the argument types. + */ +class OperandCoercionTest extends PlanTestBase { + + private static final String CREATES = + "CREATE TABLE t (d7 DECIMAL(7, 2) NOT NULL, d5 DECIMAL(5, 2) NOT NULL, " + + "i INT NOT NULL, c CHAR(10) NOT NULL)"; + + /** gte(any1, any1) binds any1 once, so decimals of two precisions both take the wider one. */ + @Test + void aComparisonOfTwoDecimalsCastsBothToOneType() throws Exception { + Expression.ScalarFunctionInvocation between = call("SELECT d7 BETWEEN 0.99 AND 1.49 FROM t"); + Expression.ScalarFunctionInvocation call = + assertInstanceOf(Expression.ScalarFunctionInvocation.class, between.arguments().get(0)); + + assertEquals("gte:any_any", call.declaration().key()); + assertEquals(List.of(R.decimal(7, 2), R.decimal(7, 2)), argumentTypes(call)); + assertBinds(call); + } + + /** + * multiply(decimal, decimal) binds each operand's own precision, so an integer + * operand becomes the decimal that holds it, not the other operand's decimal. + */ + @Test + void aDecimalTimesAnIntegerCastsTheIntegerToItsOwnDecimal() throws Exception { + Expression.ScalarFunctionInvocation call = call("SELECT d7 * i FROM t"); + + assertEquals("multiply:dec_dec", call.declaration().key()); + assertEquals(List.of(R.decimal(7, 2), R.decimal(10, 0)), argumentTypes(call)); + assertBinds(call); + } + + /** A varchar declaration does not bind a char, so a char(n) operand becomes a varchar(n). */ + @Test + void aCharOperandOfAVarcharFunctionBecomesAVarchar() throws Exception { + Expression.ScalarFunctionInvocation call = call("SELECT c LIKE 'a%' FROM t"); + + assertEquals("like:vchar_vchar", call.declaration().key()); + assertEquals(R.varChar(10), argumentTypes(call).get(0)); + assertBinds(call); + } + + private static Expression.ScalarFunctionInvocation call(String query) throws Exception { + Plan plan = + new SqlToSubstrait() + .convert( + query, SubstraitCreateStatementParser.processCreateStatementsToCatalog(CREATES)); + Expression expression = ((Project) plan.getRoots().get(0).getInput()).getExpressions().get(0); + return assertInstanceOf(Expression.ScalarFunctionInvocation.class, expression); + } + + private static List argumentTypes(Expression.ScalarFunctionInvocation call) { + return call.arguments().stream() + .filter(Expression.class::isInstance) + .map(argument -> ((Expression) argument).getType()) + .collect(Collectors.toList()); + } + + private static void assertBinds(Expression.ScalarFunctionInvocation call) { + assertDoesNotThrow(() -> call.declaration().resolveType(argumentTypes(call))); + } +} From 50a450593448c1c104ab4f5f3adba89d29397f49 Mon Sep 17 00:00:00 2001 From: Aleksandr Efimov Date: Wed, 30 Sep 2026 17:04:50 +0300 Subject: [PATCH 2/3] fix(isthmus): cast shared-parameter operands only to a type that holds them Calcite's least restrictive decimal gives up scale past precision 38, so TPC-H Q6's l_discount, a bare DECIMAL, met 0.03 - 0.01 at decimal(38,0) and the bound became 0. The common decimal is now built from the most integer digits and the largest scale, and operands are left as they are when that needs more than 38 digits. --- .../isthmus/expression/FunctionConverter.java | 42 +++++++++++++++++-- .../isthmus/OperandCoercionTest.java | 23 ++++++++++ 2 files changed, 62 insertions(+), 3 deletions(-) diff --git a/isthmus/src/main/java/io/substrait/isthmus/expression/FunctionConverter.java b/isthmus/src/main/java/io/substrait/isthmus/expression/FunctionConverter.java index 41c64d062..dfc99821d 100644 --- a/isthmus/src/main/java/io/substrait/isthmus/expression/FunctionConverter.java +++ b/isthmus/src/main/java/io/substrait/isthmus/expression/FunctionConverter.java @@ -824,12 +824,13 @@ private Optional matchByLeastRestrictive( operands.stream() .map(operand -> coerceToward(operand, type)) .collect(Collectors.toList()); - if (!binds(declaration, coercedArgs)) { + Optional common = losslessCommonType(operands, type); + if (!binds(declaration, coercedArgs) && 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 the least restrictive one exactly. + // they have one type, so they all take a type that holds each of them exactly. List exact = operands.stream() - .map(operand -> castUnlessEqual(operand, type)) + .map(operand -> castUnlessEqual(operand, common.get())) .collect(Collectors.toList()); if (binds(declaration, exact)) { coercedArgs = exact; @@ -969,6 +970,41 @@ private static Optional integerDigits(Type type) { return Optional.empty(); } + /** + * 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 losslessCommonType( + List operands, Type leastRestrictive) { + if (!(leastRestrictive instanceof Type.Decimal)) { + return Optional.of(leastRestrictive); + } + int integerDigits = 0; + int scale = 0; + for (Expression operand : operands) { + Type type = operand.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 { + Optional digits = integerDigits(type); + if (digits.isEmpty()) { + return Optional.empty(); + } + integerDigits = Math.max(integerDigits, digits.get()); + } + } + 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.withNullable(target.nullable()).equals(target)) { diff --git a/isthmus/src/test/java/io/substrait/isthmus/OperandCoercionTest.java b/isthmus/src/test/java/io/substrait/isthmus/OperandCoercionTest.java index cc42eae82..623116e61 100644 --- a/isthmus/src/test/java/io/substrait/isthmus/OperandCoercionTest.java +++ b/isthmus/src/test/java/io/substrait/isthmus/OperandCoercionTest.java @@ -58,6 +58,29 @@ void aCharOperandOfAVarcharFunctionBecomesAVarchar() throws Exception { assertBinds(call); } + /** + * TPC-H declares l_discount as a bare DECIMAL, which is decimal(38,0). No decimal holds both it + * and 0.02 exactly, so the operands are left as they are rather than cast to a type that drops + * the fraction and turns the bound into 0. + */ + @Test + void noOperandIsCastToADecimalThatLosesItsScale() throws Exception { + Plan plan = + new SqlToSubstrait() + .convert( + "SELECT l_discount BETWEEN 0.03 - 0.01 AND 0.03 + 0.01 FROM lineitem", + TPCH_CATALOG); + Expression.ScalarFunctionInvocation between = + assertInstanceOf( + Expression.ScalarFunctionInvocation.class, + ((Project) plan.getRoots().get(0).getInput()).getExpressions().get(0)); + Expression.ScalarFunctionInvocation gte = + assertInstanceOf(Expression.ScalarFunctionInvocation.class, between.arguments().get(0)); + + Type.Decimal bound = assertInstanceOf(Type.Decimal.class, argumentTypes(gte).get(1)); + assertEquals(2, bound.scale()); + } + private static Expression.ScalarFunctionInvocation call(String query) throws Exception { Plan plan = new SqlToSubstrait() From 486c612ac92e48d013a52352fb7bcd06912adaa5 Mon Sep 17 00:00:00 2001 From: Aleksandr Efimov Date: Wed, 30 Sep 2026 19:44:29 +0300 Subject: [PATCH 3/3] fix(isthmus): try each matching variant until one binds without truncating A signature match took the first variant whatever it bound. Each matching variant is now tried with the operands as they are, with char(n) as varchar(n), and with char and varchar as string where the declaration takes a string, and the first that binds is taken. Binding now also requires a concretely declared argument to have exactly its type, and a string result no shorter than the call's, so CHAR(10) || CHAR(10) binds concat:str rather than a concat:vchar that types it varchar(10). An integer operand's digits come from its Calcite type. --- .../isthmus/expression/FunctionConverter.java | 196 ++++++++++++------ .../isthmus/OperandCoercionTest.java | 48 ++++- 2 files changed, 179 insertions(+), 65 deletions(-) diff --git a/isthmus/src/main/java/io/substrait/isthmus/expression/FunctionConverter.java b/isthmus/src/main/java/io/substrait/isthmus/expression/FunctionConverter.java index dfc99821d..85cef0785 100644 --- a/isthmus/src/main/java/io/substrait/isthmus/expression/FunctionConverter.java +++ b/isthmus/src/main/java/io/substrait/isthmus/expression/FunctionConverter.java @@ -24,6 +24,7 @@ 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; @@ -48,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; @@ -529,6 +531,23 @@ private Optional signatureMatch(List inputTypes, Type outputType) { return Optional.empty(); } + private List signatureMatches(List inputTypes, Type outputType) { + List matches = new ArrayList<>(); + for (F function : functions) { + List 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); @@ -818,21 +837,25 @@ private Optional matchByLeastRestrictive( Type type = typeConverter.toSubstrait(leastRestrictive); Optional out = singularInputType.orElseThrow().tryMatch(type, outputType); + List calciteTypes = + call.getOperands().map(RexNode::getType).collect(Collectors.toList()); return out.map( declaration -> { List coercedArgs = - operands.stream() - .map(operand -> coerceToward(operand, type)) + Streams.zip( + operands.stream(), + calciteTypes.stream(), + (operand, calciteType) -> coerceToward(operand, calciteType, type)) .collect(Collectors.toList()); - Optional common = losslessCommonType(operands, type); - if (!binds(declaration, coercedArgs) && common.isPresent()) { + Optional 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 exact = operands.stream() .map(operand -> castUnlessEqual(operand, common.get())) .collect(Collectors.toList()); - if (binds(declaration, exact)) { + if (binds(declaration, exact, outputType)) { coercedArgs = exact; } } @@ -858,8 +881,8 @@ private Optional matchCoerced(C call, Type outputType, List expre .collect(Collectors.toList()); // See if all the input types can be made to match the function - Optional matchFunction = signatureMatch(operandTypes, outputType); - if (matchFunction.isEmpty()) { + List candidates = signatureMatches(operandTypes, outputType); + if (candidates.isEmpty()) { return Optional.empty(); } @@ -867,18 +890,23 @@ private Optional matchCoerced(C call, Type outputType, List expre Streams.zip( expressions.stream(), operandTypes.stream(), FunctionConverter::coerceArgument) .collect(Collectors.toList()); - if (!binds(matchFunction.get(), coercedArgs)) { - // The signature match takes a fixed-length string for a varchar one, and the declaration - // then does not bind it: a char(n) operand of a varchar declaration becomes a varchar(n). - List varchars = - coercedArgs.stream() - .map(FunctionConverter::fixedCharAsVarChar) - .collect(Collectors.toList()); - if (binds(matchFunction.get(), varchars)) { - coercedArgs = varchars; + // 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 varchars = + coercedArgs.stream() + .map(FunctionConverter::fixedCharAsVarChar) + .collect(Collectors.toList()); + for (F candidate : candidates) { + for (List 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, matchFunction.get(), coercedArgs, outputType)); + return Optional.of(generateBinding(call, candidates.get(0), coercedArgs, outputType)); } /** @@ -922,54 +950,81 @@ public interface GenericCall { } /** - * Returns whether the declaration binds the given arguments, that is, whether its parameters take - * one value each from the argument types. + * 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 boolean binds(SimpleExtension.Function declaration, List arguments) { + private static boolean binds( + SimpleExtension.Function declaration, List arguments, Type outputType) { + List types = arguments.stream().map(Expression::getType).collect(Collectors.toList()); try { - declaration.resolveType( - arguments.stream().map(Expression::getType).collect(Collectors.toList())); - return true; + TypeExpressionEvaluator.checkBindings(declaration.args(), declaration.variadic(), types); } catch (UnsupportedOperationException e) { return false; } + List 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 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; + } } - /** - * 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)}. - */ - private static Expression coerceToward(Expression operand, Type target) { - if (target instanceof Type.Decimal) { - Optional digits = integerDigits(operand.getType()); - if (digits.isPresent()) { - return ExpressionCreator.cast( - TypeCreator.of(operand.getType().nullable()).decimal(digits.get(), 0), - operand, - Expression.FailureBehavior.THROW_EXCEPTION); + private static List declaredArgumentTypes( + SimpleExtension.Function declaration) { + List 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 coerceArgument(operand, target); + return types; } - private static Optional integerDigits(Type type) { - if (type instanceof Type.I8) { - return Optional.of(3); - } - if (type instanceof Type.I16) { - return Optional.of(5); - } - if (type instanceof Type.I32) { - return Optional.of(10); + private static Optional stringLength(Type type) { + if (type instanceof Type.VarChar) { + return Optional.of(((Type.VarChar) type).length()); } - if (type instanceof Type.I64) { - return Optional.of(19); + 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 @@ -978,24 +1033,22 @@ private static Optional integerDigits(Type type) { * and there is none when that needs more than 38 digits. */ private static Optional losslessCommonType( - List operands, Type leastRestrictive) { + List operands, List calciteTypes, Type leastRestrictive) { if (!(leastRestrictive instanceof Type.Decimal)) { return Optional.of(leastRestrictive); } int integerDigits = 0; int scale = 0; - for (Expression operand : operands) { - Type type = operand.getType(); + 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 { - Optional digits = integerDigits(type); - if (digits.isEmpty()) { - return Optional.empty(); - } - integerDigits = Math.max(integerDigits, digits.get()); + return Optional.empty(); } } if (integerDigits + scale > 38) { @@ -1007,13 +1060,40 @@ private static Optional losslessCommonType( private static Expression castUnlessEqual(Expression operand, Type target) { Type type = operand.getType(); - if (type.withNullable(target.nullable()).equals(target)) { + 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 asDeclaredStrings( + SimpleExtension.Function declaration, List operands) { + List declared = declaredArgumentTypes(declaration); + List 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; diff --git a/isthmus/src/test/java/io/substrait/isthmus/OperandCoercionTest.java b/isthmus/src/test/java/io/substrait/isthmus/OperandCoercionTest.java index 623116e61..ceaec1a9e 100644 --- a/isthmus/src/test/java/io/substrait/isthmus/OperandCoercionTest.java +++ b/isthmus/src/test/java/io/substrait/isthmus/OperandCoercionTest.java @@ -20,8 +20,8 @@ class OperandCoercionTest extends PlanTestBase { private static final String CREATES = - "CREATE TABLE t (d7 DECIMAL(7, 2) NOT NULL, d5 DECIMAL(5, 2) NOT NULL, " - + "i INT NOT NULL, c CHAR(10) NOT NULL)"; + "CREATE TABLE t (d7 DECIMAL(7, 2) NOT NULL, i INT NOT NULL, c CHAR(10) NOT NULL, " + + "v VARCHAR NOT NULL, v10 VARCHAR(10) NOT NULL)"; /** gte(any1, any1) binds any1 once, so decimals of two precisions both take the wider one. */ @Test @@ -37,7 +37,7 @@ void aComparisonOfTwoDecimalsCastsBothToOneType() throws Exception { /** * multiply(decimal, decimal) binds each operand's own precision, so an integer - * operand becomes the decimal that holds it, not the other operand's decimal. + * operand becomes the decimal that holds it, not the least restrictive decimal(12,2). */ @Test void aDecimalTimesAnIntegerCastsTheIntegerToItsOwnDecimal() throws Exception { @@ -81,11 +81,45 @@ void noOperandIsCastToADecimalThatLosesItsScale() throws Exception { assertEquals(2, bound.scale()); } - private static Expression.ScalarFunctionInvocation call(String query) throws Exception { + /** + * concat:vchar declares one operand's length as its result, so two char(10) operands would come + * out as a varchar(10). The next variant, concat:str, takes them as strings. + */ + @Test + void aConcatenationNeverBindsADeclarationThatTruncatesIt() throws Exception { + Expression.ScalarFunctionInvocation call = call("SELECT c || c FROM t"); + + assertEquals("concat:str", call.declaration().key()); + assertEquals(List.of(R.STRING, R.STRING), argumentTypes(call)); + assertBinds(call); + } + + /** + * like:vchar_vchar does not bind an unbounded varchar, which is a string, so the next variant, + * like:str_str, is tried and binds once the pattern is a string too. + */ + @Test + void aVariantThatDoesNotBindGivesWayToTheNextOne() throws Exception { + Expression.ScalarFunctionInvocation call = call("SELECT v LIKE 'a%' FROM t"); + + assertEquals("like:str_str", call.declaration().key()); + assertEquals(List.of(R.STRING, R.STRING), argumentTypes(call)); + assertBinds(call); + } + + /** A declaration's concrete argument type is bound exactly: replace:str_str_str takes strings. */ + @Test + void aConcretelyDeclaredStringTakesAString() throws Exception { + Expression.ScalarFunctionInvocation call = call("SELECT REPLACE(c, v10, v10) FROM t"); + + assertEquals("replace:str_str_str", call.declaration().key()); + assertEquals(List.of(R.STRING, R.STRING, R.STRING), argumentTypes(call)); + } + + private Expression.ScalarFunctionInvocation call(String query) throws Exception { Plan plan = - new SqlToSubstrait() - .convert( - query, SubstraitCreateStatementParser.processCreateStatementsToCatalog(CREATES)); + toSubstraitPlan( + query, SubstraitCreateStatementParser.processCreateStatementsToCatalog(CREATES)); Expression expression = ((Project) plan.getRoots().get(0).getInput()).getExpressions().get(0); return assertInstanceOf(Expression.ScalarFunctionInvocation.class, expression); }