diff --git a/core/src/main/java/io/substrait/extension/FunctionBindingResolver.java b/core/src/main/java/io/substrait/extension/FunctionBindingResolver.java
index 79b7a4ee3..7a10b2ec4 100644
--- a/core/src/main/java/io/substrait/extension/FunctionBindingResolver.java
+++ b/core/src/main/java/io/substrait/extension/FunctionBindingResolver.java
@@ -27,14 +27,15 @@
*
*
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, 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
+ * 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 {
@@ -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,20 @@ 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) {
- // 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 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();
}
if (declared instanceof Type) {
// A concrete declared argument type (e.g. i32) matches ignoring nullability, except under a
@@ -520,15 +549,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..d05b5d894 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
@@ -104,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
@@ -162,7 +185,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 +194,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 +210,36 @@ 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 (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())
+ || (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 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, 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);
}
} else if (declared instanceof ParameterizedType.Decimal && actual instanceof Type.Decimal) {
ParameterizedType.Decimal declaredDecimal = (ParameterizedType.Decimal) declared;
@@ -246,43 +293,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) {
@@ -334,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
@@ -388,27 +467,126 @@ public Type visit(ParameterizedType.IntervalCompound intervalCompound) {
return intervalCompound(intervalCompound.nullable(), intervalCompound.precision());
}
+ @Override
+ public Type visit(ParameterizedType.ListType list) {
+ return list(list.nullable(), list.name());
+ }
+
+ @Override
+ public Type visit(ParameterizedType.Map map) {
+ return map(map.nullable(), map.key(), map.value());
+ }
+
+ @Override
+ public Type visit(ParameterizedType.Struct struct) {
+ return struct(struct.nullable(), struct.fields());
+ }
+
+ @Override
+ public Type visit(ParameterizedType.Func function) {
+ 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 extends TypeExpression> fields) {
+ return TypeCreator.of(nullable)
+ .struct(fields.stream().map(this::evaluateNested).collect(Collectors.toList()));
+ }
+
+ private Type func(
+ boolean nullable, List extends TypeExpression> parameters, TypeExpression returnType) {
+ return TypeCreator.of(nullable)
+ .func(
+ 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(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());
+ Type bound = local instanceof Type ? (Type) local : bindings.boundType(variable.value());
+ if (bound != null) {
+ if (!(local instanceof Type)
+ && !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());
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());
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
- // because parseIntegerLiteral reports "not a number" by throwing, and a bound type name is
- // never a numeral.
+ // 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());
+ 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());
if (literal.isPresent()) {
@@ -474,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/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..253de11ae
--- /dev/null
+++ b/core/src/test/java/io/substrait/type/ContainerReturnTypeTest.java
@@ -0,0 +1,454 @@
+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;
+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 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)) {
+ 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 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)) {
+ 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);
+ 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 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"));
+ // 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)));
+ }
+ }
+
+ @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>",
+ "t = any1\nlist<1 > 0 ? t : t>")) {
+ 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 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
+ // 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()
+ .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..c870f38c5 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",
@@ -171,10 +172,19 @@ void theReturnShapesThatAreNotDerivedYet() {
"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));
- 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/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>"));
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
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
}