From bab8b94fae97754036b21e97d6d10769dbcce2bb Mon Sep 17 00:00:00 2001 From: Aleksandr Efimov Date: Mon, 7 Sep 2026 17:30:42 +0300 Subject: [PATCH 1/6] feat(core): derive container return types from nested argument bindings Parameterized list returns currently fail even when their element type is available from the arguments. For filter, sort and transform, type variables inside list and function arguments are also never bound. Bind type and integer parameters recursively through list, map, struct and function arguments, and derive container returns from those bindings. Preserve nested nullability and enforce shared parameter, literal and variadic constraints from [spec v0.102.0](https://github.com/substrait-io/substrait/blob/v0.102.0/site/docs/expressions/scalar_functions.md#nullability-and-any-type-binding). If the available bindings leave a nested element's nullability undetermined, derivation fails. This enables the six Java-side variants in #1241: filter, sort, transform, string_split, regexp_string_split and regexp_match_substring_all. quantile remains unsupported because the pinned catalog uses an anonymous any in its return type. It needs [the spec fix](https://github.com/substrait-io/substrait/pull/1193) and a packaging update. Closes #1241 --- .../extension/FunctionBindingResolver.java | 76 ++++- .../type/TypeExpressionEvaluator.java | 168 ++++++++-- .../FunctionBindingResolverTest.java | 16 +- .../type/ContainerReturnTypeTest.java | 301 ++++++++++++++++++ .../type/ParameterizedReturnTypeTest.java | 22 +- .../extensions/binding_extensions.yaml | 4 +- 6 files changed, 518 insertions(+), 69 deletions(-) create mode 100644 core/src/test/java/io/substrait/type/ContainerReturnTypeTest.java 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 76624879d..c794ed62b 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. @@ -41,12 +45,13 @@ * 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 program referring - * to an argument's value rather than a parameter of its type still fails: this API receives only - * argument types. 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 program referring to an argument's value rather than a parameter of its type still + * fails: this API receives only argument types. 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 @@ -162,7 +167,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; } @@ -171,6 +176,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) { @@ -186,13 +192,28 @@ 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. + 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; @@ -246,43 +267,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) { @@ -388,6 +437,55 @@ public Type visit(ParameterizedType.IntervalCompound intervalCompound) { return intervalCompound(intervalCompound.nullable(), intervalCompound.precision()); } + @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())); + } + + /** + * Evaluates a type nested in a container. Unlike a top-level name, a nested name keeps the + * nullability it was bound with, and a {@code ?} marker only widens it. + */ + private Type evaluateNested(ParameterizedType expression) { + if (expression instanceof ParameterizedType.StringLiteral) { + ParameterizedType.StringLiteral variable = (ParameterizedType.StringLiteral) expression; + Object local = locals.get(variable.value()); + Type bound = local instanceof Type ? (Type) local : bindings.boundType(variable.value()); + if (bound != null) { + if (local == null + && !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 evaluate(expression, Type.class); + } + @Override public Object visit(ParameterizedType.StringLiteral stringLiteral) { Object local = locals.get(stringLiteral.value()); diff --git a/core/src/test/java/io/substrait/extension/FunctionBindingResolverTest.java b/core/src/test/java/io/substrait/extension/FunctionBindingResolverTest.java index b3cc824c1..6661bb9d0 100644 --- a/core/src/test/java/io/substrait/extension/FunctionBindingResolverTest.java +++ b/core/src/test/java/io/substrait/extension/FunctionBindingResolverTest.java @@ -429,10 +429,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, () -> @@ -441,8 +448,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..f17a5fbd3 --- /dev/null +++ b/core/src/test/java/io/substrait/type/ContainerReturnTypeTest.java @@ -0,0 +1,301 @@ +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 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))); + 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 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 8291d45a5..2bcc4312a 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 @@ -156,11 +157,11 @@ void mirrorNullabilityStillApplies() { } /** - * The census of list returns the evaluator does not derive. The catalog is owned upstream, so + * The census of list returns, which now derive recursively. The catalog is owned upstream, so * this catches declarations added by a {@code substrait-packaging} bump. */ @Test - void theReturnShapesThatAreNotDerivedYet() { + void catalogReturnShapes() { assertEquals( List.of( "filter:list_func", @@ -172,9 +173,8 @@ void theReturnShapesThatAreNotDerivedYet() { "transform:list_func"), variantsReturning(ParameterizedType.ListType.class)); - assertThrows( - UnsupportedOperationException.class, - () -> resolve("string_split:vchar_vchar", R.varChar(20), R.varChar(20))); + assertEquals( + R.list(R.varChar(20)), resolve("string_split:vchar_vchar", R.varChar(20), R.varChar(20))); } @Test 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 From 99a03e4728efd5401742a58cfcd9dad839ee37e9 Mon Sep 17 00:00:00 2001 From: Aleksandr Efimov Date: Mon, 7 Sep 2026 22:05:15 +0300 Subject: [PATCH 2/6] fix(core): leave top-level wildcard nullability unconstrained Let nested wildcard occurrences determine the variable's nullability while keeping nested-to-nested consistency checks. Cover index_in with nullable list elements and argument-order independence. Add a successful derived-output validation case and check that catalog function arguments require Type.Func. --- .../type/TypeExpressionEvaluator.java | 5 +- .../type/ContainerReturnTypeTest.java | 52 +++++++++++++++++++ 2 files changed, 55 insertions(+), 2 deletions(-) diff --git a/core/src/main/java/io/substrait/type/TypeExpressionEvaluator.java b/core/src/main/java/io/substrait/type/TypeExpressionEvaluator.java index c794ed62b..1220c647a 100644 --- a/core/src/main/java/io/substrait/type/TypeExpressionEvaluator.java +++ b/core/src/main/java/io/substrait/type/TypeExpressionEvaluator.java @@ -210,8 +210,9 @@ private void bind(ParameterizedType declared, Type actual, boolean bindNames, bo if (bindNames && literal.isNumberedWildcard()) { // 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. - boolean exactNullability = !nested || !literal.nullable(); + // 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); } diff --git a/core/src/test/java/io/substrait/type/ContainerReturnTypeTest.java b/core/src/test/java/io/substrait/type/ContainerReturnTypeTest.java index f17a5fbd3..e7099f5f5 100644 --- a/core/src/test/java/io/substrait/type/ContainerReturnTypeTest.java +++ b/core/src/test/java/io/substrait/type/ContainerReturnTypeTest.java @@ -90,6 +90,41 @@ void sharedNestedWildcardsKeepInnerNullability() { 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")); @@ -192,6 +227,7 @@ void containerOuterNullabilityFollowsTheFunctionPolicy() { 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( @@ -200,6 +236,22 @@ void validatesTheDerivedElementTypeAndNullability() { } } + @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")); From b2ac4222e02c1294b022c5f8c5526e3338d4607f Mon Sep 17 00:00:00 2001 From: Aleksandr Efimov Date: Tue, 22 Sep 2026 10:38:01 +0300 Subject: [PATCH 3/6] fix(core): derive containers whose children are return-program expressions A container whose element carries arithmetic, as in list>, parses as a TypeExpression container rather than a ParameterizedType one, and a nested name can refer to a local of the return program. --- .../type/TypeExpressionEvaluator.java | 56 +++++++++++++++---- .../substrait/type/ReturnProgramTypeTest.java | 10 ++++ 2 files changed, 55 insertions(+), 11 deletions(-) diff --git a/core/src/main/java/io/substrait/type/TypeExpressionEvaluator.java b/core/src/main/java/io/substrait/type/TypeExpressionEvaluator.java index 1220c647a..c380c5d2a 100644 --- a/core/src/main/java/io/substrait/type/TypeExpressionEvaluator.java +++ b/core/src/main/java/io/substrait/type/TypeExpressionEvaluator.java @@ -440,36 +440,70 @@ public Type visit(ParameterizedType.IntervalCompound intervalCompound) { @Override public Type visit(ParameterizedType.ListType list) { - return TypeCreator.of(list.nullable()).list(evaluateNested(list.name())); + return list(list.nullable(), list.name()); } @Override public Type visit(ParameterizedType.Map map) { - return TypeCreator.of(map.nullable()) - .map(evaluateNested(map.key()), evaluateNested(map.value())); + return map(map.nullable(), map.key(), map.value()); } @Override public Type visit(ParameterizedType.Struct struct) { - return TypeCreator.of(struct.nullable()) - .struct(struct.fields().stream().map(this::evaluateNested).collect(Collectors.toList())); + return struct(struct.nullable(), struct.fields()); } @Override public Type visit(ParameterizedType.Func function) { - return TypeCreator.of(function.nullable()) + return func(function.nullable(), function.parameterTypes(), function.returnType()); + } + + @Override + public Type visit(TypeExpression.ListType list) { + return list(list.nullable(), list.elementType()); + } + + @Override + public Type visit(TypeExpression.Map map) { + return map(map.nullable(), map.key(), map.value()); + } + + @Override + public Type visit(TypeExpression.Struct struct) { + return struct(struct.nullable(), struct.fields()); + } + + @Override + public Type visit(TypeExpression.Func function) { + return func(function.nullable(), function.parameterTypes(), function.returnType()); + } + + private Type list(boolean nullable, TypeExpression element) { + return TypeCreator.of(nullable).list(evaluateNested(element)); + } + + private Type map(boolean nullable, TypeExpression key, TypeExpression value) { + return TypeCreator.of(nullable).map(evaluateNested(key), evaluateNested(value)); + } + + private Type struct(boolean nullable, List fields) { + return TypeCreator.of(nullable) + .struct(fields.stream().map(this::evaluateNested).collect(Collectors.toList())); + } + + private Type func( + boolean nullable, List parameters, TypeExpression returnType) { + return TypeCreator.of(nullable) .func( - function.parameterTypes().stream() - .map(this::evaluateNested) - .collect(Collectors.toList()), - evaluateNested(function.returnType())); + parameters.stream().map(this::evaluateNested).collect(Collectors.toList()), + evaluateNested(returnType)); } /** * Evaluates a type nested in a container. Unlike a top-level name, a nested name keeps the * nullability it was bound with, and a {@code ?} marker only widens it. */ - private Type evaluateNested(ParameterizedType expression) { + private Type evaluateNested(TypeExpression expression) { if (expression instanceof ParameterizedType.StringLiteral) { ParameterizedType.StringLiteral variable = (ParameterizedType.StringLiteral) expression; Object local = locals.get(variable.value()); diff --git a/core/src/test/java/io/substrait/type/ReturnProgramTypeTest.java b/core/src/test/java/io/substrait/type/ReturnProgramTypeTest.java index 088b17125..c12b4a4d7 100644 --- a/core/src/test/java/io/substrait/type/ReturnProgramTypeTest.java +++ b/core/src/test/java/io/substrait/type/ReturnProgramTypeTest.java @@ -103,6 +103,16 @@ void assignmentsUseEarlierResultsAndRemainLocalToOneEvaluation() { UnsupportedOperationException.class, () -> evaluate("a = 1\nb = a = 99\ni64\nvarchar")); } + @Test + void containersEvaluateTheirChildrenWithTheProgramsLocals() { + assertEquals(R.list(R.varChar(11)), evaluate("list>")); + assertEquals(R.list(R.varChar(20)), evaluate("a = L * 2\nlist>")); + assertEquals(R.list(R.varChar(10)), evaluate("t = varchar\nlist")); + assertEquals( + R.map(R.varChar(10), N.varChar(12)), evaluate("t = varchar?\nmap, t>")); + assertThrows(UnsupportedOperationException.class, () -> evaluate("a = L\nlist")); + } + @Test void conditionsSelectOnlyTheChosenBranch() { assertEquals(R.varChar(10), evaluate("varchar 5 ? L : missing>")); From 1ef1c226f78bc381281cb6476b97ace67b5c1180 Mon Sep 17 00:00:00 2001 From: Aleksandr Efimov Date: Wed, 30 Sep 2026 10:37:37 +0300 Subject: [PATCH 4/6] fix(core): bind an unmarked top-level wildcard exactly The spec strips only an outermost argument's own nullability before binding, so index_in(i32, list) binds any1 to both i32 and i32? and does not bind. The relaxation that let it bind is reverted, which also lets a return such as f(any1) -> list derive the element's nullability from a wildcard bound only at the top level. AggregateConversion's Javadoc now gives the reason decimal avg is still rejected: its intermediate derives, but a phase that consumes the state binds it from that state. --- .../type/TypeExpressionEvaluator.java | 10 +++--- .../type/ContainerReturnTypeTest.java | 31 +++++++++++++------ .../isthmus/AggregateConversion.java | 9 +++--- 3 files changed, 32 insertions(+), 18 deletions(-) diff --git a/core/src/main/java/io/substrait/type/TypeExpressionEvaluator.java b/core/src/main/java/io/substrait/type/TypeExpressionEvaluator.java index c380c5d2a..851f72347 100644 --- a/core/src/main/java/io/substrait/type/TypeExpressionEvaluator.java +++ b/core/src/main/java/io/substrait/type/TypeExpressionEvaluator.java @@ -208,11 +208,13 @@ private void bind(ParameterizedType declared, Type actual, boolean bindNames, bo throw cannotBind(declared, actual); } if (bindNames && literal.isNumberedWildcard()) { - // An unmarked nested wildcard binds the complete type, including nullability. A '?' + // An unmarked wildcard binds exactly. Nested, it binds the complete type, including + // nullability. In an outermost argument that argument's own nullability is stripped + // first, as the spec does under MIRROR and DECLARED_OUTPUT (spec v0.102.0), so i32? + // there binds i32; signature validation checks it against DISCRETE separately. 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(); + // nullability: both i32 and i32? become i32? after substitution. + boolean exactNullability = !literal.nullable(); Type binding = nested && !literal.nullable() ? actual : actual.withNullable(false); bindType(literal.value(), binding, exactNullability); } diff --git a/core/src/test/java/io/substrait/type/ContainerReturnTypeTest.java b/core/src/test/java/io/substrait/type/ContainerReturnTypeTest.java index e7099f5f5..97bd4d094 100644 --- a/core/src/test/java/io/substrait/type/ContainerReturnTypeTest.java +++ b/core/src/test/java/io/substrait/type/ContainerReturnTypeTest.java @@ -91,32 +91,34 @@ void sharedNestedWildcardsKeepInnerNullability() { } @Test - void catalogIndexInAcceptsNullableElements() { + void catalogIndexInBindsItsElementExactly() { + // index_in(any1, list): the value's own nullability is stripped before binding, as the + // spec does for an outermost argument, while the element keeps its, so a nullable element + // binds any1 to a different type than the value does. 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))); - } + assertDerives(function, N.I64, List.of(value, R.list(R.I32))); + assertInvalid(function, value, R.list(N.I32)); } assertInvalid(function, R.FP64, R.list(R.I32)); assertInvalid(function, R.I32, R.list(N.FP64)); } @Test - void topLevelWildcardsDoNotConstrainNestedNullabilityInEitherOrder() { + void topLevelWildcardsBindExactlyInEitherOrder() { 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)); - } + Type expected = TypeCreator.of(value.nullable()).list(R.I32); + assertDerives(forward, expected, List.of(value, R.list(R.I32))); + assertDerives(reverse, expected, List.of(R.list(R.I32), value)); + assertInvalid(forward, value, R.list(N.I32)); + assertInvalid(reverse, R.list(N.I32), value); } assertInvalid(forward, R.I32, R.list(N.FP64)); assertInvalid(reverse, R.list(N.FP64), R.I32); @@ -125,6 +127,15 @@ void topLevelWildcardsDoNotConstrainNestedNullabilityInEitherOrder() { assertInvalid(function(list, list, list, ANY1), R.list(R.I32), R.list(N.I32), R.I32); } + @Test + void aWildcardBoundOnlyAtTheTopLevelDerivesANestedReturn() { + // f(any1) -> list: the element takes the argument's type without its own nullability, + // which the function's nullability handling applies to the list. + SimpleExtension.Function wrap = function(P.listE(ANY1), ANY1); + assertDerives(wrap, R.list(R.I32), List.of(R.I32)); + assertDerives(wrap, N.list(R.I32), List.of(N.I32)); + } + @Test void nullableWildcardMarkersAreSubstitutedAcrossArgumentShapes() { ParameterizedType nullableElement = P.listE(Q.parameter("any1")); diff --git a/isthmus/src/main/java/io/substrait/isthmus/AggregateConversion.java b/isthmus/src/main/java/io/substrait/isthmus/AggregateConversion.java index 6eca842c6..92b631573 100644 --- a/isthmus/src/main/java/io/substrait/isthmus/AggregateConversion.java +++ b/isthmus/src/main/java/io/substrait/isthmus/AggregateConversion.java @@ -33,10 +33,11 @@ public enum FunctionBindingValidation { *

Type derivation is fail-closed: a function whose return expression the derivation does not * yet support is rejected rather than assumed valid, so this mode is not adoptable for plans * that use such functions. Among the standard aggregates that means {@code quantile}, whose - * declared return {@code LIST?} uses a plain {@code any} carrying no identity to bind, and - * {@code avg} over a decimal, whose intermediate {@code STRUCT,i64>} is a - * parameterized struct the derivation has no case for. Because an unspecified aggregate phase - * consumes that intermediate state, an ordinary decimal {@code avg} is rejected too. + * declared return {@code LIST?} uses a plain {@code any} carrying no identity to bind. + * {@code avg} over a decimal is rejected too, in every phase that consumes the intermediate + * state, an unspecified phase among them: its intermediate {@code STRUCT,i64>} + * derives from the initial arguments, but such a phase binds it from the state it receives, and + * a struct cannot bind the declared {@code DECIMAL} argument. */ EXTENSION_DECLARATION } From bbbf03a3c36abc47175fea045ed1eb310cf54b05 Mon Sep 17 00:00:00 2001 From: Aleksandr Efimov Date: Wed, 30 Sep 2026 14:02:58 +0300 Subject: [PATCH 5/6] fix(core): validate shared bindings and keep an element's nullability A bare wildcard in a return expression replaced the nullability it was bound with by the return's own, so under DECLARED_OUTPUT f(list) -> any1 over list derived i32. A nullable return now widens it instead. Signature validation checked each argument's shape on its own, so matchesDeclaration accepted list and list for two list arguments and list for list>. It now binds the parameters through the same binder the derivation uses. Both reject the unbound type explicitly instead of binding it or reading its nullability. --- .../extension/FunctionBindingResolver.java | 36 ++++++++-- .../type/TypeExpressionEvaluator.java | 43 ++++++++--- .../type/ContainerReturnTypeTest.java | 72 +++++++++++++++++++ .../type/ParameterizedReturnTypeTest.java | 10 +++ 4 files changed, 146 insertions(+), 15 deletions(-) diff --git a/core/src/main/java/io/substrait/extension/FunctionBindingResolver.java b/core/src/main/java/io/substrait/extension/FunctionBindingResolver.java index 7cce7a06e..7a10b2ec4 100644 --- a/core/src/main/java/io/substrait/extension/FunctionBindingResolver.java +++ b/core/src/main/java/io/substrait/extension/FunctionBindingResolver.java @@ -28,8 +28,9 @@ *

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), 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} + * nullability must match; wildcard and integer parameters bind recursively, through the same + * binding the derivation uses (see {@link TypeExpressionEvaluator#checkBindings}). 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 @@ -334,6 +335,16 @@ private static void validateSignature( !repeated || bindRepeats, wildcardBindings); } + // The checks above judge each argument's shape on its own. Binding the parameters the way + // derivation does also checks what the arguments share: a nested wildcard's identity, integer + // parameters and their literal constraints. + try { + TypeExpressionEvaluator.checkBindings( + declared, declaration.variadic(), valueAndTypeArgumentTypes(arguments)); + } catch (UnsupportedOperationException e) { + throw new InvalidFunctionBindingException( + String.format("%s: %s", declaration.getAnchor(), e.getMessage()), e); + } } private static void checkArgument( @@ -381,6 +392,14 @@ private static void checkArgumentType( boolean bindWildcards, Map wildcardBindings) { Type actualType = actual.type().orElseThrow(IllegalStateException::new); + if (actualType instanceof Type.Unbound) { + // The unbound type does not unify with any declared shape, and it carries no nullability for + // the checks below to read. + throw new InvalidFunctionBindingException( + String.format( + "%s argument %d: the unbound type does not match declared %s", + declaration.getAnchor(), index, declaredType)); + } if (declaredType instanceof ParameterizedType.StringLiteral && ((ParameterizedType.StringLiteral) declaredType).isWildcard()) { checkWildcardArgument( @@ -472,10 +491,17 @@ private static void requireKind( private static boolean typeMatches( ParameterizedType declared, Type actual, boolean exactNullability) { + if (actual instanceof Type.Unbound) { + // The unbound type does not unify with any declared shape, and it carries no nullability for + // the checks below to read. + throw new InvalidFunctionBindingException( + String.format( + "Cannot validate declared argument shape %s against the unbound type", declared)); + } if (declared instanceof ParameterizedType.StringLiteral) { - // 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. + // Top-level wildcards are handled by checkWildcardArgument. Nested unmarked wildcards may + // bind a nullable type; an explicit '?' requires a nullable actual. Shared identities are + // checked when validateSignature binds the parameters. return !exactNullability || !((ParameterizedType.StringLiteral) declared).nullable() || actual.nullable(); diff --git a/core/src/main/java/io/substrait/type/TypeExpressionEvaluator.java b/core/src/main/java/io/substrait/type/TypeExpressionEvaluator.java index 851f72347..c462fbd38 100644 --- a/core/src/main/java/io/substrait/type/TypeExpressionEvaluator.java +++ b/core/src/main/java/io/substrait/type/TypeExpressionEvaluator.java @@ -109,6 +109,24 @@ public static Type evaluateExpression( } } + /** + * Binds the declaration's type parameters from the actual argument types, as {@link + * #evaluateExpression} does before it evaluates a return expression, and reports the first + * argument that does not fit: a shape that does not match, a parameter bound to two different + * values, or a literal parameter the actual type does not carry. + * + * @param declaredArguments the declared arguments of the function + * @param variadic the declaration's variadic behavior, if it is variadic + * @param actualTypes the actual argument types supplied at the call site + * @throws UnsupportedOperationException if the actual types do not bind the declaration + */ + public static void checkBindings( + List declaredArguments, + Optional variadic, + List actualTypes) { + bindParameters(declaredArguments, variadic, actualTypes); + } + /** * Binds the declaration's type parameters — both numbered wildcards (e.g. the {@code any1} of * {@code min(any1) -> any1}) and integer parameters (e.g. the {@code P} and {@code S} of {@code @@ -193,6 +211,11 @@ private Integer boundInteger(String token) { * independent) while literal constraints are still enforced. */ private void bind(ParameterizedType declared, Type actual, boolean bindNames, boolean nested) { + if (actual instanceof Type.Unbound) { + // The unbound type does not unify with any declared shape, and it has no nullability for + // the checks below to read. + throw cannotBind(declared, actual); + } if (nested && !(declared instanceof ParameterizedType.StringLiteral)) { if ((declared instanceof NullableType && ((NullableType) declared).nullable() != actual.nullable()) @@ -210,10 +233,10 @@ private void bind(ParameterizedType declared, Type actual, boolean bindNames, bo if (bindNames && literal.isNumberedWildcard()) { // An unmarked wildcard binds exactly. Nested, it binds the complete type, including // nullability. In an outermost argument that argument's own nullability is stripped - // first, as the spec does under MIRROR and DECLARED_OUTPUT (spec v0.102.0), so i32? - // there binds i32; signature validation checks it against DISCRETE separately. A '?' - // marker requires a nullable actual, but does not constrain the variable's own - // nullability: both i32 and i32? become i32? after substitution. + // first, as the spec does under MIRROR and DECLARED_OUTPUT, so i32? there binds i32; + // signature validation checks it against DISCRETE separately. A '?' marker requires a + // nullable actual, but does not constrain the variable's own nullability: both i32 and + // i32? become i32? after substitution. boolean exactNullability = !literal.nullable(); Type binding = nested && !literal.nullable() ? actual : actual.withNullable(false); bindType(literal.value(), binding, exactNullability); @@ -511,7 +534,7 @@ private Type evaluateNested(TypeExpression expression) { Object local = locals.get(variable.value()); Type bound = local instanceof Type ? (Type) local : bindings.boundType(variable.value()); if (bound != null) { - if (local == null + if (!(local instanceof Type) && !variable.nullable() && !bindings.exactTypeNullabilities.contains(variable.value())) { throw new UnsupportedOperationException( @@ -535,15 +558,15 @@ public Object visit(ParameterizedType.StringLiteral stringLiteral) { if (integer != null) { return integer.longValue(); } - // 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. This is checked before parsing the token as an integer literal + // A wildcard return (e.g. min(any1) -> any1) resolves to the bound argument type, and a + // nullable return expression makes it nullable. A type bound from a nested element keeps its + // own nullability. MIRROR policy, if any, is applied afterwards by the caller. This is + // checked before parsing the token as an integer literal // because parseIntegerLiteral reports "not a number" by throwing, and a bound type name is // never a numeral. Type bound = bindings.boundType(stringLiteral.value()); if (bound != null) { - return bound.withNullable(stringLiteral.nullable()); + return stringLiteral.nullable() ? bound.withNullable(true) : bound; } OptionalInt literal = parseIntegerLiteral(stringLiteral.value()); if (literal.isPresent()) { diff --git a/core/src/test/java/io/substrait/type/ContainerReturnTypeTest.java b/core/src/test/java/io/substrait/type/ContainerReturnTypeTest.java index 97bd4d094..ed3ea2360 100644 --- a/core/src/test/java/io/substrait/type/ContainerReturnTypeTest.java +++ b/core/src/test/java/io/substrait/type/ContainerReturnTypeTest.java @@ -1,18 +1,22 @@ package io.substrait.type; import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; import static org.junit.jupiter.api.Assertions.assertThrows; import static org.junit.jupiter.api.Assertions.assertTrue; +import io.substrait.expression.Expression; import io.substrait.extension.DefaultExtensionCatalog; import io.substrait.extension.FunctionBindingResolver; import io.substrait.extension.ImmutableSimpleExtension; import io.substrait.extension.InvalidFunctionBindingException; +import io.substrait.extension.ResolvedAggregateBinding; 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 io.substrait.type.parser.TypeStringParser; import java.util.Arrays; import java.util.List; import java.util.stream.Collectors; @@ -331,6 +335,74 @@ void variadicContainersRespectParameterConsistencyAndLiteralConstraints() { } } + @Test + void aWildcardBoundFromAnElementKeepsItsNullabilityWhereverTheReturnNamesIt() { + // Under DECLARED_OUTPUT the return's nullability is the declaration's, so a nullable element + // bound into any1 has to survive a bare any1, a local assigned from it, and a conditional. + for (String program : List.of("any1", "t = any1\nlist", "list<1 > 0 ? any1 : any1>")) { + SimpleExtension.Function function = + ImmutableSimpleExtension.ScalarFunctionVariant.builder() + .from( + function( + TypeStringParser.parseExpression(program, "extension:test"), P.listE(ANY1))) + .nullability(SimpleExtension.Nullability.DECLARED_OUTPUT) + .build(); + for (Type element : List.of(R.I32, N.I32)) { + Type expected = program.endsWith(">") ? R.list(element) : element; + assertDerives(function, expected, List.of(R.list(element))); + } + } + } + + @Test + void signatureValidationChecksWhatTheArgumentsShare() { + // matchesDeclaration validates the signature without deriving a type, so it has to bind the + // parameters the way derivation does: a shared nested wildcard, an integer parameter's literal. + ParameterizedType list = P.listE(ANY1); + assertFalse(matches(aggregate(list, list), R.list(R.I32), R.list(R.I64))); + assertTrue(matches(aggregate(list, list), R.list(R.I32), R.list(R.I32))); + ParameterizedType decimals = P.listE(P.decimalE("P", "0")); + assertFalse(matches(aggregate(decimals), R.list(R.decimal(12, 1)))); + assertTrue(matches(aggregate(decimals), R.list(R.decimal(12, 0)))); + } + + @Test + void theUnboundTypeMatchesNoDeclaredShape() { + Type unbound = Type.Unbound.builder().build(); + assertInvalid(function(ANY1, P.listE(ANY1)), R.list(unbound)); + assertInvalid(function(R.I64, ANY1, P.listE(ANY1)), R.I32, unbound); + // Reported as a binding that does not match rather than escaping from a nullability check. + assertFalse(matches(aggregate(P.listE(Q.parameter("any1"))), R.list(unbound))); + assertFalse(matches(aggregate(ANY1), unbound)); + } + + private static SimpleExtension.AggregateFunctionVariant aggregate( + ParameterizedType... parameters) { + return ImmutableSimpleExtension.AggregateFunctionVariant.builder() + .urn("extension:io.substrait:container_test") + .name("container_aggregate") + .returnType(R.I64) + .args( + Arrays.stream(parameters) + .map(p -> SimpleExtension.ValueArgument.builder().value(p).build()) + .collect(Collectors.toList())) + .build(); + } + + private static boolean matches( + SimpleExtension.AggregateFunctionVariant declaration, Type... actual) { + return FunctionBindingResolver.matchesDeclaration( + ResolvedAggregateBinding.builder() + .function( + FunctionBindingResolver.resolve( + declaration, + Arrays.stream(actual).map(ResolvedArgument::value).collect(Collectors.toList()), + List.of())) + .phase(Expression.AggregationPhase.INITIAL_TO_RESULT) + .invocation(Expression.AggregationInvocation.ALL) + .build()); + } + private static SimpleExtension.ScalarFunctionVariant function( TypeExpression result, ParameterizedType... parameters) { return ImmutableSimpleExtension.ScalarFunctionVariant.builder() diff --git a/core/src/test/java/io/substrait/type/ParameterizedReturnTypeTest.java b/core/src/test/java/io/substrait/type/ParameterizedReturnTypeTest.java index 2bcc4312a..c870f38c5 100644 --- a/core/src/test/java/io/substrait/type/ParameterizedReturnTypeTest.java +++ b/core/src/test/java/io/substrait/type/ParameterizedReturnTypeTest.java @@ -172,6 +172,16 @@ void catalogReturnShapes() { "string_split:vchar_vchar", "transform:list_func"), variantsReturning(ParameterizedType.ListType.class)); + // A container element carrying arithmetic or a program local parses to the TypeExpression + // flavour, a sibling class the assertion above cannot see, so census both. The pinned catalog + // declares neither, and no map, struct or func return in any flavour. + assertEquals(List.of(), variantsReturning(TypeExpression.ListType.class)); + assertEquals(List.of(), variantsReturning(ParameterizedType.Map.class)); + assertEquals(List.of(), variantsReturning(TypeExpression.Map.class)); + assertEquals(List.of(), variantsReturning(ParameterizedType.Struct.class)); + assertEquals(List.of(), variantsReturning(TypeExpression.Struct.class)); + assertEquals(List.of(), variantsReturning(ParameterizedType.Func.class)); + assertEquals(List.of(), variantsReturning(TypeExpression.Func.class)); assertEquals( R.list(R.varChar(20)), resolve("string_split:vchar_vchar", R.varChar(20), R.varChar(20))); From 1243caa775b0aac0e3d9584442ef49eb24374679 Mon Sep 17 00:00:00 2001 From: Aleksandr Efimov Date: Wed, 30 Sep 2026 14:31:47 +0300 Subject: [PATCH 6/6] fix(core): keep an undetermined nullability out of containers list over list leaves open whether any1 is i32 or i32?. A direct list already failed on that, but a local or a conditional carried any1 into the list as a required i32. A wildcard whose nullability the arguments leave open may now stand unmarked only as the result itself. A local named in a return also keeps its nullability, as a bound wildcard does. --- .../type/TypeExpressionEvaluator.java | 40 +++++++++++++++---- .../type/ContainerReturnTypeTest.java | 20 +++++++++- 2 files changed, 52 insertions(+), 8 deletions(-) diff --git a/core/src/main/java/io/substrait/type/TypeExpressionEvaluator.java b/core/src/main/java/io/substrait/type/TypeExpressionEvaluator.java index c462fbd38..d05b5d894 100644 --- a/core/src/main/java/io/substrait/type/TypeExpressionEvaluator.java +++ b/core/src/main/java/io/substrait/type/TypeExpressionEvaluator.java @@ -409,6 +409,10 @@ private static final class ReturnTypeEvaluator // can hold any of them; each operation checks the kind it consumes. private final Map locals = new HashMap<>(); private boolean inProgram; + // Whether the expression being evaluated is the result itself rather than a container child or + // an assignment. Only there may a wildcard whose nullability the arguments leave open stand + // unmarked: the function's nullability handling decides the result's own nullability. + private boolean resultPosition = true; private ReturnTypeEvaluator(TypeExpression returnExpression, ParameterBindings bindings) { // Rendered only on failure: the expression can be a whole return program, and most @@ -529,6 +533,16 @@ private Type func( * nullability it was bound with, and a {@code ?} marker only widens it. */ private Type evaluateNested(TypeExpression expression) { + boolean enclosing = resultPosition; + resultPosition = false; + try { + return evaluateChild(expression); + } finally { + resultPosition = enclosing; + } + } + + private Type evaluateChild(TypeExpression expression) { if (expression instanceof ParameterizedType.StringLiteral) { ParameterizedType.StringLiteral variable = (ParameterizedType.StringLiteral) expression; Object local = locals.get(variable.value()); @@ -550,8 +564,8 @@ private Type evaluateNested(TypeExpression expression) { public Object visit(ParameterizedType.StringLiteral stringLiteral) { Object local = locals.get(stringLiteral.value()); if (local != null) { - return local instanceof Type - ? ((Type) local).withNullable(stringLiteral.nullable()) + return local instanceof Type && stringLiteral.nullable() + ? ((Type) local).withNullable(true) : local; } Integer integer = bindings.boundInteger(stringLiteral.value()); @@ -561,11 +575,17 @@ public Object visit(ParameterizedType.StringLiteral stringLiteral) { // A wildcard return (e.g. min(any1) -> any1) resolves to the bound argument type, and a // nullable return expression makes it nullable. A type bound from a nested element keeps its // own nullability. MIRROR policy, if any, is applied afterwards by the caller. This is - // checked before parsing the token as an integer literal - // because parseIntegerLiteral reports "not a number" by throwing, and a bound type name is - // never a numeral. + // checked before parsing the token as an integer literal because parseIntegerLiteral reports + // "not a number" by throwing, and a bound type name is never a numeral. Type bound = bindings.boundType(stringLiteral.value()); if (bound != null) { + if (!resultPosition + && !stringLiteral.nullable() + && !bindings.exactTypeNullabilities.contains(stringLiteral.value())) { + // Reached through an assignment or a conditional, on its way into a container. + throw new UnsupportedOperationException( + "Cannot derive nullability of type parameter '" + stringLiteral.value() + "'"); + } return stringLiteral.nullable() ? bound.withNullable(true) : bound; } OptionalInt literal = parseIntegerLiteral(stringLiteral.value()); @@ -632,8 +652,14 @@ public Object visit(TypeExpression.ReturnProgram program) { "Cannot evaluate a return program nested in another: " + program); } inProgram = true; - for (TypeExpression.ReturnProgram.Assignment assignment : program.assignments()) { - locals.put(assignment.name(), evaluate(assignment.expr(), Object.class)); + boolean enclosing = resultPosition; + resultPosition = false; + try { + for (TypeExpression.ReturnProgram.Assignment assignment : program.assignments()) { + locals.put(assignment.name(), evaluate(assignment.expr(), Object.class)); + } + } finally { + resultPosition = enclosing; } return evaluate(program.finalExpression(), Type.class); } diff --git a/core/src/test/java/io/substrait/type/ContainerReturnTypeTest.java b/core/src/test/java/io/substrait/type/ContainerReturnTypeTest.java index ed3ea2360..253de11ae 100644 --- a/core/src/test/java/io/substrait/type/ContainerReturnTypeTest.java +++ b/core/src/test/java/io/substrait/type/ContainerReturnTypeTest.java @@ -339,7 +339,12 @@ void variadicContainersRespectParameterConsistencyAndLiteralConstraints() { void aWildcardBoundFromAnElementKeepsItsNullabilityWhereverTheReturnNamesIt() { // Under DECLARED_OUTPUT the return's nullability is the declaration's, so a nullable element // bound into any1 has to survive a bare any1, a local assigned from it, and a conditional. - for (String program : List.of("any1", "t = any1\nlist", "list<1 > 0 ? any1 : any1>")) { + for (String program : + List.of( + "any1", + "t = any1\nlist", + "list<1 > 0 ? any1 : any1>", + "t = any1\nlist<1 > 0 ? t : t>")) { SimpleExtension.Function function = ImmutableSimpleExtension.ScalarFunctionVariant.builder() .from( @@ -354,6 +359,19 @@ void aWildcardBoundFromAnElementKeepsItsNullabilityWhereverTheReturnNamesIt() { } } + @Test + void anOpenNullabilityCannotReachAContainerThroughALocalOrAConditional() { + // list over list leaves open whether any1 is i32 or i32?, so a required element + // named through a local or a conditional is as undetermined as a direct list. + for (String program : List.of("list", "t = any1\nlist", "list<1 > 0 ? any1 : any1>")) { + assertInvalid( + function( + TypeStringParser.parseExpression(program, "extension:test"), + P.listE(Q.parameter("any1"))), + R.list(N.I32)); + } + } + @Test void signatureValidationChecksWhatTheArgumentsShare() { // matchesDeclaration validates the signature without deriving a type, so it has to bind the