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


##########
native/spark-expr/src/kernels/temporal.rs:
##########
@@ -1260,154 +1213,109 @@ pub(crate) fn timestamp_trunc_array_fmt_dyn(
     array: &dyn Array,
     formats: &dyn Array,
 ) -> Result<ArrayRef, SparkError> {
-    match (array.data_type().clone(), formats.data_type().clone()) {
-        (DataType::Dictionary(_, _), DataType::Dictionary(_, _)) => {
-            downcast_dictionary_array!(
-                formats => {
-                    downcast_dictionary_array!(
-                        array => {
-                            timestamp_trunc_array_fmt_dict_dict(
-                                    
&array.downcast_dict::<TimestampMicrosecondArray>().unwrap(),
-                                    
&formats.downcast_dict::<StringArray>().unwrap())
-                            .map(|a| Arc::new(a) as ArrayRef)
-                        }
-                        dt => return_compute_error_with!("timestamp_trunc does 
not support", dt)
-                    )
-                }
-                fmt => return_compute_error_with!("timestamp_trunc does not 
support format type", fmt),
-            )
+    let array = unpack_dictionary(array)?;
+    let formats = unpack_dictionary(formats)?;
+    let array = match array.data_type() {
+        DataType::Timestamp(TimeUnit::Microsecond, _) => {
+            array.as_primitive::<TimestampMicrosecondType>()
         }
-        (DataType::Dictionary(_, _), DataType::Utf8) => {
-            downcast_dictionary_array!(
-                array => {
-                  timestamp_trunc_array_fmt_dict_plain(
-                        
&array.downcast_dict::<PrimitiveArray<TimestampMicrosecondType>>().unwrap(),
-                        formats.as_any().downcast_ref::<StringArray>()
-                            .expect("Unexpected value type in formats"))
-                  .map(|a| Arc::new(a) as ArrayRef)
-                }
-                dt => return_compute_error_with!("timestamp_trunc does not 
support", dt),
-            )
-        }
-        (DataType::Timestamp(TimeUnit::Microsecond, _), 
DataType::Dictionary(_, _)) => {
-            downcast_dictionary_array!(
-                formats => {
-                downcast_temporal_array!(array => {
-                        timestamp_trunc_array_fmt_plain_dict(
-                                array,
-                                
&formats.downcast_dict::<StringArray>().unwrap())
-                        .map(|a| Arc::new(a) as ArrayRef)
-                    }
-                    dt => return_compute_error_with!("timestamp_trunc does not 
support", dt),
-                    )
-                }
-                fmt => return_compute_error_with!("timestamp_trunc does not 
support format type", fmt),
-            )
-        }
-        (DataType::Timestamp(TimeUnit::Microsecond, _), DataType::Utf8) => {
-            downcast_temporal_array!(
-                array => {
-                    timestamp_trunc_array_fmt_plain_plain(array,
-                        
formats.as_any().downcast_ref::<StringArray>().expect("Unexpected value type in 
formats"))
-                    .map(|a| Arc::new(a) as ArrayRef)
-                },
-                dt => return_compute_error_with!("timestamp_trunc does not 
support", dt),
-            )
+        dt => {
+            return_compute_error_with!("Unsupported input type for function 
'timestamp_trunc'", dt)
         }
-        (dt, fmt) => Err(SparkError::Internal(format!(
-            "Unsupported datatype: {dt:}, format: {fmt:?} for function 
'timestamp_trunc'"
-        ))),
-    }
+    };
+    let formats = match formats.data_type() {
+        DataType::Utf8 => formats.as_string::<i32>(),
+        fmt => return_compute_error_with!("timestamp_trunc does not support 
format type", fmt),
+    };
+    Ok(Arc::new(timestamp_trunc_by_row_format(array, formats)?))
 }
 
-macro_rules! timestamp_trunc_array_fmt_helper {
-    ($array: ident, $formats: ident, $datatype: ident) => {{
-        let mut builder = 
TimestampMicrosecondBuilder::with_capacity($array.len());
-        let iter = $array.into_iter();
-        assert_eq!(
-            $array.len(),
-            $formats.len(),
-            "lengths of values array and format array must be the same"
-        );
-        match $datatype {
-            DataType::Timestamp(TimeUnit::Microsecond, None) => {
-                // TimestampNTZ: operate directly on naive microsecond values
-                for (index, val) in iter.enumerate() {
-                    let micros_val = val.map(|v| i64::from(v));
-                    let trunc_fn = 
ntz_trunc_fn_for_format($formats.value(index))?;
-                    timestamp_trunc_ntz_single(micros_val, &mut builder, 
trunc_fn)?;
-                }
-                Ok(builder.finish())
-            }
-            DataType::Timestamp(TimeUnit::Microsecond, Some(tz_str)) => {
-                let tz: Tz = tz_str.parse()?;
-                for (index, val) in iter.enumerate() {
-                    let trunc_fn = 
tz_trunc_fn_for_format($formats.value(index))?;
-                    as_timestamp_tz_with_op_single::<T, _>(val, &mut builder, 
&tz, |dt| {
-                        as_micros_from_unix_epoch_utc(trunc_fn(dt))
-                    })?;
-                }
-                Ok(builder.finish().with_timezone(tz_str.as_ref()))
-            }
-            dt => {
-                return_compute_error_with!(
-                    "Unsupported input type '{:?}' for function 
'timestamp_trunc'",
-                    dt
-                )
-            }
-        }
-    }};
+fn unpack_dictionary(array: &dyn Array) -> Result<ArrayRef, SparkError> {
+    match array.data_type() {
+        DataType::Dictionary(_, value_type) => arrow::compute::cast(array, 
value_type)
+            .map_err(|error| SparkError::Internal(error.to_string())),
+        _ => Ok(make_array(array.to_data())),
+    }
 }
 
-fn timestamp_trunc_array_fmt_plain_plain<T>(
-    array: &PrimitiveArray<T>,
+/// Truncates each row with the format in the same row.
+///
+/// Rows are grouped by format, and each group goes through the literal-format 
kernel with the
+/// other rows set to NULL. A row is therefore truncated by exactly the rules 
a literal format
+/// applies, including the DST handling, and a value in one group cannot make 
another group fail.
+/// A NULL format gives a NULL result, as in Spark.
+fn timestamp_trunc_by_row_format(
+    array: &TimestampMicrosecondArray,
     formats: &StringArray,
-) -> Result<TimestampMicrosecondArray, SparkError>
-where
-    T: ArrowTemporalType + ArrowNumericType,
-    i64: From<T::Native>,
-{
-    let data_type = array.data_type();
-    timestamp_trunc_array_fmt_helper!(array, formats, data_type)
-}
-fn timestamp_trunc_array_fmt_plain_dict<T, K>(
-    array: &PrimitiveArray<T>,
-    formats: &TypedDictionaryArray<K, StringArray>,
-) -> Result<TimestampMicrosecondArray, SparkError>
-where
-    T: ArrowTemporalType + ArrowNumericType,
-    i64: From<T::Native>,
-    K: ArrowDictionaryKeyType,
-{
-    let data_type = array.data_type();
-    timestamp_trunc_array_fmt_helper!(array, formats, data_type)
-}
+) -> Result<TimestampMicrosecondArray, SparkError> {
+    if array.len() != formats.len() {
+        return Err(SparkError::Internal(format!(
+            "timestamp_trunc has {} values but {} formats",
+            array.len(),
+            formats.len()
+        )));
+    }
 
-fn timestamp_trunc_array_fmt_dict_plain<T, K>(
-    array: &TypedDictionaryArray<K, PrimitiveArray<T>>,
-    formats: &StringArray,
-) -> Result<TimestampMicrosecondArray, SparkError>
-where
-    T: ArrowTemporalType + ArrowNumericType,
-    i64: From<T::Native>,
-    K: ArrowDictionaryKeyType,
-{
-    let data_type = array.values().data_type();
-    timestamp_trunc_array_fmt_helper!(array, formats, data_type)
-}
+    let mut granularities: Vec<&'static str> = Vec::new();
+    // A batch rarely holds more than a few spellings, so parse each one once.
+    let mut spellings: Vec<(&str, usize)> = Vec::new();
+    let mut row_groups: Vec<Option<usize>> = Vec::with_capacity(array.len());
+    let mut null_format_hides_value = false;
+    for index in 0..array.len() {
+        if array.is_null(index) {
+            row_groups.push(None);
+            continue;
+        }
+        if formats.is_null(index) {
+            null_format_hides_value = true;
+            row_groups.push(None);
+            continue;
+        }
+        let spelling = formats.value(index);
+        let group = match spellings.iter().find(|(seen, _)| *seen == spelling) 
{
+            Some((_, group)) => *group,
+            None => {
+                let granularity = normalize_timestamp_trunc_format(spelling)?;
+                let group = match granularities.iter().position(|g| *g == 
granularity) {
+                    Some(group) => group,
+                    None => {
+                        granularities.push(granularity);
+                        granularities.len() - 1
+                    }
+                };
+                spellings.push((spelling, group));
+                group
+            }
+        };
+        row_groups.push(Some(group));
+    }
 
-fn timestamp_trunc_array_fmt_dict_dict<T, K, F>(
-    array: &TypedDictionaryArray<K, PrimitiveArray<T>>,
-    formats: &TypedDictionaryArray<F, StringArray>,
-) -> Result<TimestampMicrosecondArray, SparkError>
-where
-    T: ArrowTemporalType + ArrowNumericType,
-    i64: From<T::Native>,
-    K: ArrowDictionaryKeyType,
-    F: ArrowDictionaryKeyType,
-{
-    let data_type = array.values().data_type();
-    timestamp_trunc_array_fmt_helper!(array, formats, data_type)
+    if granularities.len() == 1 && !null_format_hides_value {
+        return timestamp_trunc_upstream(array, granularities[0]);

Review Comment:
   [P2] Preserve Spark 4.2's checked overflow behavior when routing column 
formats here. With a UTC session, 
`spark.comet.expression.TruncTimestamp.allowIncompatible=true`, and columns 
containing `us=-9223372036854775808` and `fmt='SECOND'`, 
`unix_micros(date_trunc(fmt, timestamp_micros(us)))` now yields 
`9223372036854551616`. Spark 4.2 raises `ArithmeticException: long overflow`, 
and the base column path also errors. `MILLISECOND` similarly returns 
`9223372036854775616`. Although the scalar mismatch predates this PR, this new 
delegation materially worsens the column path by turning an error into a 
silently wrapped timestamp. Could you thread the version-aware overflow policy 
described for #6740 through both the single-group and masked-group calls, 
retaining wrapping for Spark 3.4–4.1, and add a Spark 4.2 column-format 
regression?
   
   Evidence: Compiled the unchanged base/head temporal kernels against locked 
Arrow 59.3.0, chrono 0.4.45 and DataFusion 55.1.0. For 
TimestampMicrosecondArray([i64::MIN]).with_timezone("UTC") and 
StringArray(["SECOND"]), the base returns Err("Unable to read value as 
datetime") and head returns 9223372036854551616. MILLISECOND returns 
9223372036854775616 at head. Spark 4.2.0 DateTimeUtils.truncTimestamp and 
TruncTimestamp.eval with BoundReference column inputs both throw 
ArithmeticException: long overflow for both formats. Spark 4.2 source uses 
Math.subtractExact, while 3.4–4.1 use wrapping subtraction.



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