sunchao commented on code in PR #4727:
URL: https://github.com/apache/datafusion-comet/pull/4727#discussion_r4124906434


##########
spark/src/main/scala/org/apache/spark/sql/comet/operators.scala:
##########
@@ -2349,6 +2366,48 @@ object CometObjectHashAggregateExec
       op.child,
       SerializedPlan(None))
   }
+
+  /**
+   * For intermediate aggregates containing TypedImperativeAggregate functions 
(like CollectSet or
+   * CollectList), Spark declares buffer columns as BinaryType because it 
serializes the JVM
+   * state. Native Comet keeps the actual state type instead: 
ArrayType(elementType) with
+   * containsNull true for CollectSet/CollectList. Rewrite the Spark-side 
output attributes for
+   * Partial, PartialMerge, and mixed {Partial, PartialMerge} stages so 
shuffle and downstream
+   * native aggregate stages see the schema that native execution really 
produces.
+   *
+   * Final aggregates output user-visible values rather than intermediate 
state, so their Spark
+   * result schema is left unchanged.
+   */
+  private def adjustOutputForNativeState(op: ObjectHashAggregateExec): 
Seq[Attribute] = {

Review Comment:
   [P2] Preserve the shared helper’s `Mode` schema adjustment. This new 
overload is selected by `CometObjectHashAggregateExec.createExec` instead of 
the inherited `CometBaseAggregate` helper, but it omits that helper’s `Mode` 
case. Consequently, native `mode` partials advertise `BinaryType` while 
emitting a struct of values and counts. With Comet shuffle enabled, ordinary 
supported queries such as `SELECT mode(v) FROM mode_int` now abort during 
schema reconciliation instead of returning the mode. Please reuse the shared 
helper, which already handles `PartialMerge`, or preserve its complete state 
mappings.
   
   Evidence: Exact-head CI job 
https://github.com/apache/datafusion-comet/actions/runs/36243582628/job/108409634203
 reproduces failures in `mode.sql:47` and `mode_within_group.sql:42`: 
`CometSchemaAlignExec cannot reconcile ... expected Binary, found 
Struct("values": non-null List(Int32), "counts": non-null List(Int64))`. The 
first fixture expects `5`. The base revision uses the inherited helper 
containing the `Mode` mapping, while this added overload bypasses it.



##########
native/core/src/execution/spark_aggregate_state.rs:
##########
@@ -0,0 +1,688 @@
+// Licensed to the Apache Software Foundation (ASF) under one
+// or more contributor license agreements.  See the NOTICE file
+// distributed with this work for additional information
+// regarding copyright ownership.  The ASF licenses this file
+// to you under the Apache License, Version 2.0 (the
+// "License"); you may not use this file except in compliance
+// with the License.  You may obtain a copy of the License at
+//
+//   http://www.apache.org/licenses/LICENSE-2.0
+//
+// Unless required by applicable law or agreed to in writing,
+// software distributed under the License is distributed on an
+// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+// KIND, either express or implied.  See the License for the
+// specific language governing permissions and limitations
+// under the License.
+
+//! Decoders for Spark JVM aggregate state consumed by native PartialMerge.
+
+use std::{borrow::Cow, sync::Arc};
+
+use arrow::array::{
+    builder::{make_builder, ArrayBuilder, ListBuilder},
+    Array, ArrayRef, BinaryArray, GenericByteArray, LargeBinaryArray, 
OffsetSizeTrait,
+};
+use arrow::datatypes::{DataType, FieldRef, GenericBinaryType};
+use datafusion::common::{DataFusionError, Result};
+use datafusion::physical_expr::aggregate::AggregateFunctionExpr;
+use datafusion_comet_shuffle::spark_unsafe::list::{append_to_builder, 
SparkUnsafeArray};
+
+/// Decoder used by `MergeAsPartial` before forwarding state to the inner 
accumulator.
+#[derive(Clone, Debug)]
+pub(crate) enum PartialMergeStateDecoder {
+    PassThrough,
+    SparkCollect(SparkCollectStateDecoder),
+}
+
+impl PartialMergeStateDecoder {
+    pub(crate) fn try_new(
+        inner_expr: &AggregateFunctionExpr,
+        state_fields: &[FieldRef],
+    ) -> Result<Self> {
+        Ok(
+            match SparkCollectStateDecoder::try_new(inner_expr, state_fields)? 
{
+                Some(decoder) => Self::SparkCollect(decoder),
+                None => Self::PassThrough,
+            },
+        )
+    }
+
+    pub(crate) fn decode<'a>(&self, values: &'a [ArrayRef]) -> Result<Cow<'a, 
[ArrayRef]>> {
+        match self {
+            Self::PassThrough => Ok(Cow::Borrowed(values)),
+            Self::SparkCollect(decoder) => decoder.decode(values),
+        }
+    }
+}
+
+/// Decodes Spark JVM collect aggregate buffers into DataFusion collect state.
+///
+/// Spark's `CollectList` / `CollectSet` are `TypedImperativeAggregate`s. When 
Spark runs the
+/// lower Partial aggregate, each buffer is serialized as a `BinaryType` value 
containing a
+/// single-field `UnsafeRow`; field 0 is the `UnsafeArrayData` with the 
collected elements.
+/// DataFusion's collect accumulators expect the merge input to be a 
list-typed state column, so
+/// mixed Spark-Partial -> Comet-PartialMerge plans must materialize those 
unsafe bytes into an
+/// Arrow `ListArray` before calling the inner accumulator's `merge_batch`. 
The single-field
+/// `UnsafeRow` and nested `UnsafeArrayData` layouts used here are unchanged 
across Spark 3.4-4.2.
+#[derive(Clone, Debug)]
+pub(crate) struct SparkCollectStateDecoder {
+    item_field: FieldRef,
+}
+
+impl SparkCollectStateDecoder {
+    fn try_new(
+        inner_expr: &AggregateFunctionExpr,
+        state_fields: &[FieldRef],
+    ) -> Result<Option<Self>> {
+        if !matches!(inner_expr.fun().name(), "collect_list" | "collect_set") {
+            return Ok(None);
+        }
+
+        let [state_field] = state_fields else {
+            return Err(DataFusionError::Internal(format!(
+                "Spark collect state decoder expected one state field, got {}",
+                state_fields.len()
+            )));
+        };
+        let DataType::List(item_field) = state_field.data_type() else {
+            return Err(DataFusionError::Internal(format!(
+                "Spark collect state decoder expected List state, got {}",
+                state_field.data_type()
+            )));
+        };
+
+        Ok(Some(Self {
+            item_field: Arc::clone(item_field),
+        }))
+    }
+
+    fn decode<'a>(&self, values: &'a [ArrayRef]) -> Result<Cow<'a, 
[ArrayRef]>> {
+        if values.len() != 1 {
+            return Err(DataFusionError::Internal(format!(
+                "Spark collect state decoder expected one state column, got 
{}",
+                values.len()
+            )));
+        }
+
+        match values[0].data_type() {
+            DataType::Binary => {
+                let array = values[0]
+                    .as_any()
+                    .downcast_ref::<BinaryArray>()
+                    .ok_or_else(|| {
+                        Self::decode_error("expected BinaryArray for Binary 
collect state")
+                    })?;
+                Ok(Cow::Owned(vec![self.decode_binary_array(array)?]))
+            }
+            DataType::LargeBinary => {
+                let array = values[0]
+                    .as_any()
+                    .downcast_ref::<LargeBinaryArray>()
+                    .ok_or_else(|| {
+                        Self::decode_error(
+                            "expected LargeBinaryArray for LargeBinary collect 
state",
+                        )
+                    })?;
+                Ok(Cow::Owned(vec![self.decode_binary_array(array)?]))
+            }
+            _ => Ok(Cow::Borrowed(values)),
+        }
+    }
+
+    fn decode_binary_array<O: OffsetSizeTrait>(
+        &self,
+        array: &GenericByteArray<GenericBinaryType<O>>,
+    ) -> Result<ArrayRef> {
+        let mut builder = self.new_list_builder(array.len());
+
+        for row_idx in 0..array.len() {
+            if array.is_null(row_idx) {
+                builder.append_null();
+            } else {
+                self.append_unsafe_row_array(array.value(row_idx), &mut 
builder)?;
+            }
+        }
+
+        Ok(Arc::new(builder.finish()))
+    }
+
+    fn new_list_builder(&self, capacity: usize) -> ListBuilder<Box<dyn 
ArrayBuilder>> {
+        let value_builder = make_builder(self.item_field.data_type(), 
capacity);
+        ListBuilder::with_capacity(value_builder, 
capacity).with_field(Arc::clone(&self.item_field))
+    }
+
+    fn append_unsafe_row_array(
+        &self,
+        row_bytes: &[u8],
+        builder: &mut ListBuilder<Box<dyn ArrayBuilder>>,
+    ) -> Result<()> {
+        match self.spark_array_from_single_field_unsafe_row(row_bytes)? {
+            Some(array) => {
+                append_to_builder::<true>(self.item_field.data_type(), 
builder.values(), &array)

Review Comment:
   [P2] Decode binary `CollectSet` elements using Spark’s actual buffer type. 
Across the supported Spark versions, `CollectSet(BinaryType)` serializes each 
value as `UnsafeArrayData(ArrayType(ByteType))`, then extracts its byte payload 
during evaluation. This decoder instead uses the native `Binary` item type and 
copies the entire nested array, including its header and padding. With Comet 
shuffle enabled and a Spark `LocalTableScan` partial, `SELECT x, count(DISTINCT 
y), collect_set(b) FROM VALUES (1,1,X'ABCD'), (1,2,X'ABCD') AS t(x,y,b) GROUP 
BY x` therefore returns a corrupted binary value. Preserve the 
aggregate-specific serialized element type and unwrap the byte array, or retain 
fallback for this boundary.
   
   Evidence: A bounded Rust test compiled from the exact-head decoder expected 
`[171, 205]` but returned `[2, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 
171, 205, 0, 0, 0, 0, 0, 0]`. Spark 4.1.3 `CollectSet.serialize` produced 
exactly the test’s 64-byte buffer, and `CollectSet.eval` returned the original 
`ABCD`. Spark’s `bufferElementType`, `convertToBufferElement`, and `eval` 
implement this special binary representation in every inspected version. The 
four existing decoder tests passed.



-- 
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.

To unsubscribe, e-mail: [email protected]

For queries about this service, please contact Infrastructure at:
[email protected]


---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]

Reply via email to