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))