Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
33 changes: 29 additions & 4 deletions core/src/main/java/io/substrait/expression/Expression.java
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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.
*
* <p>When CONSISTENT, all variadic arguments must have the same type (ignoring nullability).
* When INCONSISTENT, arguments can have different types.
Expand Down Expand Up @@ -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 {
/**
Expand All @@ -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<SortDirection> direction();

/**
* Returns the custom comparison function, if this sort field uses one rather than a direction.
*
* @return the comparison function declaration
*/
public abstract Optional<SimpleExtension.ScalarFunctionVariant> 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.
Expand Down
16 changes: 13 additions & 3 deletions core/src/main/java/io/substrait/expression/WindowBound.java
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Comment thread
anasik marked this conversation as resolved.
*
* @param boundsType the window's bounds type
* @param lowerBound the window's lower bound
Expand All @@ -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,
Expand All @@ -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");
}
}

/**
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -778,14 +778,7 @@ public Expression visit(
List<Expression> partitionExprs = toProto(expr.partitionBy());

List<SortField> 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());
Expand All @@ -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)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -745,10 +745,22 @@ private static List<FunctionArg> 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");
}
}

/**
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -44,8 +44,9 @@ public abstract class ConsistentPartitionWindow extends SingleInputRel implement
public abstract List<SortField> 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() {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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()) {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -189,13 +188,7 @@ protected io.substrait.proto.Type toProto(io.substrait.type.Type type) {

private List<io.substrait.proto.SortField> toProtoS(List<Expression.SortField> 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());
}

Expand Down
50 changes: 50 additions & 0 deletions core/src/test/java/io/substrait/expression/SortFieldTest.java
Original file line number Diff line number Diff line change
@@ -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());
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -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 =
Expand Down Expand Up @@ -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 =
Expand Down
Loading
Loading