-
Notifications
You must be signed in to change notification settings - Fork 96
feat(core)!: derive container return types from nested argument bindings #1289
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from all commits
bab8b94
99a03e4
b2ac422
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -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. | ||
| * | ||
| * <p>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<varchar<L1>>}, 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. | ||
| * <p>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<varchar<L1>>} and type parameters such as {@code list<any1>}. 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. | ||
| * | ||
| * <p>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<P,0>) 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<String, Type> types = new HashMap<>(); | ||
| private final Set<String> exactTypeNullabilities = new HashSet<>(); | ||
| private final Map<String, Integer> integers = new HashMap<>(); | ||
|
|
||
| private Type boundType(String name) { | ||
|
|
@@ -186,13 +192,29 @@ private Integer boundInteger(String token) { | |
| * an INCONSISTENT variadic repetition — named parameters are left unbound (each repetition is | ||
| * independent) while literal constraints are still enforced. | ||
| */ | ||
| private void bind(ParameterizedType declared, Type actual, boolean bindNames) { | ||
| private void bind(ParameterizedType declared, Type actual, boolean bindNames, boolean nested) { | ||
| if (nested && !(declared instanceof ParameterizedType.StringLiteral)) { | ||
| if ((declared instanceof NullableType | ||
| && ((NullableType) declared).nullable() != actual.nullable()) | ||
| || (declared instanceof Type && !declared.equals(actual))) { | ||
| throw cannotBind(declared, actual); | ||
| } | ||
| } | ||
| if (declared instanceof ParameterizedType.StringLiteral) { | ||
| ParameterizedType.StringLiteral literal = (ParameterizedType.StringLiteral) declared; | ||
| // Only a numbered wildcard names a parameter that a return expression can refer to and that | ||
| // has to stay consistent across the call; a plain "any" binds independently each time. | ||
| if (nested && literal.nullable() && !actual.nullable()) { | ||
| throw cannotBind(declared, actual); | ||
| } | ||
| if (bindNames && literal.isNumberedWildcard()) { | ||
| bindType(literal.value(), actual); | ||
| // An unmarked nested wildcard binds the complete type, including nullability. A '?' | ||
| // marker requires a nullable actual, but does not constrain the variable's own | ||
| // nullability: both i32 and i32? become i32? after substitution. Outermost argument | ||
| // nullability is excluded from binding and also leaves the variable's nullability open. | ||
| boolean exactNullability = nested && !literal.nullable(); | ||
| Type binding = nested && !literal.nullable() ? actual : actual.withNullable(false); | ||
| bindType(literal.value(), binding, exactNullability); | ||
| } | ||
| } else if (declared instanceof ParameterizedType.Decimal && actual instanceof Type.Decimal) { | ||
| ParameterizedType.Decimal declaredDecimal = (ParameterizedType.Decimal) declared; | ||
|
|
@@ -246,43 +268,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<ParameterizedType> declared, List<Type> 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<any1>} | ||
| * 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 +438,89 @@ 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) { | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Please say in the description how this derives an aggregate's parameterized |
||
| 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) { | ||
| 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()); | ||
|
|
||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Drop
nested &&here so an unmarked wildcard in a top-level argument binds exactly too.index_in(i32, list<i32?>)then stops binding, which matches the spec's rules and what substrait-io/substrait#1080 (still open) says about the same shape. It also fixes a return that needs a wildcard bound only at the top level:f(any1) -> list<any1>currently throws even for a requiredi32, and that is the shape quantile has after #1307. Two tests assert the relaxation and need flipping for the nullable-element cases (catalogIndexInAcceptsNullableElementsandtopLevelWildcardsDoNotConstrainNestedNullabilityInEitherOrder), and the description paragraph saying the spec table doesn't cover this case can go.