From bd5fd19a26124ebf019a366e703f04cd4e76f22d Mon Sep 17 00:00:00 2001 From: Roman Kondakov Date: Sun, 26 Jul 2026 16:45:51 +0700 Subject: [PATCH 1/2] feat(core): resolve extension function bindings and derive their output types MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Adds a resolved-binding model for extension function invocations, and makes TypeExpressionEvaluator actually evaluate parameterized return types instead of throwing. ResolvedArgument keeps an argument's kind (value, type or enum) and its selected enum option, so two invocations differing only by an enum argument — e.g. std_dev(POPULATION, fp32) vs std_dev(SAMPLE, fp32) — are not conflated, and an enum the plan left unspecified stays distinct from any specified option. ResolvedFunctionBinding captures the semantic identity of an invocation: anchor, ordered arguments and options. ResolvedAggregateBinding adds the aggregate phase and invocation semantics, and selects the declaration's intermediate type for a phase that stops at the intermediate state rather than its return type. FunctionBindingResolver keeps the two concerns apart. resolve() only captures identity and performs no validation, so an invocation that merely differs from its declaration still resolves. validate() opts into checking arity, argument kinds and types, enum options, function options and the plan-declared output type against the declaration. TypeExpressionEvaluator previously threw UnsupportedOperationException("NYI") for anything but an already concrete return type, so SimpleExtension.Function's resolveType could not evaluate a parameterized declaration at all. It now binds numbered wildcards (the any1 of min(any1) -> any1) and integer type parameters (the P and S of DECIMAL) from the actual argument types and substitutes them, taking variadic parameter consistency and the declaration's MIRROR nullability policy into account. Occurrences of one numbered wildcard must agree on a single type, while each plain any binds independently. Derivations that are not supported — arithmetic derivations, if/then, return programs — stay fail-closed: they raise rather than fall back to a caller-supplied type, which is what makes a derived type trustworthy enough to validate a plan against. Co-Authored-By: Claude Opus 5 (1M context) --- .../extension/FunctionBindingResolver.java | 524 ++++++++++++++++++ .../InvalidFunctionBindingException.java | 40 ++ .../extension/ResolvedAggregateBinding.java | 154 +++++ .../substrait/extension/ResolvedArgument.java | 142 +++++ .../extension/ResolvedFunctionBinding.java | 114 ++++ .../substrait/extension/SimpleExtension.java | 3 +- .../substrait/function/ParameterizedType.java | 25 + .../type/TypeExpressionEvaluator.java | 217 +++++++- .../FunctionBindingResolverTest.java | 410 ++++++++++++++ .../ResolvedAggregateBindingTest.java | 178 ++++++ .../type/TypeExpressionEvaluatorTest.java | 80 +++ .../extensions/binding_extensions.yaml | 56 ++ 12 files changed, 1938 insertions(+), 5 deletions(-) create mode 100644 core/src/main/java/io/substrait/extension/FunctionBindingResolver.java create mode 100644 core/src/main/java/io/substrait/extension/InvalidFunctionBindingException.java create mode 100644 core/src/main/java/io/substrait/extension/ResolvedAggregateBinding.java create mode 100644 core/src/main/java/io/substrait/extension/ResolvedArgument.java create mode 100644 core/src/main/java/io/substrait/extension/ResolvedFunctionBinding.java create mode 100644 core/src/test/java/io/substrait/extension/FunctionBindingResolverTest.java create mode 100644 core/src/test/java/io/substrait/extension/ResolvedAggregateBindingTest.java create mode 100644 core/src/test/java/io/substrait/type/TypeExpressionEvaluatorTest.java create mode 100644 core/src/test/resources/extensions/binding_extensions.yaml diff --git a/core/src/main/java/io/substrait/extension/FunctionBindingResolver.java b/core/src/main/java/io/substrait/extension/FunctionBindingResolver.java new file mode 100644 index 000000000..8f16d0ac0 --- /dev/null +++ b/core/src/main/java/io/substrait/extension/FunctionBindingResolver.java @@ -0,0 +1,524 @@ +package io.substrait.extension; + +import io.substrait.expression.FunctionOption; +import io.substrait.function.NullableType; +import io.substrait.function.ParameterizedType; +import io.substrait.function.TypeExpression; +import io.substrait.type.Type; +import io.substrait.type.TypeExpressionEvaluator; +import java.util.ArrayList; +import java.util.HashMap; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Locale; +import java.util.Map; +import java.util.Optional; + +/** + * Resolves an extension function invocation to a {@link ResolvedFunctionBinding} and, on demand, + * validates and derives its output type. + * + *

{@link #resolve} only captures identity (anchor, arguments, options); it performs no + * validation, so a plan that merely differs from the declaration still resolves. {@link #validate} + * (and {@link #resolveAndValidate}) opt into checking the signature (arity, argument kinds and + * types), the options and the plan-declared output type against the declaration. {@link + * #deriveOutputType} derives the return type; unsupported derivations fail closed (see {@link + * TypeExpressionEvaluator}), never falling back to a plan-supplied type. + * + *

Signature type matching covers value-argument kinds, wildcards and the common + * concrete/decimal/char/binary type classes; other parameterized argument shapes are accepted + * without deep structural checking. 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, and only requires the repetitions to agree when its parameters are + * {@code CONSISTENT}. Enum options and option preferences are matched case-insensitively; an + * unspecified enum option is rejected only where the declaration requires one. + */ +public final class FunctionBindingResolver { + + private FunctionBindingResolver() {} + + /** + * Resolves a function declaration and arguments into a binding, capturing its identity. This + * performs no validation against the declaration; see {@link #validate}. + * + * @param declaration the function declaration + * @param arguments the ordered resolved arguments + * @param options the function options + * @return the resolved binding + */ + public static ResolvedFunctionBinding resolve( + SimpleExtension.Function declaration, + List arguments, + List options) { + // Resolving only captures identity; it does not validate. Signature/options/output validation + // is a separate, opt-in concern (see #validate), so a plan that merely differs from the + // declaration still converts under the default policy. + // Options are kept exactly as the plan spelled them, so a binding always reproduces the + // invocation it was resolved from. + return ResolvedFunctionBinding.builder() + .anchor(declaration.getAnchor()) + .declaration(declaration) + .arguments(arguments) + .options(options) + .build(); + } + + /** + * Resolves a function declaration and validates that the plan-declared output type matches the + * derived one. + * + * @param declaration the function declaration + * @param arguments the ordered resolved arguments + * @param options the function options + * @param declaredOutputType the output type declared by the plan + * @return the resolved binding + * @throws InvalidFunctionBindingException if the signature, options or output type are invalid + */ + public static ResolvedFunctionBinding resolveAndValidate( + SimpleExtension.Function declaration, + List arguments, + List options, + Type declaredOutputType) { + ResolvedFunctionBinding binding = resolve(declaration, arguments, options); + validate(binding, declaredOutputType); + return binding; + } + + /** + * Validates a resolved binding against its extension declaration: the signature (arity, argument + * kinds and types), the options, and — against the given plan-declared type — the output type. + * This is the opt-in extension-declaration check, separate from resolving the binding. + * + * @param binding the resolved binding + * @param declaredOutputType the output type declared by the plan + * @throws InvalidFunctionBindingException if the binding is inconsistent with the declaration + */ + public static void validate(ResolvedFunctionBinding binding, Type declaredOutputType) { + validateSignature(binding.declaration(), binding.arguments()); + validateOptions(binding.declaration(), binding.options()); + validateOutputType(binding, declaredOutputType); + } + + /** + * Validates a resolved aggregate binding against its extension declaration. Identical to {@link + * #validate(ResolvedFunctionBinding, Type)} except that the type a partial aggregate produces is + * its declaration's intermediate type rather than its return type, which {@link + * ResolvedAggregateBinding#outputType()} selects by phase. + * + * @param binding the resolved aggregate binding + * @param declaredOutputType the output type declared by the plan + * @throws InvalidFunctionBindingException if the binding is inconsistent with the declaration + */ + public static void validate(ResolvedAggregateBinding binding, Type declaredOutputType) { + validateAggregate(binding); + if (!binding.outputType().equals(declaredOutputType)) { + throw InvalidFunctionBindingException.outputTypeMismatch( + binding.function().anchor(), declaredOutputType, binding.outputType()); + } + } + + /** + * Reports whether a resolved aggregate binding is consistent with its declaration — its signature + * (as its phase defines it) and its options — without comparing any plan-declared output type and + * without throwing. Note that consistency is not the same question as "does this binding still + * describe that invocation": a phase whose state is parameterized by the initial arguments cannot + * be re-validated from the state alone, so comparing the recorded arguments answers that better. + * + * @param binding the resolved aggregate binding + * @return {@code true} if the binding satisfies its declaration + */ + public static boolean matchesDeclaration(ResolvedAggregateBinding binding) { + try { + validateAggregate(binding); + return true; + } catch (InvalidFunctionBindingException e) { + return false; + } + } + + private static void validateAggregate(ResolvedAggregateBinding binding) { + ResolvedFunctionBinding function = binding.function(); + if (binding.consumesIntermediateState()) { + validateIntermediateSignature(binding); + } else { + validateSignature(function.declaration(), function.arguments()); + } + validateOptions(function.declaration(), function.options()); + } + + /** + * Validates the signature of a phase that consumes an intermediate state: by contract its single + * argument is the declaration's intermediate value, not one of its declared arguments, and only a + * decomposable declaration has such a state at all. + */ + private static void validateIntermediateSignature(ResolvedAggregateBinding binding) { + ResolvedFunctionBinding function = binding.function(); + SimpleExtension.Function declaration = function.declaration(); + if (!(declaration instanceof SimpleExtension.AggregateFunctionVariant)) { + throw new InvalidFunctionBindingException( + String.format("%s is not an aggregate declaration", function.anchor())); + } + List arguments = function.arguments(); + if (arguments.size() != 1) { + throw new InvalidFunctionBindingException( + String.format( + "%s in phase %s takes a single intermediate-state argument but got %d", + function.anchor(), binding.phase(), arguments.size())); + } + ResolvedArgument state = arguments.get(0); + if (state.kind() != ResolvedArgument.Kind.VALUE) { + // The state is a value the previous phase produced. A type or enum argument in its place + // would carry no operand at all, and conversion drops it silently. + throw new InvalidFunctionBindingException( + String.format( + "%s in phase %s takes its intermediate state as a value argument but got %s", + function.anchor(), binding.phase(), state.kind())); + } + Type intermediate = + deriveIntermediateType( + (SimpleExtension.AggregateFunctionVariant) declaration, function.arguments()); + Type stateType = state.type().orElseThrow(IllegalStateException::new); + if (!stateType.equals(intermediate)) { + throw new InvalidFunctionBindingException( + String.format( + "%s in phase %s takes its intermediate state %s but got %s", + function.anchor(), binding.phase(), intermediate, stateType)); + } + } + + /** + * Validates that a plan-declared output type exactly matches the type derived from the binding's + * declaration. + * + * @param binding the resolved binding + * @param declaredOutputType the output type declared by the plan + * @throws InvalidFunctionBindingException if the declared type differs from the derived type + */ + public static void validateOutputType(ResolvedFunctionBinding binding, Type declaredOutputType) { + if (!binding.outputType().equals(declaredOutputType)) { + throw InvalidFunctionBindingException.outputTypeMismatch( + binding.anchor(), declaredOutputType, binding.outputType()); + } + } + + /** + * Derives the output type of a function from its declaration and arguments, applying the + * declaration's nullability policy. + * + * @param declaration the function declaration + * @param arguments the ordered resolved arguments + * @return the derived output type + */ + public static Type deriveOutputType( + SimpleExtension.Function declaration, List arguments) { + return derive(declaration.returnType(), declaration, arguments); + } + + /** + * Derives the intermediate type of a decomposable aggregate — the type its partial phases produce + * — from its declaration and arguments, applying the declaration's nullability policy. + * + *

The type parameters are bound from the given arguments, so an intermediate expression that + * refers to a parameter of the initial arguments cannot be derived from a phase that + * only sees the intermediate value; that fails closed rather than guessing. + * + * @param declaration the aggregate function declaration + * @param arguments the ordered resolved arguments + * @return the derived intermediate type + * @throws InvalidFunctionBindingException if the declaration is not decomposable or its + * intermediate expression cannot be derived + */ + public static Type deriveIntermediateType( + SimpleExtension.AggregateFunctionVariant declaration, List arguments) { + TypeExpression intermediate = declaration.intermediate(); + if (declaration.decomposability() == SimpleExtension.Decomposability.NONE + || intermediate == null) { + throw new InvalidFunctionBindingException( + String.format( + "%s is not decomposable and has no intermediate state", declaration.getAnchor())); + } + return derive(intermediate, declaration, arguments); + } + + private static Type derive( + TypeExpression expression, + SimpleExtension.Function declaration, + List arguments) { + List valueTypes = valueAndTypeArgumentTypes(arguments); + Type base; + try { + base = + TypeExpressionEvaluator.evaluateExpression( + expression, declaration.args(), declaration.variadic(), valueTypes); + } catch (UnsupportedOperationException e) { + // Surface unresolved/inconsistent type expressions as a binding error rather than an + // unchecked "not implemented" exception, so callers can handle them uniformly. + throw new InvalidFunctionBindingException( + String.format("Cannot derive type for %s: %s", declaration.getAnchor(), e.getMessage())); + } + if (declaration.nullability() == SimpleExtension.Nullability.MIRROR) { + // MIRROR: the output is nullable iff any value argument is nullable. + return base.withNullable(anyValueArgumentNullable(arguments)); + } + // DECLARED_OUTPUT / DISCRETE: the declared type's nullability is authoritative. + return base; + } + + private static List valueAndTypeArgumentTypes(List arguments) { + List types = new ArrayList<>(); + for (ResolvedArgument argument : arguments) { + if (argument.kind() != ResolvedArgument.Kind.ENUM) { + argument.type().ifPresent(types::add); + } + } + return types; + } + + private static boolean anyValueArgumentNullable(List arguments) { + for (ResolvedArgument argument : arguments) { + if (argument.kind() == ResolvedArgument.Kind.VALUE + && argument.type().map(Type::nullable).orElse(false)) { + return true; + } + } + return false; + } + + private static void validateSignature( + SimpleExtension.Function declaration, List arguments) { + List declared = declaration.args(); + // getRange() already accounts for variadic min/max and required arguments. + if (!declaration.getRange().within(arguments.size())) { + throw new InvalidFunctionBindingException( + String.format( + "%s does not accept %d argument(s)", declaration.getAnchor(), arguments.size())); + } + if (declared.isEmpty()) { + return; + } + boolean discrete = declaration.nullability() == SimpleExtension.Nullability.DISCRETE; + // A variadic declaration whose parameters are INCONSISTENT repeats its trailing argument + // without requiring the repetitions to agree with each other. + boolean bindRepeats = + declaration + .variadic() + .map( + behavior -> + behavior.parameterConsistency() + == SimpleExtension.VariadicBehavior.ParameterConsistency.CONSISTENT) + .orElse(false); + Map wildcardBindings = new HashMap<>(); + for (int i = 0; i < arguments.size(); i++) { + // For variadic functions the trailing declared argument repeats. + SimpleExtension.Argument declaredArgument = declared.get(Math.min(i, declared.size() - 1)); + boolean repeated = i >= declared.size(); + checkArgument( + declaration, + i, + declaredArgument, + arguments.get(i), + discrete, + !repeated || bindRepeats, + wildcardBindings); + } + } + + private static void checkArgument( + SimpleExtension.Function declaration, + int index, + SimpleExtension.Argument declared, + ResolvedArgument actual, + boolean discrete, + boolean bindWildcards, + Map wildcardBindings) { + if (declared instanceof SimpleExtension.ValueArgument) { + requireKind(declaration, index, ResolvedArgument.Kind.VALUE, actual); + Type actualType = actual.type().orElseThrow(IllegalStateException::new); + ParameterizedType declaredType = ((SimpleExtension.ValueArgument) declared).value(); + if (declaredType instanceof ParameterizedType.StringLiteral + && ((ParameterizedType.StringLiteral) declaredType).isWildcard()) { + checkWildcardArgument( + declaration, + index, + (ParameterizedType.StringLiteral) declaredType, + actualType, + discrete, + bindWildcards, + wildcardBindings); + } else if (!typeMatches(declaredType, actualType, discrete)) { + throw new InvalidFunctionBindingException( + String.format( + "%s argument %d: type %s is not compatible with declared %s", + declaration.getAnchor(), index, actualType, declared.toTypeString())); + } + } else if (declared instanceof SimpleExtension.TypeArgument) { + requireKind(declaration, index, ResolvedArgument.Kind.TYPE, actual); + } else if (declared instanceof SimpleExtension.EnumArgument) { + requireKind(declaration, index, ResolvedArgument.Kind.ENUM, actual); + checkEnumOption(declaration, index, (SimpleExtension.EnumArgument) declared, actual); + } + } + + private static void checkWildcardArgument( + SimpleExtension.Function declaration, + int index, + ParameterizedType.StringLiteral declared, + Type actualType, + boolean discrete, + boolean bindWildcards, + Map wildcardBindings) { + if (discrete && declared.nullable() != actualType.nullable()) { + // Under DISCRETE the declared nullability is part of the signature, wildcards included. + throw new InvalidFunctionBindingException( + String.format( + "%s argument %d: type %s is not compatible with declared %s", + declaration.getAnchor(), index, actualType, declared.value())); + } + if (!declared.isNumberedWildcard() || !bindWildcards) { + // A plain "any" matches each argument independently; only a numbered wildcard (any1) has to + // bind to a single type across the invocation. + return; + } + String name = declared.value(); + Type existing = wildcardBindings.putIfAbsent(name, actualType); + if (existing != null && !existing.equalsIgnoringNullability(actualType)) { + throw new InvalidFunctionBindingException( + String.format( + "%s argument %d: wildcard '%s' is bound to both %s and %s", + declaration.getAnchor(), index, name, existing, actualType)); + } + } + + private static void checkEnumOption( + SimpleExtension.Function declaration, + int index, + SimpleExtension.EnumArgument declared, + ResolvedArgument actual) { + Optional option = actual.enumValue(); + if (!option.isPresent()) { + // A plan may leave an enum argument unspecified; that is only valid where the declaration + // does not require an option. + if (declared.required()) { + throw new InvalidFunctionBindingException( + String.format( + "%s argument %d: enum option is required but was left unspecified (expected one of" + + " %s)", + declaration.getAnchor(), index, declared.options())); + } + return; + } + // Substrait matches enum option symbols case-insensitively. + for (String allowed : declared.options()) { + if (allowed.equalsIgnoreCase(option.get())) { + return; + } + } + throw new InvalidFunctionBindingException( + String.format( + "%s argument %d: enum option '%s' is not one of %s", + declaration.getAnchor(), index, option.get(), declared.options())); + } + + private static void requireKind( + SimpleExtension.Function declaration, + int index, + ResolvedArgument.Kind expected, + ResolvedArgument actual) { + if (actual.kind() != expected) { + throw new InvalidFunctionBindingException( + String.format( + "%s argument %d: expected %s argument but got %s", + declaration.getAnchor(), index, expected, actual.kind())); + } + } + + 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; + } + if (declared instanceof Type) { + // A concrete declared argument type (e.g. i32) matches ignoring nullability, except under a + // DISCRETE declaration where argument nullability is part of the signature. + return exactNullability + ? actual.equals(declared) + : actual.equalsIgnoringNullability((Type) declared); + } + if (declared instanceof ParameterizedType.Decimal) { + return actual instanceof Type.Decimal + && nullabilityMatches(declared, actual, exactNullability); + } + if (declared instanceof ParameterizedType.FixedChar) { + return actual instanceof Type.FixedChar + && nullabilityMatches(declared, actual, exactNullability); + } + if (declared instanceof ParameterizedType.VarChar) { + return actual instanceof Type.VarChar + && nullabilityMatches(declared, actual, exactNullability); + } + if (declared instanceof ParameterizedType.FixedBinary) { + return actual instanceof Type.FixedBinary + && nullabilityMatches(declared, actual, exactNullability); + } + if (declared instanceof ParameterizedType.PrecisionTimestamp) { + return actual instanceof Type.PrecisionTimestamp + && nullabilityMatches(declared, actual, exactNullability); + } + if (declared instanceof ParameterizedType.PrecisionTimestampTZ) { + return actual instanceof Type.PrecisionTimestampTZ + && nullabilityMatches(declared, actual, exactNullability); + } + // Other parameterized shapes (structs, lists, maps, ...) are accepted without deep checking. + 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 + * nullability of its own has nothing to compare, so it is accepted. + */ + private static boolean nullabilityMatches( + ParameterizedType declared, Type actual, boolean exactNullability) { + if (!exactNullability || !(declared instanceof NullableType)) { + return true; + } + return ((NullableType) declared).nullable() == actual.nullable(); + } + + private static void validateOptions( + SimpleExtension.Function declaration, List options) { + // A binding keeps option names and values exactly as the plan spelled them, so lower-case both + // sides here: the Substrait spec matches options case-insensitively. + Map> declaredOptions = new LinkedHashMap<>(); + for (Map.Entry entry : declaration.options().entrySet()) { + List values = new ArrayList<>(); + for (String value : entry.getValue().getValues()) { + values.add(value.toLowerCase(Locale.ROOT)); + } + declaredOptions.put(entry.getKey().toLowerCase(Locale.ROOT), values); + } + for (FunctionOption option : options) { + String name = option.getName().toLowerCase(Locale.ROOT); + List allowed = declaredOptions.get(name); + if (allowed == null) { + throw new InvalidFunctionBindingException( + String.format("%s has no option named '%s'", declaration.getAnchor(), name)); + } + if (option.values().isEmpty()) { + throw new InvalidFunctionBindingException( + String.format( + "%s option '%s' has an empty preference list", declaration.getAnchor(), name)); + } + for (String value : option.values()) { + if (!allowed.contains(value.toLowerCase(Locale.ROOT))) { + throw new InvalidFunctionBindingException( + String.format( + "%s option '%s' value '%s' is not one of %s", + declaration.getAnchor(), name, value, allowed)); + } + } + } + } +} diff --git a/core/src/main/java/io/substrait/extension/InvalidFunctionBindingException.java b/core/src/main/java/io/substrait/extension/InvalidFunctionBindingException.java new file mode 100644 index 000000000..535d35e3e --- /dev/null +++ b/core/src/main/java/io/substrait/extension/InvalidFunctionBindingException.java @@ -0,0 +1,40 @@ +package io.substrait.extension; + +import io.substrait.type.Type; + +/** + * Thrown when an extension function invocation cannot be resolved, or when the output type declared + * by a plan is inconsistent with the type derived from the function's extension declaration. + * + *

Resolution is fail-closed: rather than trusting a plan-supplied output type, the resolver + * derives the type from the declaration and raises this exception when they disagree. + */ +public class InvalidFunctionBindingException extends RuntimeException { + + private static final long serialVersionUID = 1L; + + /** + * Creates an exception with the given message. + * + * @param message the detail message + */ + public InvalidFunctionBindingException(String message) { + super(message); + } + + /** + * Creates an exception describing a mismatch between the declared and the derived output type. + * + * @param anchor the function anchor being resolved + * @param declared the output type declared by the plan + * @param derived the output type derived from the declaration + * @return the exception + */ + public static InvalidFunctionBindingException outputTypeMismatch( + SimpleExtension.FunctionAnchor anchor, Type declared, Type derived) { + return new InvalidFunctionBindingException( + String.format( + "Declared output type %s for %s does not match the type %s derived from its declaration", + declared, anchor, derived)); + } +} diff --git a/core/src/main/java/io/substrait/extension/ResolvedAggregateBinding.java b/core/src/main/java/io/substrait/extension/ResolvedAggregateBinding.java new file mode 100644 index 000000000..27858d200 --- /dev/null +++ b/core/src/main/java/io/substrait/extension/ResolvedAggregateBinding.java @@ -0,0 +1,154 @@ +package io.substrait.extension; + +import io.substrait.expression.AggregateFunctionInvocation; +import io.substrait.expression.EnumArg; +import io.substrait.expression.Expression; +import io.substrait.expression.FunctionArg; +import io.substrait.type.Type; +import java.util.ArrayList; +import java.util.List; +import java.util.Optional; +import org.immutables.value.Value; + +/** + * A fully-resolved binding of an aggregate function invocation. + * + *

Wraps a {@link ResolvedFunctionBinding} with the aggregate-specific phase and invocation + * semantics. The output type is derived from the function declaration (never from a plan-supplied + * value) and depends on the phase, so it does not participate in the binding's identity. + */ +@Value.Immutable +public abstract class ResolvedAggregateBinding { + + /** + * Returns the resolved function binding. + * + * @return the function binding + */ + public abstract ResolvedFunctionBinding function(); + + /** + * Returns the aggregation phase. + * + * @return the aggregation phase + */ + public abstract Expression.AggregationPhase phase(); + + /** + * Returns the aggregation invocation semantics (all vs. distinct). + * + * @return the aggregation invocation + */ + public abstract Expression.AggregationInvocation invocation(); + + /** + * Returns the intermediate type, when it has been resolved explicitly. When empty, {@link + * #outputType()} derives it from the declaration on demand. Not part of the binding identity. + * + * @return the intermediate type, if known + */ + @Value.Auxiliary + public abstract Optional intermediateType(); + + /** + * Returns the type this aggregate produces in its phase: a phase that stops at the + * intermediate state produces the declaration's intermediate type, while a phase that runs to the + * result produces its return type. Derived from the declaration, never from a plan-supplied + * value. + * + * @return the output type of this phase + * @throws InvalidFunctionBindingException if an intermediate type is required but the declaration + * is not an aggregate variant, or its type expression cannot be derived + */ + public Type outputType() { + return producesIntermediateState() + ? intermediateType().orElseGet(this::deriveIntermediateType) + : function().outputType(); + } + + /** + * Returns whether this phase stops at the declaration's intermediate state instead of producing + * its result. + * + * @return {@code true} for the initial-to-intermediate and intermediate-to-intermediate phases + */ + public boolean producesIntermediateState() { + return phase() == Expression.AggregationPhase.INITIAL_TO_INTERMEDIATE + || phase() == Expression.AggregationPhase.INTERMEDIATE_TO_INTERMEDIATE; + } + + /** + * Returns whether this phase consumes the declaration's intermediate state instead of its + * declared arguments. The arguments of such an invocation are intermediate values, so they are + * not expected to match the declaration's argument list. + * + * @return {@code true} for the intermediate-to-intermediate and intermediate-to-result phases + */ + public boolean consumesIntermediateState() { + return phase() == Expression.AggregationPhase.INTERMEDIATE_TO_INTERMEDIATE + || phase() == Expression.AggregationPhase.INTERMEDIATE_TO_RESULT; + } + + private Type deriveIntermediateType() { + SimpleExtension.Function declaration = function().declaration(); + if (!(declaration instanceof SimpleExtension.AggregateFunctionVariant)) { + throw new InvalidFunctionBindingException( + String.format( + "%s is not an aggregate declaration and has no intermediate type", + function().anchor())); + } + return FunctionBindingResolver.deriveIntermediateType( + (SimpleExtension.AggregateFunctionVariant) declaration, function().arguments()); + } + + /** + * Creates a builder for {@link ResolvedAggregateBinding}. + * + * @return a new builder + */ + public static ImmutableResolvedAggregateBinding.Builder builder() { + return ImmutableResolvedAggregateBinding.builder(); + } + + /** + * Resolves an aggregate function invocation into a binding, capturing its semantic identity + * (anchor, arguments and options). This does not validate the invocation against the + * declaration; signature, options and output-type validation is a separate, opt-in step (see + * {@link FunctionBindingResolver#validate}). The output type is derived on demand via {@link + * #outputType()}. + * + * @param invocation the aggregate invocation to resolve + * @return the resolved aggregate binding + */ + public static ResolvedAggregateBinding resolve(AggregateFunctionInvocation invocation) { + List arguments = resolvedArguments(invocation.arguments()); + ResolvedFunctionBinding function = + FunctionBindingResolver.resolve(invocation.declaration(), arguments, invocation.options()); + return builder() + .function(function) + .phase(invocation.aggregationPhase()) + .invocation(invocation.invocation()) + .intermediateType(Optional.empty()) + .build(); + } + + private static List resolvedArguments(List arguments) { + List resolved = new ArrayList<>(); + for (FunctionArg arg : arguments) { + if (arg instanceof Expression) { + resolved.add(ResolvedArgument.value(((Expression) arg).getType())); + } else if (arg instanceof Type) { + resolved.add(ResolvedArgument.type((Type) arg)); + } else if (arg instanceof EnumArg) { + // An enum argument may carry no option; keep that distinct from an empty one so the + // identity is faithful and validation can report it against the declaration. + resolved.add( + ((EnumArg) arg) + .value() + .map(ResolvedArgument::enumOption) + .orElseGet(ResolvedArgument::unspecifiedEnumOption)); + } + } + return resolved; + } +} diff --git a/core/src/main/java/io/substrait/extension/ResolvedArgument.java b/core/src/main/java/io/substrait/extension/ResolvedArgument.java new file mode 100644 index 000000000..4cfffb6de --- /dev/null +++ b/core/src/main/java/io/substrait/extension/ResolvedArgument.java @@ -0,0 +1,142 @@ +package io.substrait.extension; + +import io.substrait.type.Type; +import java.util.Objects; +import java.util.Optional; +import org.jspecify.annotations.Nullable; + +/** + * An ordered, kind-aware argument of a resolved function binding. + * + *

Unlike a bare {@link Type} list, this preserves the argument kind (value, type or + * enum) and the selected enum option, so that two invocations that differ only by an enum argument + * — e.g. {@code std_dev(POPULATION, fp32)} vs {@code std_dev(SAMPLE, fp32)} — are not conflated. + * + *

An enum argument may also carry no option, for a plan that leaves it unspecified; + * that is distinct from any specified option. An option is kept exactly as the plan spelled it, so + * identity is case-sensitive and merely conservative: two invocations differing only in the case of + * an enum symbol stay distinct, which costs a missed deduplication and never rewrites the plan's + * data. Matching an option against its declaration is case-insensitive, as the spec requires. + */ +public final class ResolvedArgument { + + /** The kind of a function argument. */ + public enum Kind { + /** A value argument, carrying a data {@link Type}. */ + VALUE, + /** A type argument, carrying a {@link Type}. */ + TYPE, + /** An enum argument, carrying a selected option. */ + ENUM + } + + private final Kind kind; + private final @Nullable Type type; + private final @Nullable String enumValue; + + private ResolvedArgument(Kind kind, @Nullable Type type, @Nullable String enumValue) { + this.kind = kind; + this.type = type; + this.enumValue = enumValue; + } + + /** + * Creates a value argument. + * + * @param type the value type + * @return the resolved argument + */ + public static ResolvedArgument value(Type type) { + return new ResolvedArgument(Kind.VALUE, Objects.requireNonNull(type), null); + } + + /** + * Creates a type argument. + * + * @param type the argument type + * @return the resolved argument + */ + public static ResolvedArgument type(Type type) { + return new ResolvedArgument(Kind.TYPE, Objects.requireNonNull(type), null); + } + + /** + * Creates an enum argument. + * + * @param option the selected enum option + * @return the resolved argument + */ + public static ResolvedArgument enumOption(String option) { + return new ResolvedArgument(Kind.ENUM, null, Objects.requireNonNull(option)); + } + + /** + * Creates an enum argument whose option the plan left unspecified. This is distinct from an + * argument carrying an empty option, and is only valid where the declaration does not require an + * option. + * + * @return the resolved argument + */ + public static ResolvedArgument unspecifiedEnumOption() { + return new ResolvedArgument(Kind.ENUM, null, null); + } + + /** + * Returns the argument kind. + * + * @return the kind + */ + public Kind kind() { + return kind; + } + + /** + * Returns the data type of a value or type argument. + * + * @return the type, if this is a value or type argument + */ + public Optional type() { + return Optional.ofNullable(type); + } + + /** + * Returns the selected option of an enum argument. + * + * @return the enum option, or empty if this is not an enum argument or its option was left + * unspecified + */ + public Optional enumValue() { + return Optional.ofNullable(enumValue); + } + + @Override + public boolean equals(@Nullable Object o) { + if (this == o) { + return true; + } + if (!(o instanceof ResolvedArgument)) { + return false; + } + ResolvedArgument other = (ResolvedArgument) o; + return kind == other.kind + && Objects.equals(type, other.type) + && Objects.equals(enumValue, other.enumValue); + } + + @Override + public int hashCode() { + return Objects.hash(kind, type, enumValue); + } + + @Override + public String toString() { + switch (kind) { + case VALUE: + return "value(" + type + ")"; + case TYPE: + return "type(" + type + ")"; + default: + return "enum(" + (enumValue == null ? "unspecified" : enumValue) + ")"; + } + } +} diff --git a/core/src/main/java/io/substrait/extension/ResolvedFunctionBinding.java b/core/src/main/java/io/substrait/extension/ResolvedFunctionBinding.java new file mode 100644 index 000000000..d65db3370 --- /dev/null +++ b/core/src/main/java/io/substrait/extension/ResolvedFunctionBinding.java @@ -0,0 +1,114 @@ +package io.substrait.extension; + +import io.substrait.expression.FunctionOption; +import io.substrait.type.Type; +import java.util.List; +import java.util.Optional; +import org.immutables.value.Value; + +/** + * A fully-resolved binding of an extension function declaration to concrete arguments. + * + *

The binding captures the semantic identity of a function invocation: the function + * anchor, the ordered arguments and the options. The {@linkplain #outputType() output type} is + * derived from the declaration on demand (see {@link FunctionBindingResolver}) and is + * deliberately not part of this identity — it belongs to the type-carrying Calcite wrapper, not to + * the semantic identity of the function. Two bindings are equal when they agree on anchor, + * arguments and options, which is what lets a consumer distinguish two functions that happen to + * lower to the same engine operator. + */ +@Value.Immutable +public abstract class ResolvedFunctionBinding { + + /** + * Returns the canonical function anchor (extension urn + compound signature). + * + * @return the function anchor + */ + public abstract SimpleExtension.FunctionAnchor anchor(); + + /** + * Returns the ordered, kind-aware arguments the function was resolved against. + * + * @return the resolved arguments + */ + public abstract List arguments(); + + /** + * Returns the options exactly as the invocation carries them: named preference lists, in order, + * spelled the way the plan spelled them. Selecting a single value from a preference list is a + * consumer-specific concern and is intentionally not done here. + * + *

This is a list rather than a map by name because that is what a plan carries — an option + * name may legitimately appear more than once, and a map would silently drop all but one. + * + * @return the function options + */ + public abstract List options(); + + /** + * Returns the preference list of the named option, matched case-insensitively as the Substrait + * spec requires. + * + * @param name the option name + * @return the preferred values of the first option with that name, if any + */ + public Optional> option(String name) { + for (FunctionOption option : options()) { + if (option.getName().equalsIgnoreCase(name)) { + return Optional.of(option.values()); + } + } + return Optional.empty(); + } + + /** + * Returns the function declaration this binding resolves. Not part of the binding identity. + * + * @return the function declaration + */ + @Value.Auxiliary + public abstract SimpleExtension.Function declaration(); + + /** + * Derives the output type from the declaration and arguments, applying the declaration's + * nullability policy. Computed on demand (may throw for as-yet-unsupported derivations); it is + * neither stored nor part of the binding's identity. + * + * @return the derived output type + */ + public Type outputType() { + return FunctionBindingResolver.deriveOutputType(declaration(), arguments()); + } + + /** + * Enforces that the anchor is consistent with the declaration. + * + *

Options need no normalization: {@link FunctionOption} is itself immutable, so the binding is + * immutable all the way down and its identity cannot change after it was built. Names and values + * are kept exactly as the plan spelled them — a binding is also the record of what to convert + * back, so it must not rewrite the plan's data. Identity is therefore case-sensitive and merely + * conservative: two invocations that differ only in the case of an option stay distinct, which + * costs a missed deduplication and never loses information. Validation against the declaration is + * case-insensitive, as the spec requires. + */ + @Value.Check + protected void checkAnchor() { + if (!anchor().equals(declaration().getAnchor())) { + throw new IllegalArgumentException( + String.format( + "anchor %s does not match declaration anchor %s", + anchor(), declaration().getAnchor())); + } + } + + /** + * Creates a builder for {@link ResolvedFunctionBinding}. Prefer {@link FunctionBindingResolver}, + * which validates the signature and options before constructing a binding. + * + * @return a new builder + */ + public static ImmutableResolvedFunctionBinding.Builder builder() { + return ImmutableResolvedFunctionBinding.builder(); + } +} diff --git a/core/src/main/java/io/substrait/extension/SimpleExtension.java b/core/src/main/java/io/substrait/extension/SimpleExtension.java index c2c2298b0..b48b5c782 100644 --- a/core/src/main/java/io/substrait/extension/SimpleExtension.java +++ b/core/src/main/java/io/substrait/extension/SimpleExtension.java @@ -655,7 +655,8 @@ public String key() { * @return the resolve Type */ public io.substrait.type.Type resolveType(List argumentTypes) { - return TypeExpressionEvaluator.evaluateExpression(returnType(), args(), argumentTypes); + return TypeExpressionEvaluator.evaluateExpression( + returnType(), args(), variadic(), argumentTypes); } } diff --git a/core/src/main/java/io/substrait/function/ParameterizedType.java b/core/src/main/java/io/substrait/function/ParameterizedType.java index 2b38d0930..6c1fabece 100644 --- a/core/src/main/java/io/substrait/function/ParameterizedType.java +++ b/core/src/main/java/io/substrait/function/ParameterizedType.java @@ -56,6 +56,17 @@ default boolean isWildcard() { return false; } + /** + * Returns whether this type is a numbered wildcard ({@code any1}, {@code any2}, ...), as + * opposed to a plain {@code any}. Occurrences of the same numbered wildcard within one signature + * must bind to the same type, while each plain {@code any} binds independently. + * + * @return {@code true} if this type is a numbered wildcard + */ + default boolean isNumberedWildcard() { + return false; + } + /** Base class for parameterized types that dispatch to a {@link ParameterizedTypeVisitor}. */ abstract class BaseParameterizedType implements ParameterizedType { @Override @@ -453,6 +464,20 @@ public boolean isWildcard() { return value().toLowerCase(Locale.ROOT).startsWith("any"); } + @Override + public boolean isNumberedWildcard() { + String literal = value().toLowerCase(Locale.ROOT); + if (!literal.startsWith("any") || literal.length() == "any".length()) { + return false; + } + for (int i = "any".length(); i < literal.length(); i++) { + if (!Character.isDigit(literal.charAt(i))) { + return false; + } + } + return true; + } + @Override R accept(final ParameterizedTypeVisitor parameterizedTypeVisitor) throws E { diff --git a/core/src/main/java/io/substrait/type/TypeExpressionEvaluator.java b/core/src/main/java/io/substrait/type/TypeExpressionEvaluator.java index b5ae26318..5171e86fa 100644 --- a/core/src/main/java/io/substrait/type/TypeExpressionEvaluator.java +++ b/core/src/main/java/io/substrait/type/TypeExpressionEvaluator.java @@ -1,11 +1,26 @@ package io.substrait.type; import io.substrait.extension.SimpleExtension; +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.List; +import java.util.Map; +import java.util.Optional; /** * Evaluates a {@link TypeExpression} to a concrete {@link Type} given a set of actual arguments. + * + *

A declaration's {@code return} expression is either already concrete (e.g. {@code i64?}) or + * parameterized in terms of its argument types (e.g. {@code DECIMAL?<38,S>}, where {@code S} is the + * scale of the argument). This evaluator resolves the latter by binding the integer type parameters + * ({@code P}, {@code S}, ...) from the actual argument types and substituting them. + * + *

Expression shapes that are not yet supported (arithmetic derivations, {@code if/then}, return + * programs, ...) fail closed with an {@link UnsupportedOperationException}. The evaluator never + * falls back to a caller-supplied type: an unresolved expression is an error, not a default. */ public class TypeExpressionEvaluator { @@ -13,19 +28,213 @@ public class TypeExpressionEvaluator { * Evaluates a return-type expression to a concrete {@link Type}. * * @param returnExpression the type expression to evaluate - * @param parameterizedTypeList the declared parameter types of the function + * @param declaredArguments the declared arguments of the function (used to bind type parameters) * @param actualTypes the actual argument types supplied at the call site * @return the resolved concrete type - * @throws UnsupportedOperationException if the expression cannot yet be evaluated + * @throws UnsupportedOperationException if the expression cannot be evaluated */ public static Type evaluateExpression( TypeExpression returnExpression, - List parameterizedTypeList, + List declaredArguments, List actualTypes) { + return evaluateExpression(returnExpression, declaredArguments, Optional.empty(), actualTypes); + } + /** + * Evaluates a return-type expression to a concrete {@link Type}, taking the declaration's + * variadic behavior into account. + * + * @param returnExpression the type expression to evaluate + * @param declaredArguments the declared arguments of the function (used to bind type parameters) + * @param variadic the declaration's variadic behavior, if it is variadic + * @param actualTypes the actual argument types supplied at the call site + * @return the resolved concrete type + * @throws UnsupportedOperationException if the expression cannot be evaluated + */ + public static Type evaluateExpression( + TypeExpression returnExpression, + List declaredArguments, + Optional variadic, + List actualTypes) { + // Bind the parameters even when the return type is concrete: binding is what catches a + // signature that uses one parameter name inconsistently — f(DECIMAL, DECIMAL) called + // with two different scales — and that is an error whatever the function returns. + ParameterBindings bindings = bindParameters(declaredArguments, variadic, actualTypes); if (returnExpression instanceof Type) { + // The declared return type is already concrete; nothing to derive. return (Type) returnExpression; } - throw new UnsupportedOperationException("NYI"); + return returnExpression.accept(new ReturnTypeEvaluator(returnExpression, bindings)); + } + + /** + * 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 + * DECIMAL}) — by matching each declared value argument against the corresponding actual + * argument type in a single pass. + * + *

Binding the same parameter name to two different values is a signature error and is rejected + * rather than silently overwritten, for wildcards and integer parameters alike. A variadic + * declaration states its trailing argument once but accepts it repeatedly, so every actual + * argument is bound against that trailing declaration — unless the declaration marks its + * parameters {@code INCONSISTENT}, in which case each repetition is independent and only the + * first one binds. + */ + private static ParameterBindings bindParameters( + List declaredArguments, + Optional variadic, + List actualTypes) { + // Enum arguments select an overload and carry no value type, so they do not consume an actual + // type; dropping them aligns the declared and actual arguments positionally. A null entry is a + // type argument: it consumes an actual type but binds no parameter. + List declaredTypes = new ArrayList<>(); + for (SimpleExtension.Argument declared : declaredArguments) { + if (declared instanceof SimpleExtension.EnumArgument) { + continue; + } + declaredTypes.add( + declared instanceof SimpleExtension.ValueArgument + ? ((SimpleExtension.ValueArgument) declared).value() + : null); + } + + ParameterBindings bindings = new ParameterBindings(); + if (declaredTypes.isEmpty()) { + return bindings; + } + boolean bindRepeats = + variadic + .map( + behavior -> + behavior.parameterConsistency() + == SimpleExtension.VariadicBehavior.ParameterConsistency.CONSISTENT) + .orElse(false); + for (int index = 0; index < actualTypes.size(); index++) { + boolean repeated = index >= declaredTypes.size(); + if (repeated && !bindRepeats) { + // Either the arity is wrong (validated by the resolver, not here) or the repetitions are + // declared independent; in both cases there is nothing further to bind. + break; + } + ParameterizedType declared = + repeated ? declaredTypes.get(declaredTypes.size() - 1) : declaredTypes.get(index); + if (declared != null) { + bindings.bind(declared, actualTypes.get(index)); + } + } + return bindings; + } + + /** The wildcard and integer type parameters bound from a call site's actual argument types. */ + private static final class ParameterBindings { + + private final Map types = new HashMap<>(); + private final Map integers = new HashMap<>(); + + private Type boundType(String name) { + return types.get(name); + } + + private Integer boundInteger(String token) { + return integers.get(token); + } + + private void bind(ParameterizedType declared, Type 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 (literal.isNumberedWildcard()) { + bindType(literal.value(), actual); + } + } else if (declared instanceof ParameterizedType.Decimal && actual instanceof Type.Decimal) { + ParameterizedType.Decimal declaredDecimal = (ParameterizedType.Decimal) declared; + Type.Decimal actualDecimal = (Type.Decimal) actual; + bindInteger(declaredDecimal.precision().value(), actualDecimal.precision()); + bindInteger(declaredDecimal.scale().value(), actualDecimal.scale()); + } + } + + 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. + Type existing = types.putIfAbsent(name, actual); + if (existing != null && !existing.equalsIgnoringNullability(actual)) { + throw new UnsupportedOperationException( + String.format( + "Inconsistent binding for type parameter '%s': %s vs %s", name, existing, actual)); + } + } + + private void bindInteger(String token, int value) { + if (isInteger(token)) { + // A numeric token is a literal, not a parameter name; nothing to bind. + return; + } + Integer existing = integers.putIfAbsent(token, value); + if (existing != null && existing != value) { + throw new UnsupportedOperationException( + String.format( + "Inconsistent binding for type parameter '%s': %d vs %d", token, existing, value)); + } + } + + private static boolean isInteger(String token) { + try { + Integer.parseInt(token.trim()); + return true; + } catch (NumberFormatException e) { + return false; + } + } + } + + /** + * Evaluates the supported return-type expression shapes. Everything else falls through to the + * throwing base, keeping unsupported derivations fail-closed. + */ + private static final class ReturnTypeEvaluator + extends TypeExpressionVisitor.TypeExpressionThrowsVisitor { + + private final ParameterBindings bindings; + + private ReturnTypeEvaluator(TypeExpression returnExpression, ParameterBindings bindings) { + super("Cannot evaluate return-type expression: " + returnExpression); + this.bindings = bindings; + } + + @Override + public Type visit(ParameterizedType.Decimal decimal) { + int precision = resolveInteger(decimal.precision().value()); + int scale = resolveInteger(decimal.scale().value()); + return TypeCreator.of(decimal.nullable()).decimal(precision, scale); + } + + @Override + public Type visit(ParameterizedType.StringLiteral stringLiteral) { + // A wildcard return (e.g. min(any1) -> any1) resolves to the bound argument type, taking the + // nullability declared on the return expression in both directions (a required return forces + // the type non-null, a nullable one forces it nullable). MIRROR policy, if any, is applied + // afterwards by the caller. + Type bound = bindings.boundType(stringLiteral.value()); + if (bound == null) { + throw new UnsupportedOperationException( + "Unbound type parameter '" + stringLiteral.value() + "' in return-type expression"); + } + return bound.withNullable(stringLiteral.nullable()); + } + + private int resolveInteger(String token) { + Integer bound = bindings.boundInteger(token); + if (bound != null) { + return bound; + } + try { + return Integer.parseInt(token.trim()); + } catch (NumberFormatException e) { + throw new UnsupportedOperationException( + "Unbound type parameter '" + token + "' in return-type expression"); + } + } } } diff --git a/core/src/test/java/io/substrait/extension/FunctionBindingResolverTest.java b/core/src/test/java/io/substrait/extension/FunctionBindingResolverTest.java new file mode 100644 index 000000000..92410c1f8 --- /dev/null +++ b/core/src/test/java/io/substrait/extension/FunctionBindingResolverTest.java @@ -0,0 +1,410 @@ +package io.substrait.extension; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertNotEquals; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import com.google.common.io.Resources; +import io.substrait.expression.FunctionOption; +import io.substrait.type.TypeCreator; +import java.io.IOException; +import java.io.UncheckedIOException; +import java.nio.charset.StandardCharsets; +import java.util.ArrayList; +import java.util.List; +import java.util.Optional; +import org.junit.jupiter.api.Test; + +class FunctionBindingResolverTest { + + static final TypeCreator R = TypeCreator.REQUIRED; + static final TypeCreator N = TypeCreator.NULLABLE; + + /** Declarations that the standard extensions do not provide (plain wildcards, DISCRETE). */ + static final SimpleExtension.ExtensionCollection TEST_EXTENSIONS = loadTestExtensions(); + + final SimpleExtension.ExtensionCollection extensions = DefaultExtensionCatalog.DEFAULT_COLLECTION; + + private static SimpleExtension.ExtensionCollection loadTestExtensions() { + try { + return SimpleExtension.load( + Resources.toString( + Resources.getResource("extensions/binding_extensions.yaml"), StandardCharsets.UTF_8)); + } catch (IOException e) { + throw new UncheckedIOException(e); + } + } + + private SimpleExtension.AggregateFunctionVariant aggregate(String urn, String key) { + return extensions.getAggregateFunction(SimpleExtension.FunctionAnchor.of(urn, key)); + } + + private SimpleExtension.ScalarFunctionVariant scalar(String urn, String key) { + return extensions.getScalarFunction(SimpleExtension.FunctionAnchor.of(urn, key)); + } + + private ResolvedFunctionBinding binding( + SimpleExtension.Function declaration, List arguments) { + return ResolvedFunctionBinding.builder() + .anchor(declaration.getAnchor()) + .declaration(declaration) + .arguments(arguments) + .options(List.of()) + .build(); + } + + @Test + void resolvesConcreteIntegerSum() { + SimpleExtension.AggregateFunctionVariant sum = + aggregate(DefaultExtensionCatalog.FUNCTIONS_ARITHMETIC, "sum:i32"); + ResolvedFunctionBinding binding = + FunctionBindingResolver.resolveAndValidate( + sum, List.of(ResolvedArgument.value(R.I32)), List.of(), N.I64); + assertEquals(N.I64, binding.outputType()); + assertEquals(sum.getAnchor(), binding.anchor()); + } + + @Test + void resolvesDecimalSumWidth() { + SimpleExtension.AggregateFunctionVariant sum = + aggregate(DefaultExtensionCatalog.FUNCTIONS_ARITHMETIC_DECIMAL, "sum:dec"); + ResolvedFunctionBinding binding = + FunctionBindingResolver.resolveAndValidate( + sum, List.of(ResolvedArgument.value(R.decimal(10, 2))), List.of(), N.decimal(38, 2)); + assertEquals(N.decimal(38, 2), binding.outputType()); + } + + @Test + void appliesMirrorNullabilityForScalars() { + // add:i32_i32 has MIRROR nullability: a nullable operand makes the result nullable. + SimpleExtension.ScalarFunctionVariant add = + scalar(DefaultExtensionCatalog.FUNCTIONS_ARITHMETIC, "add:i32_i32"); + ResolvedFunctionBinding required = + FunctionBindingResolver.resolve( + add, List.of(ResolvedArgument.value(R.I32), ResolvedArgument.value(R.I32)), List.of()); + assertEquals(false, required.outputType().nullable()); + ResolvedFunctionBinding nullable = + FunctionBindingResolver.resolve( + add, List.of(ResolvedArgument.value(N.I32), ResolvedArgument.value(R.I32)), List.of()); + assertEquals(true, nullable.outputType().nullable()); + } + + @Test + void derivesWildcardReturnFromArgument() { + // any_value(any1) -> any1?: the wildcard return resolves to the argument type, made nullable. + SimpleExtension.AggregateFunctionVariant anyValue = + aggregate(DefaultExtensionCatalog.FUNCTIONS_AGGREGATE_GENERIC, "any_value:any"); + assertEquals( + N.I32, + FunctionBindingResolver.deriveOutputType(anyValue, List.of(ResolvedArgument.value(R.I32)))); + } + + @Test + void rejectsDeclaredOutputTypeThatDivergesFromDeclaration() { + SimpleExtension.AggregateFunctionVariant sum = + aggregate(DefaultExtensionCatalog.FUNCTIONS_ARITHMETIC, "sum:i32"); + // The standard declaration derives i64? for sum(i32); a plan declaring required i64 is invalid. + assertThrows( + InvalidFunctionBindingException.class, + () -> + FunctionBindingResolver.resolveAndValidate( + sum, List.of(ResolvedArgument.value(R.I32)), List.of(), R.I64)); + } + + @Test + void rejectsIncompatibleArgumentType() { + SimpleExtension.AggregateFunctionVariant sum = + aggregate(DefaultExtensionCatalog.FUNCTIONS_ARITHMETIC, "sum:i32"); + // Validation is opt-in: resolve() alone captures identity and does not check the signature. + assertThrows( + InvalidFunctionBindingException.class, + () -> + FunctionBindingResolver.resolveAndValidate( + sum, List.of(ResolvedArgument.value(R.FP64)), List.of(), N.I64)); + } + + @Test + void rejectsInconsistentNumberedWildcard() { + // nullif(any1, any1): the numbered wildcard must bind to one type; i32 + string is invalid. + SimpleExtension.ScalarFunctionVariant nullif = + scalar(DefaultExtensionCatalog.FUNCTIONS_COMPARISON, "nullif:any_any"); + assertThrows( + InvalidFunctionBindingException.class, + () -> + FunctionBindingResolver.resolveAndValidate( + nullif, + List.of(ResolvedArgument.value(R.I32), ResolvedArgument.value(R.STRING)), + List.of(), + R.I32)); + } + + @Test + void rejectsWrongArity() { + SimpleExtension.AggregateFunctionVariant sum = + aggregate(DefaultExtensionCatalog.FUNCTIONS_ARITHMETIC, "sum:i32"); + assertThrows( + InvalidFunctionBindingException.class, + () -> + FunctionBindingResolver.resolveAndValidate( + sum, + List.of(ResolvedArgument.value(R.I32), ResolvedArgument.value(R.I32)), + List.of(), + N.I64)); + } + + @Test + void outputTypeIsDerivedNotSettable() { + SimpleExtension.AggregateFunctionVariant sum = + aggregate(DefaultExtensionCatalog.FUNCTIONS_ARITHMETIC, "sum:i32"); + // There is no builder setter for outputType; it is always derived from the declaration. + assertEquals(N.I64, binding(sum, List.of(ResolvedArgument.value(R.I32))).outputType()); + } + + @Test + void identityIncludesArgumentsAndEnumValues() { + SimpleExtension.AggregateFunctionVariant sum = + aggregate(DefaultExtensionCatalog.FUNCTIONS_ARITHMETIC, "sum:i32"); + assertEquals( + binding(sum, List.of(ResolvedArgument.value(R.I32))), + binding(sum, List.of(ResolvedArgument.value(R.I32)))); + // Different value-argument types are distinct identities. + assertNotEquals( + binding(sum, List.of(ResolvedArgument.value(R.I32))), + binding(sum, List.of(ResolvedArgument.value(R.FP64)))); + // Different enum options are distinct identities: std_dev declares a leading "distribution" + // enum argument, so POPULATION and SAMPLE are two different functions on the same anchor. + SimpleExtension.AggregateFunctionVariant stdDev = stdDev(); + assertNotEquals( + binding(stdDev, stdDevArguments("POPULATION")), binding(stdDev, stdDevArguments("SAMPLE"))); + } + + @Test + void identityKeepsEnumOptionsAsSpelled() { + SimpleExtension.AggregateFunctionVariant stdDev = stdDev(); + // Identity is conservative rather than spec-exact: the plan's spelling is preserved verbatim + // (see acceptsEnumOptionRegardlessOfCase for the case-insensitive matching the spec requires), + // so two spellings of one option stay distinct instead of one of them being rewritten. + assertNotEquals( + binding(stdDev, stdDevArguments("SAMPLE")), binding(stdDev, stdDevArguments("sample"))); + // An unspecified option is not the same as any specified one either. + assertNotEquals( + binding( + stdDev, + List.of(ResolvedArgument.unspecifiedEnumOption(), ResolvedArgument.value(R.FP32))), + binding(stdDev, stdDevArguments("SAMPLE"))); + } + + @Test + void optionsAreKeptAsSpelledButMatchedCaseInsensitively() { + SimpleExtension.AggregateFunctionVariant count = + aggregate(DefaultExtensionCatalog.FUNCTIONS_AGGREGATE_GENERIC, "count:any"); + ResolvedFunctionBinding binding = + FunctionBindingResolver.resolveAndValidate( + count, + List.of(ResolvedArgument.value(R.I32)), + List.of(FunctionOption.builder().name("Overflow").addValues("ERROR").build()), + R.I64); + // The declaration spells the option "overflow" with values [SILENT, SATURATE, ERROR]; matching + // ignores case, but the binding still reports what the plan asked for. + assertEquals( + List.of(FunctionOption.builder().name("Overflow").addValues("ERROR").build()), + binding.options()); + } + + @Test + void acceptsEnumOptionRegardlessOfCase() { + SimpleExtension.AggregateFunctionVariant stdDev = stdDev(); + ResolvedFunctionBinding binding = + FunctionBindingResolver.resolveAndValidate( + stdDev, stdDevArguments("sample"), List.of(), N.FP32); + assertEquals(N.FP32, binding.outputType()); + } + + @Test + void rejectsUnknownEnumOption() { + SimpleExtension.AggregateFunctionVariant stdDev = stdDev(); + assertThrows( + InvalidFunctionBindingException.class, + () -> + FunctionBindingResolver.resolveAndValidate( + stdDev, stdDevArguments("MEDIAN"), List.of(), N.FP32)); + } + + @Test + void rejectsUnspecifiedRequiredEnumOption() { + SimpleExtension.AggregateFunctionVariant stdDev = stdDev(); + // std_dev's distribution argument is required, so leaving it unspecified does not satisfy the + // declaration — but it is reported as a missing option, not as an unknown one. + InvalidFunctionBindingException e = + assertThrows( + InvalidFunctionBindingException.class, + () -> + FunctionBindingResolver.resolveAndValidate( + stdDev, + List.of( + ResolvedArgument.unspecifiedEnumOption(), ResolvedArgument.value(R.FP32)), + List.of(), + N.FP32)); + assertTrue(e.getMessage().contains("left unspecified"), e.getMessage()); + } + + @Test + void plainWildcardArgumentsBindIndependently() { + // pair(any, any): plain wildcards are matched per argument, so two different types are valid. + SimpleExtension.ScalarFunctionVariant pair = testScalar("pair:any_any"); + FunctionBindingResolver.resolveAndValidate( + pair, + List.of(ResolvedArgument.value(R.I32), ResolvedArgument.value(R.STRING)), + List.of(), + R.BOOLEAN); + } + + @Test + void numberedWildcardArgumentsMustAgree() { + // same_pair(any1, any1): the numbered wildcard has to bind to one type across the invocation. + SimpleExtension.ScalarFunctionVariant samePair = testScalar("same_pair:any_any"); + assertThrows( + InvalidFunctionBindingException.class, + () -> + FunctionBindingResolver.resolveAndValidate( + samePair, + List.of(ResolvedArgument.value(R.I32), ResolvedArgument.value(R.STRING)), + List.of(), + R.BOOLEAN)); + } + + @Test + void discreteDeclarationFixesWildcardNullability() { + // discrete_wildcard(any1) -> any1 with DISCRETE nullability: the declared required argument + // does not accept a nullable one, and the derived type keeps the declared nullability. + SimpleExtension.ScalarFunctionVariant discrete = testScalar("discrete_wildcard:any"); + assertEquals( + R.I32, + FunctionBindingResolver.resolveAndValidate( + discrete, List.of(ResolvedArgument.value(R.I32)), List.of(), R.I32) + .outputType()); + assertThrows( + InvalidFunctionBindingException.class, + () -> + FunctionBindingResolver.resolveAndValidate( + discrete, List.of(ResolvedArgument.value(N.I32)), List.of(), N.I32)); + } + + @Test + void discreteDeclarationFixesNullabilityOfParameterizedTypes() { + // discrete_decimal(DECIMAL?) with DISCRETE nullability: the declared nullable argument is + // part of the signature, so a required decimal does not satisfy it either. + SimpleExtension.ScalarFunctionVariant discrete = testScalar("discrete_decimal:dec"); + FunctionBindingResolver.resolveAndValidate( + discrete, List.of(ResolvedArgument.value(N.decimal(10, 2))), List.of(), R.BOOLEAN); + assertThrows( + InvalidFunctionBindingException.class, + () -> + FunctionBindingResolver.resolveAndValidate( + discrete, List.of(ResolvedArgument.value(R.decimal(10, 2))), List.of(), R.BOOLEAN)); + } + + @Test + void variadicArgumentsBindBeyondTheDeclaredTail() { + // coalesce(any1...) declares its argument once but accepts it repeatedly; every actual + // argument has to agree on the wildcard, including the ones past the declared tail. + SimpleExtension.ScalarFunctionVariant coalesce = + scalar(DefaultExtensionCatalog.FUNCTIONS_COMPARISON, "coalesce:any"); + assertEquals( + R.I32, + FunctionBindingResolver.deriveOutputType( + coalesce, + List.of( + ResolvedArgument.value(R.I32), + ResolvedArgument.value(R.I32), + ResolvedArgument.value(R.I32)))); + assertThrows( + InvalidFunctionBindingException.class, + () -> + FunctionBindingResolver.deriveOutputType( + coalesce, + List.of( + ResolvedArgument.value(R.I32), + ResolvedArgument.value(R.I32), + ResolvedArgument.value(R.STRING)))); + } + + @Test + void sharedTypeParametersMustAgreeEvenWithAConcreteReturn() { + // decimal_pair(DECIMAL, DECIMAL) -> boolean: nothing about the return depends on P + // and S, but one signature cannot bind them to two different values. + SimpleExtension.ScalarFunctionVariant pair = testScalar("decimal_pair:dec_dec"); + FunctionBindingResolver.resolveAndValidate( + pair, + List.of(ResolvedArgument.value(R.decimal(10, 2)), ResolvedArgument.value(R.decimal(10, 2))), + List.of(), + R.BOOLEAN); + assertThrows( + InvalidFunctionBindingException.class, + () -> + FunctionBindingResolver.resolveAndValidate( + pair, + List.of( + ResolvedArgument.value(R.decimal(10, 2)), + ResolvedArgument.value(R.decimal(20, 5))), + List.of(), + R.BOOLEAN)); + } + + @Test + void optionsCannotChangeAfterTheBindingIsBuilt() { + SimpleExtension.AggregateFunctionVariant count = + aggregate(DefaultExtensionCatalog.FUNCTIONS_AGGREGATE_GENERIC, "count:any"); + List options = + new ArrayList<>( + List.of(FunctionOption.builder().name("overflow").addValues("error").build())); + ResolvedFunctionBinding binding = + FunctionBindingResolver.resolve(count, List.of(ResolvedArgument.value(R.I32)), options); + // Mutating the list the caller handed over must not change the binding's identity: the option + // list is copied and every FunctionOption in it is immutable in its own right. + options.add(FunctionOption.builder().name("overflow").addValues("silent").build()); + assertEquals(1, binding.options().size()); + // Lookup by name matches case-insensitively, as the spec requires, while the stored option + // keeps the plan's spelling. + assertEquals(Optional.of(List.of("error")), binding.option("OVERFLOW")); + } + + @Test + void repeatedOptionNamesAreKept() { + // A plan may repeat an option name; a map keyed by name would silently drop all but one. + SimpleExtension.AggregateFunctionVariant count = + aggregate(DefaultExtensionCatalog.FUNCTIONS_AGGREGATE_GENERIC, "count:any"); + List options = + List.of( + FunctionOption.builder().name("overflow").addValues("error").build(), + FunctionOption.builder().name("overflow").addValues("silent").build()); + ResolvedFunctionBinding binding = + FunctionBindingResolver.resolve(count, List.of(ResolvedArgument.value(R.I32)), options); + assertEquals(options, binding.options()); + } + + private SimpleExtension.AggregateFunctionVariant stdDev() { + return aggregate(DefaultExtensionCatalog.FUNCTIONS_ARITHMETIC, "std_dev:req_fp32"); + } + + private SimpleExtension.ScalarFunctionVariant testScalar(String key) { + return TEST_EXTENSIONS.getScalarFunction( + SimpleExtension.FunctionAnchor.of("extension:test:binding_extensions", key)); + } + + private List stdDevArguments(String distribution) { + return List.of(ResolvedArgument.enumOption(distribution), ResolvedArgument.value(R.FP32)); + } + + @Test + void identityDistinguishesExtensionsByAnchor() { + SimpleExtension.AggregateFunctionVariant sumI32 = + aggregate(DefaultExtensionCatalog.FUNCTIONS_ARITHMETIC, "sum:i32"); + SimpleExtension.AggregateFunctionVariant sumI64 = + aggregate(DefaultExtensionCatalog.FUNCTIONS_ARITHMETIC, "sum:i64"); + assertNotEquals( + binding(sumI32, List.of(ResolvedArgument.value(R.I32))), + binding(sumI64, List.of(ResolvedArgument.value(R.I64)))); + } +} diff --git a/core/src/test/java/io/substrait/extension/ResolvedAggregateBindingTest.java b/core/src/test/java/io/substrait/extension/ResolvedAggregateBindingTest.java new file mode 100644 index 000000000..07557db7b --- /dev/null +++ b/core/src/test/java/io/substrait/extension/ResolvedAggregateBindingTest.java @@ -0,0 +1,178 @@ +package io.substrait.extension; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertNotEquals; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import io.substrait.TestBase; +import io.substrait.expression.AggregateFunctionInvocation; +import io.substrait.expression.Expression; +import io.substrait.relation.Rel; +import io.substrait.type.Type; +import java.util.List; +import java.util.Optional; +import org.junit.jupiter.api.Test; + +class ResolvedAggregateBindingTest extends TestBase { + + /** The intermediate state {@code avg} accumulates: a sum and a count. */ + static final Type AVG_INTERMEDIATE = R.struct(R.I64, R.I64); + + @Test + void finalPhaseProducesTheDeclaredReturnType() { + ResolvedAggregateBinding binding = + ResolvedAggregateBinding.resolve(avg(Expression.AggregationPhase.INITIAL_TO_RESULT)); + assertEquals(N.I32, binding.outputType()); + } + + @Test + void partialPhaseProducesTheDeclaredIntermediateType() { + // A phase that stops at the intermediate state produces avg's accumulator, not its average. + for (Expression.AggregationPhase phase : + new Expression.AggregationPhase[] { + Expression.AggregationPhase.INITIAL_TO_INTERMEDIATE, + Expression.AggregationPhase.INTERMEDIATE_TO_INTERMEDIATE + }) { + assertEquals(AVG_INTERMEDIATE, ResolvedAggregateBinding.resolve(avg(phase)).outputType()); + } + } + + @Test + void validationComparesAgainstThePhaseOutputType() { + ResolvedAggregateBinding partial = + ResolvedAggregateBinding.resolve(avg(Expression.AggregationPhase.INITIAL_TO_INTERMEDIATE)); + // A partial aggregate that declares the intermediate type is valid ... + FunctionBindingResolver.validate(partial, AVG_INTERMEDIATE); + // ... while one declaring the final return type is not. + assertThrows( + InvalidFunctionBindingException.class, + () -> FunctionBindingResolver.validate(partial, N.I32)); + } + + @Test + void explicitIntermediateTypeWins() { + ResolvedAggregateBinding resolved = + ResolvedAggregateBinding.resolve(avg(Expression.AggregationPhase.INITIAL_TO_INTERMEDIATE)); + ResolvedAggregateBinding overridden = + ResolvedAggregateBinding.builder().from(resolved).intermediateType(R.struct(R.I64)).build(); + assertEquals(R.struct(R.I64), overridden.outputType()); + } + + @Test + void phaseIsPartOfTheIdentity() { + // Identity has to tell a partial aggregate apart from a full one: they produce different types + // and cannot be substituted for one another. + ResolvedAggregateBinding partial = + ResolvedAggregateBinding.resolve(avg(Expression.AggregationPhase.INITIAL_TO_INTERMEDIATE)); + ResolvedAggregateBinding full = + ResolvedAggregateBinding.resolve(avg(Expression.AggregationPhase.INITIAL_TO_RESULT)); + assertNotEquals(partial, full); + } + + @Test + void intermediateConsumingPhaseTakesTheIntermediateState() { + // sum:i32 accumulates into i64?, so its final phase consumes an i64? — validating that against + // the declaration's i32 argument would reject a perfectly valid partial-aggregation plan. + ResolvedAggregateBinding finalPhase = + sum(Expression.AggregationPhase.INTERMEDIATE_TO_RESULT, N.I64); + FunctionBindingResolver.validate(finalPhase, N.I64); + assertTrue(FunctionBindingResolver.matchesDeclaration(finalPhase)); + + // The same argument in an initial phase is not valid: there it is the declared i32 that is + // expected. + assertFalse( + FunctionBindingResolver.matchesDeclaration( + sum(Expression.AggregationPhase.INITIAL_TO_RESULT, N.I64))); + } + + @Test + void intermediateConsumingPhaseRejectsAnotherState() { + assertThrows( + InvalidFunctionBindingException.class, + () -> + FunctionBindingResolver.validate( + sum(Expression.AggregationPhase.INTERMEDIATE_TO_RESULT, R.I32), N.I64)); + } + + @Test + void bindingStopsMatchingWhenTheArgumentsChange() { + // What the reverse conversion relies on: a binding that no longer describes the call it is + // attached to (a rule replaced the i32 operand with an i64 one) must not be restored. + assertFalse( + FunctionBindingResolver.matchesDeclaration( + sum(Expression.AggregationPhase.INITIAL_TO_RESULT, R.I64))); + assertTrue( + FunctionBindingResolver.matchesDeclaration( + sum(Expression.AggregationPhase.INITIAL_TO_RESULT, R.I32))); + } + + @Test + void intermediateStateMustArriveAsAValueArgument() { + // A type argument carries no operand: conversion drops it and leaves the aggregate with no + // arguments at all, so accepting one here would validate a plan that cannot be converted. + ResolvedAggregateBinding binding = + ResolvedAggregateBinding.builder() + .function( + FunctionBindingResolver.resolve( + sumDeclaration(), List.of(ResolvedArgument.type(N.I64)), List.of())) + .phase(Expression.AggregationPhase.INTERMEDIATE_TO_RESULT) + .invocation(Expression.AggregationInvocation.ALL) + .intermediateType(Optional.empty()) + .build(); + assertFalse(FunctionBindingResolver.matchesDeclaration(binding)); + } + + @Test + void nonDecomposableDeclarationHasNoIntermediateState() { + // mode declares no decomposability, so there is no intermediate state for a partial phase to + // produce — deriving one fails closed rather than inventing it. + SimpleExtension.AggregateFunctionVariant mode = + extensions.getAggregateFunction( + SimpleExtension.FunctionAnchor.of( + DefaultExtensionCatalog.FUNCTIONS_ARITHMETIC, "mode:i32")); + ResolvedAggregateBinding binding = + ResolvedAggregateBinding.resolve( + AggregateFunctionInvocation.builder() + .declaration(mode) + .outputType(N.I32) + .aggregationPhase(Expression.AggregationPhase.INITIAL_TO_INTERMEDIATE) + .invocation(Expression.AggregationInvocation.ALL) + .addArguments(sb.i32(1)) + .build()); + assertThrows(InvalidFunctionBindingException.class, binding::outputType); + } + + private SimpleExtension.AggregateFunctionVariant sumDeclaration() { + return extensions.getAggregateFunction( + SimpleExtension.FunctionAnchor.of(DefaultExtensionCatalog.FUNCTIONS_ARITHMETIC, "sum:i32")); + } + + private ResolvedAggregateBinding sum(Expression.AggregationPhase phase, Type argumentType) { + SimpleExtension.AggregateFunctionVariant declaration = sumDeclaration(); + Rel input = sb.namedScan(List.of("t"), List.of("x"), List.of(argumentType)); + return ResolvedAggregateBinding.resolve( + AggregateFunctionInvocation.builder() + .declaration(declaration) + .outputType(N.I64) + .aggregationPhase(phase) + .invocation(Expression.AggregationInvocation.ALL) + .addArguments(sb.fieldReference(input, 0)) + .build()); + } + + private AggregateFunctionInvocation avg(Expression.AggregationPhase phase) { + SimpleExtension.AggregateFunctionVariant declaration = + extensions.getAggregateFunction( + SimpleExtension.FunctionAnchor.of( + DefaultExtensionCatalog.FUNCTIONS_ARITHMETIC, "avg:i32")); + return AggregateFunctionInvocation.builder() + .declaration(declaration) + .outputType(N.I32) + .aggregationPhase(phase) + .invocation(Expression.AggregationInvocation.ALL) + .addArguments(sb.i32(1)) + .build(); + } +} diff --git a/core/src/test/java/io/substrait/type/TypeExpressionEvaluatorTest.java b/core/src/test/java/io/substrait/type/TypeExpressionEvaluatorTest.java new file mode 100644 index 000000000..873e71396 --- /dev/null +++ b/core/src/test/java/io/substrait/type/TypeExpressionEvaluatorTest.java @@ -0,0 +1,80 @@ +package io.substrait.type; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertThrows; + +import io.substrait.extension.DefaultExtensionCatalog; +import io.substrait.extension.SimpleExtension; +import java.util.List; +import org.junit.jupiter.api.Test; + +class TypeExpressionEvaluatorTest { + + final SimpleExtension.ExtensionCollection extensions = DefaultExtensionCatalog.DEFAULT_COLLECTION; + + private SimpleExtension.AggregateFunctionVariant aggregate(String urn, String key) { + return extensions.getAggregateFunction(SimpleExtension.FunctionAnchor.of(urn, key)); + } + + private SimpleExtension.ScalarFunctionVariant scalar(String urn, String key) { + return extensions.getScalarFunction(SimpleExtension.FunctionAnchor.of(urn, key)); + } + + @Test + void concreteIntegerSumPassesThrough() { + SimpleExtension.AggregateFunctionVariant sum = + aggregate(DefaultExtensionCatalog.FUNCTIONS_ARITHMETIC, "sum:i32"); + // sum:i32 declares a concrete "i64?" return, independent of the argument. + assertEquals(TypeCreator.NULLABLE.I64, sum.resolveType(List.of(TypeCreator.REQUIRED.I32))); + } + + @Test + void decimalSumDerivesWidthAndScale() { + SimpleExtension.AggregateFunctionVariant sum = + aggregate(DefaultExtensionCatalog.FUNCTIONS_ARITHMETIC_DECIMAL, "sum:dec"); + // sum:dec declares DECIMAL -> DECIMAL?<38,S>: precision widens to 38, scale mirrors S. + assertEquals( + TypeCreator.NULLABLE.decimal(38, 2), + sum.resolveType(List.of(TypeCreator.REQUIRED.decimal(10, 2)))); + assertEquals( + TypeCreator.NULLABLE.decimal(38, 7), + sum.resolveType(List.of(TypeCreator.REQUIRED.decimal(20, 7)))); + } + + @Test + void rejectsInconsistentWildcardBinding() { + SimpleExtension.ScalarFunctionVariant nullif = + scalar(DefaultExtensionCatalog.FUNCTIONS_COMPARISON, "nullif:any_any"); + // nullif(any1, any1) -> any1?: binding the same numbered wildcard to two different types is a + // signature error, and must be rejected rather than silently resolved to the first binding. + assertThrows( + UnsupportedOperationException.class, + () -> nullif.resolveType(List.of(TypeCreator.REQUIRED.I32, TypeCreator.REQUIRED.STRING))); + } + + @Test + void variadicWildcardBindsEveryArgument() { + SimpleExtension.ScalarFunctionVariant coalesce = + scalar(DefaultExtensionCatalog.FUNCTIONS_COMPARISON, "coalesce:any"); + // coalesce(any1...) states its argument once and accepts it repeatedly, so an argument past + // the declared tail still has to agree with the wildcard. + assertThrows( + UnsupportedOperationException.class, + () -> coalesce.resolveType(List.of(TypeCreator.REQUIRED.I32, TypeCreator.REQUIRED.STRING))); + assertEquals( + TypeCreator.REQUIRED.I32, + coalesce.resolveType( + List.of(TypeCreator.REQUIRED.I32, TypeCreator.REQUIRED.I32, TypeCreator.REQUIRED.I32))); + } + + @Test + void wildcardBindingIgnoresNullability() { + SimpleExtension.ScalarFunctionVariant nullif = + scalar(DefaultExtensionCatalog.FUNCTIONS_COMPARISON, "nullif:any_any"); + // Nullability is not part of a wildcard's identity: i32 and i32? bind the same any1. The + // declared "any1?" return then decides the result's nullability. + assertEquals( + TypeCreator.NULLABLE.I32, + nullif.resolveType(List.of(TypeCreator.REQUIRED.I32, TypeCreator.NULLABLE.I32))); + } +} diff --git a/core/src/test/resources/extensions/binding_extensions.yaml b/core/src/test/resources/extensions/binding_extensions.yaml new file mode 100644 index 000000000..8041e92d9 --- /dev/null +++ b/core/src/test/resources/extensions/binding_extensions.yaml @@ -0,0 +1,56 @@ +%YAML 1.2 +--- +urn: extension:test:binding_extensions +scalar_functions: + - name: "pair" + description: >- + Two plain (unnumbered) wildcards. Each argument binds independently, so the two may + be of different types. + impls: + - args: + - name: x + value: any + - name: y + value: any + return: boolean + - name: "same_pair" + description: >- + Two occurrences of one numbered wildcard, which must agree on a single type. + impls: + - args: + - name: x + value: any1 + - name: y + value: any1 + return: boolean + - name: "decimal_pair" + description: >- + Two decimals sharing one set of type parameters, with a concrete return type. The shared + parameters must agree even though nothing about the return depends on them. + impls: + - args: + - name: x + value: DECIMAL + - name: y + value: DECIMAL + return: boolean + - name: "discrete_decimal" + description: >- + A DISCRETE signature over a parameterized decimal, so the declared argument nullability is + part of the signature just as it is for a concrete type. + impls: + - args: + - name: x + value: "DECIMAL?" + nullability: DISCRETE + return: boolean + - name: "discrete_wildcard" + description: >- + A DISCRETE signature over a numbered wildcard, so the declared argument nullability is + part of the signature. + impls: + - args: + - name: x + value: any1 + nullability: DISCRETE + return: any1 From 4d7aa8f4dd097b4dd015734cd7ad30345a57a2fd Mon Sep 17 00:00:00 2001 From: Roman Kondakov Date: Sun, 26 Jul 2026 16:47:30 +0700 Subject: [PATCH 2/2] feat(isthmus)!: preserve aggregate output types and semantics through Calcite MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit A Substrait aggregate used to lose part of itself on the way to Calcite. The output type the plan declared was handed to AggregateCall.create, but RelBuilder re-infers a call's type from its operator, so the plan's type was replaced by Calcite's inference — sum(DECIMAL(10,2)), which the standard sum:dec declaration types as DECIMAL(38,2), came back as the argument's own width. Function options and the aggregation phase have no place in a Calcite aggregate call at all, so they were dropped, and converting back re-matched the operator against the extension catalog: among equally-shaped declarations it could only guess, it always produced a full INITIAL_TO_RESULT aggregation without options, and for a phase consuming an intermediate state it could not match at all, because the call's operands are then accumulator state rather than the declaration's arguments. Each converted measure now resolves a ResolvedAggregateBinding. Where Calcite cannot hold what the binding carries — the chosen output type differs from its inference, or the invocation has options or a phase other than INITIAL_TO_RESULT — the operator is wrapped so that both the binding and the type travel with it, and AggregateFunctionConverter rebuilds the invocation from that binding instead of re-matching. The binding records the plan as it was converted, so it is dropped again when a planner rule has since changed the call's arguments, and DISTINCT is always read off the Calcite call because a rule may legitimately have removed a redundant one. Two further conversion fixes come with it. Measures that are equal to Calcite but carry different types no longer collapse into a single column: AggregateCall equality ignores the stored type and RelBuilder deduplicates, so deduplication is now switched off for exactly those aggregates. And an aggregate with no groupings is built as one empty grouping set rather than none at all, so its measures are no longer re-inferred as if there were no empty group; converting such a plan back therefore spells the same global aggregation the way Calcite spells it, as a single empty grouping. BREAKING CHANGE: converting a Substrait aggregate to Calcite now preserves the plan's declared output type instead of Calcite's inferred one, and the resulting AggregateCall may carry a wrapper operator rather than the plain SqlAggFunction. Consumers that compare the operator by identity should call AggregateFunctions.unwrapBound first; rollup and splitting are disabled for a wrapped call, since both re-infer the type from the underlying function and would discard the preserved one. Pass AggregateConversion with OutputTypeSource.CALCITE_INFERENCE to SubstraitToCalcite to restore the previous behavior. A ConverterProvider subclass that overrides getSubstraitRelNodeConverter(RelBuilder) should override the new getSubstraitRelNodeConverter(RelBuilder, AggregateConversion) as well, otherwise its customization is skipped for a non-default configuration. Co-Authored-By: Claude Opus 5 (1M context) --- .../isthmus/AggregateConversion.java | 109 ++++ .../substrait/isthmus/AggregateFunctions.java | 243 ++++++++ .../substrait/isthmus/ConverterProvider.java | 25 +- .../isthmus/SubstraitRelNodeConverter.java | 118 +++- .../substrait/isthmus/SubstraitToCalcite.java | 24 +- .../AggregateFunctionConverter.java | 151 ++++- .../SubstraitRelNodeConverterTest.java | 564 +++++++++++++++++- 7 files changed, 1207 insertions(+), 27 deletions(-) create mode 100644 isthmus/src/main/java/io/substrait/isthmus/AggregateConversion.java diff --git a/isthmus/src/main/java/io/substrait/isthmus/AggregateConversion.java b/isthmus/src/main/java/io/substrait/isthmus/AggregateConversion.java new file mode 100644 index 000000000..604b1813a --- /dev/null +++ b/isthmus/src/main/java/io/substrait/isthmus/AggregateConversion.java @@ -0,0 +1,109 @@ +package io.substrait.isthmus; + +import java.util.Objects; + +/** + * Configures how the output type of a converted Substrait aggregate is chosen and validated. + * + *

The two settings are independent: {@link OutputTypeSource} decides which type the resulting + * Calcite {@code AggregateCall} carries, and {@link FunctionBindingValidation} decides whether the + * plan-declared type is checked against the extension declaration. + * + *

The {@link #DEFAULT} is {@link OutputTypeSource#PLAN_OUTPUT} with {@link + * FunctionBindingValidation#NONE}: the plan's declared output type is preserved (so conversion + * never silently changes a type) without asserting that the plan is spec-compliant. + */ +public final class AggregateConversion { + + /** Where the Calcite output type of a converted aggregate comes from. */ + public enum OutputTypeSource { + /** Use the type Calcite infers for the operator (ignores the plan's declared type). */ + CALCITE_INFERENCE, + /** Preserve the {@code AggregateFunction.output_type} declared by the plan. */ + PLAN_OUTPUT + } + + /** Whether the plan-declared output type is validated against the extension declaration. */ + public enum FunctionBindingValidation { + /** Do not validate the declared output type. */ + NONE, + /** + * Require the declared output type to match the type derived from the extension declaration. + * + *

Type derivation is fail-closed: a function whose return expression the derivation does not + * yet support (arithmetic derivations, {@code if/then}, return programs) is rejected rather + * than assumed valid, so this mode is not adoptable for plans that use such functions. + */ + EXTENSION_DECLARATION + } + + /** Preserve the plan's output type without validating it against the declaration. */ + public static final AggregateConversion DEFAULT = + new AggregateConversion(OutputTypeSource.PLAN_OUTPUT, FunctionBindingValidation.NONE); + + private final OutputTypeSource outputTypeSource; + private final FunctionBindingValidation bindingValidation; + + /** + * Creates a configuration. + * + * @param outputTypeSource where the Calcite output type comes from + * @param bindingValidation whether the declared type is validated against the declaration + */ + public AggregateConversion( + OutputTypeSource outputTypeSource, FunctionBindingValidation bindingValidation) { + this.outputTypeSource = Objects.requireNonNull(outputTypeSource); + this.bindingValidation = Objects.requireNonNull(bindingValidation); + } + + /** + * Returns the output-type source. + * + * @return the output-type source + */ + public OutputTypeSource outputTypeSource() { + return outputTypeSource; + } + + /** + * Returns the binding-validation policy. + * + * @return the binding-validation policy + */ + public FunctionBindingValidation bindingValidation() { + return bindingValidation; + } + + /** + * Returns whether this configuration is equivalent to {@link #DEFAULT}. + * + * @return {@code true} if this preserves the plan output type without validation + */ + public boolean isDefault() { + return outputTypeSource == OutputTypeSource.PLAN_OUTPUT + && bindingValidation == FunctionBindingValidation.NONE; + } + + @Override + public boolean equals(Object o) { + if (this == o) { + return true; + } + if (!(o instanceof AggregateConversion)) { + return false; + } + AggregateConversion other = (AggregateConversion) o; + return outputTypeSource == other.outputTypeSource + && bindingValidation == other.bindingValidation; + } + + @Override + public int hashCode() { + return Objects.hash(outputTypeSource, bindingValidation); + } + + @Override + public String toString() { + return "AggregateConversion[" + outputTypeSource + ", " + bindingValidation + "]"; + } +} diff --git a/isthmus/src/main/java/io/substrait/isthmus/AggregateFunctions.java b/isthmus/src/main/java/io/substrait/isthmus/AggregateFunctions.java index 8ea324948..fafa5514e 100644 --- a/isthmus/src/main/java/io/substrait/isthmus/AggregateFunctions.java +++ b/isthmus/src/main/java/io/substrait/isthmus/AggregateFunctions.java @@ -1,15 +1,25 @@ package io.substrait.isthmus; +import io.substrait.extension.ResolvedAggregateBinding; +import java.util.Objects; import java.util.Optional; +import org.apache.calcite.rel.core.AggregateCall; import org.apache.calcite.rel.type.RelDataType; import org.apache.calcite.sql.SqlAggFunction; +import org.apache.calcite.sql.SqlCall; import org.apache.calcite.sql.SqlKind; +import org.apache.calcite.sql.SqlOperandCountRange; import org.apache.calcite.sql.SqlOperatorBinding; +import org.apache.calcite.sql.SqlSplittableAggFunction; +import org.apache.calcite.sql.SqlSyntax; +import org.apache.calcite.sql.SqlWriter; import org.apache.calcite.sql.fun.SqlAvgAggFunction; import org.apache.calcite.sql.fun.SqlMinMaxAggFunction; import org.apache.calcite.sql.fun.SqlSumAggFunction; import org.apache.calcite.sql.fun.SqlSumEmptyIsZeroAggFunction; import org.apache.calcite.sql.type.ReturnTypes; +import org.apache.calcite.util.Optionality; +import org.checkerframework.checker.nullness.qual.Nullable; /** * Provides Substrait-specific variants of Calcite aggregate functions to ensure type inference @@ -59,6 +69,100 @@ public class AggregateFunctions { /** Substrait-specific SUM0 aggregate function (non-null BIGINT return type). */ public static final SqlAggFunction SUM0 = new SubstraitSumEmptyIsZeroAggFunction(); + /** + * Wraps a configured aggregate function so it carries a resolved Substrait binding and infers the + * given output type. + * + *

This is a transport adapter: Calcite validates an {@code AggregateCall}'s stored type + * against the function's inference rule and {@code RelBuilder} re-infers it when copying a call, + * so the chosen output type has to travel with the function. The operator's identity is the + * binding alone (see {@link #boundBinding}); the carried type is readable via {@link + * #declaredOutputType}. + * + * @param aggFunction configured aggregate function to adapt + * @param binding resolved Substrait aggregate binding + * @param outputType Calcite output type the wrapper should infer (e.g. the plan-declared type) + * @return an aggregate function that carries the binding and infers the given output type + */ + public static SqlAggFunction bind( + SqlAggFunction aggFunction, ResolvedAggregateBinding binding, RelDataType outputType) { + return new BoundSqlAggFunction(unwrapBound(aggFunction), binding, outputType); + } + + /** + * Returns the configured aggregate function underlying a {@link #bind bound} function, or the + * input unchanged. + * + * @param aggFunction aggregate function to inspect + * @return the underlying configured aggregate function + */ + public static SqlAggFunction unwrapBound(SqlAggFunction aggFunction) { + return aggFunction instanceof BoundSqlAggFunction + ? ((BoundSqlAggFunction) aggFunction).delegate + : aggFunction; + } + + /** + * Returns the given call with a {@link #bind bound} function replaced by the configured function + * it wraps, keeping the type the call already carries. + * + *

Use this before executing a converted plan. Calcite dispatches aggregate + * implementations by operator identity ({@code RexImpTable} keyed on the {@link SqlAggFunction}, + * which compares class, name and kind), so a bound function has no implementor and Enumerable or + * Bindable execution rejects the aggregate. The same already holds for the Substrait variants in + * this class — {@link #SUM}, {@link #MIN}, {@link #AVG} and the rest are distinct classes from + * the stock operators — so executability was never a property of a converted aggregate plan; + * unwrapping restores it for the calls that can have it, at the cost of the carried binding. + * + * @param call aggregate call to unwrap + * @return the call on the underlying configured function, or the call unchanged if it carries no + * binding + */ + public static AggregateCall unwrapBound(AggregateCall call) { + SqlAggFunction unwrapped = unwrapBound(call.getAggregation()); + if (unwrapped == call.getAggregation()) { + return call; + } + return AggregateCall.create( + unwrapped, + call.isDistinct(), + call.isApproximate(), + call.ignoreNulls(), + call.rexList, + call.getArgList(), + call.filterArg, + call.distinctKeys, + call.getCollation(), + call.getType(), + call.getName()); + } + + /** + * Returns the resolved Substrait aggregate binding carried by the given function, if any. + * + * @param aggFunction aggregate function to inspect + * @return the resolved binding, or empty if the function does not carry one + */ + public static Optional boundBinding(SqlAggFunction aggFunction) { + return aggFunction instanceof BoundSqlAggFunction + ? Optional.of(((BoundSqlAggFunction) aggFunction).binding) + : Optional.empty(); + } + + /** + * Returns the output type a {@link #bind bound} function carries, if any. This is the type the + * wrapper infers — under {@code PLAN_OUTPUT} the type declared by the plan, which may differ from + * the one derived from the extension declaration when the plan is not validated against it. + * + * @param aggFunction aggregate function to inspect + * @return the carried output type, or empty if the function does not carry one + */ + public static Optional declaredOutputType(SqlAggFunction aggFunction) { + return aggFunction instanceof BoundSqlAggFunction + ? Optional.of(((BoundSqlAggFunction) aggFunction).outputType) + : Optional.empty(); + } + /** * Converts default Calcite aggregate functions to Substrait-specific variants when needed. * @@ -154,4 +258,143 @@ public RelDataType inferReturnType(SqlOperatorBinding opBinding) { return ReturnTypes.BIGINT.inferReturnType(opBinding); } } + + /** + * Aggregate function that carries a resolved Substrait binding and infers a fixed output type. + * All non-type behavior is delegated to the configured function. + * + *

Identity is the {@link ResolvedAggregateBinding} alone, so that operator equality stays a + * question of which function is being called and never of which type it produces: two + * calls of the same Substrait function remain the same operator even when their plan-declared + * types differ, while calls resolving to different functions (different extension, options or + * enum arguments) stay distinct. Keeping two such calls apart in a single aggregate — {@code + * AggregateCall} equality ignores the stored type, so {@code RelBuilder} would otherwise + * deduplicate them — is the converter's job, not this operator's. + * + *

A bound function is a plan-interchange device, not an executable operator: Calcite looks up + * aggregate implementations by operator identity, so a plan containing one cannot be run by the + * Enumerable or Bindable convention. {@link #unwrapBound(AggregateCall)} trades the binding back + * for executability. (The Substrait variants above are already distinct classes from the stock + * operators, so the same restriction applies to them.) + * + *

Type-rederiving transformations that expose a {@link SqlAggFunction} hook are disabled until + * they become binding-aware: rollup ({@link #getRollup()}) and splitting ({@link + * SqlSplittableAggFunction}) both recompute the transformed call's type from the underlying + * function, which would discard the carried type. Rules that instead match directly on {@link + * org.apache.calcite.sql.SqlKind} (e.g. {@code AGGREGATE_REDUCE_FUNCTIONS}) have no such hook and + * are retained on the delegate's kind; a caller optimizing a converted plan with a bound + * aggregate should either avoid such rules or make them binding-aware, otherwise they may rewrite + * the call and drop the carried type. + */ + private static final class BoundSqlAggFunction extends SqlAggFunction { + private final SqlAggFunction delegate; + private final ResolvedAggregateBinding binding; + private final RelDataType outputType; + + private BoundSqlAggFunction( + SqlAggFunction delegate, ResolvedAggregateBinding binding, RelDataType outputType) { + super( + delegate.getName(), + delegate.getSqlIdentifier(), + delegate.getKind(), + ReturnTypes.explicit(outputType), + delegate.getOperandTypeInference(), + delegate.getOperandTypeChecker(), + delegate.getFunctionType(), + delegate.requiresOrder(), + delegate.requiresOver(), + delegate.requiresGroupOrder()); + this.delegate = delegate; + this.binding = binding; + this.outputType = outputType; + } + + @Override + public SqlOperandCountRange getOperandCountRange() { + return delegate.getOperandCountRange(); + } + + @Override + public SqlSyntax getSyntax() { + return delegate.getSyntax(); + } + + @Override + public void unparse(SqlWriter writer, SqlCall call, int leftPrec, int rightPrec) { + delegate.unparse(writer, call, leftPrec, rightPrec); + } + + @Override + public Optionality getDistinctOptionality() { + return delegate.getDistinctOptionality(); + } + + @Override + public boolean allowsFilter() { + return delegate.allowsFilter(); + } + + @Override + public boolean allowsNullTreatment() { + return delegate.allowsNullTreatment(); + } + + @Override + public boolean skipsNullInputs() { + return delegate.skipsNullInputs(); + } + + @Override + public @Nullable SqlAggFunction getRollup() { + // A rollup re-infers the top aggregate's return type from the underlying function, which + // would discard the canonical Substrait output type. Opt out until rollup is binding-aware. + return null; + } + + @Override + public boolean isPercentile() { + return delegate.isPercentile(); + } + + @Override + public boolean allowsFraming() { + return delegate.allowsFraming(); + } + + @Override + public @Nullable T unwrap(Class clazz) { + if (clazz == SqlSplittableAggFunction.class) { + // Splitting (e.g. pushing an aggregate through a join) re-infers the split calls' types + // from the underlying function, discarding the canonical type. Opt out until it is + // binding-aware. + return null; + } + if (clazz.isInstance(binding)) { + // Expose the resolved binding so consumers (and the reverse converter) can recover the + // original Substrait function without re-matching on the Calcite operator. + return clazz.cast(binding); + } + T unwrapped = super.unwrap(clazz); + return unwrapped != null ? unwrapped : delegate.unwrap(clazz); + } + + @Override + public boolean equals(@Nullable Object obj) { + if (obj == this) { + return true; + } + if (!(obj instanceof BoundSqlAggFunction)) { + return false; + } + BoundSqlAggFunction other = (BoundSqlAggFunction) obj; + // Identity is the semantic binding only: the carried output type is deliberately excluded so + // that matching on the operator stays type-agnostic. + return binding.equals(other.binding); + } + + @Override + public int hashCode() { + return Objects.hash(getClass(), binding); + } + } } diff --git a/isthmus/src/main/java/io/substrait/isthmus/ConverterProvider.java b/isthmus/src/main/java/io/substrait/isthmus/ConverterProvider.java index f8b4b130e..864d5e2fc 100644 --- a/isthmus/src/main/java/io/substrait/isthmus/ConverterProvider.java +++ b/isthmus/src/main/java/io/substrait/isthmus/ConverterProvider.java @@ -309,11 +309,34 @@ public Function getSchemaResolver() { * A {@link SubstraitRelNodeConverter} is used when converting from Substrait {@link Rel}s to * Calcite {@link org.apache.calcite.rel.RelNode}s. * + *

This overload is used for the default {@link AggregateConversion} only; a subclass that + * overrides it to customize conversion should override {@link + * #getSubstraitRelNodeConverter(RelBuilder, AggregateConversion)} as well, otherwise the + * customization is skipped whenever a caller asks for a non-default configuration. + * * @param relBuilder the RelBuilder to use for creating Calcite RelNodes * @return a new SubstraitRelNodeConverter instance */ public SubstraitRelNodeConverter getSubstraitRelNodeConverter(RelBuilder relBuilder) { - return new SubstraitRelNodeConverter(relBuilder, this); + return getSubstraitRelNodeConverter(relBuilder, AggregateConversion.DEFAULT); + } + + /** + * A {@link SubstraitRelNodeConverter} is used when converting from Substrait {@link Rel}s to + * Calcite {@link org.apache.calcite.rel.RelNode}s. + * + *

{@link SubstraitToCalcite} calls this overload only for a non-default {@link + * AggregateConversion}, and {@link #getSubstraitRelNodeConverter(RelBuilder)} otherwise — that + * dispatch keeps subclasses which only override the long-standing single-argument factory + * working. A subclass that customizes conversion should therefore override both. + * + * @param relBuilder the RelBuilder to use for creating Calcite RelNodes + * @param aggregateConversion controls how aggregate output types are chosen and validated + * @return a new SubstraitRelNodeConverter instance + */ + public SubstraitRelNodeConverter getSubstraitRelNodeConverter( + RelBuilder relBuilder, AggregateConversion aggregateConversion) { + return new SubstraitRelNodeConverter(relBuilder, this, aggregateConversion); } /** diff --git a/isthmus/src/main/java/io/substrait/isthmus/SubstraitRelNodeConverter.java b/isthmus/src/main/java/io/substrait/isthmus/SubstraitRelNodeConverter.java index 26f063c7b..7f8e49641 100644 --- a/isthmus/src/main/java/io/substrait/isthmus/SubstraitRelNodeConverter.java +++ b/isthmus/src/main/java/io/substrait/isthmus/SubstraitRelNodeConverter.java @@ -4,6 +4,8 @@ import io.substrait.expression.Expression; import io.substrait.expression.Expression.SortDirection; import io.substrait.expression.FunctionArg; +import io.substrait.extension.FunctionBindingResolver; +import io.substrait.extension.ResolvedAggregateBinding; import io.substrait.extension.SimpleExtension; import io.substrait.isthmus.calcite.rel.CreateTable; import io.substrait.isthmus.calcite.rel.CreateView; @@ -111,6 +113,9 @@ public class SubstraitRelNodeConverter /** Type converter to translate between Calcite and Substrait type systems. */ private final TypeConverter typeConverter; + /** Controls how aggregate output types are chosen and validated. */ + private final AggregateConversion aggregateConversion; + /** * Creates a new SubstraitRelNodeConverter with the specified extensions, type factory, and * relation builder. @@ -136,6 +141,20 @@ public SubstraitRelNodeConverter( * @param converterProvider the converter provider containing configuration and converters */ public SubstraitRelNodeConverter(RelBuilder relBuilder, ConverterProvider converterProvider) { + this(relBuilder, converterProvider, AggregateConversion.DEFAULT); + } + + /** + * Creates a new SubstraitRelNodeConverter with an explicit aggregate-conversion configuration. + * + * @param relBuilder the Calcite relation builder + * @param converterProvider the converter provider containing configuration and converters + * @param aggregateConversion controls how aggregate output types are chosen and validated + */ + public SubstraitRelNodeConverter( + RelBuilder relBuilder, + ConverterProvider converterProvider, + AggregateConversion aggregateConversion) { this.typeFactory = converterProvider.getTypeFactory(); this.typeConverter = converterProvider.getTypeConverter(); this.relBuilder = relBuilder; @@ -143,6 +162,7 @@ public SubstraitRelNodeConverter(RelBuilder relBuilder, ConverterProvider conver this.scalarFunctionConverter = converterProvider.getScalarFunctionConverter(); this.aggregateFunctionConverter = converterProvider.getAggregateFunctionConverter(); this.expressionRexConverter = converterProvider.getExpressionRexConverter(this); + this.aggregateConversion = aggregateConversion; } /** @@ -320,11 +340,23 @@ public RelNode visit(Aggregate aggregate, Context context) throws RuntimeExcepti .collect(java.util.stream.Collectors.toList()); List groupExprs = groupExprLists.stream().flatMap(Collection::stream).collect(Collectors.toList()); - RelBuilder.GroupKey groupKey = relBuilder.groupKey(groupExprs, groupExprLists); + // An aggregate with no groupings at all is a global aggregation, i.e. a single empty grouping + // set. Handing RelBuilder an empty list of grouping sets instead would leave the Calcite + // aggregate with no grouping set at all, and RelBuilder would then re-infer every measure as if + // there were no empty group — which is the opposite of what a global aggregation means. + boolean globalAggregation = groupExprLists.isEmpty(); + RelBuilder.GroupKey groupKey = + globalAggregation + ? relBuilder.groupKey(groupExprs) + : relBuilder.groupKey(groupExprs, groupExprLists); + // Mirrors how RelBuilder derives an aggregate call's hasEmptyGroup from the grouping sets it + // builds out of this group key; the two must agree or the inferred type below is not the one + // the call ends up with. + boolean hasEmptyGroup = globalAggregation || groupExprLists.stream().anyMatch(List::isEmpty); List aggregateCalls = aggregate.getMeasures().stream() - .map(measure -> fromMeasure(measure, context)) + .map(measure -> fromMeasure(measure, context, child, hasEmptyGroup)) .collect(java.util.stream.Collectors.toList()); Optional remap = aggregate.getRemap(); @@ -361,11 +393,49 @@ public RelNode visit(Aggregate aggregate, Context context) throws RuntimeExcepti } } - RelNode node = relBuilder.push(child).aggregate(groupKey, aggregateCalls).build(); + // RelBuilder deduplicates equal aggregate calls, and AggregateCall equality ignores the stored + // type: two measures of the same function that differ only by their declared output type would + // collapse into one column. Opt out of deduplication for exactly those aggregates — narrowly, + // because transform() yields a plain RelBuilder carrying the same cluster and schema but not a + // custom builder subclass. + RelBuilder aggregateBuilder = + hasTypeDistinctDuplicates(aggregateCalls) + ? relBuilder.transform(config -> config.withDedupAggregateCalls(false)) + : relBuilder; + + RelNode node = aggregateBuilder.push(child).aggregate(groupKey, aggregateCalls).build(); return applyRemap(node, remap); } - private AggregateCall fromMeasure(Aggregate.Measure measure, Context context) { + /** + * Returns whether the binding carries semantics no stock Calcite aggregate operator can express, + * so that two invocations differing only in those would be indistinguishable once converted: the + * function's options and any phase other than a full initial-to-result aggregation. Enum + * arguments are excluded — they select the Calcite operator itself (e.g. {@code STDDEV_POP} vs + * {@code STDDEV_SAMP}) and are recovered from it on the way back. + */ + private static boolean carriesOpaqueSemantics(ResolvedAggregateBinding binding) { + return !binding.function().options().isEmpty() + || binding.phase() != Expression.AggregationPhase.INITIAL_TO_RESULT; + } + + /** + * Returns whether two aggregate calls are equal to Calcite yet carry different types, i.e. + * whether deduplicating them would silently give one of them the other's type. + */ + private static boolean hasTypeDistinctDuplicates(List aggregateCalls) { + Map typesByCall = new HashMap<>(); + for (AggregateCall call : aggregateCalls) { + RelDataType existing = typesByCall.putIfAbsent(call, call.getType()); + if (existing != null && !existing.equals(call.getType())) { + return true; + } + } + return false; + } + + private AggregateCall fromMeasure( + Aggregate.Measure measure, Context context, RelNode input, boolean hasEmptyGroup) { List eArgs = measure.getFunction().arguments(); // Only value (Expression) arguments map to Calcite aggregate operands. Enum arguments such as // the std_dev/variance "distribution" are used to disambiguate the operator, not as operands. @@ -403,7 +473,14 @@ private AggregateCall fromMeasure(Aggregate.Measure measure, Context context) { measure.getFunction().invocation().equals(Expression.AggregationInvocation.DISTINCT); SqlAggFunction aggFunction; - RelDataType returnType = typeConverter.toCalcite(typeFactory, measure.getFunction().getType()); + // Resolve the Substrait binding (semantic identity). Validating the signature, options and + // declared output type against the extension declaration is a separate, opt-in concern + // (EXTENSION_DECLARATION). + ResolvedAggregateBinding binding = ResolvedAggregateBinding.resolve(measure.getFunction()); + if (aggregateConversion.bindingValidation() + == AggregateConversion.FunctionBindingValidation.EXTENSION_DECLARATION) { + FunctionBindingResolver.validate(binding, measure.getFunction().getType()); + } if (operator.get() instanceof SqlAggFunction) { aggFunction = (SqlAggFunction) operator.get(); @@ -430,6 +507,37 @@ private AggregateCall fromMeasure(Aggregate.Measure measure, Context context) { .collect(Collectors.toList())); } + AggregateCall inferredCall = + AggregateCall.create( + aggFunction, + distinct, + false, + false, + Collections.emptyList(), + argIndex, + filterArg, + null, + relCollation, + hasEmptyGroup, + input, + null, + null); + RelDataType inferredType = inferredCall.getType(); + boolean preservePlanType = + aggregateConversion.outputTypeSource() == AggregateConversion.OutputTypeSource.PLAN_OUTPUT; + // Convert the declared type only where it is used: CALCITE_INFERENCE promises to ignore it, and + // a type this converter cannot represent must not fail a conversion that never needed it. + RelDataType returnType = + preservePlanType + ? typeConverter.toCalcite(typeFactory, measure.getFunction().getType()) + : inferredType; + if ((preservePlanType && !returnType.equals(inferredType)) || carriesOpaqueSemantics(binding)) { + // Calcite would either re-infer a different type or lose semantics its operator cannot + // express; carry the binding and the chosen type on a transport wrapper so both survive + // Calcite's re-inference (RelBuilder / planner rules) and its call deduplication. + aggFunction = AggregateFunctions.bind(aggFunction, binding, returnType); + } + return AggregateCall.create( aggFunction, distinct, diff --git a/isthmus/src/main/java/io/substrait/isthmus/SubstraitToCalcite.java b/isthmus/src/main/java/io/substrait/isthmus/SubstraitToCalcite.java index bea3c318b..8a0ef7966 100644 --- a/isthmus/src/main/java/io/substrait/isthmus/SubstraitToCalcite.java +++ b/isthmus/src/main/java/io/substrait/isthmus/SubstraitToCalcite.java @@ -35,6 +35,9 @@ public class SubstraitToCalcite { /** The converter provider containing configuration. */ protected ConverterProvider converterProvider; + /** Controls how aggregate output types are chosen and validated. */ + protected final AggregateConversion aggregateConversion; + /** * Creates a Substrait-to-Calcite converter using configuration from the provider. * @@ -52,9 +55,24 @@ public SubstraitToCalcite(ConverterProvider converterProvider) { */ public SubstraitToCalcite( ConverterProvider converterProvider, Prepare.CatalogReader catalogReader) { + this(converterProvider, catalogReader, AggregateConversion.DEFAULT); + } + + /** + * Creates a Substrait-to-Calcite converter with an explicit aggregate-conversion configuration. + * + * @param converterProvider the converter provider containing configuration and converters + * @param catalogReader Calcite catalog reader for schema resolution + * @param aggregateConversion controls how aggregate output types are chosen and validated + */ + public SubstraitToCalcite( + ConverterProvider converterProvider, + Prepare.CatalogReader catalogReader, + AggregateConversion aggregateConversion) { this.converterProvider = converterProvider; this.typeFactory = converterProvider.getTypeFactory(); this.catalogReader = catalogReader; + this.aggregateConversion = aggregateConversion; } /** @@ -76,8 +94,12 @@ public RelNode convert(Rel rel) { CalciteSchema rootSchema = converterProvider.getSchemaResolver().apply(rel); relBuilder = converterProvider.getRelBuilder(rootSchema); } + // Under the default policy, dispatch through the long-standing single-argument factory so a + // ConverterProvider subclass overriding it still customizes conversion. SubstraitRelNodeConverter converter = - converterProvider.getSubstraitRelNodeConverter(relBuilder); + aggregateConversion.isDefault() + ? converterProvider.getSubstraitRelNodeConverter(relBuilder) + : converterProvider.getSubstraitRelNodeConverter(relBuilder, aggregateConversion); return rel.accept(converter, Context.newContext()); } diff --git a/isthmus/src/main/java/io/substrait/isthmus/expression/AggregateFunctionConverter.java b/isthmus/src/main/java/io/substrait/isthmus/expression/AggregateFunctionConverter.java index 0b9b7e63d..b06288615 100644 --- a/isthmus/src/main/java/io/substrait/isthmus/expression/AggregateFunctionConverter.java +++ b/isthmus/src/main/java/io/substrait/isthmus/expression/AggregateFunctionConverter.java @@ -2,10 +2,12 @@ import com.google.common.collect.ImmutableList; import io.substrait.expression.AggregateFunctionInvocation; +import io.substrait.expression.EnumArg; import io.substrait.expression.Expression; -import io.substrait.expression.ExpressionCreator; import io.substrait.expression.FunctionArg; import io.substrait.expression.StatisticalDistribution; +import io.substrait.extension.ResolvedAggregateBinding; +import io.substrait.extension.ResolvedArgument; import io.substrait.extension.SimpleExtension; import io.substrait.isthmus.AggregateFunctions; import io.substrait.isthmus.SubstraitRelVisitor; @@ -103,24 +105,127 @@ protected AggregateFunctionInvocation generateBinding( Type outputType) { AggregateCall agg = call.getUnderlying(); - List sorts = - agg.getCollation() != null - ? agg.getCollation().getFieldCollations().stream() - .map(r -> SubstraitRelVisitor.toSortField(r, call.inputType)) - .collect(java.util.stream.Collectors.toList()) - : Collections.emptyList(); - Expression.AggregationInvocation invocation = - agg.isDistinct() - ? Expression.AggregationInvocation.DISTINCT - : Expression.AggregationInvocation.ALL; + return AggregateFunctionInvocation.builder() + .declaration(function) + .outputType(outputType) + .aggregationPhase(Expression.AggregationPhase.INITIAL_TO_RESULT) + .sort(sortFields(agg, call.inputType)) + .invocation( + agg.isDistinct() + ? Expression.AggregationInvocation.DISTINCT + : Expression.AggregationInvocation.ALL) + .addAllArguments(arguments) + .build(); + } + + /** + * Rebuilds the invocation from the binding a converted call carries, when that binding still + * describes the call. + * + *

A binding names the exact declaration the plan used and the semantics Calcite cannot express + * — the function's options and its aggregation phase — so it beats re-matching: matching can only + * guess a declaration among equally-shaped candidates, always produces a full initial-to-result + * aggregation with no options, and for a phase that consumes an intermediate state cannot match + * at all, because the call's operands are then intermediate values rather than the declaration's + * arguments. + * + *

The binding describes the plan as it was converted, and a planner rule may have rewritten + * the call since. Staleness is therefore decided by comparing the arguments the binding recorded + * with the ones the call now has — not by re-validating against the declaration, which cannot + * judge a phase whose state is parameterized by the initial arguments it no longer sees (the + * intermediate {@code STRUCT,i64>} of a decimal average, say). When the arguments + * differ, this returns empty and the call is matched on its own terms, binding and all. + */ + private Optional fromCarriedBinding( + RelNode input, + Type.Struct inputType, + AggregateCall call, + Function topLevelConverter) { + Optional bound = + AggregateFunctions.boundBinding(call.getAggregation()); + if (!bound.isPresent() + || !(bound.get().function().declaration() + instanceof SimpleExtension.AggregateFunctionVariant)) { + return Optional.empty(); + } + ResolvedAggregateBinding binding = bound.get(); + SimpleExtension.AggregateFunctionVariant declaration = + (SimpleExtension.AggregateFunctionVariant) binding.function().declaration(); + + List operands = + call.getArgList().stream() + .map(index -> topLevelConverter.apply(rexBuilder.makeInputRef(input, index))) + .collect(java.util.stream.Collectors.toList()); + Optional> arguments = alignArguments(binding, declaration, operands); + if (!arguments.isPresent()) { + return Optional.empty(); + } - return ExpressionCreator.aggregateFunction( - function, - outputType, - Expression.AggregationPhase.INITIAL_TO_RESULT, - sorts, - invocation, - arguments); + AggregateFunctionInvocation invocation = + AggregateFunctionInvocation.builder() + .declaration(declaration) + .outputType(typeConverter.toSubstrait(call.getType())) + .aggregationPhase(binding.phase()) + .sort(sortFields(call, inputType)) + // The invocation comes from the Calcite call rather than from the binding: a planner + // rule may legitimately have dropped a redundant DISTINCT since the plan was converted. + .invocation( + call.isDistinct() + ? Expression.AggregationInvocation.DISTINCT + : Expression.AggregationInvocation.ALL) + .addAllArguments(arguments.get()) + .options(binding.function().options()) + .build(); + + List rebuilt = + ResolvedAggregateBinding.resolve(invocation).function().arguments(); + return rebuilt.equals(binding.function().arguments()) + ? Optional.of(invocation) + : Optional.empty(); + } + + /** + * Puts the call's operands back into the declaration's argument order, restoring the enum + * arguments — which are not Calcite operands — from the binding. Returns empty when the operands + * no longer fit the declaration, which again means the binding does not describe this call. + */ + private static Optional> alignArguments( + ResolvedAggregateBinding binding, + SimpleExtension.AggregateFunctionVariant declaration, + List operands) { + List resolved = binding.function().arguments(); + List arguments = new ArrayList<>(); + int operandIndex = 0; + for (int index = 0; index < declaration.args().size(); index++) { + if (declaration.args().get(index) instanceof SimpleExtension.EnumArgument) { + if (index >= resolved.size() || resolved.get(index).kind() != ResolvedArgument.Kind.ENUM) { + return Optional.empty(); + } + arguments.add( + resolved.get(index).enumValue().map(EnumArg::of).orElse(EnumArg.UNSPECIFIED_ENUM_ARG)); + } else { + if (operandIndex >= operands.size()) { + return Optional.empty(); + } + arguments.add(operands.get(operandIndex++)); + } + } + // A variadic declaration states its trailing argument once but accepts it repeatedly. + while (operandIndex < operands.size()) { + if (!declaration.variadic().isPresent()) { + return Optional.empty(); + } + arguments.add(operands.get(operandIndex++)); + } + return Optional.of(arguments); + } + + private static List sortFields(AggregateCall call, Type.Struct inputType) { + return call.getCollation() != null + ? call.getCollation().getFieldCollations().stream() + .map(collation -> SubstraitRelVisitor.toSortField(collation, inputType)) + .collect(java.util.stream.Collectors.toList()) + : Collections.emptyList(); } /** @@ -138,6 +243,14 @@ public Optional convert( AggregateCall call, Function topLevelConverter) { + // A call converted from Substrait may carry its resolved binding; that is a better source than + // re-matching the operator, and the only source for a phase Calcite cannot represent. + Optional carried = + fromCarriedBinding(input, inputType, call, topLevelConverter); + if (carried.isPresent()) { + return carried; + } + FunctionFinder m = getFunctionFinder(call); if (m == null) { return Optional.empty(); @@ -195,7 +308,7 @@ private List leadingEnumArgs(AggregateCall call) { protected FunctionFinder getFunctionFinder(AggregateCall call) { // replace COUNT() + distinct == true and approximate == true with APPROX_COUNT_DISTINCT // before converting into substrait function - SqlAggFunction aggFunction = call.getAggregation(); + SqlAggFunction aggFunction = AggregateFunctions.unwrapBound(call.getAggregation()); if (aggFunction == SqlStdOperatorTable.COUNT && call.isDistinct() && call.isApproximate()) { aggFunction = SqlStdOperatorTable.APPROX_COUNT_DISTINCT; } diff --git a/isthmus/src/test/java/io/substrait/isthmus/SubstraitRelNodeConverterTest.java b/isthmus/src/test/java/io/substrait/isthmus/SubstraitRelNodeConverterTest.java index 4a5ed1f9e..473965535 100644 --- a/isthmus/src/test/java/io/substrait/isthmus/SubstraitRelNodeConverterTest.java +++ b/isthmus/src/test/java/io/substrait/isthmus/SubstraitRelNodeConverterTest.java @@ -1,17 +1,43 @@ package io.substrait.isthmus; +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertNotEquals; +import static org.junit.jupiter.api.Assertions.assertSame; import static org.junit.jupiter.api.Assertions.assertThrows; - +import static org.junit.jupiter.api.Assertions.assertTrue; + +import io.substrait.expression.AggregateFunctionInvocation; +import io.substrait.expression.EnumArg; +import io.substrait.expression.Expression; +import io.substrait.expression.FunctionArg; +import io.substrait.expression.FunctionOption; +import io.substrait.expression.ImmutableAggregateFunctionInvocation; +import io.substrait.extension.DefaultExtensionCatalog; +import io.substrait.extension.InvalidFunctionBindingException; +import io.substrait.extension.ResolvedAggregateBinding; +import io.substrait.extension.ResolvedArgument; +import io.substrait.extension.SimpleExtension; import io.substrait.plan.Plan; import io.substrait.relation.Join.JoinType; import io.substrait.relation.Rel; import io.substrait.relation.Set.SetOp; import io.substrait.type.Type; +import java.util.ArrayList; import java.util.List; import java.util.Optional; import java.util.stream.Collectors; import java.util.stream.Stream; +import org.apache.calcite.plan.RelOptRule; +import org.apache.calcite.plan.hep.HepPlanner; +import org.apache.calcite.plan.hep.HepProgram; +import org.apache.calcite.plan.hep.HepProgramBuilder; import org.apache.calcite.rel.RelNode; +import org.apache.calcite.rel.core.AggregateCall; +import org.apache.calcite.rel.rules.CoreRules; +import org.apache.calcite.rel.type.RelDataType; +import org.apache.calcite.sql.SqlAggFunction; +import org.apache.calcite.tools.RelBuilder; +import org.apache.calcite.util.ImmutableBitSet; import org.junit.jupiter.api.Nested; import org.junit.jupiter.api.Test; @@ -38,6 +64,14 @@ void direct() { RelNode relNode = substraitToCalcite.convert(root.getInput()); assertRowMatch(relNode.getRowType(), R.I32, N.STRING, R.I64); + + // The declared COUNT type matches Calcite's inference, so the call keeps the stock operator + // rather than a binding-carrying wrapper. + org.apache.calcite.rel.core.Aggregate calciteAgg = + (org.apache.calcite.rel.core.Aggregate) relNode; + assertEquals( + Optional.empty(), + AggregateFunctions.boundBinding(calciteAgg.getAggCallList().get(0).getAggregation())); } @Test @@ -53,6 +87,534 @@ void emit() { RelNode relNode = substraitToCalcite.convert(root.getInput()); assertRowMatch(relNode.getRowType(), N.STRING, R.I64); } + + @Test + void declaredMeasureOutputTypes() { + Rel aggregate = + sb.aggregate( + input -> sb.grouping(input, 2), + input -> + List.of( + withOutputType(sb.sum(input, 0), R.I64), + withOutputType(sb.avg(input, 1), R.FP32)), + commonTable); + + RelNode relNode = substraitToCalcite.convert(aggregate); + assertRowMatch(relNode.getRowType(), N.STRING, R.I64, R.FP32); + + // Both declared types diverge from Calcite's inference, so both calls carry a resolved + // binding on a wrapper (PLAN_OUTPUT preserves the plan's declared type). + org.apache.calcite.rel.core.Aggregate calciteAgg = + (org.apache.calcite.rel.core.Aggregate) relNode; + for (AggregateCall call : calciteAgg.getAggCallList()) { + assertTrue(AggregateFunctions.boundBinding(call.getAggregation()).isPresent()); + } + + assertFullRoundTrip(aggregate); + } + + @Test + void declaredDecimalWidthIsPreserved() { + Rel input = + sb.namedScan(List.of("example"), List.of("d", "g"), List.of(R.decimal(10, 2), R.STRING)); + // The standard extension declarations for decimal sum and avg return DECIMAL<38,S>, while + // Calcite's inference keeps the argument's precision. The declared width must survive. + Rel aggregate = + sb.aggregate( + i -> sb.grouping(i, 1), + i -> + List.of( + sb.measure( + sb.aggregateFn( + DefaultExtensionCatalog.FUNCTIONS_ARITHMETIC_DECIMAL, + "sum:dec", + N.decimal(38, 2), + sb.fieldReference(i, 0))), + sb.measure( + sb.aggregateFn( + DefaultExtensionCatalog.FUNCTIONS_ARITHMETIC_DECIMAL, + "avg:dec", + N.decimal(38, 2), + sb.fieldReference(i, 0)))), + input); + + RelNode relNode = substraitToCalcite.convert(aggregate); + assertRowMatch(relNode.getRowType(), R.STRING, N.decimal(38, 2), N.decimal(38, 2)); + assertFullRoundTrip(aggregate); + } + + @Test + void declaredTypeOnGlobalAggregateWithoutGroupings() { + Rel input = sb.namedScan(List.of("example"), List.of("a"), List.of(R.I32)); + // An aggregate with no groupings at all (as opposed to one empty grouping) is a global + // aggregation; the declared type must be preserved for it as well. + Rel aggregate = + io.substrait.relation.Aggregate.builder() + .input(input) + .measures(List.of(withOutputType(sb.sum(input, 0), R.I64))) + .build(); + + RelNode relNode = substraitToCalcite.convert(aggregate); + assertRowMatch(relNode.getRowType(), R.I64); + + // A global aggregation is one empty grouping set, not zero grouping sets. RelBuilder derives + // each measure's hasEmptyGroup from the grouping sets, so anything else would re-infer the + // measures against a group key the aggregate does not have. + org.apache.calcite.rel.core.Aggregate calciteAgg = + (org.apache.calcite.rel.core.Aggregate) relNode; + assertEquals(List.of(ImmutableBitSet.of()), calciteAgg.getGroupSets()); + assertEquals(org.apache.calcite.rel.core.Aggregate.Group.SIMPLE, calciteAgg.getGroupType()); + + // Converting back therefore normalizes "no groupings at all" to the equivalent single empty + // grouping — the same global aggregation, spelled the way Calcite spells it. + io.substrait.relation.Aggregate roundTripped = + (io.substrait.relation.Aggregate) SubstraitRelVisitor.convert(relNode, converterProvider); + assertEquals(1, roundTripped.getGroupings().size()); + assertTrue(roundTripped.getGroupings().get(0).getExpressions().isEmpty()); + } + + @Test + void boundOperatorIdentityIgnoresDeclaredType() { + ResolvedAggregateBinding binding = + ResolvedAggregateBinding.resolve(sb.sum(commonTable, 0).getFunction()); + RelDataType i64 = converterProvider.getTypeConverter().toCalcite(typeFactory, N.I64); + RelDataType fp64 = converterProvider.getTypeConverter().toCalcite(typeFactory, N.FP64); + + SqlAggFunction boundToI64 = AggregateFunctions.bind(AggregateFunctions.SUM, binding, i64); + SqlAggFunction boundToFp64 = AggregateFunctions.bind(AggregateFunctions.SUM, binding, fp64); + + // Operator identity answers "which Substrait function is this?", never "which type does it + // produce?", so matching on the operator stays type-agnostic. + assertEquals(boundToI64, boundToFp64); + assertEquals(boundToI64.hashCode(), boundToFp64.hashCode()); + + // The carried type is still readable, it is just not part of the identity. + assertEquals(Optional.of(i64), AggregateFunctions.declaredOutputType(boundToI64)); + assertEquals(Optional.of(fp64), AggregateFunctions.declaredOutputType(boundToFp64)); + assertEquals(Optional.empty(), AggregateFunctions.declaredOutputType(AggregateFunctions.SUM)); + + // A different Substrait function is a different operator. + ResolvedAggregateBinding countBinding = + ResolvedAggregateBinding.resolve(sb.count(commonTable, 0).getFunction()); + assertNotEquals( + boundToI64, AggregateFunctions.bind(AggregateFunctions.SUM, countBinding, i64)); + } + + @Test + void enumArgumentDistinguishesMeasuresOfTheSameFunction() { + Rel input = sb.namedScan(List.of("example"), List.of("x", "g"), List.of(R.FP32, R.STRING)); + // Both measures call std_dev:req_fp32 on the same column; only the leading "distribution" + // enum argument differs, so only the resolved binding tells them apart. + Rel aggregate = + sb.aggregate( + i -> sb.grouping(i, 1), + i -> List.of(stdDev(i, "POPULATION"), stdDev(i, "SAMPLE")), + input); + + RelNode relNode = substraitToCalcite.convert(aggregate); + assertRowMatch(relNode.getRowType(), R.STRING, R.FP32, R.FP32); + + org.apache.calcite.rel.core.Aggregate calciteAgg = + (org.apache.calcite.rel.core.Aggregate) relNode; + List calls = calciteAgg.getAggCallList(); + assertEquals(2, calls.size()); + ResolvedAggregateBinding population = boundBinding(calls.get(0)); + ResolvedAggregateBinding sample = boundBinding(calls.get(1)); + assertNotEquals(population, sample); + assertEquals(Optional.of("POPULATION"), enumArgument(population)); + assertEquals(Optional.of("SAMPLE"), enumArgument(sample)); + } + + @Test + void functionOptionsSurviveConversion() { + Rel input = sb.namedScan(List.of("example"), List.of("a", "g"), List.of(R.I32, R.STRING)); + // count(i32) -> i64 matches Calcite's inference exactly, so nothing about the type forces a + // wrapper — but the plan's "overflow" option has no place in a Calcite aggregate call, and + // dropping it would silently change what the plan asks for. + Rel aggregate = sb.aggregate(i -> sb.grouping(i, 1), i -> List.of(count(i, "ERROR")), input); + + RelNode relNode = substraitToCalcite.convert(aggregate); + + org.apache.calcite.rel.core.Aggregate calciteAgg = + (org.apache.calcite.rel.core.Aggregate) relNode; + ResolvedAggregateBinding binding = boundBinding(calciteAgg.getAggCallList().get(0)); + assertEquals(Optional.of(List.of("ERROR")), binding.function().option("overflow")); + + // ... and the option comes back out on the way to Substrait, instead of being re-matched + // away. + assertFullRoundTrip(aggregate); + } + + @Test + void measuresDifferingOnlyByOptionsStayDistinct() { + Rel input = sb.namedScan(List.of("example"), List.of("a", "g"), List.of(R.I32, R.STRING)); + Rel aggregate = + sb.aggregate( + i -> sb.grouping(i, 1), i -> List.of(count(i, "ERROR"), count(i, "SILENT")), input); + + RelNode relNode = substraitToCalcite.convert(aggregate); + + // Both measures are COUNT(a) of type i64 and differ only in an option Calcite cannot express, + // so only the carried binding keeps them from being deduplicated into one column. + org.apache.calcite.rel.core.Aggregate calciteAgg = + (org.apache.calcite.rel.core.Aggregate) relNode; + assertEquals(2, calciteAgg.getAggCallList().size()); + assertNotEquals( + boundBinding(calciteAgg.getAggCallList().get(0)), + boundBinding(calciteAgg.getAggCallList().get(1))); + assertRowMatch(relNode.getRowType(), R.STRING, R.I64, R.I64); + } + + @Test + void intermediatePhaseSurvivesBothDirections() { + Rel input = sb.namedScan(List.of("example"), List.of("partial"), List.of(R.I64)); + // count accumulates into i64, and a final phase consumes that state. Calcite has no notion of + // aggregation phases, so only the carried binding can bring this back: matching would see + // COUNT(i64) and rebuild a full initial-to-result aggregation instead. + Rel aggregate = + sb.aggregate( + i -> sb.grouping(i), + i -> List.of(count(i, null, Expression.AggregationPhase.INTERMEDIATE_TO_RESULT)), + input); + + RelNode relNode = substraitToCalcite.convert(aggregate); + + org.apache.calcite.rel.core.Aggregate calciteAgg = + (org.apache.calcite.rel.core.Aggregate) relNode; + assertEquals( + Expression.AggregationPhase.INTERMEDIATE_TO_RESULT, + boundBinding(calciteAgg.getAggCallList().get(0)).phase()); + + assertFullRoundTrip(aggregate); + } + + @Test + void parameterizedIntermediateStateSurvivesBothDirections() { + // avg:dec accumulates into STRUCT,i64>. A phase that consumes that state sees + // only the state, so S cannot be bound and the intermediate type cannot be re-derived — the + // binding has to be trusted as recorded instead of re-validated against the declaration. + Type state = R.struct(R.decimal(38, 2), R.I64); + // A NamedStruct names nested fields too, so the state column contributes three names. + Rel input = + sb.namedScan(List.of("example"), List.of("partial", "total", "count"), List.of(state)); + Rel aggregate = + sb.aggregate( + i -> sb.grouping(i), + i -> + List.of( + io.substrait.relation.Aggregate.Measure.builder() + .function( + AggregateFunctionInvocation.builder() + .declaration( + extensions.getAggregateFunction( + SimpleExtension.FunctionAnchor.of( + DefaultExtensionCatalog.FUNCTIONS_ARITHMETIC_DECIMAL, + "avg:dec"))) + .outputType(N.decimal(38, 2)) + .aggregationPhase( + Expression.AggregationPhase.INTERMEDIATE_TO_RESULT) + .invocation(Expression.AggregationInvocation.ALL) + .addArguments(sb.fieldReference(i, 0)) + .build()) + .build()), + input); + + RelNode relNode = substraitToCalcite.convert(aggregate); + + org.apache.calcite.rel.core.Aggregate calciteAgg = + (org.apache.calcite.rel.core.Aggregate) relNode; + assertEquals( + Expression.AggregationPhase.INTERMEDIATE_TO_RESULT, + boundBinding(calciteAgg.getAggCallList().get(0)).phase()); + + assertFullRoundTrip(aggregate); + } + + @Test + void calciteInferenceIgnoresAnUnrepresentablePlanType() { + // The declared type is a user-defined type this converter has no mapping for. Ignoring it is + // exactly what CALCITE_INFERENCE promises, so the conversion must not fail on it. + Rel aggregate = + sb.aggregate( + input -> sb.grouping(input, 2), + input -> + List.of( + withOutputType( + sb.sum(input, 0), + sb.userDefinedType("extension:test:unmapped", "opaque"))), + commonTable); + SubstraitToCalcite calciteInference = + new SubstraitToCalcite( + converterProvider, + null, + new AggregateConversion( + AggregateConversion.OutputTypeSource.CALCITE_INFERENCE, + AggregateConversion.FunctionBindingValidation.NONE)); + + assertRowMatch(calciteInference.convert(aggregate).getRowType(), N.STRING, N.I32); + + // Preserving it, on the other hand, is impossible — and says so instead of substituting. + assertThrows( + UnsupportedOperationException.class, () -> substraitToCalcite.convert(aggregate)); + } + + @Test + void unwrappingRestoresTheDelegateAndKeepsTheType() { + Rel aggregate = + sb.aggregate( + i -> sb.grouping(i, 2), + i -> List.of(withOutputType(sb.sum(i, 0), R.I64)), + commonTable); + org.apache.calcite.rel.core.Aggregate calciteAgg = + (org.apache.calcite.rel.core.Aggregate) substraitToCalcite.convert(aggregate); + AggregateCall bound = calciteAgg.getAggCallList().get(0); + assertTrue(AggregateFunctions.boundBinding(bound.getAggregation()).isPresent()); + + // What a consumer does before executing a converted plan: Calcite dispatches aggregate + // implementations by operator identity, and a bound operator has none. + AggregateCall unwrapped = AggregateFunctions.unwrapBound(bound); + assertEquals(AggregateFunctions.SUM, unwrapped.getAggregation()); + assertEquals(bound.getType(), unwrapped.getType()); + assertEquals(Optional.empty(), AggregateFunctions.boundBinding(unwrapped.getAggregation())); + // A call that carries no binding is returned untouched. + assertSame(unwrapped, AggregateFunctions.unwrapBound(unwrapped)); + } + + @Test + void converterFactoryDispatchDependsOnTheConfiguration() { + List factories = new ArrayList<>(); + ConverterProvider provider = + new ConverterProvider() { + @Override + public SubstraitRelNodeConverter getSubstraitRelNodeConverter(RelBuilder relBuilder) { + factories.add("single-argument"); + return super.getSubstraitRelNodeConverter(relBuilder); + } + + @Override + public SubstraitRelNodeConverter getSubstraitRelNodeConverter( + RelBuilder relBuilder, AggregateConversion aggregateConversion) { + factories.add("aggregate-aware"); + return super.getSubstraitRelNodeConverter(relBuilder, aggregateConversion); + } + }; + Rel aggregate = + sb.aggregate( + input -> sb.grouping(input, 2), input -> List.of(sb.count(input, 0)), commonTable); + + // The default configuration goes through the long-standing single-argument factory, so a + // subclass that only overrides that one still customizes conversion. + new SubstraitToCalcite(provider).convert(aggregate); + assertEquals(List.of("single-argument", "aggregate-aware"), factories); + + // A non-default configuration cannot: it has nowhere to pass the configuration but the + // two-argument factory. + factories.clear(); + new SubstraitToCalcite( + provider, + null, + new AggregateConversion( + AggregateConversion.OutputTypeSource.CALCITE_INFERENCE, + AggregateConversion.FunctionBindingValidation.NONE)) + .convert(aggregate); + assertEquals(List.of("aggregate-aware"), factories); + } + + private io.substrait.relation.Aggregate.Measure count(Rel input, String overflow) { + return count(input, overflow, Expression.AggregationPhase.INITIAL_TO_RESULT); + } + + private io.substrait.relation.Aggregate.Measure count( + Rel input, String overflow, Expression.AggregationPhase phase) { + SimpleExtension.AggregateFunctionVariant declaration = + extensions.getAggregateFunction( + SimpleExtension.FunctionAnchor.of( + DefaultExtensionCatalog.FUNCTIONS_AGGREGATE_GENERIC, "count:any")); + ImmutableAggregateFunctionInvocation.Builder function = + AggregateFunctionInvocation.builder() + .declaration(declaration) + .outputType(R.I64) + .aggregationPhase(phase) + .invocation(Expression.AggregationInvocation.ALL) + .addArguments(sb.fieldReference(input, 0)); + if (overflow != null) { + function.addOptions(FunctionOption.builder().name("overflow").addValues(overflow).build()); + } + return io.substrait.relation.Aggregate.Measure.builder().function(function.build()).build(); + } + + private ResolvedAggregateBinding boundBinding(AggregateCall call) { + return AggregateFunctions.boundBinding(call.getAggregation()) + .orElseThrow(() -> new AssertionError("expected a bound aggregate function: " + call)); + } + + private Optional enumArgument(ResolvedAggregateBinding binding) { + return binding.function().arguments().stream() + .filter(argument -> argument.kind() == ResolvedArgument.Kind.ENUM) + .findFirst() + .flatMap(ResolvedArgument::enumValue); + } + + private io.substrait.relation.Aggregate.Measure stdDev(Rel input, String distribution) { + SimpleExtension.AggregateFunctionVariant declaration = + extensions.getAggregateFunction( + SimpleExtension.FunctionAnchor.of( + DefaultExtensionCatalog.FUNCTIONS_ARITHMETIC, "std_dev:req_fp32")); + return io.substrait.relation.Aggregate.Measure.builder() + .function( + AggregateFunctionInvocation.builder() + .declaration(declaration) + .arguments( + List.of(EnumArg.of(distribution), sb.fieldReference(input, 0))) + // The declaration returns fp32?; declaring it required diverges from Calcite's + // inference, which is what makes the binding travel on a wrapper. + .outputType(R.FP32) + .aggregationPhase(Expression.AggregationPhase.INITIAL_TO_RESULT) + .invocation(Expression.AggregationInvocation.ALL) + .build()) + .build(); + } + + @Test + void differentDeclaredTypesOnDuplicateMeasuresArePreserved() { + Rel input = sb.namedScan(List.of("example"), List.of("a", "g"), List.of(R.I32, R.STRING)); + Rel aggregate = + sb.aggregate( + i -> sb.grouping(i, 1), + i -> + List.of( + withOutputType(sb.sum(i, 0), R.I64), withOutputType(sb.sum(i, 0), R.FP64)), + input); + + RelNode relNode = substraitToCalcite.convert(aggregate); + + // The two SUM(a) measures are identical apart from their declared output type, and + // AggregateCall equality ignores that type. They survive as distinct columns because the + // converter opts such an aggregate out of RelBuilder's aggregate-call deduplication. + assertRowMatch(relNode.getRowType(), R.STRING, R.I64, R.FP64); + assertEquals(2, ((org.apache.calcite.rel.core.Aggregate) relNode).getAggCallList().size()); + } + + @Test + void declaredTypeSurvivesAggregateRollupRule() { + Rel input = sb.namedScan(List.of("example"), List.of("a", "g"), List.of(R.I32, R.STRING)); + Rel filtered = sb.filter(i -> sb.equal(sb.fieldReference(i, 0), sb.i32(1)), input); + Rel aggregate = + sb.aggregate( + i -> sb.grouping(i, 1), i -> List.of(withOutputType(sb.sum(i, 0), R.I64)), filtered); + + RelNode converted = substraitToCalcite.convert(aggregate); + RelNode optimized = optimize(converted, CoreRules.AGGREGATE_FILTER_TRANSPOSE); + + // The filter references the non-grouped column, so the rule would roll the aggregate up. The + // wrapper opts out of rollup, so the declared type is left intact instead of being rederived + // as a nullable SUM on the rolled-up top aggregate. + assertRowMatch(optimized.getRowType(), R.STRING, R.I64); + } + + @Test + void declaredTypeSurvivesAggregateSplitRule() { + Rel left = sb.namedScan(List.of("left"), List.of("a", "k"), List.of(R.I32, R.I32)); + Rel right = sb.namedScan(List.of("right"), List.of("k"), List.of(R.I32)); + Rel joined = + sb.innerJoin( + ji -> sb.equal(sb.fieldReference(ji, 1), sb.fieldReference(ji, 2)), left, right); + Rel aggregate = + sb.aggregate( + i -> sb.grouping(), i -> List.of(withOutputType(sb.sum(i, 0), R.I64)), joined); + + RelNode converted = substraitToCalcite.convert(aggregate); + RelNode optimized = optimize(converted, CoreRules.AGGREGATE_JOIN_TRANSPOSE_EXTENDED); + + // The wrapper opts out of splitting, so the rule leaves the declared type intact instead of + // rederiving a nullable SUM on the split top aggregate. + assertRowMatch(optimized.getRowType(), R.I64); + } + + @Test + void strictValidationRejectsNonSpecOutputType() { + // sum(i32) derives i64? from its declaration; a plan declaring i32? is not spec-compliant. + Rel aggregate = + sb.aggregate( + input -> sb.grouping(input, 2), + input -> List.of(withOutputType(sb.sum(input, 0), N.I32)), + commonTable); + SubstraitToCalcite strict = + new SubstraitToCalcite( + converterProvider, + null, + new AggregateConversion( + AggregateConversion.OutputTypeSource.PLAN_OUTPUT, + AggregateConversion.FunctionBindingValidation.EXTENSION_DECLARATION)); + assertThrows(InvalidFunctionBindingException.class, () -> strict.convert(aggregate)); + } + + @Test + void strictValidationAcceptsAndPreservesSpecOutputType() { + // sum(i32) -> i64? is spec-compliant; STRICT accepts it and PLAN_OUTPUT preserves it. + Rel aggregate = + sb.aggregate( + input -> sb.grouping(input, 2), + input -> List.of(withOutputType(sb.sum(input, 0), N.I64)), + commonTable); + SubstraitToCalcite strict = + new SubstraitToCalcite( + converterProvider, + null, + new AggregateConversion( + AggregateConversion.OutputTypeSource.PLAN_OUTPUT, + AggregateConversion.FunctionBindingValidation.EXTENSION_DECLARATION)); + RelNode relNode = strict.convert(aggregate); + assertRowMatch(relNode.getRowType(), N.STRING, N.I64); + } + + @Test + void calciteInferenceModeIgnoresPlanType() { + // With CALCITE_INFERENCE the plan's declared type is ignored and Calcite's inference wins. + Rel aggregate = + sb.aggregate( + input -> sb.grouping(input, 2), + input -> List.of(withOutputType(sb.sum(input, 0), N.I64)), + commonTable); + SubstraitToCalcite calciteInference = + new SubstraitToCalcite( + converterProvider, + null, + new AggregateConversion( + AggregateConversion.OutputTypeSource.CALCITE_INFERENCE, + AggregateConversion.FunctionBindingValidation.NONE)); + + RelNode relNode = calciteInference.convert(aggregate); + + // Calcite's SUM infers nullable i32 for sum(i32), not the plan's i64, and no wrapper is + // added. + assertRowMatch(relNode.getRowType(), N.STRING, N.I32); + org.apache.calcite.rel.core.Aggregate calciteAgg = + (org.apache.calcite.rel.core.Aggregate) relNode; + assertEquals( + Optional.empty(), + AggregateFunctions.boundBinding(calciteAgg.getAggCallList().get(0).getAggregation())); + } + + private RelNode optimize(RelNode relNode, RelOptRule rule) { + HepProgram program = new HepProgramBuilder().addRuleInstance(rule).build(); + HepPlanner planner = new HepPlanner(program); + planner.setRoot(relNode); + return planner.findBestExp(); + } + + private io.substrait.relation.Aggregate.Measure withOutputType( + io.substrait.relation.Aggregate.Measure measure, Type outputType) { + AggregateFunctionInvocation function = + AggregateFunctionInvocation.builder() + .from(measure.getFunction()) + .outputType(outputType) + .build(); + return io.substrait.relation.Aggregate.Measure.builder() + .from(measure) + .function(function) + .build(); + } } @Nested