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
15 changes: 13 additions & 2 deletions spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala
Original file line number Diff line number Diff line change
Expand Up @@ -741,6 +741,10 @@ object QueryPlanSerde extends Logging with CometExprShim with CometTypeShim {
}
}

/**
* This method does not promote the aggregate tree. Its arguments and filters are independent
* roots and must be serialized through [[exprToProto]] to receive decimal promotion.
*/
def aggExprToProto(
aggExpr: AggregateExpression,
inputs: Seq[Attribute],
Expand Down Expand Up @@ -845,7 +849,9 @@ object QueryPlanSerde extends Logging with CometExprShim with CometTypeShim {
* expression.
*
* This method performs a transformation on the plan to handle decimal promotion and then calls
* into the recursive method [[exprToProtoInternal]].
* into the recursive method [[exprToProtoInternal]]. Use this entry point for independent roots
* (including aggregate arguments and filters) and synthesized trees needing decimal promotion.
* Serdes must use [[exprToProtoInternal]] for children of the already-promoted tree.
*
* @param expr
* The input expression
Expand Down Expand Up @@ -914,6 +920,11 @@ object QueryPlanSerde extends Logging with CometExprShim with CometTypeShim {
* Convert a Spark expression to a protocol-buffer representation of a native Comet/DataFusion
* expression.
*
* The caller owns decimal promotion: this method serializes children of an already-promoted
Comment thread
peterxcli marked this conversation as resolved.
* root without traversing them again. Literals and wrappers that introduce no decimal
* arithmetic can also use this path. Newly synthesized arithmetic must enter through
* [[exprToProto]].
*
* @param expr
* The input expression
* @param inputs
Expand Down Expand Up @@ -1063,7 +1074,7 @@ object QueryPlanSerde extends Logging with CometExprShim with CometTypeShim {
binding: Boolean,
f: (ExprOuterClass.Expr.Builder, ExprOuterClass.UnaryExpr) => ExprOuterClass.Expr.Builder)
: Option[ExprOuterClass.Expr] = {
val childExpr = exprToProtoInternal(child, inputs, binding) // TODO review
val childExpr = exprToProtoInternal(child, inputs, binding)
if (childExpr.isDefined) {
// create the generic UnaryExpr message
val inner = ExprOuterClass.UnaryExpr
Expand Down
77 changes: 40 additions & 37 deletions spark/src/main/scala/org/apache/comet/serde/arrays.scala
Original file line number Diff line number Diff line change
Expand Up @@ -44,8 +44,8 @@ object CometArrayRemove
expr: ArrayRemove,
inputs: Seq[Attribute],
binding: Boolean): Option[ExprOuterClass.Expr] = {
val arrayExprProto = exprToProto(expr.left, inputs, binding)
val keyExprProto = exprToProto(expr.right, inputs, binding)
val arrayExprProto = exprToProtoInternal(expr.left, inputs, binding)
val keyExprProto = exprToProtoInternal(expr.right, inputs, binding)

scalarFunctionExprToProto("array_remove_all", arrayExprProto, keyExprProto)
}
Expand All @@ -60,8 +60,8 @@ object CometArrayAppend extends CometExpressionSerde[ArrayAppend] with ArraysBas
val (srcChild, itemChild) = widenElementInLockstep(expr.children.head, expr.children(1))
val elementType = srcChild.dataType.asInstanceOf[ArrayType].elementType

val arrayExprProto = exprToProto(srcChild, inputs, binding)
val keyExprProto = exprToProto(itemChild, inputs, binding)
val arrayExprProto = exprToProtoInternal(srcChild, inputs, binding)
val keyExprProto = exprToProtoInternal(itemChild, inputs, binding)

// DataFusion's array_append always returns a list with nullable elements,
// so we must promise ArrayType(elementType, containsNull = true) here even if
Expand All @@ -83,7 +83,8 @@ object CometArrayAppend extends CometExpressionSerde[ArrayAppend] with ArraysBas
binding,
(builder, unaryExpr) => builder.setIsNotNull(unaryExpr))

val nullLiteralProto = exprToProto(Literal(null, elementType), Seq.empty)
val nullLiteralProto =
exprToProtoInternal(Literal(null, elementType), Seq.empty, binding = true)

if (arrayAppendScalarExpr.isDefined && isNotNullExpr.isDefined && nullLiteralProto.isDefined) {
val caseWhenExpr = ExprOuterClass.CaseWhen
Expand Down Expand Up @@ -130,8 +131,8 @@ object CometArrayContains
expr: ArrayContains,
inputs: Seq[Attribute],
binding: Boolean): Option[ExprOuterClass.Expr] = {
val arrayExprProto = exprToProto(expr.children.head, inputs, binding)
val keyExprProto = exprToProto(expr.children(1), inputs, binding)
val arrayExprProto = exprToProtoInternal(expr.children.head, inputs, binding)
val keyExprProto = exprToProtoInternal(expr.children(1), inputs, binding)

scalarFunctionExprToProto("array_contains", arrayExprProto, keyExprProto)
}
Expand Down Expand Up @@ -230,8 +231,8 @@ object CometArrayIntersect
expr: ArrayIntersect,
inputs: Seq[Attribute],
binding: Boolean): Option[ExprOuterClass.Expr] = {
val leftArrayExprProto = exprToProto(expr.children.head, inputs, binding)
val rightArrayExprProto = exprToProto(expr.children(1), inputs, binding)
val leftArrayExprProto = exprToProtoInternal(expr.children.head, inputs, binding)
val rightArrayExprProto = exprToProtoInternal(expr.children(1), inputs, binding)

val arraysIntersectScalarExpr =
scalarFunctionExprToProto("array_intersect", leftArrayExprProto, rightArrayExprProto)
Expand All @@ -244,7 +245,7 @@ object CometArrayMax extends CometExpressionSerde[ArrayMax] {
expr: ArrayMax,
inputs: Seq[Attribute],
binding: Boolean): Option[ExprOuterClass.Expr] = {
val arrayExprProto = exprToProto(expr.children.head, inputs, binding)
val arrayExprProto = exprToProtoInternal(expr.children.head, inputs, binding)

val arrayMaxScalarExpr =
scalarFunctionExprToProto("array_max", arrayExprProto)
Expand All @@ -257,7 +258,7 @@ object CometArrayMin extends CometExpressionSerde[ArrayMin] {
expr: ArrayMin,
inputs: Seq[Attribute],
binding: Boolean): Option[ExprOuterClass.Expr] = {
val arrayExprProto = exprToProto(expr.children.head, inputs, binding)
val arrayExprProto = exprToProtoInternal(expr.children.head, inputs, binding)

val arrayMinScalarExpr = scalarFunctionExprToProto("array_min", arrayExprProto)
arrayMinScalarExpr
Expand All @@ -269,8 +270,8 @@ object CometArraysOverlap extends CometExpressionSerde[ArraysOverlap] {
expr: ArraysOverlap,
inputs: Seq[Attribute],
binding: Boolean): Option[ExprOuterClass.Expr] = {
val leftArrayExprProto = exprToProto(expr.left, inputs, binding)
val rightArrayExprProto = exprToProto(expr.right, inputs, binding)
val leftArrayExprProto = exprToProtoInternal(expr.left, inputs, binding)
val rightArrayExprProto = exprToProtoInternal(expr.right, inputs, binding)

val arraysOverlapScalarExpr = scalarFunctionExprToProtoWithReturnType(
"spark_arrays_overlap",
Expand All @@ -289,7 +290,7 @@ object CometArrayCompact extends CometExpressionSerde[Expression] {
inputs: Seq[Attribute],
binding: Boolean): Option[ExprOuterClass.Expr] = {
val child = expr.children.head
val arrayExprProto = exprToProto(child, inputs, binding)
val arrayExprProto = exprToProtoInternal(child, inputs, binding)

val arrayCompactScalarExpr = scalarFunctionExprToProto("array_compact", arrayExprProto)
arrayCompactScalarExpr
Expand Down Expand Up @@ -345,8 +346,8 @@ object CometArrayExcept
return None
case None =>
}
val leftArrayExprProto = exprToProto(expr.left, inputs, binding)
val rightArrayExprProto = exprToProto(expr.right, inputs, binding)
val leftArrayExprProto = exprToProtoInternal(expr.left, inputs, binding)
val rightArrayExprProto = exprToProtoInternal(expr.right, inputs, binding)

val arrayExceptScalarExpr =
scalarFunctionExprToProto("array_except", leftArrayExprProto, rightArrayExprProto)
Expand Down Expand Up @@ -402,16 +403,16 @@ object CometArrayJoin
expr: ArrayJoin,
inputs: Seq[Attribute],
binding: Boolean): Option[ExprOuterClass.Expr] = {
val arrayExprProto = exprToProto(expr.array, inputs, binding)
val delimiterExprProto = exprToProto(expr.delimiter, inputs, binding)
val arrayExprProto = exprToProtoInternal(expr.array, inputs, binding)
val delimiterExprProto = exprToProtoInternal(expr.delimiter, inputs, binding)

val joined = expr.nullReplacement match {
case Some(nullReplacementExpr) =>
scalarFunctionExprToProto(
"array_to_string",
arrayExprProto,
delimiterExprProto,
exprToProto(nullReplacementExpr, inputs, binding))
exprToProtoInternal(nullReplacementExpr, inputs, binding))
case None =>
scalarFunctionExprToProto("array_to_string", arrayExprProto, delimiterExprProto)
}
Expand All @@ -422,8 +423,8 @@ object CometArrayJoin
case Some(nullReplacementExpr) =>
for {
innerProto <- joined
replacementIsNull <- exprToProto(IsNull(nullReplacementExpr), inputs, binding)
nullLiteral <- exprToProto(Literal(null, expr.dataType), inputs, binding)
replacementIsNull <- exprToProtoInternal(IsNull(nullReplacementExpr), inputs, binding)
nullLiteral <- exprToProtoInternal(Literal(null, expr.dataType), inputs, binding)
} yield ExprOuterClass.Expr
.newBuilder()
.setIf(
Expand Down Expand Up @@ -478,9 +479,9 @@ object CometSlice extends CometExpressionSerde[Slice] {
expr: Slice,
inputs: Seq[Attribute],
binding: Boolean): Option[ExprOuterClass.Expr] = {
val arrayExprProto = exprToProto(expr.x, inputs, binding)
val startExprProto = exprToProto(Cast(expr.start, LongType), inputs, binding)
val lengthExprProto = exprToProto(Cast(expr.length, LongType), inputs, binding)
val arrayExprProto = exprToProtoInternal(expr.x, inputs, binding)
val startExprProto = exprToProtoInternal(Cast(expr.start, LongType), inputs, binding)
val lengthExprProto = exprToProtoInternal(Cast(expr.length, LongType), inputs, binding)
// No serialized return type: native `spark_array_slice` reuses its input's list field for the
// output, so only `return_field_from_args` is guaranteed to match. Spark's `expr.dataType` is
// not: `CometCreateArray` may have widened the input to a deeply-nullable element type, and
Expand All @@ -498,8 +499,8 @@ object CometArrayUnion extends CometExpressionSerde[ArrayUnion] {
expr: ArrayUnion,
inputs: Seq[Attribute],
binding: Boolean): Option[ExprOuterClass.Expr] = {
val leftArrayExprProto = exprToProto(expr.children.head, inputs, binding)
val rightArrayExprProto = exprToProto(expr.children(1), inputs, binding)
val leftArrayExprProto = exprToProtoInternal(expr.children.head, inputs, binding)
val rightArrayExprProto = exprToProtoInternal(expr.children(1), inputs, binding)

val arraysUnionScalarExpr =
scalarFunctionExprToProto("array_union", leftArrayExprProto, rightArrayExprProto)
Expand Down Expand Up @@ -606,7 +607,7 @@ object CometArrayReverse extends CometExpressionSerde[Reverse] with ArraysBase {
withFallbackReason(expr, s"child data type not supported: ${expr.child.dataType}")
return None
}
val reverseExprProto = exprToProto(expr.child, inputs, binding)
val reverseExprProto = exprToProtoInternal(expr.child, inputs, binding)
val reverseScalarExpr = scalarFunctionExprToProto("array_reverse", reverseExprProto)
reverseScalarExpr
}
Expand Down Expand Up @@ -694,7 +695,8 @@ object CometElementAt extends CometExpressionSerde[ElementAt] {
inputs,
binding,
(builder, unaryExpr) => builder.setIsNotNull(unaryExpr))
val nullLiteralProto = exprToProto(Literal(null, expr.dataType), Seq.empty)
val nullLiteralProto =
exprToProtoInternal(Literal(null, expr.dataType), Seq.empty, binding = true)
for {
base <- baseExpr
notNull <- isNotNullExpr
Expand Down Expand Up @@ -732,7 +734,7 @@ object CometFlatten extends CometExpressionSerde[Flatten] with ArraysBase {
expr: Flatten,
inputs: Seq[Attribute],
binding: Boolean): Option[ExprOuterClass.Expr] = {
val flattenExprProto = exprToProto(expr.child, inputs, binding)
val flattenExprProto = exprToProtoInternal(expr.child, inputs, binding)
val flattenScalarExpr = scalarFunctionExprToProto("flatten", flattenExprProto)
flattenScalarExpr
}
Expand Down Expand Up @@ -777,7 +779,7 @@ object CometSize extends CometExpressionSerde[Size] {
expr: Size,
inputs: Seq[Attribute],
binding: Boolean): Option[ExprOuterClass.Expr] = {
val arrayExprProto = exprToProto(expr.child, inputs, binding)
val arrayExprProto = exprToProtoInternal(expr.child, inputs, binding)
for {
isNotNullExprProto <- createIsNotNullExprProto(expr, inputs, binding)
sizeScalarExprProto <- scalarFunctionExprToProto("size", arrayExprProto)
Expand Down Expand Up @@ -810,7 +812,7 @@ object CometSize extends CometExpressionSerde[Size] {

private def createLiteralExprProto(legacySizeOfNull: Boolean): Option[ExprOuterClass.Expr] = {
val value = if (legacySizeOfNull) -1 else null
exprToProto(Literal(value, IntegerType), Seq.empty)
exprToProtoInternal(Literal(value, IntegerType), Seq.empty, binding = true)
}

}
Expand All @@ -830,8 +832,8 @@ object CometArrayPosition extends CometExpressionSerde[ArrayPosition] with Array
expr: ArrayPosition,
inputs: Seq[Attribute],
binding: Boolean): Option[ExprOuterClass.Expr] = {
val arrayExprProto = exprToProto(expr.left, inputs, binding)
val elementExprProto = exprToProto(expr.right, inputs, binding)
val arrayExprProto = exprToProtoInternal(expr.left, inputs, binding)
val elementExprProto = exprToProtoInternal(expr.right, inputs, binding)

// Use spark_array_position which returns Int64 and 0 when not found
// (matching Spark's behavior)
Expand Down Expand Up @@ -881,7 +883,8 @@ object CometArraysZip extends CometExpressionSerde[ArraysZip] {
// mimic Spark's ArraysZip behavior: returns NULL if any argument is NULL
val combinedNullCheck = expr.children.map(child => IsNotNull(child)).reduce(And)
val isNotNullExpr = exprToProtoInternal(combinedNullCheck, inputs, binding)
val nullLiteralProto = exprToProto(Literal(null, expr.dataType), Seq.empty)
val nullLiteralProto =
exprToProtoInternal(Literal(null, expr.dataType), Seq.empty, binding = true)

if (exprChildren.forall(
_.isDefined) && isNotNullExpr.isDefined && nullLiteralProto.isDefined) {
Expand Down Expand Up @@ -1017,12 +1020,12 @@ object CometSequence extends CometExpressionSerde[Sequence] with CodegenDispatch
expr: Sequence,
inputs: Seq[Attribute],
binding: Boolean): Option[ExprOuterClass.Expr] = {
val startExprProto = exprToProto(expr.start, inputs, binding)
val stopExprProto = exprToProto(expr.stop, inputs, binding)
val startExprProto = exprToProtoInternal(expr.start, inputs, binding)
Comment thread
peterxcli marked this conversation as resolved.
val stopExprProto = exprToProtoInternal(expr.stop, inputs, binding)
// With no step argument the native kernel computes Spark's per-row default,
// `start <= stop ? 1 : -1`, which cannot be expressed as a plan-time literal.
val argProtos = Seq(startExprProto, stopExprProto) ++
expr.stepOpt.map(exprToProto(_, inputs, binding))
expr.stepOpt.map(exprToProtoInternal(_, inputs, binding))
scalarFunctionExprToProtoWithReturnType(
"spark_sequence",
expr.dataType,
Expand Down
6 changes: 3 additions & 3 deletions spark/src/main/scala/org/apache/comet/serde/bitwise.scala
Original file line number Diff line number Diff line change
Expand Up @@ -50,7 +50,7 @@ object CometBitwiseNot extends CometExpressionSerde[BitwiseNot] {
expr: BitwiseNot,
inputs: Seq[Attribute],
binding: Boolean): Option[ExprOuterClass.Expr] = {
val childProto = exprToProto(expr.child, inputs, binding)
val childProto = exprToProtoInternal(expr.child, inputs, binding)
val bitNotScalarExpr =
scalarFunctionExprToProto("bitwise_not", childProto)
bitNotScalarExpr
Expand Down Expand Up @@ -144,8 +144,8 @@ object CometBitwiseGet extends CometExpressionSerde[BitwiseGet] {
expr: BitwiseGet,
inputs: Seq[Attribute],
binding: Boolean): Option[ExprOuterClass.Expr] = {
val argProto = exprToProto(expr.left, inputs, binding)
val posProto = exprToProto(expr.right, inputs, binding)
val argProto = exprToProtoInternal(expr.left, inputs, binding)
val posProto = exprToProtoInternal(expr.right, inputs, binding)
val bitGetScalarExpr =
scalarFunctionExprToProtoWithReturnType("bit_get", ByteType, false, argProto, posProto)
bitGetScalarExpr
Expand Down
Loading
Loading