diff --git a/core/src/main/java/io/substrait/extension/FunctionBindingResolver.java b/core/src/main/java/io/substrait/extension/FunctionBindingResolver.java index 79b7a4ee3..7cce7a06e 100644 --- a/core/src/main/java/io/substrait/extension/FunctionBindingResolver.java +++ b/core/src/main/java/io/substrait/extension/FunctionBindingResolver.java @@ -27,14 +27,14 @@ * *

Signature type matching is fail-closed. It checks value- and type-argument patterns alike: * wildcards, concrete types and the scalar-parameterized classes (decimal, char, binary, precision - * time/timestamp, intervals); a declared shape carrying nested types (lists, maps, structs, - * function types) is rejected rather than accepted unchecked. Occurrences of one numbered wildcard - * ({@code any1}) must agree on a single type, while each plain {@code any} matches independently; a - * variadic declaration repeats its trailing argument, requiring the repetitions to agree only when - * its parameters are {@code CONSISTENT} — a literal integer parameter (the {@code 0} of {@code - * DECIMAL}) constrains every repetition regardless. Enum options and option preferences are - * matched case-insensitively; an unspecified enum option is always rejected, since the extension - * schema cannot declare an optional one. + * time/timestamp, intervals), and nested list, map, struct and function types. Nested structure and + * nullability must match; wildcard and integer parameters bind recursively. Occurrences of one + * numbered wildcard ({@code any1}) must agree on a single type, while each plain {@code any} + * matches independently; a variadic declaration repeats its trailing argument, requiring the + * repetitions to agree only when its parameters are {@code CONSISTENT} — a literal integer + * parameter (the {@code 0} of {@code DECIMAL}) constrains every repetition regardless. Enum + * options and option preferences are matched case-insensitively; an unspecified enum option is + * always rejected, since the extension schema cannot declare an optional one. */ public final class FunctionBindingResolver { @@ -473,9 +473,12 @@ private static void requireKind( private static boolean typeMatches( ParameterizedType declared, Type actual, boolean exactNullability) { if (declared instanceof ParameterizedType.StringLiteral) { - // Non-wildcard extension parameter names at the top level are accepted; numbered wildcards - // are handled by the caller for cross-argument consistency. - return true; + // Top-level wildcards are handled by checkWildcard. Nested unmarked wildcards may bind a + // nullable type; an explicit '?' requires a nullable actual. The evaluator checks shared + // variable identities while deriving the return type, even when that return is concrete. + return !exactNullability + || !((ParameterizedType.StringLiteral) declared).nullable() + || actual.nullable(); } if (declared instanceof Type) { // A concrete declared argument type (e.g. i32) matches ignoring nullability, except under a @@ -520,15 +523,56 @@ private static boolean typeMatches( return actual instanceof Type.IntervalCompound && nullabilityMatches(declared, actual, exactNullability); } - // The remaining declared shapes — lists, maps, structs and function types — carry nested types - // this validator cannot yet check structurally, and the spec requires nested structure and - // nullability to match exactly: h(list, list) invoked as h(list, list) - // must not bind (spec v0.99.0, scalar binding rules). A validator that advertises strictness - // must fail closed on a shape it cannot judge rather than silently accept it. + if (declared instanceof ParameterizedType.ListType) { + return actual instanceof Type.ListType + && nullabilityMatches(declared, actual, exactNullability) + && typeMatches( + ((ParameterizedType.ListType) declared).name(), + ((Type.ListType) actual).elementType(), + true); + } + if (declared instanceof ParameterizedType.Map) { + if (!(actual instanceof Type.Map) + || !nullabilityMatches(declared, actual, exactNullability)) { + return false; + } + ParameterizedType.Map pattern = (ParameterizedType.Map) declared; + Type.Map map = (Type.Map) actual; + return typeMatches(pattern.key(), map.key(), true) + && typeMatches(pattern.value(), map.value(), true); + } + if (declared instanceof ParameterizedType.Struct) { + return actual instanceof Type.Struct + && nullabilityMatches(declared, actual, exactNullability) + && typeListMatches( + ((ParameterizedType.Struct) declared).fields(), ((Type.Struct) actual).fields()); + } + if (declared instanceof ParameterizedType.Func) { + if (!(actual instanceof Type.Func) + || !nullabilityMatches(declared, actual, exactNullability)) { + return false; + } + ParameterizedType.Func pattern = (ParameterizedType.Func) declared; + Type.Func function = (Type.Func) actual; + return typeListMatches(pattern.parameterTypes(), function.parameterTypes()) + && typeMatches(pattern.returnType(), function.returnType(), true); + } throw new InvalidFunctionBindingException( String.format("Validation of the declared argument shape %s is not supported", declared)); } + private static boolean typeListMatches(List declared, List actual) { + if (declared.size() != actual.size()) { + return false; + } + for (int index = 0; index < declared.size(); index++) { + if (!typeMatches(declared.get(index), actual.get(index), true)) { + return false; + } + } + return true; + } + /** * Under a DISCRETE declaration the declared nullability is part of the signature, for a * parameterized argument as much as for a concrete one. A declared shape that carries no diff --git a/core/src/main/java/io/substrait/type/TypeExpressionEvaluator.java b/core/src/main/java/io/substrait/type/TypeExpressionEvaluator.java index 45f7e5e98..4346ad939 100644 --- a/core/src/main/java/io/substrait/type/TypeExpressionEvaluator.java +++ b/core/src/main/java/io/substrait/type/TypeExpressionEvaluator.java @@ -1,15 +1,19 @@ package io.substrait.type; import io.substrait.extension.SimpleExtension; +import io.substrait.function.NullableType; import io.substrait.function.ParameterizedType; import io.substrait.function.TypeExpression; import io.substrait.function.TypeExpressionVisitor; import java.util.ArrayList; import java.util.HashMap; +import java.util.HashSet; import java.util.List; import java.util.Map; import java.util.Optional; import java.util.OptionalInt; +import java.util.Set; +import java.util.stream.Collectors; /** * Evaluates a {@link TypeExpression} to a concrete {@link Type} given a set of actual arguments. @@ -31,12 +35,12 @@ * {@code interval_compound} at all, as an argument or as a return -- those two are supported for * symmetry, and pinned against hand-written declarations rather than the catalog. * - *

A {@code list} return still fails whatever its element, because the evaluator does not descend - * into a container -- so an element parameter it would otherwise substitute, as in {@code - * list>}, is out of reach just as an element type to evaluate is. A multi-line return - * program still fails because evaluating one needs integer arithmetic over the bound parameters - * rather than substitution. And a plain {@code any} cannot be derived at all: unlike {@code any1} - * it names nothing, so there is no identity to bind. + *

List, map, struct and function declarations bind their element, field, parameter and return + * types recursively. Container returns also evaluate their children, including integer parameters + * such as {@code List>} and type parameters such as {@code list}. Nested + * nullability is preserved; only the outermost argument nullability is excluded from wildcard + * identity. A multi-line return program still needs arithmetic rather than substitution. A plain + * {@code any} has no identity to bind, so it cannot be derived as a return type. * *

Which shipped variants those cover is pinned by {@code ParameterizedReturnTypeTest} against * the declarations the catalog ships, and deliberately not repeated here -- the catalog is owned @@ -146,7 +150,7 @@ private static ParameterBindings bindParameters( } // An INCONSISTENT variadic repetition binds no named parameters — each repetition is // independent — but a literal constraint (the 0 of DECIMAL) still applies to it. - bindings.bind(declared, actualTypes.get(index), !repeated || bindRepeats); + bindings.bind(declared, actualTypes.get(index), !repeated || bindRepeats, false); } return bindings; } @@ -155,6 +159,7 @@ private static ParameterBindings bindParameters( private static final class ParameterBindings { private final Map types = new HashMap<>(); + private final Set exactTypeNullabilities = new HashSet<>(); private final Map integers = new HashMap<>(); private Type boundType(String name) { @@ -170,13 +175,29 @@ private Integer boundInteger(String token) { * an INCONSISTENT variadic repetition — named parameters are left unbound (each repetition is * independent) while literal constraints are still enforced. */ - private void bind(ParameterizedType declared, Type actual, boolean bindNames) { + private void bind(ParameterizedType declared, Type actual, boolean bindNames, boolean nested) { + if (nested && !(declared instanceof ParameterizedType.StringLiteral)) { + if ((declared instanceof NullableType + && ((NullableType) declared).nullable() != actual.nullable()) + || (declared instanceof Type && !declared.equals(actual))) { + throw cannotBind(declared, actual); + } + } if (declared instanceof ParameterizedType.StringLiteral) { ParameterizedType.StringLiteral literal = (ParameterizedType.StringLiteral) declared; // Only a numbered wildcard names a parameter that a return expression can refer to and that // has to stay consistent across the call; a plain "any" binds independently each time. + if (nested && literal.nullable() && !actual.nullable()) { + throw cannotBind(declared, actual); + } if (bindNames && literal.isNumberedWildcard()) { - bindType(literal.value(), actual); + // An unmarked nested wildcard binds the complete type, including nullability. A '?' + // marker requires a nullable actual, but does not constrain the variable's own + // nullability: both i32 and i32? become i32? after substitution. Outermost argument + // nullability is excluded from binding and also leaves the variable's nullability open. + boolean exactNullability = nested && !literal.nullable(); + Type binding = nested && !literal.nullable() ? actual : actual.withNullable(false); + bindType(literal.value(), binding, exactNullability); } } else if (declared instanceof ParameterizedType.Decimal && actual instanceof Type.Decimal) { ParameterizedType.Decimal declaredDecimal = (ParameterizedType.Decimal) declared; @@ -230,43 +251,71 @@ private void bind(ParameterizedType declared, Type actual, boolean bindNames) { ((ParameterizedType.IntervalCompound) declared).precision().value(), ((Type.IntervalCompound) actual).precision(), bindNames); - } else if (!(declared instanceof Type) && !isContainer(declared)) { - // A shape one of the arms above should have taken: the declaration carries a parameter and - // the actual type is not the class that would bind it. Binding nothing here would enforce - // the shared-parameter rule for some calls and skip it for others. + } else if (declared instanceof ParameterizedType.ListType + && actual instanceof Type.ListType) { + bind( + ((ParameterizedType.ListType) declared).name(), + ((Type.ListType) actual).elementType(), + bindNames, + true); + } else if (declared instanceof ParameterizedType.Map && actual instanceof Type.Map) { + ParameterizedType.Map pattern = (ParameterizedType.Map) declared; + Type.Map map = (Type.Map) actual; + bind(pattern.key(), map.key(), bindNames, true); + bind(pattern.value(), map.value(), bindNames, true); + } else if (declared instanceof ParameterizedType.Struct && actual instanceof Type.Struct) { + bindFields( + ((ParameterizedType.Struct) declared).fields(), + ((Type.Struct) actual).fields(), + bindNames); + } else if (declared instanceof ParameterizedType.Func && actual instanceof Type.Func) { + ParameterizedType.Func pattern = (ParameterizedType.Func) declared; + Type.Func function = (Type.Func) actual; + bindFields(pattern.parameterTypes(), function.parameterTypes(), bindNames); + bind(pattern.returnType(), function.returnType(), bindNames, true); + } else if (!(declared instanceof Type)) { + throw cannotBind(declared, actual); + } + } + + private void bindFields( + List declared, List actual, boolean bindNames) { + if (declared.size() != actual.size()) { throw new UnsupportedOperationException( - String.format( - "Cannot bind parameters from declared argument type %s to actual type %s", - declared, actual)); + "Cannot bind container fields: expected " + + declared.size() + + " types but got " + + actual.size()); + } + for (int index = 0; index < declared.size(); index++) { + bind(declared.get(index), actual.get(index), bindNames, true); } } - /** - * Whether the declared type holds other types rather than an integer parameter. Binding does - * not descend into these, so their parameters bind nothing and a mismatch cannot be told from a - * shape this method simply does not reach yet -- unlike the classes above, refusing here would - * reject declarations that resolve today without binding anything, such as a {@code list} - * argument to a function returning a concrete type. - * - * @param declared the declared argument type - * @return {@code true} if the type is a list, map, struct or function declaration - */ - private boolean isContainer(ParameterizedType declared) { - return declared instanceof ParameterizedType.ListType - || declared instanceof ParameterizedType.Map - || declared instanceof ParameterizedType.Struct - || declared instanceof ParameterizedType.Func; + private static UnsupportedOperationException cannotBind( + ParameterizedType declared, Type actual) { + return new UnsupportedOperationException( + String.format( + "Cannot bind parameters from declared argument type %s to actual type %s", + declared, actual)); } - private void bindType(String name, Type actual) { - // Nullability is not part of a wildcard's identity: any1 binds to i32 and i32? alike, and the - // return expression's own nullability (or the MIRROR policy) decides the result's. + private void bindType(String name, Type actual, boolean exactNullability) { Type existing = types.putIfAbsent(name, actual); - if (existing != null && !existing.equalsIgnoringNullability(actual)) { + boolean existingExact = exactTypeNullabilities.contains(name); + if (existing != null + && (!existing.equalsIgnoringNullability(actual) + || (existingExact && exactNullability && !existing.equals(actual)))) { throw new UnsupportedOperationException( String.format( "Inconsistent binding for type parameter '%s': %s vs %s", name, existing, actual)); } + if (exactNullability) { + exactTypeNullabilities.add(name); + if (!existingExact) { + types.put(name, actual); + } + } } private void bindInteger(String token, int value, boolean bindNames) { @@ -374,18 +423,65 @@ public Type visit(ParameterizedType.IntervalCompound intervalCompound) { .intervalCompound(resolveInteger(intervalCompound.precision().value())); } + @Override + public Type visit(ParameterizedType.ListType list) { + return TypeCreator.of(list.nullable()).list(evaluateNested(list.name())); + } + + @Override + public Type visit(ParameterizedType.Map map) { + return TypeCreator.of(map.nullable()) + .map(evaluateNested(map.key()), evaluateNested(map.value())); + } + + @Override + public Type visit(ParameterizedType.Struct struct) { + return TypeCreator.of(struct.nullable()) + .struct(struct.fields().stream().map(this::evaluateNested).collect(Collectors.toList())); + } + + @Override + public Type visit(ParameterizedType.Func function) { + return TypeCreator.of(function.nullable()) + .func( + function.parameterTypes().stream() + .map(this::evaluateNested) + .collect(Collectors.toList()), + evaluateNested(function.returnType())); + } + + private Type evaluateNested(ParameterizedType expression) { + if (expression instanceof Type) { + return (Type) expression; + } + if (expression instanceof ParameterizedType.StringLiteral) { + ParameterizedType.StringLiteral variable = (ParameterizedType.StringLiteral) expression; + Type bound = boundType(variable); + if (!variable.nullable() && !bindings.exactTypeNullabilities.contains(variable.value())) { + throw new UnsupportedOperationException( + "Cannot derive nullability of type parameter '" + variable.value() + "'"); + } + return bound.withNullable(bound.nullable() || variable.nullable()); + } + return expression.accept(this); + } + @Override public Type visit(ParameterizedType.StringLiteral stringLiteral) { // A wildcard return (e.g. min(any1) -> any1) resolves to the bound argument type, taking the // nullability declared on the return expression in both directions (a required return forces // the type non-null, a nullable one forces it nullable). MIRROR policy, if any, is applied // afterwards by the caller. + return boundType(stringLiteral).withNullable(stringLiteral.nullable()); + } + + private Type boundType(ParameterizedType.StringLiteral stringLiteral) { Type bound = bindings.boundType(stringLiteral.value()); if (bound == null) { throw new UnsupportedOperationException( "Unbound type parameter '" + stringLiteral.value() + "' in return-type expression"); } - return bound.withNullable(stringLiteral.nullable()); + return bound; } private int resolveInteger(String token) { diff --git a/core/src/test/java/io/substrait/extension/FunctionBindingResolverTest.java b/core/src/test/java/io/substrait/extension/FunctionBindingResolverTest.java index af3e60601..e1919e6db 100644 --- a/core/src/test/java/io/substrait/extension/FunctionBindingResolverTest.java +++ b/core/src/test/java/io/substrait/extension/FunctionBindingResolverTest.java @@ -405,10 +405,17 @@ void inconsistentVariadicKeepsLiteralConstraints() { } @Test - void failsClosedOnANestedShapeItCannotCheck() { + void checksNestedShapeAndNullability() { SimpleExtension.ScalarFunctionVariant listPair = testScalar("list_pair:list_list"); - // A declared list against a non-list actual used to be accepted silently; a strict - // validator must reject a shape it cannot check rather than pass it. + assertDoesNotThrow( + () -> + FunctionBindingResolver.resolveAndValidate( + listPair, + List.of( + ResolvedArgument.value(R.list(N.I32)), ResolvedArgument.value(R.list(N.I32))), + List.of(), + R.BOOLEAN)); + // The container shape must match before its element can bind. assertThrows( InvalidFunctionBindingException.class, () -> @@ -417,8 +424,7 @@ void failsClosedOnANestedShapeItCannotCheck() { List.of(ResolvedArgument.value(R.I32), ResolvedArgument.value(R.I32)), List.of(), R.BOOLEAN)); - // list vs list must not bind either: nested nullability is part of the structural - // match, which is exactly the check this validator cannot do yet — so it fails closed here too. + // Inner nullability is part of the shared wildcard binding. assertThrows( InvalidFunctionBindingException.class, () -> diff --git a/core/src/test/java/io/substrait/type/ContainerReturnTypeTest.java b/core/src/test/java/io/substrait/type/ContainerReturnTypeTest.java new file mode 100644 index 000000000..e7099f5f5 --- /dev/null +++ b/core/src/test/java/io/substrait/type/ContainerReturnTypeTest.java @@ -0,0 +1,353 @@ +package io.substrait.type; + +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.extension.DefaultExtensionCatalog; +import io.substrait.extension.FunctionBindingResolver; +import io.substrait.extension.ImmutableSimpleExtension; +import io.substrait.extension.InvalidFunctionBindingException; +import io.substrait.extension.ResolvedArgument; +import io.substrait.extension.SimpleExtension; +import io.substrait.function.ParameterizedType; +import io.substrait.function.ParameterizedTypeCreator; +import io.substrait.function.TypeExpression; +import java.util.Arrays; +import java.util.List; +import java.util.stream.Collectors; +import java.util.stream.Stream; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.Arguments; +import org.junit.jupiter.params.provider.MethodSource; + +class ContainerReturnTypeTest { + private static final TypeCreator R = TypeCreator.REQUIRED; + private static final TypeCreator N = TypeCreator.NULLABLE; + private static final ParameterizedTypeCreator P = ParameterizedTypeCreator.REQUIRED; + private static final ParameterizedTypeCreator Q = ParameterizedTypeCreator.NULLABLE; + private static final ParameterizedType ANY1 = P.parameter("any1"); + + static Stream catalogReturns() { + return Stream.of( + Arguments.of( + "string_split:vchar_vchar", + R.list(R.varChar(20)), + List.of(R.varChar(20), R.varChar(20))), + Arguments.of( + "regexp_string_split:vchar_vchar", + R.list(R.varChar(20)), + List.of(R.varChar(20), R.varChar(20))), + Arguments.of( + "regexp_match_substring_all:vchar_vchar_i64_i64", + R.list(R.varChar(20)), + List.of(R.varChar(20), R.varChar(20), R.I64, R.I64)), + Arguments.of("sort:list", R.list(N.I32), List.of(R.list(N.I32))), + Arguments.of("sort:list", N.list(R.I32), List.of(N.list(R.I32))), + Arguments.of( + "filter:list_func", + R.list(N.I32), + List.of(R.list(N.I32), R.func(List.of(N.I32), N.BOOLEAN))), + Arguments.of( + "transform:list_func", + R.list(N.varChar(30)), + List.of(R.list(N.I32), R.func(List.of(N.I32), N.varChar(30))))); + } + + @ParameterizedTest + @MethodSource("catalogReturns") + void derivesCatalogListReturns(String key, Type expected, List actual) { + SimpleExtension.Function function = + DefaultExtensionCatalog.DEFAULT_COLLECTION.scalarFunctions().stream() + .filter(f -> f.key().equals(key)) + .findFirst() + .orElseThrow(); + assertDerives(function, expected, actual); + } + + @Test + void recursesThroughMapsStructsAndFunctions() { + ParameterizedType declaration = + P.mapE( + P.varCharE("L"), + P.structE(P.listE(ANY1), P.funcE(List.of(ANY1), P.decimalE("P", "S")))); + Type actual = + R.map(R.varChar(12), R.struct(R.list(N.I64), R.func(List.of(N.I64), R.decimal(15, 3)))); + assertDerives(function(declaration, declaration), actual, List.of(actual)); + } + + @Test + void concreteNestedReturnTypesNeedNoParameters() { + assertDerives(function(P.listE(R.I32)), R.list(R.I32), List.of()); + } + + @Test + void sharedNestedWildcardsKeepInnerNullability() { + SimpleExtension.Function pair = function(P.listE(ANY1), P.listE(ANY1), P.listE(ANY1)); + assertDerives(pair, N.list(N.I32), List.of(N.list(N.I32), R.list(N.I32))); + assertInvalid(pair, R.list(R.I32), R.list(N.I32)); + assertInvalid(pair, R.list(R.I32), R.list(R.I64)); + } + + @Test + void catalogIndexInAcceptsNullableElements() { + SimpleExtension.Function function = + DefaultExtensionCatalog.DEFAULT_COLLECTION.scalarFunctions().stream() + .filter(f -> f.key().equals("index_in:any_list")) + .findFirst() + .orElseThrow(); + for (Type value : List.of(R.I32, N.I32)) { + for (Type element : List.of(R.I32, N.I32)) { + assertDerives(function, N.I64, List.of(value, R.list(element))); + } + } + assertInvalid(function, R.FP64, R.list(R.I32)); + assertInvalid(function, R.I32, R.list(N.FP64)); + } + + @Test + void topLevelWildcardsDoNotConstrainNestedNullabilityInEitherOrder() { + ParameterizedType list = P.listE(ANY1); + SimpleExtension.Function forward = function(list, ANY1, list); + SimpleExtension.Function reverse = function(list, list, ANY1); + for (Type value : List.of(R.I32, N.I32)) { + for (Type element : List.of(R.I32, N.I32)) { + Type expected = TypeCreator.of(value.nullable()).list(element); + assertDerives(forward, expected, List.of(value, R.list(element))); + assertDerives(reverse, expected, List.of(R.list(element), value)); + } + } + assertInvalid(forward, R.I32, R.list(N.FP64)); + assertInvalid(reverse, R.list(N.FP64), R.I32); + assertInvalid(function(list, ANY1, list, list), R.I32, R.list(R.I32), R.list(N.I32)); + assertInvalid(function(list, list, ANY1, list), R.list(R.I32), R.I32, R.list(N.I32)); + assertInvalid(function(list, list, list, ANY1), R.list(R.I32), R.list(N.I32), R.I32); + } + + @Test + void nullableWildcardMarkersAreSubstitutedAcrossArgumentShapes() { + ParameterizedType nullableElement = P.listE(Q.parameter("any1")); + // These are the scalar-binding examples for j(any1, list), in both argument orders. + SimpleExtension.Function forward = function(nullableElement, ANY1, nullableElement); + SimpleExtension.Function reverse = function(nullableElement, nullableElement, ANY1); + assertDerives(forward, R.list(N.I32), List.of(R.I32, R.list(N.I32))); + assertDerives(reverse, R.list(N.I32), List.of(R.list(N.I32), R.I32)); + assertInvalid(forward, R.I32, R.list(R.I32)); + assertInvalid(forward, R.I32, R.list(N.I64)); + assertInvalid(reverse, R.list(N.I64), R.I32); + // A nullable marker does not remove nullability already bound by an unmarked nested wildcard. + assertDerives( + function(P.listE(ANY1), nullableElement, P.listE(ANY1)), + R.list(N.I32), + List.of(R.list(N.I32), R.list(N.I32))); + } + + @Test + void nestedIntegerParametersAndLiteralsAreChecked() { + ParameterizedType list = P.listE(P.decimalE("P", "0")); + SimpleExtension.Function pair = function(P.listE(P.decimalE("P", "0")), list, list); + assertDerives( + pair, + R.list(R.decimal(12, 0)), + List.of(R.list(R.decimal(12, 0)), R.list(R.decimal(12, 0)))); + assertInvalid(pair, R.list(R.decimal(12, 0)), R.list(R.decimal(13, 0))); + assertInvalid(pair, R.list(R.decimal(12, 0)), R.list(R.decimal(12, 1))); + } + + @Test + void rejectsWrongContainerShapesAndArity() { + assertInvalid(function(R.I64, P.listE(ANY1)), R.I64); + assertInvalid(function(R.I64, P.mapE(ANY1, ANY1)), R.list(R.I32)); + assertInvalid(function(R.I64, P.structE(ANY1, ANY1)), R.struct(R.I32)); + assertInvalid( + function(R.I64, P.funcE(List.of(ANY1), ANY1)), R.func(List.of(R.I32, R.I32), R.I32)); + assertInvalid(function(R.I64, P.listE(P.listE(ANY1))), R.list(N.list(R.I32))); + assertInvalid(function(R.I64, P.listE(R.I32)), R.list(N.I32)); + } + + @Test + void concreteReturnsStillRejectIncorrectNestedMembers() { + List patterns = + List.of( + P.mapE(R.STRING, ANY1), + P.mapE(ANY1, R.I32), + P.structE(R.I32, ANY1), + P.structE(ANY1, R.I32), + P.funcE(List.of(R.I32), ANY1), + P.funcE(List.of(ANY1), R.BOOLEAN)); + List actual = + List.of( + R.map(R.I64, R.I32), + R.map(R.STRING, N.I32), + R.struct(R.I64, R.I32), + R.struct(R.I32, N.I32), + R.func(List.of(N.I32), R.I64), + R.func(List.of(R.I32), N.BOOLEAN)); + for (int index = 0; index < patterns.size(); index++) { + SimpleExtension.Function function = function(R.I64, patterns.get(index)); + Type argument = actual.get(index); + assertThrows( + UnsupportedOperationException.class, () -> function.resolveType(List.of(argument))); + assertInvalid(function, argument); + assertThrows( + InvalidFunctionBindingException.class, + () -> + FunctionBindingResolver.resolveAndValidate( + function, List.of(ResolvedArgument.value(argument)), List.of(), R.I64)); + } + } + + @Test + void containerOuterNullabilityFollowsTheFunctionPolicy() { + for (SimpleExtension.Nullability policy : SimpleExtension.Nullability.values()) { + SimpleExtension.Function function = + ImmutableSimpleExtension.ScalarFunctionVariant.builder() + .from(function(P.listE(ANY1), P.listE(ANY1))) + .nullability(policy) + .build(); + assertDerives(function, R.list(N.I32), List.of(R.list(N.I32))); + if (policy == SimpleExtension.Nullability.DISCRETE) { + assertThrows( + InvalidFunctionBindingException.class, + () -> + FunctionBindingResolver.resolveAndValidate( + function, + List.of(ResolvedArgument.value(N.list(N.I32))), + List.of(), + R.list(N.I32))); + } else { + Type expected = TypeCreator.of(policy == SimpleExtension.Nullability.MIRROR).list(N.I32); + assertDerives(function, expected, List.of(N.list(N.I32))); + } + } + } + + @Test + void validatesTheDerivedElementTypeAndNullability() { + SimpleExtension.Function function = function(P.listE(P.varCharE("L")), P.varCharE("L")); + List arguments = List.of(ResolvedArgument.value(R.varChar(20))); + assertDerives(function, R.list(R.varChar(20)), List.of(R.varChar(20))); + for (Type wrong : + List.of(R.list(R.varChar(19)), R.list(N.varChar(20)), N.list(R.varChar(20)))) { + assertThrows( + InvalidFunctionBindingException.class, + () -> FunctionBindingResolver.resolveAndValidate(function, arguments, List.of(), wrong)); + } + } + + @Test + void catalogFunctionArgumentsRequireFunctionTypes() { + for (String key : List.of("all_match:list_func", "any_match:list_func")) { + SimpleExtension.Function function = + DefaultExtensionCatalog.DEFAULT_COLLECTION.scalarFunctions().stream() + .filter(f -> f.key().equals(key)) + .findFirst() + .orElseThrow(); + assertDerives(function, N.BOOLEAN, List.of(R.list(R.I64), R.func(List.of(R.I64), N.BOOLEAN))); + assertThrows( + UnsupportedOperationException.class, + () -> function.resolveType(List.of(R.list(R.I64), N.BOOLEAN))); + assertInvalid(function, R.list(R.I64), N.BOOLEAN); + } + } + + @Test + void aPlainAnyStillHasNoReturnBinding() { + SimpleExtension.Function function = function(P.listE(P.parameter("any")), P.parameter("any")); + assertThrows(UnsupportedOperationException.class, () -> function.resolveType(List.of(R.I32))); + assertInvalid(function, R.I32); + } + + @Test + void catalogQuantileStillHasAnUnboundElementType() { + SimpleExtension.Function quantile = + DefaultExtensionCatalog.DEFAULT_COLLECTION.aggregateFunctions().stream() + .filter(f -> f.key().equals("quantile:req_req_i64_any")) + .findFirst() + .orElseThrow(); + UnsupportedOperationException error = + assertThrows( + UnsupportedOperationException.class, () -> quantile.resolveType(List.of(R.I64, R.I32))); + assertTrue(error.getMessage().contains("Unbound type parameter 'any'"), error.getMessage()); + } + + @Test + void aNullableMarkerAloneCannotDetermineTheVariablesOwnNullability() { + ParameterizedType nullableElement = P.listE(Q.parameter("any1")); + assertDerives( + function(nullableElement, nullableElement), R.list(N.I32), List.of(R.list(N.I32))); + // Both any1=i32 and any1=i32? satisfy list. Without another occurrence, the + // nullability of an unmarked return element is not determined by the argument. + assertInvalid(function(P.listE(ANY1), nullableElement), R.list(N.I32)); + } + + @Test + void containerTypeArgumentsAlsoBindParameters() { + ParameterizedType list = P.listE(P.varCharE("L")); + SimpleExtension.Function function = + ImmutableSimpleExtension.ScalarFunctionVariant.builder() + .from(function(list)) + .args(List.of(SimpleExtension.TypeArgument.builder().type(list).build())) + .build(); + assertEquals( + R.list(R.varChar(17)), + FunctionBindingResolver.deriveOutputType( + function, List.of(ResolvedArgument.type(R.list(R.varChar(17)))))); + } + + @Test + void variadicContainersRespectParameterConsistencyAndLiteralConstraints() { + ParameterizedType list = P.listE(P.decimalE("P", "0")); + for (SimpleExtension.VariadicBehavior.ParameterConsistency consistency : + SimpleExtension.VariadicBehavior.ParameterConsistency.values()) { + SimpleExtension.Function function = + ImmutableSimpleExtension.ScalarFunctionVariant.builder() + .from(function(list, list)) + .variadic( + ImmutableSimpleExtension.VariadicBehavior.builder() + .min(1) + .parameterConsistency(consistency) + .build()) + .build(); + List actual = List.of(R.list(R.decimal(12, 0)), R.list(R.decimal(15, 0))); + if (consistency == SimpleExtension.VariadicBehavior.ParameterConsistency.CONSISTENT) { + assertInvalid(function, actual.toArray(new Type[0])); + } else { + assertDerives(function, actual.get(0), actual); + } + assertInvalid(function, actual.get(0), R.list(R.decimal(15, 1))); + } + } + + private static SimpleExtension.ScalarFunctionVariant function( + TypeExpression result, ParameterizedType... parameters) { + return ImmutableSimpleExtension.ScalarFunctionVariant.builder() + .urn("extension:io.substrait:container_test") + .name("container") + .returnType(result) + .args( + Arrays.stream(parameters) + .map(p -> SimpleExtension.ValueArgument.builder().value(p).build()) + .collect(Collectors.toList())) + .build(); + } + + private static void assertDerives( + SimpleExtension.Function function, Type expected, List actual) { + assertEquals(expected, function.resolveType(actual)); + List arguments = + actual.stream().map(ResolvedArgument::value).collect(Collectors.toList()); + assertEquals(expected, FunctionBindingResolver.deriveOutputType(function, arguments)); + FunctionBindingResolver.resolveAndValidate(function, arguments, List.of(), expected); + } + + private static void assertInvalid(SimpleExtension.Function function, Type... actual) { + assertThrows( + InvalidFunctionBindingException.class, + () -> + FunctionBindingResolver.deriveOutputType( + function, + Arrays.stream(actual).map(ResolvedArgument::value).collect(Collectors.toList()))); + } +} diff --git a/core/src/test/java/io/substrait/type/ParameterizedReturnTypeTest.java b/core/src/test/java/io/substrait/type/ParameterizedReturnTypeTest.java index ba4566df7..73848834b 100644 --- a/core/src/test/java/io/substrait/type/ParameterizedReturnTypeTest.java +++ b/core/src/test/java/io/substrait/type/ParameterizedReturnTypeTest.java @@ -86,14 +86,15 @@ void aParameterizedDeclarationRejectsAnotherActualShape() { } @Test - void aContainerDeclarationIsNotRefusedForABindingItNeverMakes() { - // Binding descends into none of the container declarations, so a `list` or a - // `func boolean?>` argument binds nothing. All four of these declare a concrete return - // and need no binding at all, so refusing the shape would reject calls that resolve today. + void concreteReturnsStillBindContainerArguments() { assertEquals(R.I64, resolve("cardinality:list", R.list(R.I64))); assertEquals(N.I64, resolve("index_in:any_list", R.I64, R.list(R.I64))); - assertEquals(N.BOOLEAN, resolve("all_match:list_func", R.list(R.I64), N.BOOLEAN)); - assertEquals(N.BOOLEAN, resolve("any_match:list_func", R.list(R.I64), N.BOOLEAN)); + assertEquals( + N.BOOLEAN, + resolve("all_match:list_func", R.list(R.I64), R.func(List.of(R.I64), N.BOOLEAN))); + assertEquals( + N.BOOLEAN, + resolve("any_match:list_func", R.list(R.I64), R.func(List.of(R.I64), N.BOOLEAN))); } @Test @@ -155,14 +156,9 @@ void mirrorNullabilityStillApplies() { assertEquals(N.intervalDay(6), resolve("multiply:i8_iday", R.I8, N.intervalDay(6))); } - /** - * The census of what the evaluator does not derive: a {@code list} return is the first shape, a - * multi-line return program the second. {@link TypeExpressionEvaluator}'s Javadoc describes those - * shapes and points here rather than naming variants, so this test is the only place a {@code - * substrait-packaging} bump can make the two disagree. - */ + /** Pins the catalog's list-return and return-program shapes across packaging updates. */ @Test - void theReturnShapesThatAreNotDerivedYet() { + void catalogReturnShapes() { assertEquals( List.of( "filter:list_func", @@ -193,11 +189,9 @@ void theReturnShapesThatAreNotDerivedYet() { "subtract:dec_dec"), variantsReturning(TypeExpression.ReturnProgram.class)); - // The lists above pin which variants carry each shape; these pin that the shapes actually fail, - // so making one derivable cannot leave the census passing and the Javadoc stale. - assertThrows( - UnsupportedOperationException.class, - () -> resolve("string_split:vchar_vchar", R.varChar(20), R.varChar(20))); + // List returns now derive recursively; return programs remain a separate expression shape. + assertEquals( + R.list(R.varChar(20)), resolve("string_split:vchar_vchar", R.varChar(20), R.varChar(20))); assertThrows( UnsupportedOperationException.class, () -> resolve("add:dec_dec", R.decimal(10, 2), R.decimal(10, 2))); diff --git a/core/src/test/resources/extensions/binding_extensions.yaml b/core/src/test/resources/extensions/binding_extensions.yaml index 321830249..85a3a1b1f 100644 --- a/core/src/test/resources/extensions/binding_extensions.yaml +++ b/core/src/test/resources/extensions/binding_extensions.yaml @@ -109,8 +109,8 @@ scalar_functions: return: boolean - name: "list_pair" description: >- - A numbered wildcard nested inside a list. Nested shapes cannot be checked structurally - yet, so strict validation fails closed on them. + A numbered wildcard nested inside a list. Both elements must bind to the same type, + including their nullability. impls: - args: - name: x