diff --git a/gluten-substrait/src/main/java/org/apache/gluten/substrait/expression/ExpressionBuilder.java b/gluten-substrait/src/main/java/org/apache/gluten/substrait/expression/ExpressionBuilder.java index 9d5b5c9a10..8a790a5319 100644 --- a/gluten-substrait/src/main/java/org/apache/gluten/substrait/expression/ExpressionBuilder.java +++ b/gluten-substrait/src/main/java/org/apache/gluten/substrait/expression/ExpressionBuilder.java @@ -256,6 +256,11 @@ public static SingularOrListNode makeSingularOrListNode( return new SingularOrListNode(value, expressionNodes); } + public static SingularOrListNode makeSingularOrListNode( + ExpressionNode value, List rawValues, org.apache.spark.sql.types.DataType dataType) { + return new SingularOrListNode(value, rawValues, dataType); + } + public static WindowFunctionNode makeWindowFunction( Integer functionId, List expressionNodes, diff --git a/gluten-substrait/src/main/java/org/apache/gluten/substrait/expression/SingularOrListNode.java b/gluten-substrait/src/main/java/org/apache/gluten/substrait/expression/SingularOrListNode.java index c55791a085..acdd4eca1f 100644 --- a/gluten-substrait/src/main/java/org/apache/gluten/substrait/expression/SingularOrListNode.java +++ b/gluten-substrait/src/main/java/org/apache/gluten/substrait/expression/SingularOrListNode.java @@ -17,6 +17,7 @@ package org.apache.gluten.substrait.expression; import io.substrait.proto.Expression; +import org.apache.spark.sql.types.DataType; import java.io.Serializable; import java.util.ArrayList; @@ -25,18 +26,39 @@ public class SingularOrListNode implements ExpressionNode, Serializable { private final ExpressionNode value; private final List listNodes = new ArrayList<>(); + // rawValues and dataType allow delaying literal node construction until toProtobuf() + private final List rawValues; + private final DataType dataType; SingularOrListNode(ExpressionNode value, List listNodes) { this.value = value; this.listNodes.addAll(listNodes); + this.rawValues = null; + this.dataType = null; + } + + SingularOrListNode(ExpressionNode value, List rawValues, DataType dataType) { + this.value = value; + this.rawValues = new ArrayList<>(rawValues); + this.dataType = dataType; } @Override public Expression toProtobuf() { Expression.SingularOrList.Builder builder = Expression.SingularOrList.newBuilder(); builder.setValue(value.toProtobuf()); - for (ExpressionNode expressionNode : listNodes) { - builder.addOptions(expressionNode.toProtobuf()); + if (!listNodes.isEmpty()) { + for (ExpressionNode expressionNode : listNodes) { + builder.addOptions(expressionNode.toProtobuf()); + } + } else if (rawValues != null) { + for (Object obj : rawValues) { + // construct a temporary LiteralNode and convert to protobuf to avoid keeping + // many LiteralNode objects in memory at once. Use per-value nullability. + LiteralNode literalNode = + (LiteralNode) ExpressionBuilder.makeLiteral(obj, dataType, obj == null); + builder.addOptions(literalNode.toProtobuf()); + } } Expression.Builder expressionBuilder = Expression.newBuilder(); expressionBuilder.setSingularOrList(builder); diff --git a/gluten-substrait/src/main/scala/org/apache/gluten/expression/PredicateExpressionTransformer.scala b/gluten-substrait/src/main/scala/org/apache/gluten/expression/PredicateExpressionTransformer.scala index 9f443973a9..3f74cab69a 100644 --- a/gluten-substrait/src/main/scala/org/apache/gluten/expression/PredicateExpressionTransformer.scala +++ b/gluten-substrait/src/main/scala/org/apache/gluten/expression/PredicateExpressionTransformer.scala @@ -40,10 +40,18 @@ case class InSetTransformer( original: InSet) extends UnaryExpressionTransformer { override def doTransform(context: SubstraitContext): ExpressionNode = { - InExpressionTransformer.toTransformer( - child.doTransform(context), - original.hset, - original.child.dataType) + val leftNode = child.doTransform(context) + val values = original.hset + val valueType = original.child.dataType + + // Keep raw literal values and defer building LiteralNode objects until toProtobuf(). + val rawValues = new java.util.ArrayList[Object]( + values.toSeq + // Sort elements for deterministic behaviours. + .sortBy(Literal(_, valueType).toString()) + .map(_.asInstanceOf[Object]) + .asJava) + ExpressionBuilder.makeSingularOrListNode(leftNode, rawValues, valueType) } }