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]