diff --git a/core/src/main/java/io/substrait/expression/Expression.java b/core/src/main/java/io/substrait/expression/Expression.java index 767003285..cc77902dc 100644 --- a/core/src/main/java/io/substrait/expression/Expression.java +++ b/core/src/main/java/io/substrait/expression/Expression.java @@ -12,6 +12,7 @@ import java.nio.ByteBuffer; import java.util.List; import java.util.Map; +import java.util.Optional; import java.util.UUID; import java.util.stream.Collectors; import org.immutables.value.Value; @@ -1631,7 +1632,8 @@ public Type getType() { /** * Validates that variadic arguments satisfy the parameter consistency requirement, that {@code * bounds_type} is set whenever a window bound requires it, and that a RANGE bound with a - * Preceding or Following side has exactly one, non-CLUSTERED ordering expression. + * Preceding or Following side has exactly one ordering expression, which must not use + * SORT_DIRECTION_CLUSTERED or a custom comparison function. * *

When CONSISTENT, all variadic arguments must have the same type (ignoring nullability). * When INCONSISTENT, arguments can have different types. @@ -1962,7 +1964,10 @@ public static ImmutableExpression.MultiOrListRecord.Builder builder() { } } - /** Represents a sort field with an expression and sort direction. */ + /** + * Represents a sort field with an expression and a sort kind, which is either a direction or a + * reference to a custom comparison function. + */ @Value.Immutable abstract class SortField { /** @@ -1973,11 +1978,31 @@ abstract class SortField { public abstract Expression expr(); /** - * Returns the sort direction. + * Returns the sort direction, if this sort field uses one rather than a custom comparison + * function. * * @return the sort direction */ - public abstract SortDirection direction(); + public abstract Optional direction(); + + /** + * Returns the custom comparison function, if this sort field uses one rather than a direction. + * + * @return the comparison function declaration + */ + public abstract Optional comparisonFunction(); + + /** + * Validates that exactly one of {@link #direction()} and {@link #comparisonFunction()} is set. + */ + @Value.Check + protected void check() { + if (this.direction().isPresent() == this.comparisonFunction().isPresent()) { + throw new IllegalArgumentException( + "SortField must set exactly one of direction or comparisonFunction, but " + + (this.direction().isPresent() ? "both were set" : "neither was set")); + } + } /** * Creates a new builder for constructing a SortField. diff --git a/core/src/main/java/io/substrait/expression/WindowBound.java b/core/src/main/java/io/substrait/expression/WindowBound.java index 6ec2e994f..2af6b4dde 100644 --- a/core/src/main/java/io/substrait/expression/WindowBound.java +++ b/core/src/main/java/io/substrait/expression/WindowBound.java @@ -66,7 +66,7 @@ static void checkBoundsType( /** * Validates a RANGE window's ordering against its bounds, per the spec's rule that a RANGE frame * with a {@link Preceding} or {@link Following} bound must have exactly one ordering expression, - * which must not use {@code SORT_DIRECTION_CLUSTERED}. + * which must not use {@code SORT_DIRECTION_CLUSTERED} or a custom comparison function. * * @param boundsType the window's bounds type * @param lowerBound the window's lower bound @@ -75,7 +75,8 @@ static void checkBoundsType( * @param function identifies the window function being validated, for the exception message * @throws IllegalArgumentException if {@code boundsType} is {@code RANGE} and either bound is * {@link Preceding} or {@link Following}, and {@code sorts} does not hold exactly one - * ordering expression whose direction is not {@code SORT_DIRECTION_CLUSTERED} + * ordering expression whose direction is not {@code SORT_DIRECTION_CLUSTERED} and which does + * not use a custom comparison function */ static void checkRangeOrdering( Expression.WindowBoundsType boundsType, @@ -99,12 +100,21 @@ static void checkRangeOrdering( + " expression, but found " + sorts.size()); } - if (sorts.get(0).direction() == Expression.SortDirection.CLUSTERED) { + Expression.SortField sort = sorts.get(0); + if (sort.direction() + .filter(direction -> direction == Expression.SortDirection.CLUSTERED) + .isPresent()) { throw new IllegalArgumentException( function + ": a RANGE bound with a Preceding or Following side cannot use" + " SORT_DIRECTION_CLUSTERED for its ordering expression"); } + if (sort.comparisonFunction().isPresent()) { + throw new IllegalArgumentException( + function + + ": a RANGE bound with a Preceding or Following side cannot use a custom" + + " comparison function for its ordering expression"); + } } /** diff --git a/core/src/main/java/io/substrait/expression/proto/ExpressionProtoConverter.java b/core/src/main/java/io/substrait/expression/proto/ExpressionProtoConverter.java index e7ef7a444..379689827 100644 --- a/core/src/main/java/io/substrait/expression/proto/ExpressionProtoConverter.java +++ b/core/src/main/java/io/substrait/expression/proto/ExpressionProtoConverter.java @@ -778,14 +778,7 @@ public Expression visit( List partitionExprs = toProto(expr.partitionBy()); List sortFields = - expr.sort().stream() - .map( - s -> - SortField.newBuilder() - .setDirection(s.direction().toProto()) - .setExpr(toProto(s.expr())) - .build()) - .collect(java.util.stream.Collectors.toList()); + expr.sort().stream().map(this::toProto).collect(java.util.stream.Collectors.toList()); Expression.WindowFunction.Bound lowerBound = toProto(expr.lowerBound()); Expression.WindowFunction.Bound upperBound = toProto(expr.upperBound()); @@ -810,6 +803,36 @@ public Expression visit( .build(); } + /** + * Converts a sort field to its protobuf representation. + * + * @param s the sort field to convert + * @return the proto sort field + */ + public SortField toProto(io.substrait.expression.Expression.SortField s) { + return toProto(s, toProto(s.expr())); + } + + /** + * Converts a sort field to its protobuf representation, using an already-converted proto + * expression. Lets a caller that overrides {@link #toProto(io.substrait.expression.Expression)} + * route the sort field's expression through that override. + * + * @param s the sort field to convert + * @param expr the already-converted proto expression for {@code s.expr()} + * @return the proto sort field + */ + public SortField toProto(io.substrait.expression.Expression.SortField s, Expression expr) { + SortField.Builder builder = SortField.newBuilder().setExpr(expr); + if (s.comparisonFunction().isPresent()) { + builder.setComparisonFunctionReference( + extensionCollector.getFunctionReference(s.comparisonFunction().get())); + } else { + builder.setDirection(s.direction().get().toProto()); + } + return builder.build(); + } + @Override public Expression visit( io.substrait.expression.Expression.DynamicParameter expr, EmptyVisitationContext context) diff --git a/core/src/main/java/io/substrait/expression/proto/ProtoExpressionConverter.java b/core/src/main/java/io/substrait/expression/proto/ProtoExpressionConverter.java index f1e11c0b0..138a77362 100644 --- a/core/src/main/java/io/substrait/expression/proto/ProtoExpressionConverter.java +++ b/core/src/main/java/io/substrait/expression/proto/ProtoExpressionConverter.java @@ -745,10 +745,22 @@ private static List fromFunctionArgumentList( * @return the converted sort field */ public Expression.SortField fromSortField(SortField s) { - return Expression.SortField.builder() - .direction(Expression.SortDirection.fromProto(s.getDirection())) - .expr(from(s.getExpr())) - .build(); + Expression expr = from(s.getExpr()); + switch (s.getSortKindCase()) { + case DIRECTION: + return Expression.SortField.builder() + .expr(expr) + .direction(Expression.SortDirection.fromProto(s.getDirection())) + .build(); + case COMPARISON_FUNCTION_REFERENCE: + return Expression.SortField.builder() + .expr(expr) + .comparisonFunction( + lookup.getScalarFunction(s.getComparisonFunctionReference(), extensions)) + .build(); + default: + throw new IllegalArgumentException("SortField has no sort_kind set"); + } } /** diff --git a/core/src/main/java/io/substrait/relation/ConsistentPartitionWindow.java b/core/src/main/java/io/substrait/relation/ConsistentPartitionWindow.java index 98ca0c2e2..abd9510e3 100644 --- a/core/src/main/java/io/substrait/relation/ConsistentPartitionWindow.java +++ b/core/src/main/java/io/substrait/relation/ConsistentPartitionWindow.java @@ -44,8 +44,9 @@ public abstract class ConsistentPartitionWindow extends SingleInputRel implement public abstract List getSorts(); /** - * Validates that a RANGE bound with a Preceding or Following side has exactly one, non-CLUSTERED - * ordering expression, for every window function invocation. + * Validates that a RANGE bound with a Preceding or Following side has exactly one ordering + * expression, which must not use SORT_DIRECTION_CLUSTERED or a custom comparison function, for + * every window function invocation. */ @Value.Check protected void check() { diff --git a/core/src/main/java/io/substrait/relation/ProtoRelConverter.java b/core/src/main/java/io/substrait/relation/ProtoRelConverter.java index 5f43ace2b..716df2291 100644 --- a/core/src/main/java/io/substrait/relation/ProtoRelConverter.java +++ b/core/src/main/java/io/substrait/relation/ProtoRelConverter.java @@ -954,12 +954,7 @@ protected Sort newSort(SortRel rel) { .input(input) .sortFields( rel.getSortsList().stream() - .map( - field -> - Expression.SortField.builder() - .direction(Expression.SortDirection.fromProto(field.getDirection())) - .expr(converter.from(field.getExpr())) - .build()) + .map(converter::fromSortField) .collect(java.util.stream.Collectors.toList())); if (rel.hasAdvancedExtension()) { diff --git a/core/src/main/java/io/substrait/relation/RelProtoConverter.java b/core/src/main/java/io/substrait/relation/RelProtoConverter.java index 885cd87c9..a782f8877 100644 --- a/core/src/main/java/io/substrait/relation/RelProtoConverter.java +++ b/core/src/main/java/io/substrait/relation/RelProtoConverter.java @@ -40,7 +40,6 @@ import io.substrait.proto.RelCommon.Hint.Stats; import io.substrait.proto.RelRoot; import io.substrait.proto.SetRel; -import io.substrait.proto.SortField; import io.substrait.proto.SortRel; import io.substrait.proto.TopNRel; import io.substrait.proto.UpdateRel; @@ -189,13 +188,7 @@ protected io.substrait.proto.Type toProto(io.substrait.type.Type type) { private List toProtoS(List sorts) { return sorts.stream() - .map( - s -> { - return SortField.newBuilder() - .setDirection(s.direction().toProto()) - .setExpr(toProto(s.expr())) - .build(); - }) + .map(s -> exprProtoConverter.toProto(s, toProto(s.expr()))) .collect(Collectors.toList()); } diff --git a/core/src/test/java/io/substrait/expression/SortFieldTest.java b/core/src/test/java/io/substrait/expression/SortFieldTest.java new file mode 100644 index 000000000..8e1affee8 --- /dev/null +++ b/core/src/test/java/io/substrait/expression/SortFieldTest.java @@ -0,0 +1,50 @@ +package io.substrait.expression; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertThrows; + +import io.substrait.TestBase; +import io.substrait.extension.DefaultExtensionCatalog; +import io.substrait.extension.SimpleExtension; +import org.junit.jupiter.api.Test; + +class SortFieldTest extends TestBase { + + @Test + void neitherDirectionNorComparisonFunctionIsRejected() { + assertThrows( + IllegalArgumentException.class, + () -> Expression.SortField.builder().expr(sb.i64(1)).build()); + } + + @Test + void bothDirectionAndComparisonFunctionIsRejected() { + SimpleExtension.ScalarFunctionVariant comparisonFunction = + extensions.getScalarFunction( + SimpleExtension.FunctionAnchor.of( + DefaultExtensionCatalog.FUNCTIONS_COMPARISON, "nullif:any_any")); + + assertThrows( + IllegalArgumentException.class, + () -> + Expression.SortField.builder() + .expr(sb.i64(1)) + .direction(Expression.SortDirection.ASC_NULLS_FIRST) + .comparisonFunction(comparisonFunction) + .build()); + } + + @Test + void protoWithNoSortKindSetIsRejected() { + io.substrait.proto.SortField protoSortField = + io.substrait.proto.SortField.newBuilder() + .setExpr(expressionProtoConverter.toProto(sb.i64(1))) + .build(); + + IllegalArgumentException e = + assertThrows( + IllegalArgumentException.class, + () -> protoExpressionConverter.fromSortField(protoSortField)); + assertEquals("SortField has no sort_kind set", e.getMessage()); + } +} diff --git a/core/src/test/java/io/substrait/type/proto/ConsistentPartitionWindowRelRoundtripTest.java b/core/src/test/java/io/substrait/type/proto/ConsistentPartitionWindowRelRoundtripTest.java index 7b2c0c71e..626ed8aec 100644 --- a/core/src/test/java/io/substrait/type/proto/ConsistentPartitionWindowRelRoundtripTest.java +++ b/core/src/test/java/io/substrait/type/proto/ConsistentPartitionWindowRelRoundtripTest.java @@ -75,6 +75,41 @@ void consistentPartitionWindowRoundtripSingle() { assertEquals(rel2.getRecordType().fields(), Arrays.asList(R.I64, R.I16, R.I32, R.I64)); } + @Test + void windowFunctionInvocationWithCustomComparisonFunctionRoundtrips() { + // Expression.WindowFunctionInvocation goes through ExpressionProtoConverter, a separate + // POJO->proto implementation from ConsistentPartitionWindow's RelProtoConverter path. + SimpleExtension.WindowFunctionVariant windowFunctionDeclaration = + extensions.getWindowFunction( + SimpleExtension.FunctionAnchor.of( + DefaultExtensionCatalog.FUNCTIONS_ARITHMETIC, "lead:any")); + SimpleExtension.ScalarFunctionVariant comparisonFunction = + extensions.getScalarFunction( + SimpleExtension.FunctionAnchor.of( + DefaultExtensionCatalog.FUNCTIONS_COMPARISON, "nullif:any_any")); + Expression.SortField sortField = + Expression.SortField.builder() + .expr(sb.i64(1)) + .comparisonFunction(comparisonFunction) + .build(); + + Expression.WindowFunctionInvocation windowFunction = + Expression.WindowFunctionInvocation.builder() + .declaration(windowFunctionDeclaration) + .arguments(Arrays.asList(sb.i64(1))) + .partitionBy(Collections.emptyList()) + .sort(Arrays.asList(sortField)) + .outputType(R.I64) + .aggregationPhase(Expression.AggregationPhase.INITIAL_TO_RESULT) + .invocation(Expression.AggregationInvocation.ALL) + .lowerBound(WindowBound.UNBOUNDED) + .upperBound(WindowBound.CURRENT_ROW) + .boundsType(Expression.WindowBoundsType.RANGE) + .build(); + + verifyRoundTrip(windowFunction); + } + @Test void consistentPartitionWindowRoundtripMulti() { SimpleExtension.WindowFunctionVariant windowFunctionLeadDeclaration = @@ -354,6 +389,44 @@ void rangePrecedingWithClusteredOrderingIsRejected() { assertThrows(IllegalArgumentException.class, relBuilder::build); } + @Test + void rangePrecedingWithCustomComparisonFunctionOrderingIsRejected() { + SimpleExtension.WindowFunctionVariant windowFunctionDeclaration = + extensions.getWindowFunction( + SimpleExtension.FunctionAnchor.of( + DefaultExtensionCatalog.FUNCTIONS_ARITHMETIC, "lead:any")); + SimpleExtension.ScalarFunctionVariant comparisonFunction = + extensions.getScalarFunction( + SimpleExtension.FunctionAnchor.of( + DefaultExtensionCatalog.FUNCTIONS_COMPARISON, "nullif:any_any")); + Rel input = sb.namedScan(Arrays.asList("test"), Arrays.asList("a"), Arrays.asList(R.I64)); + // A RANGE bound with a Preceding side cannot use a custom comparison function for its + // ordering expression. + ImmutableConsistentPartitionWindow.Builder relBuilder = + ConsistentPartitionWindow.builder() + .input(input) + .windowFunctions( + Arrays.asList( + ConsistentPartitionWindow.WindowRelFunctionInvocation.builder() + .declaration(windowFunctionDeclaration) + .arguments(Arrays.asList(sb.fieldReference(input, 0))) + .outputType(R.I64) + .aggregationPhase(Expression.AggregationPhase.INITIAL_TO_RESULT) + .invocation(Expression.AggregationInvocation.ALL) + .lowerBound(WindowBound.Preceding.of(5)) + .upperBound(WindowBound.CURRENT_ROW) + .boundsType(Expression.WindowBoundsType.RANGE) + .build())) + .sorts( + Arrays.asList( + Expression.SortField.builder() + .expr(sb.fieldReference(input, 0)) + .comparisonFunction(comparisonFunction) + .build())); + + assertThrows(IllegalArgumentException.class, relBuilder::build); + } + @Test void rangePrecedingWithTwoOrderingExpressionsOnAnInvocationIsRejected() { SimpleExtension.WindowFunctionVariant declaration = diff --git a/core/src/test/java/io/substrait/type/proto/SortRelRoundtripTest.java b/core/src/test/java/io/substrait/type/proto/SortRelRoundtripTest.java index d3bebd41d..35d69bb92 100644 --- a/core/src/test/java/io/substrait/type/proto/SortRelRoundtripTest.java +++ b/core/src/test/java/io/substrait/type/proto/SortRelRoundtripTest.java @@ -2,6 +2,8 @@ import io.substrait.TestBase; import io.substrait.expression.Expression; +import io.substrait.extension.DefaultExtensionCatalog; +import io.substrait.extension.SimpleExtension; import io.substrait.relation.Rel; import io.substrait.relation.Sort; import java.util.Arrays; @@ -222,6 +224,24 @@ void sortAllDirections() { verifyRoundTrip(sort); } + @Test + void sortByCustomComparisonFunction() { + // A sort field can reference a custom comparison function instead of a direction. + SimpleExtension.ScalarFunctionVariant comparisonFunction = + extensions.getScalarFunction( + SimpleExtension.FunctionAnchor.of( + DefaultExtensionCatalog.FUNCTIONS_COMPARISON, "nullif:any_any")); + Expression.SortField sortField = + Expression.SortField.builder() + .expr(sb.fieldReference(baseTable, 2)) + .comparisonFunction(comparisonFunction) + .build(); + + Rel sort = Sort.builder().input(baseTable).addSortFields(sortField).build(); + + verifyRoundTrip(sort); + } + @Test void nestedSort() { // Sort on top of another sort diff --git a/examples/substrait-spark/src/main/java/io/substrait/examples/util/SubstraitStringify.java b/examples/substrait-spark/src/main/java/io/substrait/examples/util/SubstraitStringify.java index 62f9b3798..13d120edc 100644 --- a/examples/substrait-spark/src/main/java/io/substrait/examples/util/SubstraitStringify.java +++ b/examples/substrait-spark/src/main/java/io/substrait/examples/util/SubstraitStringify.java @@ -1,5 +1,6 @@ package io.substrait.examples.util; +import io.substrait.expression.Expression; import io.substrait.relation.Aggregate; import io.substrait.relation.ConsistentPartitionWindow; import io.substrait.relation.Cross; @@ -73,6 +74,13 @@ public SubstraitStringify() { super(0); } + private static String sortKind(Expression.SortField sortField) { + return sortField + .direction() + .map(Object::toString) + .orElseGet(() -> "comparisonFunction=" + sortField.comparisonFunction().get()); + } + /** * Explains the Substrait plan * @@ -300,7 +308,7 @@ public String visit(Sort sort, EmptyVisitationContext context) throws RuntimeExc .forEach( sf -> { ExpressionStringify expr = new ExpressionStringify(indent); - sb.append(sf.expr().accept(expr, context)).append(" ").append(sf.direction()); + sb.append(sf.expr().accept(expr, context)).append(" ").append(sortKind(sf)); }); List inputs = sort.getInputs(); inputs.forEach( @@ -457,7 +465,7 @@ public String visit(TopN topN, EmptyVisitationContext context) throws RuntimeExc .forEach( sf -> { ExpressionStringify expr = new ExpressionStringify(indent); - sb.append(sf.expr().accept(expr, context)).append(" ").append(sf.direction()); + sb.append(sf.expr().accept(expr, context)).append(" ").append(sortKind(sf)); }); topN.getInputs() .forEach( diff --git a/isthmus/src/main/java/io/substrait/isthmus/SubstraitRelNodeConverter.java b/isthmus/src/main/java/io/substrait/isthmus/SubstraitRelNodeConverter.java index 31e39117e..88d730d4a 100644 --- a/isthmus/src/main/java/io/substrait/isthmus/SubstraitRelNodeConverter.java +++ b/isthmus/src/main/java/io/substrait/isthmus/SubstraitRelNodeConverter.java @@ -812,9 +812,9 @@ public RelNode visit(Sort sort, Context context) throws RuntimeException { } private RexNode directedRexNode(Expression.SortField sortField, Context context) { + SortDirection sortDirection = requireDirection(sortField); Expression expression = sortField.expr(); RexNode rexNode = expression.accept(expressionRexConverter, context); - SortDirection sortDirection = sortField.direction(); if (sortDirection == Expression.SortDirection.ASC_NULLS_FIRST) { return relBuilder.nullsFirst(rexNode); @@ -836,6 +836,15 @@ private RexNode directedRexNode(Expression.SortField sortField, Context context) throw new IllegalArgumentException("Unsupported sort direction: " + sortDirection); } + private static SortDirection requireDirection(Expression.SortField sortField) { + return sortField + .direction() + .orElseThrow( + () -> + new UnsupportedOperationException( + "A sort field using a custom comparison function is not supported")); + } + @Override public RelNode visit(Fetch fetch, Context context) throws RuntimeException { RelNode child = fetch.getInput().accept(this, context); @@ -853,9 +862,9 @@ public RelNode visit(Fetch fetch, Context context) throws RuntimeException { } private RelFieldCollation toRelFieldCollation(Expression.SortField sortField, Context context) { + SortDirection sortDirection = requireDirection(sortField); Expression expression = sortField.expr(); RexNode rex = expression.accept(expressionRexConverter, context); - SortDirection sortDirection = sortField.direction(); RexSlot rexSlot = (RexSlot) rex; int fieldIndex = rexSlot.getIndex(); diff --git a/isthmus/src/main/java/io/substrait/isthmus/expression/ExpressionRexConverter.java b/isthmus/src/main/java/io/substrait/isthmus/expression/ExpressionRexConverter.java index 446fef62c..897211007 100644 --- a/isthmus/src/main/java/io/substrait/isthmus/expression/ExpressionRexConverter.java +++ b/isthmus/src/main/java/io/substrait/isthmus/expression/ExpressionRexConverter.java @@ -670,7 +670,14 @@ public RexNode visit(Expression.WindowFunctionInvocation expr, Context context) expr.sort().stream() .map( sf -> { - Set direction = asSqlKind(sf.direction()); + Expression.SortDirection sortDirection = + sf.direction() + .orElseThrow( + () -> + new UnsupportedOperationException( + "A sort field using a custom comparison function is not" + + " supported")); + Set direction = asSqlKind(sortDirection); return new RexFieldCollation(sf.expr().accept(this, context), direction); }) .collect(ImmutableList.toImmutableList()); diff --git a/isthmus/src/test/java/io/substrait/isthmus/SubstraitExpressionConverterTest.java b/isthmus/src/test/java/io/substrait/isthmus/SubstraitExpressionConverterTest.java index 84ae622e6..4bf08d25c 100644 --- a/isthmus/src/test/java/io/substrait/isthmus/SubstraitExpressionConverterTest.java +++ b/isthmus/src/test/java/io/substrait/isthmus/SubstraitExpressionConverterTest.java @@ -14,6 +14,7 @@ import io.substrait.expression.LambdaBuilder; import io.substrait.expression.WindowBound; import io.substrait.extension.DefaultExtensionCatalog; +import io.substrait.extension.SimpleExtension; import io.substrait.isthmus.SubstraitRelNodeConverter.AnchoredInput; import io.substrait.isthmus.SubstraitRelNodeConverter.Context; import io.substrait.isthmus.expression.ExpressionRexConverter; @@ -649,6 +650,35 @@ void reportWindowInferenceFailureWithoutFailingConversion() { assertEquals("controlled window inference failure", failure.getMessage()); } + @Test + void rejectsWindowFunctionSortedByCustomComparisonFunction() { + // RexFieldCollation's Set flags have no representation for a custom comparator, so a + // sort field using one cannot be converted. + Expression.SortField sortField = + Expression.SortField.builder() + .expr(sb.i32(1)) + .comparisonFunction( + extensions.getScalarFunction( + SimpleExtension.FunctionAnchor.of( + DefaultExtensionCatalog.FUNCTIONS_COMPARISON, "nullif:any_any"))) + .build(); + Expression.WindowFunctionInvocation expr = + sb.windowFn( + DefaultExtensionCatalog.FUNCTIONS_ARITHMETIC, + "row_number:", + R.I64, + Expression.AggregationPhase.INITIAL_TO_RESULT, + Expression.AggregationInvocation.ALL, + List.of(sortField), + Expression.WindowBoundsType.RANGE, + WindowBound.UNBOUNDED, + WindowBound.UNBOUNDED); + + assertThrows( + UnsupportedOperationException.class, + () -> expr.accept(expressionRexConverter, Context.newContext())); + } + @Test void propagateWindowObserverException() { ExpressionRexConverter observingConverter = diff --git a/isthmus/src/test/java/io/substrait/isthmus/SubstraitRelNodeConverterTest.java b/isthmus/src/test/java/io/substrait/isthmus/SubstraitRelNodeConverterTest.java index 6235c51ec..e27ba35cd 100644 --- a/isthmus/src/test/java/io/substrait/isthmus/SubstraitRelNodeConverterTest.java +++ b/isthmus/src/test/java/io/substrait/isthmus/SubstraitRelNodeConverterTest.java @@ -1402,5 +1402,26 @@ void emit() { RelNode relNode = substraitToCalcite.convert(root.getInput()); assertRowMatch(relNode.getRowType(), R.I32, N.STRING); } + + @Test + void rejectsCustomComparisonFunction() { + // Calcite's RelFieldCollation has no representation for a custom comparator, so a sort field + // using one cannot be converted. + SimpleExtension.ScalarFunctionVariant comparisonFunction = + extensions.getScalarFunction( + SimpleExtension.FunctionAnchor.of( + DefaultExtensionCatalog.FUNCTIONS_COMPARISON, "nullif:any_any")); + io.substrait.relation.Sort sort = + io.substrait.relation.Sort.builder() + .input(commonTable) + .addSortFields( + Expression.SortField.builder() + .expr(sb.fieldReference(commonTable, 0)) + .comparisonFunction(comparisonFunction) + .build()) + .build(); + + assertThrows(UnsupportedOperationException.class, () -> substraitToCalcite.convert(sort)); + } } } diff --git a/spark/src/main/scala/io/substrait/spark/logical/ToLogicalPlan.scala b/spark/src/main/scala/io/substrait/spark/logical/ToLogicalPlan.scala index c3fa30dba..712972859 100644 --- a/spark/src/main/scala/io/substrait/spark/logical/ToLogicalPlan.scala +++ b/spark/src/main/scala/io/substrait/spark/logical/ToLogicalPlan.scala @@ -110,6 +110,11 @@ class ToLogicalPlan(val spark: AnyRef = SparkCompat.instance.getOrCreateSparkSes throw new IllegalArgumentException(msg) }) + if (function.sort().asScala.exists(!_.direction().isPresent)) { + throw new UnsupportedOperationException( + "A sort field using a custom comparison function is not supported") + } + val filter = Option(measure.getPreMeasureFilter.orElse(null)) .map(_.accept(expressionConverter, EmptyVisitationContext.INSTANCE)) @@ -243,8 +248,12 @@ class ToLogicalPlan(val spark: AnyRef = SparkCompat.instance.getOrCreateSparkSes } private def toSortOrder(sortField: SExpression.SortField): SortOrder = { + if (!sortField.direction().isPresent) { + throw new UnsupportedOperationException( + "A sort field using a custom comparison function is not supported") + } val expression = sortField.expr().accept(expressionConverter, EmptyVisitationContext.INSTANCE) - val (direction, nullOrdering) = sortField.direction() match { + val (direction, nullOrdering) = sortField.direction().get() match { case SExpression.SortDirection.ASC_NULLS_FIRST => (Ascending, NullsFirst) case SExpression.SortDirection.DESC_NULLS_FIRST => (Descending, NullsFirst) case SExpression.SortDirection.ASC_NULLS_LAST => (Ascending, NullsLast) diff --git a/spark/src/test/scala/io/substrait/spark/AggregateWithJoinSuite.scala b/spark/src/test/scala/io/substrait/spark/AggregateWithJoinSuite.scala index e2d84f5b6..c3fdeec83 100644 --- a/spark/src/test/scala/io/substrait/spark/AggregateWithJoinSuite.scala +++ b/spark/src/test/scala/io/substrait/spark/AggregateWithJoinSuite.scala @@ -10,10 +10,11 @@ import org.apache.spark.sql.test.SharedSparkSession import io.substrait.`type`.{NamedStruct, Type, TypeCreator} import io.substrait.dsl.SubstraitBuilder -import io.substrait.expression.{Expression, ExpressionCreator} -import io.substrait.extension.DefaultExtensionCatalog +import io.substrait.expression.{AggregateFunctionInvocation, Expression, ExpressionCreator} +import io.substrait.extension.{DefaultExtensionCatalog, SimpleExtension} import io.substrait.plan.Plan import io.substrait.relation.{Aggregate, Join, Project, Rel, VirtualTableScan} +import io.substrait.util.EmptyVisitationContext import java.util import java.util.Arrays @@ -310,6 +311,37 @@ class AggregateWithJoinSuite assertRow(rows(2), "Diesel", "GU", 35000) } + test("aggregate measure sorted by custom comparison function is rejected") { + // Aggregate measure ordering has no representation in Spark's AggregateExpression, so a + // measure sorted by a custom comparison function must be rejected rather than silently + // converted as unordered. + val testsTable = createTestsTable() + val comparisonFunction = extensions.getScalarFunction( + SimpleExtension.FunctionAnchor + .of(DefaultExtensionCatalog.FUNCTIONS_COMPARISON, "nullif:any_any")) + val sortField = Expression.SortField + .builder() + .expr(sb.fieldReference(testsTable, 0)) + .comparisonFunction(comparisonFunction) + .build() + val baseMeasure = sb.sum(testsTable, 6) // sum of test_mileage + val sortedFunction = AggregateFunctionInvocation + .builder() + .from(baseMeasure.getFunction) + .addSort(sortField) + .build() + val aggregate = Aggregate + .builder() + .input(testsTable) + .addGroupings(Aggregate.Grouping.builder().build()) + .addMeasures(Aggregate.Measure.builder().function(sortedFunction).build()) + .build() + + intercept[UnsupportedOperationException] { + new ToLogicalPlan(spark).visit(aggregate, EmptyVisitationContext.INSTANCE) + } + } + def assertRow(row: Row, fuelType: String, postcodeArea: String, totalTestMileage: Long): Unit = { assertResult(fuelType)(row.getString(0)) assertResult(postcodeArea)(row.getString(1))