diff --git a/core/src/main/java/io/substrait/extendedexpression/ExtendedExpressionProtoConverter.java b/core/src/main/java/io/substrait/extendedexpression/ExtendedExpressionProtoConverter.java index 1edbd5508..45d6f8e68 100644 --- a/core/src/main/java/io/substrait/extendedexpression/ExtendedExpressionProtoConverter.java +++ b/core/src/main/java/io/substrait/extendedexpression/ExtendedExpressionProtoConverter.java @@ -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()) { @@ -52,9 +56,7 @@ 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 { @@ -62,8 +64,7 @@ public ExtendedExpression toProto( "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); diff --git a/core/src/main/java/io/substrait/relation/AggregateFunctionProtoConverter.java b/core/src/main/java/io/substrait/relation/AggregateFunctionProtoConverter.java index e8ba02c61..b098f9c8a 100644 --- a/core/src/main/java/io/substrait/relation/AggregateFunctionProtoConverter.java +++ b/core/src/main/java/io/substrait/relation/AggregateFunctionProtoConverter.java @@ -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; } /** @@ -56,8 +73,16 @@ public AggregateFunction toProto(Aggregate.Measure measure) { args.get(i) .accept(aggFuncDef, i, argVisitor, EmptyVisitationContext.INSTANCE)) .collect(Collectors.toList())) + .addAllSorts( + measure.getFunction().sort().stream() + .map(exprProtoConverter::toProto) + .collect(Collectors.toList())) .setFunctionReference( functionCollector.getFunctionReference(measure.getFunction().declaration())) + .addAllOptions( + measure.getFunction().options().stream() + .map(ExpressionProtoConverter::from) + .collect(Collectors.toList())) .build(); } } diff --git a/core/src/main/java/io/substrait/relation/RelProtoConverter.java b/core/src/main/java/io/substrait/relation/RelProtoConverter.java index a782f8877..28f90e898 100644 --- a/core/src/main/java/io/substrait/relation/RelProtoConverter.java +++ b/core/src/main/java/io/substrait/relation/RelProtoConverter.java @@ -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; @@ -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}. * @@ -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); } /** @@ -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 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(); diff --git a/core/src/test/java/io/substrait/extendedexpression/ExtendedExpressionRoundTripTest.java b/core/src/test/java/io/substrait/extendedexpression/ExtendedExpressionRoundTripTest.java index 9b6ae95a8..f17609d96 100644 --- a/core/src/test/java/io/substrait/extendedexpression/ExtendedExpressionRoundTripTest.java +++ b/core/src/test/java/io/substrait/extendedexpression/ExtendedExpressionRoundTripTest.java @@ -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; @@ -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()) + .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 =