Skip to content
Open
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
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,10 @@ public ExtendedExpression toProto(

final ExpressionProtoConverter expressionProtoConverter =
new ExpressionProtoConverter(functionCollector, null);
final TypeProtoConverter typeProtoConverter = new TypeProtoConverter(functionCollector);
final AggregateFunctionProtoConverter aggregateFunctionProtoConverter =
new AggregateFunctionProtoConverter(
functionCollector, expressionProtoConverter, typeProtoConverter);

for (io.substrait.extendedexpression.ExtendedExpression.ExpressionReferenceBase
expressionReference : extendedExpression.getReferredExpressions()) {
Expand All @@ -52,18 +56,15 @@ public ExtendedExpression toProto(
expressionReference;
ExpressionReference.Builder expressionReferenceBuilder =
ExpressionReference.newBuilder()
.setMeasure(
new AggregateFunctionProtoConverter(functionCollector)
.toProto(aft.getMeasure()))
.setMeasure(aggregateFunctionProtoConverter.toProto(aft.getMeasure()))
.addAllOutputNames(expressionReference.getOutputNames());
builder.addReferredExpr(expressionReferenceBuilder);
} else {
throw new UnsupportedOperationException(
"Only Expression or Aggregate Function type are supported in conversion to proto Extended Expressions");
}
}
builder.setBaseSchema(
extendedExpression.getBaseSchema().toProto(new TypeProtoConverter(functionCollector)));
builder.setBaseSchema(extendedExpression.getBaseSchema().toProto(typeProtoConverter));

// the process of adding simple extensions (URNs and declarations) is handled on the fly
functionCollector.addExtensionsToExtendedExpression(builder);
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -28,9 +28,26 @@ public class AggregateFunctionProtoConverter {
* @param functionCollector the extension collector for tracking function references
*/
public AggregateFunctionProtoConverter(ExtensionCollector functionCollector) {
this(
functionCollector,
new ExpressionProtoConverter(functionCollector, null),
new TypeProtoConverter(functionCollector));
}

/**
* Constructs a converter using the caller's expression and type converters.
*
* @param functionCollector the extension collector shared by the converters
* @param exprProtoConverter the converter for arguments and sort expressions
* @param typeProtoConverter the converter for argument and output types
*/
public AggregateFunctionProtoConverter(
ExtensionCollector functionCollector,
ExpressionProtoConverter exprProtoConverter,
TypeProtoConverter typeProtoConverter) {
this.functionCollector = functionCollector;
this.exprProtoConverter = new ExpressionProtoConverter(functionCollector, null);
this.typeProtoConverter = new TypeProtoConverter(functionCollector);
this.exprProtoConverter = exprProtoConverter;
this.typeProtoConverter = typeProtoConverter;
}

/**
Expand All @@ -56,8 +73,16 @@ public AggregateFunction toProto(Aggregate.Measure measure) {
args.get(i)
.accept(aggFuncDef, i, argVisitor, EmptyVisitationContext.INSTANCE))
.collect(Collectors.toList()))
.addAllSorts(
Comment thread
nielspardon marked this conversation as resolved.
measure.getFunction().sort().stream()
.map(exprProtoConverter::toProto)
.collect(Collectors.toList()))
Comment thread
nielspardon marked this conversation as resolved.
.setFunctionReference(
functionCollector.getFunctionReference(measure.getFunction().declaration()))
.addAllOptions(
measure.getFunction().options().stream()
.map(ExpressionProtoConverter::from)
.collect(Collectors.toList()))
.build();
}
}
36 changes: 8 additions & 28 deletions core/src/main/java/io/substrait/relation/RelProtoConverter.java
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,6 @@
import io.substrait.extension.ExtensionProtoConverter;
import io.substrait.extension.SimpleExtension;
import io.substrait.plan.Plan;
import io.substrait.proto.AggregateFunction;
import io.substrait.proto.AggregateRel;
import io.substrait.proto.ConsistentPartitionWindowRel;
import io.substrait.proto.CrossRel;
Expand Down Expand Up @@ -82,6 +81,8 @@ public class RelProtoConverter
/** Collects function and type references encountered during conversion. */
@NonNull protected final ExtensionCollector extensionCollector;

private final AggregateFunctionProtoConverter aggregateFunctionProtoConverter;

/**
* Constructor with custom {@link ExtensionCollector}.
*
Expand Down Expand Up @@ -112,6 +113,9 @@ public RelProtoConverter(
this.exprProtoConverter = new ExpressionProtoConverter(extensionCollector, this);
this.typeProtoConverter = new TypeProtoConverter(extensionCollector);
this.extensionProtoConverter = extensionProtoConverter;
this.aggregateFunctionProtoConverter =
new AggregateFunctionProtoConverter(
extensionCollector, exprProtoConverter, typeProtoConverter);
}

/**
Expand Down Expand Up @@ -252,33 +256,9 @@ public Rel visit(Aggregate aggregate, EmptyVisitationContext context) throws Run
}

private AggregateRel.Measure toProto(Aggregate.Measure measure) {
FunctionArg.FuncArgVisitor<
io.substrait.proto.FunctionArgument, EmptyVisitationContext, RuntimeException>
argVisitor = FunctionArg.toProto(typeProtoConverter, exprProtoConverter);
List<FunctionArg> args = measure.getFunction().arguments();
SimpleExtension.AggregateFunctionVariant aggFuncDef = measure.getFunction().declaration();

AggregateFunction.Builder func =
AggregateFunction.newBuilder()
.setPhase(measure.getFunction().aggregationPhase().toProto())
.setInvocation(measure.getFunction().invocation().toProto())
.setOutputType(toProto(measure.getFunction().getType()))
.addAllArguments(
IntStream.range(0, args.size())
.mapToObj(
i ->
args.get(i)
.accept(aggFuncDef, i, argVisitor, EmptyVisitationContext.INSTANCE))
.collect(Collectors.toList()))
.addAllSorts(toProtoS(measure.getFunction().sort()))
.setFunctionReference(
extensionCollector.getFunctionReference(measure.getFunction().declaration()))
.addAllOptions(
measure.getFunction().options().stream()
.map(ExpressionProtoConverter::from)
.collect(Collectors.toList()));

AggregateRel.Measure.Builder builder = AggregateRel.Measure.newBuilder().setMeasure(func);
AggregateRel.Measure.Builder builder =
AggregateRel.Measure.newBuilder()
.setMeasure(aggregateFunctionProtoConverter.toProto(measure));

measure.getPreMeasureFilter().ifPresent(f -> builder.setFilter(toProto(f)));
return builder.build();
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -6,8 +6,10 @@
import io.substrait.expression.Expression;
import io.substrait.expression.ExpressionCreator;
import io.substrait.expression.FieldReference;
import io.substrait.expression.FunctionOption;
import io.substrait.expression.ImmutableFieldReference;
import io.substrait.extension.DefaultExtensionCatalog;
import io.substrait.extension.SimpleExtension;
import io.substrait.relation.Aggregate;
import io.substrait.type.NamedStruct;
import io.substrait.type.Type;
Expand Down Expand Up @@ -41,6 +43,65 @@ void testRoundTrip(ExtendedExpression.ExpressionReferenceBase expressionReferenc
assertExtendedExpressionOperation(expressionReferences, namedStruct);
}

@Test
void preservesAggregateOrderingAndFunctionsUsedOnlyInSorts() {
FieldReference value = FieldReference.newRootStructReference(0, R.STRING);
FieldReference sortKey = FieldReference.newRootStructReference(1, R.I64);
AggregateFunctionInvocation function =
AggregateFunctionInvocation.builder()
.declaration(
extensions.getAggregateFunction(
SimpleExtension.FunctionAnchor.of(
DefaultExtensionCatalog.FUNCTIONS_STRING, "string_agg:str_str")))
.addArguments(value, ExpressionCreator.string(false, ","))
.outputType(R.STRING)
.aggregationPhase(Expression.AggregationPhase.INITIAL_TO_RESULT)
.invocation(Expression.AggregationInvocation.ALL)
.addSort(
Expression.SortField.builder()
.expr(sb.add(sortKey, sb.i64(1)))
.direction(Expression.SortDirection.DESC_NULLS_LAST)
.build(),
Expression.SortField.builder()
.expr(value)
.direction(Expression.SortDirection.ASC_NULLS_FIRST)
.build(),
Expression.SortField.builder()
.expr(sortKey)
.comparisonFunction(
extensions.getScalarFunction(
SimpleExtension.FunctionAnchor.of(
DefaultExtensionCatalog.FUNCTIONS_COMPARISON, "nullif:any_any")))
.build())
.build();

assertExtendedExpressionOperation(
List.of(
ImmutableAggregateFunctionReference.builder()
.measure(Aggregate.Measure.builder().function(function).build())
.addOutputNames("concatenated")
.build()),
NamedStruct.of(List.of("value", "sort_key"), R.struct(R.STRING, R.I64)));
}

@Test
void preservesAggregateOptionPreferences() {
AggregateFunctionInvocation function =
AggregateFunctionInvocation.builder()
.from(sb.sum(FieldReference.newRootStructReference(0, R.I64)).getFunction())
.addOptions(
FunctionOption.builder().name("overflow").addValues("ERROR", "SATURATE").build())

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Reverse these so the test can tell a preserved order from a sorted one: ERROR, SATURATE is already alphabetical, so a converter that sorts preferences would still pass.

Suggested change
FunctionOption.builder().name("overflow").addValues("ERROR", "SATURATE").build())
FunctionOption.builder().name("overflow").addValues("SATURATE", "ERROR").build())

.build();

assertExtendedExpressionOperation(
List.of(
ImmutableAggregateFunctionReference.builder()
.measure(Aggregate.Measure.builder().function(function).build())
.addOutputNames("total")
.build()),
NamedStruct.of(List.of("value"), R.struct(R.I64)));
}

@Test
void getNoExpressionDefined() {
IllegalStateException illegalStateException =
Expand Down
Loading