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


##########
spark/src/main/scala/org/apache/comet/serde/aggregates.scala:
##########
@@ -1116,6 +1116,66 @@ object CometApproxCountDistinct extends 
CometAggregateExpressionSerde[HyperLogLo
   }
 }
 
+object CometPivotFirst extends CometAggregateExpressionSerde[PivotFirst] {
+
+  // Delegate to Spark's own PivotFirst.supportsDataType so the two lists 
cannot drift if a
+  // future Spark version adds a value type to the fast-path gate.
+  private def unsupportedValueTypeReason(dt: DataType): String =
+    s"Unsupported value data type: $dt"
+
+  private val emptyPivotValuesReason = "Pivot values list is empty"
+
+  override def getUnsupportedReasons(): Seq[String] = Seq(
+    "Value data type outside PivotFirst.supportsDataType " +
+      "(Boolean, Byte, Short, Int, Long, Float, Double, Decimal)",
+    emptyPivotValuesReason)
+
+  override def getSupportLevel(expr: PivotFirst): SupportLevel = {
+    if (!PivotFirst.supportsDataType(expr.valueDataType)) {
+      Unsupported(Some(unsupportedValueTypeReason(expr.valueDataType)))
+    } else if (expr.pivotColumnValues.isEmpty) {
+      Unsupported(Some(emptyPivotValuesReason))
+    } else {
+      Compatible()

Review Comment:
   [P2] Gate binary pivot columns before returning `Compatible()`. For a 
Parquet row `(g=1, k=X'61', v=10)`, `SELECT * FROM t PIVOT (sum(v) FOR k IN 
(X'61', X'62'))` returns `(1,NULL,NULL)` in Spark, but the new native aggregate 
populates the first slot with 10. Spark’s atomic-key `HashMap[Any, Int]` 
compares these byte arrays by identity, while `ScalarValue::Binary` compares 
their contents. Enabling this path changes query results by default. Preserve 
Spark matching, or make binary pivot columns fall back.
   
   Evidence: Ran the Parquet-backed reference query on Spark 4.1.3 and obtained 
`[Row(g=1, a=None, b=None)]`, with `pivotfirst` in its physical plan. The 
exact-source native probe using pivot literal 
`ScalarValue::Binary(Some(vec![97]))` and an independently constructed binary 
input containing `b"a"` produced `[Int32(10)]`. The checked Spark 
implementations select `HashMap` for atomic pivot-column types, and the new 
serde has no binary-key restriction.



##########
native/spark-expr/src/agg_funcs/pivot_first.rs:
##########
@@ -0,0 +1,497 @@
+// 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.
+
+//! Spark's `PivotFirst` aggregate. Used only by the second phase of the 
optimized pivot plan
+//! generated by `PivotTransformer`. For each group, `PivotFirst` maintains an 
array of
+//! `pivot_values.len()` slots; on each input row it evaluates the pivot 
column, looks up its
+//! index in `pivot_values`, and writes the value column into that slot when a 
match is found
+//! and the value is non-null. Rows with unmatched pivot values are ignored; 
matched rows with
+//! a null value column leave the slot unchanged (matches Spark).
+//!
+//! State layout is one column per pivot slot, matching Spark's 
`aggBufferAttributes` (which
+//! declares `indexSize` `AttributeReference`s, one per pivot value). This 
keeps the shuffle
+//! schema between Partial and Final consistent with what Spark catalyst 
declared; otherwise
+//! the shuffle exchange rejects the batch. `evaluate()` reassembles the slots 
into a
+//! `ListArray` matching `PivotFirst.dataType = ArrayType(value_type)`.
+
+use arrow::array::{Array, ArrayRef};
+use arrow::datatypes::{DataType, Field, FieldRef};
+use datafusion::common::utils::SingleRowListArrayBuilder;
+use datafusion::common::{DataFusionError, Result as DFResult, ScalarValue};
+use datafusion::logical_expr::function::{AccumulatorArgs, StateFieldsArgs};
+use datafusion::logical_expr::Volatility::Immutable;
+use datafusion::logical_expr::{Accumulator, AggregateUDFImpl, Signature};
+use datafusion::physical_expr::expressions::format_state_name;
+use std::collections::HashMap;
+use std::sync::Arc;
+
+/// UDAF implementation of Spark's `PivotFirst`.
+///
+/// `pivot_values` is a fixed, plan-time list of the pivot column values that 
occupy each
+/// output slot; `pivot_index[v] = i` means an input row whose pivot column 
equals `v` writes
+/// into slot `i`. Both the vector and the map are wrapped in `Arc` because 
`accumulator()`
+/// fires once per group in a grouped aggregate and we want that path to bump 
a refcount
+/// rather than deep-clone.
+#[derive(Debug)]
+pub struct SparkPivotFirst {
+    signature: Signature,
+    value_type: DataType,
+    // Kept for `PartialEq`/`Hash` (identity of the aggregate for plan 
comparison) and for the
+    // deterministic slot ordering `state_fields` needs. `HashMap` alone would 
give us the map
+    // but not a stable order or a `Hash` impl.
+    pivot_values: Arc<Vec<ScalarValue>>,
+    pivot_index: Arc<HashMap<ScalarValue, usize>>,
+}
+
+impl PartialEq for SparkPivotFirst {
+    fn eq(&self, other: &Self) -> bool {
+        self.value_type == other.value_type && self.pivot_values == 
other.pivot_values
+    }
+}
+
+impl Eq for SparkPivotFirst {}
+
+impl std::hash::Hash for SparkPivotFirst {
+    fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
+        self.value_type.hash(state);
+        self.pivot_values.hash(state);
+    }
+}
+
+impl SparkPivotFirst {
+    pub fn new(value_type: DataType, pivot_values: Vec<ScalarValue>) -> Self {
+        let mut pivot_index = HashMap::with_capacity(pivot_values.len());
+        // Spark's PivotFirst uses the FIRST occurrence's index 
(HashMap/TreeMap semantics), so
+        // when duplicates are somehow present we mirror that by only 
inserting the first one.
+        // `pivot_key` can fold two distinct pivot values (`0.0` and `-0.0`) 
onto one key, and can
+        // drop one entirely (NaN), so the index is not necessarily the same 
length as the slot
+        // vector - the slot count is always `pivot_values.len()`.
+        for (i, v) in pivot_values.iter().enumerate() {
+            if let Some(key) = pivot_key(v.clone()) {
+                pivot_index.entry(key).or_insert(i);
+            }
+        }
+        Self {
+            signature: Signature::user_defined(Immutable),
+            value_type,
+            pivot_values: Arc::new(pivot_values),
+            pivot_index: Arc::new(pivot_index),
+        }
+    }
+}
+
+/// Rewrite a pivot column value into the key Spark would match it on, or 
`None` when Spark can
+/// never match it.
+///
+/// Spark's `PivotFirst` looks pivot values up in a Scala `HashMap[Any, Int]`, 
so matching goes
+/// through `BoxesRunTime.equals` / `Statics.anyHash` on the boxed Catalyst 
value rather than
+/// through `ScalarValue`'s own equality. The two disagree on floats in 
opposite directions:
+///
+/// * `-0.0` and `0.0` are one key for Spark (`-0.0 == 0.0` numerically, and 
`doubleHash` folds
+///   both onto the hash of `0L`), while `ScalarValue` keeps them apart.
+/// * `NaN` matches nothing for Spark, not even another `NaN`, because Scala's 
`==` on `Double`
+///   is IEEE. `ScalarValue` treats `NaN` as equal to itself.
+///
+/// Nulls are left alone: a null pivot column value does match a null entry in 
the pivot list,
+/// which is what Spark's `pivotIndex.getOrElse(null, -1)` does.
+fn pivot_key(v: ScalarValue) -> Option<ScalarValue> {
+    match v {
+        ScalarValue::Float32(Some(f)) => {
+            if f.is_nan() {
+                None
+            } else if f == 0.0 {
+                Some(ScalarValue::Float32(Some(0.0)))
+            } else {
+                Some(ScalarValue::Float32(Some(f)))
+            }
+        }
+        ScalarValue::Float64(Some(f)) => {
+            if f.is_nan() {
+                None
+            } else if f == 0.0 {
+                Some(ScalarValue::Float64(Some(0.0)))
+            } else {
+                Some(ScalarValue::Float64(Some(f)))
+            }
+        }
+        other => Some(other),
+    }
+}
+
+impl AggregateUDFImpl for SparkPivotFirst {
+    fn name(&self) -> &str {
+        "pivot_first"
+    }
+
+    fn signature(&self) -> &Signature {
+        &self.signature
+    }
+
+    fn return_type(&self, _arg_types: &[DataType]) -> DFResult<DataType> {
+        Ok(DataType::List(Arc::new(Field::new_list_field(
+            self.value_type.clone(),
+            true,
+        ))))
+    }
+
+    fn state_fields(&self, args: StateFieldsArgs) -> DFResult<Vec<FieldRef>> {
+        // One field per pivot slot, matching Spark's aggBufferAttributes so 
the shuffle
+        // exchange sees the same schema catalyst declared. 
`format_state_name` is the same
+        // helper other aggregates in this crate use (see `avg.rs`, 
`stddev.rs`).
+        Ok((0..self.pivot_values.len())
+            .map(|i| {
+                Arc::new(Field::new(
+                    format_state_name(args.name, &i.to_string()),
+                    self.value_type.clone(),
+                    true,
+                ))
+            })
+            .collect())
+    }
+
+    fn accumulator(&self, _acc_args: AccumulatorArgs) -> DFResult<Box<dyn 
Accumulator>> {
+        Ok(Box::new(PivotFirstAccumulator::new(
+            self.value_type.clone(),
+            Arc::clone(&self.pivot_index),
+            self.pivot_values.len(),
+        )))
+    }
+}
+
+/// Per-group state: `slots[i]` holds the latest non-null value assigned to 
pivot slot `i`, or
+/// `None` when nothing has written to that slot yet.
+#[derive(Debug)]
+struct PivotFirstAccumulator {
+    value_type: DataType,
+    pivot_index: Arc<HashMap<ScalarValue, usize>>,
+    slots: Vec<Option<ScalarValue>>,
+}
+
+impl PivotFirstAccumulator {
+    /// `num_slots` is the pivot list's length, which is what `state_fields` 
declares. It can
+    /// exceed `pivot_index.len()` when pivot values collide under `pivot_key` 
or are unmatchable
+    /// (NaN); those slots exist in the output and stay null.
+    fn new(
+        value_type: DataType,
+        pivot_index: Arc<HashMap<ScalarValue, usize>>,
+        num_slots: usize,
+    ) -> Self {
+        let slots = vec![None; num_slots];
+        Self {
+            value_type,
+            pivot_index,
+            slots,
+        }
+    }
+
+    /// Turn slot `i` into a `ScalarValue`, substituting a typed null when the 
slot is empty.
+    fn slot_or_null(&self, i: usize) -> DFResult<ScalarValue> {
+        Ok(match &self.slots[i] {
+            Some(v) => v.clone(),
+            None => ScalarValue::try_from(&self.value_type)?,
+        })
+    }
+}
+
+impl Accumulator for PivotFirstAccumulator {
+    fn update_batch(&mut self, values: &[ArrayRef]) -> DFResult<()> {
+        if values.len() != 2 {
+            return Err(DataFusionError::Internal(format!(
+                "PivotFirst expects 2 inputs (pivot, value); got {}",
+                values.len()
+            )));
+        }
+        let pivot_arr = &values[0];
+        let value_arr = &values[1];
+        if pivot_arr.len() != value_arr.len() {
+            return Err(DataFusionError::Internal(
+                "PivotFirst pivot and value arrays have different 
lengths".into(),
+            ));
+        }
+        for row in 0..pivot_arr.len() {
+            // Spark ignores the row entirely if either the pivot value is 
unmatched (index<0)
+            // or the value is null. Matching Spark exactly here is important 
because
+            // `PivotFirst.update` never writes for a null value, so a 
mid-batch null does not
+            // clobber an earlier non-null.
+            let pivot_scalar = ScalarValue::try_from_array(pivot_arr, row)?;
+            let Some(key) = pivot_key(pivot_scalar) else {
+                // A NaN pivot value matches no slot in Spark, so the row is 
ignored.
+                continue;
+            };
+            if let Some(&slot_idx) = self.pivot_index.get(&key) {

Review Comment:
   [P1] Match array pivot keys using Spark value semantics. `ScalarValue::List` 
equality includes Arrow field metadata, whereas Spark’s `TreeMap` uses 
interpreted value ordering. Pivot literals have nullable list children, but 
input arrays can have non-nullable children or different field names. Equal 
values therefore miss this lookup and silently produce null totals. This 
already breaks `FOR a IN (array(1, 1), array(2, 2))` in both Spark SQL CI 
shards. Nested signed zeros also fail to match. Please implement 
Spark-compatible complex-key comparison or make these pivot-column types fall 
back until supported.
   
   Evidence: Head-associated CI jobs 
https://github.com/apache/datafusion-comet/actions/runs/34362701066/job/102532509588
 and 
https://github.com/apache/datafusion-comet/actions/runs/34362701066/job/102530993096
 report pivot.sql query #25 expecting `(2012,35000,NULL)` and 
`(2013,NULL,78000)`, but receiving nulls in every pivot column. An exact-source 
Rust probe with literal `[1,1]` and identical input values returned 
`Int32(NULL)` when child nullability or field name differed, and `Int32(10)` 
when metadata matched. A `[0.0]` input also missed a `[-0.0]` pivot key, while 
Spark 4.1.3 returned 10.



##########
native/spark-expr/src/agg_funcs/pivot_first.rs:
##########
@@ -0,0 +1,497 @@
+// 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.
+
+//! Spark's `PivotFirst` aggregate. Used only by the second phase of the 
optimized pivot plan
+//! generated by `PivotTransformer`. For each group, `PivotFirst` maintains an 
array of
+//! `pivot_values.len()` slots; on each input row it evaluates the pivot 
column, looks up its
+//! index in `pivot_values`, and writes the value column into that slot when a 
match is found
+//! and the value is non-null. Rows with unmatched pivot values are ignored; 
matched rows with
+//! a null value column leave the slot unchanged (matches Spark).
+//!
+//! State layout is one column per pivot slot, matching Spark's 
`aggBufferAttributes` (which
+//! declares `indexSize` `AttributeReference`s, one per pivot value). This 
keeps the shuffle
+//! schema between Partial and Final consistent with what Spark catalyst 
declared; otherwise
+//! the shuffle exchange rejects the batch. `evaluate()` reassembles the slots 
into a
+//! `ListArray` matching `PivotFirst.dataType = ArrayType(value_type)`.
+
+use arrow::array::{Array, ArrayRef};
+use arrow::datatypes::{DataType, Field, FieldRef};
+use datafusion::common::utils::SingleRowListArrayBuilder;
+use datafusion::common::{DataFusionError, Result as DFResult, ScalarValue};
+use datafusion::logical_expr::function::{AccumulatorArgs, StateFieldsArgs};
+use datafusion::logical_expr::Volatility::Immutable;
+use datafusion::logical_expr::{Accumulator, AggregateUDFImpl, Signature};
+use datafusion::physical_expr::expressions::format_state_name;
+use std::collections::HashMap;
+use std::sync::Arc;
+
+/// UDAF implementation of Spark's `PivotFirst`.
+///
+/// `pivot_values` is a fixed, plan-time list of the pivot column values that 
occupy each
+/// output slot; `pivot_index[v] = i` means an input row whose pivot column 
equals `v` writes
+/// into slot `i`. Both the vector and the map are wrapped in `Arc` because 
`accumulator()`
+/// fires once per group in a grouped aggregate and we want that path to bump 
a refcount
+/// rather than deep-clone.
+#[derive(Debug)]
+pub struct SparkPivotFirst {
+    signature: Signature,
+    value_type: DataType,
+    // Kept for `PartialEq`/`Hash` (identity of the aggregate for plan 
comparison) and for the
+    // deterministic slot ordering `state_fields` needs. `HashMap` alone would 
give us the map
+    // but not a stable order or a `Hash` impl.
+    pivot_values: Arc<Vec<ScalarValue>>,
+    pivot_index: Arc<HashMap<ScalarValue, usize>>,
+}
+
+impl PartialEq for SparkPivotFirst {
+    fn eq(&self, other: &Self) -> bool {
+        self.value_type == other.value_type && self.pivot_values == 
other.pivot_values
+    }
+}
+
+impl Eq for SparkPivotFirst {}
+
+impl std::hash::Hash for SparkPivotFirst {
+    fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
+        self.value_type.hash(state);
+        self.pivot_values.hash(state);
+    }
+}
+
+impl SparkPivotFirst {
+    pub fn new(value_type: DataType, pivot_values: Vec<ScalarValue>) -> Self {
+        let mut pivot_index = HashMap::with_capacity(pivot_values.len());
+        // Spark's PivotFirst uses the FIRST occurrence's index 
(HashMap/TreeMap semantics), so
+        // when duplicates are somehow present we mirror that by only 
inserting the first one.
+        // `pivot_key` can fold two distinct pivot values (`0.0` and `-0.0`) 
onto one key, and can
+        // drop one entirely (NaN), so the index is not necessarily the same 
length as the slot
+        // vector - the slot count is always `pivot_values.len()`.
+        for (i, v) in pivot_values.iter().enumerate() {
+            if let Some(key) = pivot_key(v.clone()) {
+                pivot_index.entry(key).or_insert(i);
+            }
+        }
+        Self {
+            signature: Signature::user_defined(Immutable),
+            value_type,
+            pivot_values: Arc::new(pivot_values),
+            pivot_index: Arc::new(pivot_index),
+        }
+    }
+}
+
+/// Rewrite a pivot column value into the key Spark would match it on, or 
`None` when Spark can
+/// never match it.
+///
+/// Spark's `PivotFirst` looks pivot values up in a Scala `HashMap[Any, Int]`, 
so matching goes
+/// through `BoxesRunTime.equals` / `Statics.anyHash` on the boxed Catalyst 
value rather than
+/// through `ScalarValue`'s own equality. The two disagree on floats in 
opposite directions:
+///
+/// * `-0.0` and `0.0` are one key for Spark (`-0.0 == 0.0` numerically, and 
`doubleHash` folds
+///   both onto the hash of `0L`), while `ScalarValue` keeps them apart.
+/// * `NaN` matches nothing for Spark, not even another `NaN`, because Scala's 
`==` on `Double`
+///   is IEEE. `ScalarValue` treats `NaN` as equal to itself.
+///
+/// Nulls are left alone: a null pivot column value does match a null entry in 
the pivot list,
+/// which is what Spark's `pivotIndex.getOrElse(null, -1)` does.
+fn pivot_key(v: ScalarValue) -> Option<ScalarValue> {
+    match v {
+        ScalarValue::Float32(Some(f)) => {
+            if f.is_nan() {
+                None
+            } else if f == 0.0 {
+                Some(ScalarValue::Float32(Some(0.0)))
+            } else {
+                Some(ScalarValue::Float32(Some(f)))
+            }
+        }
+        ScalarValue::Float64(Some(f)) => {
+            if f.is_nan() {
+                None
+            } else if f == 0.0 {
+                Some(ScalarValue::Float64(Some(0.0)))
+            } else {
+                Some(ScalarValue::Float64(Some(f)))
+            }
+        }
+        other => Some(other),
+    }
+}
+
+impl AggregateUDFImpl for SparkPivotFirst {
+    fn name(&self) -> &str {
+        "pivot_first"
+    }
+
+    fn signature(&self) -> &Signature {
+        &self.signature
+    }
+
+    fn return_type(&self, _arg_types: &[DataType]) -> DFResult<DataType> {
+        Ok(DataType::List(Arc::new(Field::new_list_field(
+            self.value_type.clone(),
+            true,
+        ))))
+    }
+
+    fn state_fields(&self, args: StateFieldsArgs) -> DFResult<Vec<FieldRef>> {
+        // One field per pivot slot, matching Spark's aggBufferAttributes so 
the shuffle
+        // exchange sees the same schema catalyst declared. 
`format_state_name` is the same
+        // helper other aggregates in this crate use (see `avg.rs`, 
`stddev.rs`).
+        Ok((0..self.pivot_values.len())

Review Comment:
   [P2] Preserve Catalyst’s buffer width for duplicate pivot values. Spark 
sizes `aggBufferAttributes` from `pivotIndex.size`, while this code declares 
one state field per original list entry. With ANSI disabled, pivoting a row 
`(1,'a',10)` over `IN ('a','b','b')` succeeds in Spark with `(1,10,NULL,NULL)` 
and two aggregate-buffer fields. Native execution produces three buffer fields, 
making the partial aggregate incompatible with Spark’s declared shuffle/FFI 
schema and triggering column-count checks. Please fall back for duplicate lists 
until both the buffer layout and Spark’s last-occurrence index mapping are 
reproduced.
   
   Evidence: Spark 4.1.3 inspection reported `list 3`, `HashMap(b -> 2, a -> 
0)`, and `buffer_fields 2`, and returned `(1,10,NULL,NULL)`. The exact-source 
Rust probe built the aggregate expression and observed three state fields with 
`[Int64(10), Int64(NULL), Int64(NULL)]`. Combining its output with Catalyst’s 
schema failed with `number of columns(4) must match number of fields(3) in 
schema`. Comet emits raw partial buffers, and both `export_batch` and shuffle 
decoding enforce column counts.



-- 
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