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


##########
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:
   Fixed in 6964f9237, after #6740 landed. The column-format path now passes 
`wrap_second_millisecond_overflow` to both calls into 
`timestamp_trunc_upstream`, so like the literal path it wraps before Spark 4.2 
and raises `long overflow` from 4.2 on. 
`row_formats_follow_the_overflow_policy` covers a single-format column and a 
mixed one in both modes, and both of #6740's overflow fixtures now have a 
column-format case. On Spark 4.2 the column queries raise `long overflow`, and 
forcing the column path to wrap makes that fixture fail.



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