peterxcli commented on code in PR #6595:
URL: https://github.com/apache/datafusion-comet/pull/6595#discussion_r4187166769


##########
spark/src/main/spark-4.x/org/apache/comet/serde/variant.scala:
##########
@@ -0,0 +1,115 @@
+/*
+ * 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.
+ */
+
+package org.apache.comet.serde
+
+import org.apache.spark.SparkThrowable
+import org.apache.spark.sql.catalyst.expressions.{Attribute, 
AttributeReference, Literal}
+import org.apache.spark.sql.catalyst.expressions.variant.{ArrayExtraction, 
ObjectExtraction, VariantGet}
+import org.apache.spark.sql.types._
+
+import org.apache.comet.serde.ExprOuterClass.Expr
+import org.apache.comet.serde.QueryPlanSerde.serializeDataType
+
+object CometVariantGet extends CometExpressionSerde[VariantGet] {
+  private val pathReason = "Variant extraction requires a non-null foldable 
path."
+  private val invalidPathReason =
+    "Invalid Variant paths require Spark's null and evaluation semantics."
+  private val inputReason = "Variant extraction requires a top-level Variant 
column or literal."
+  private val targetReason =
+    "Variant extraction supports Boolean, numeric, binary, date and timestamp 
targets only."
+  private val stringReason =
+    "Variant STRING extraction requires Spark-compatible JSON and scalar 
formatting " +
+      "(https://github.com/apache/datafusion-comet/issues/5424)."
+  private val decimalReason =
+    "Floating-point to decimal rounding can differ from Spark on JDK 17 " +
+      "(https://github.com/apache/datafusion-comet/issues/5424)."
+  private val temporalReason =
+    "Date/time parsing and timezone conversion support a narrower year range 
than Spark " +
+      "(https://github.com/apache/datafusion-comet/issues/5424)."
+
+  override def getUnsupportedReasons(): Seq[String] =
+    Seq(
+      pathReason,
+      invalidPathReason,
+      inputReason,
+      targetReason,
+      stringReason,
+      "Timezones with second-resolution offsets cannot be represented in 
native code.")
+  override def getIncompatibleReasons(): Seq[String] = Seq(decimalReason, 
temporalReason)
+
+  override def getSupportLevel(expr: VariantGet): SupportLevel = {
+    if (!expr.path.foldable || expr.path.eval() == null) {
+      Unsupported(Some(pathReason))
+    } else if (!hasValidPath(expr)) {
+      Unsupported(Some(invalidPathReason))
+    } else if (!(expr.child.isInstanceOf[AttributeReference] || expr.child
+        .isInstanceOf[Literal]) ||
+      expr.child.dataType != VariantType) {
+      Unsupported(Some(inputReason))
+    } else if (CometTimeZone.nativeId(expr.timeZoneId).isEmpty) {
+      CometTimeZone.supportLevel(expr.timeZoneId)
+    } else {
+      expr.dataType match {
+        case BooleanType | ByteType | ShortType | IntegerType | LongType | 
FloatType |
+            DoubleType | BinaryType =>
+          Compatible()

Review Comment:
   Fixed in 9ddf60b19. The shared cast now parses hexadecimal strings at the 
target width, including FLOAT rounding without an intermediate DOUBLE. Added 
stored-column SQL parity tests for both extraction modes and dictionary 
settings, plus native bit checks for rounding, signed zero, subnormals and 
overflow. These pass on Spark 4.1. The temporary binary expansion links to the 
parser bug, Alexhuszagh/rust-lexical#87.



##########
native/core/src/execution/expressions/variant_get.rs:
##########
@@ -0,0 +1,632 @@
+// 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.
+
+use std::{
+    collections::HashMap,
+    fmt::{Display, Formatter},
+    hash::{Hash, Hasher},
+    sync::Arc,
+};
+
+use arrow::{
+    array::{new_null_array, timezone::Tz, Array, ArrayRef, AsArray},
+    compute::interleave,
+    datatypes::{i256, DataType, Schema, TimeUnit},
+    record_batch::RecordBatch,
+};
+use base64::{engine::general_purpose::STANDARD, Engine};
+use datafusion::{
+    common::{exec_err, plan_err, DataFusionError, Result, ScalarValue},
+    logical_expr::ColumnarValue,
+    physical_expr::PhysicalExpr,
+};
+use datafusion_comet_common::{decode_utf8_spark_lossy, SparkError};
+use datafusion_comet_proto::spark_expression::{self, 
variant_path_segment::Segment};
+use datafusion_comet_spark_expr::{spark_cast, EvalMode, SparkCastOptions};
+use parquet::variant::{Variant, VariantType};
+
+use crate::execution::serde::to_arrow_datatype;
+
+#[derive(Debug, Clone, PartialEq, Eq, Hash)]
+enum PathSegment {
+    Key(String),
+    Index(usize),
+}
+
+#[derive(Debug, Eq)]
+pub struct VariantGet {
+    child: Arc<dyn PhysicalExpr>,
+    path: Vec<PathSegment>,
+    path_sql: String,
+    target: DataType,
+    target_sql: String,
+    fail_on_error: bool,
+    size_limit: usize,
+    cast_options: SparkCastOptions,
+}
+
+impl PartialEq for VariantGet {
+    fn eq(&self, other: &Self) -> bool {
+        self.child.eq(&other.child)
+            && self.path == other.path
+            && self.path_sql == other.path_sql
+            && self.target == other.target
+            && self.target_sql == other.target_sql
+            && self.fail_on_error == other.fail_on_error
+            && self.size_limit == other.size_limit
+            && self.cast_options == other.cast_options
+    }
+}
+
+impl Hash for VariantGet {
+    fn hash<H: Hasher>(&self, state: &mut H) {
+        self.child.hash(state);
+        self.path.hash(state);
+        self.path_sql.hash(state);
+        self.target.hash(state);
+        self.target_sql.hash(state);
+        self.fail_on_error.hash(state);
+        self.size_limit.hash(state);
+        self.cast_options.hash(state);
+    }
+}
+
+impl VariantGet {
+    pub fn try_new(
+        child: Arc<dyn PhysicalExpr>,
+        expr: &spark_expression::VariantGet,
+        schema: &Schema,
+    ) -> Result<Self> {
+        if !child
+            .return_field(schema)?
+            .has_valid_extension_type::<VariantType>()
+        {
+            return plan_err!("variant_get requires a Variant extension field");
+        }
+        let target =
+            to_arrow_datatype(expr.datatype.as_ref().ok_or_else(|| {
+                DataFusionError::Plan("variant_get requires a target 
type".into())
+            })?);
+        if !matches!(
+            target,
+            DataType::Boolean
+                | DataType::Int8
+                | DataType::Int16
+                | DataType::Int32
+                | DataType::Int64
+                | DataType::Float32
+                | DataType::Float64
+                | DataType::Decimal128(_, _)
+                | DataType::Binary
+                | DataType::Date32
+                | DataType::Timestamp(TimeUnit::Microsecond, _)
+        ) {
+            return plan_err!("Unsupported variant_get target: {target}");
+        }
+        if expr.size_limit == 0 {
+            return plan_err!("variant_get requires a positive Variant size 
limit");
+        }
+        let path = expr
+            .path
+            .iter()
+            .map(|part| match &part.segment {
+                Some(Segment::Key(key)) => Ok(PathSegment::Key(key.clone())),
+                Some(Segment::Index(index)) => Ok(PathSegment::Index(*index as 
usize)),
+                None => plan_err!("Missing variant_get path segment"),
+            })
+            .collect::<Result<Vec<_>>>()?;
+        // Validate the zone during planning. Spark's analyzer has already 
resolved it.
+        expr.timezone.parse::<Tz>()?;
+        Ok(Self {
+            child,
+            path,
+            path_sql: expr.path_sql.clone(),
+            target,
+            target_sql: expr.target_sql.clone(),
+            fail_on_error: expr.fail_on_error,
+            size_limit: expr.size_limit as usize,
+            cast_options: SparkCastOptions::new_with_version(
+                EvalMode::Try,
+                &expr.timezone,
+                true,
+                true,
+            ),
+        })
+    }
+
+    fn extract<'v>(&self, metadata: &[u8], value: &'v [u8]) -> 
Result<Option<&'v [u8]>> {
+        if metadata.first().is_none_or(|header| header & 15 != 1) {
+            return Err(SparkError::MalformedVariant.into());
+        }
+        if metadata.len() > self.size_limit || value.len() > self.size_limit {
+            return Err(SparkError::VariantConstructorSizeLimit.into());
+        }
+        let mut position = 0_i32;
+        for part in &self.path {
+            let header = checked_header(
+                value
+                    .get(checked_index(position)?..)
+                    .ok_or(SparkError::MalformedVariant)?,
+            )?;
+            let info = header >> 2;
+            let (data_start, offset) =
+                match (part, header & 3) {
+                    (PathSegment::Key(key), 2) => {
+                        let size_bytes = if info & 16 == 0 { 1 } else { 4 };
+                        let size =
+                            unsigned(value, 
checked_index(position.wrapping_add(1))?, size_bytes)?;
+                        let id_width = ((info >> 2) & 3) as usize + 1;
+                        let offset_width = (info & 3) as usize + 1;
+                        // Spark uses absolute positions and Java int 
arithmetic, including wrap.
+                        // All accesses below still check the resulting index 
against the buffer.
+                        let ids = position.wrapping_add(1 + size_bytes as i32);
+                        let offsets = ids.wrapping_add((size as 
i32).wrapping_mul(id_width as i32));
+                        let data = offsets.wrapping_add(
+                            (size as i32)
+                                .wrapping_add(1)
+                                .wrapping_mul(offset_width as i32),
+                        );
+                        let key_at = |i: usize| {
+                            metadata_key(
+                                metadata,
+                                unsigned(
+                                    value,
+                                    checked_index(
+                                        ids.wrapping_add((i as 
i32).wrapping_mul(id_width as i32)),
+                                    )?,
+                                    id_width,
+                                )?,
+                            )
+                        };
+                        // Match Spark 4.0/4.1/4.2's small-object linear 
lookup and UTF-16 binary
+                        // lookup. Do not decode or validate values outside 
the requested path.
+                        let index = if size < 32 {
+                            let mut found = None;
+                            for i in 0..size {
+                                if key_at(i)?.as_ref() == key {
+                                    found = Some(i);
+                                    break;
+                                }
+                            }
+                            found
+                        } else {
+                            let (mut low, mut high) = (0, size);
+                            let mut found = None;
+                            while low < high {
+                                let mid = low + (high - low - 1) / 2;
+                                match 
key_at(mid)?.encode_utf16().cmp(key.encode_utf16()) {
+                                    std::cmp::Ordering::Less => low = mid + 1,
+                                    std::cmp::Ordering::Greater => high = mid,
+                                    std::cmp::Ordering::Equal => {
+                                        found = Some(mid);
+                                        break;
+                                    }
+                                }
+                            }
+                            found
+                        };
+                        let Some(index) = index else {
+                            return Ok(None);
+                        };
+                        (
+                            data,
+                            unsigned(
+                                value,
+                                checked_index(offsets.wrapping_add(
+                                    (index as i32).wrapping_mul(offset_width 
as i32),
+                                ))?,
+                                offset_width,
+                            )?,
+                        )
+                    }
+                    (PathSegment::Index(index), 3) => {
+                        let size_bytes = if info & 4 == 0 { 1 } else { 4 };
+                        let size =
+                            unsigned(value, 
checked_index(position.wrapping_add(1))?, size_bytes)?;
+                        if *index >= size {
+                            return Ok(None);
+                        }
+                        let offset_width = (info & 3) as usize + 1;
+                        let offsets = position.wrapping_add(1 + size_bytes as 
i32);
+                        (
+                            offsets.wrapping_add(
+                                (size as i32)
+                                    .wrapping_add(1)
+                                    .wrapping_mul(offset_width as i32),
+                            ),
+                            unsigned(
+                                value,
+                                checked_index(offsets.wrapping_add(
+                                    (*index as i32).wrapping_mul(offset_width 
as i32),
+                                ))?,
+                                offset_width,
+                            )?,
+                        )
+                    }
+                    _ => return Ok(None),
+                };
+            position = data_start.wrapping_add(offset as i32);
+        }
+        let selected = value
+            .get(checked_index(position)?..)
+            .ok_or(SparkError::MalformedVariant)?;
+        checked_header(selected)?;
+        Ok(Some(selected))
+    }
+
+    fn evaluate_array(&self, input: &ArrayRef) -> Result<ArrayRef> {
+        let array = input.as_struct_opt().ok_or_else(|| {
+            DataFusionError::Execution("Variant input must use struct 
storage".into())
+        })?;
+        let values = array
+            .column_by_name("value")
+            .and_then(|a| a.as_binary_opt::<i32>());
+        let metadata = array
+            .column_by_name("metadata")
+            .and_then(|a| a.as_binary_opt::<i32>());
+        let (Some(values), Some(metadata)) = (values, metadata) else {
+            return exec_err!("Variant input must contain Binary value and 
metadata children");
+        };
+        // Cast together values with the same source type. This avoids 
invoking Arrow's cast
+        // machinery and allocating an array for each row in a heterogeneous 
Variant column.
+        let mut groups: Vec<Vec<ScalarValue>> = vec![vec![]];
+        let mut group_types = HashMap::new();
+        let mut positions = Vec::with_capacity(array.len());
+        let mut extracted = Vec::with_capacity(array.len());
+        for row in 0..array.len() {
+            let value = if array.is_null(row) {
+                None
+            } else {
+                if values.is_null(row) || metadata.is_null(row) {
+                    return Err(SparkError::MalformedVariant.into());
+                }
+                // Use shallow decoding: Spark validates only the selected 
path. Full Arrow
+                // validation also rejects legacy Spark key order and empty 
dictionary keys.
+                self.extract(metadata.value(row), values.value(row))?
+            };
+            let scalar = match value.as_ref() {
+                None | Some(&[0, ..]) => None,
+                Some(value) => self.scalar_for_cast(value)?,
+            };
+            if let Some(scalar) = scalar {
+                let data_type = scalar.data_type();
+                let group = *group_types.entry(data_type).or_insert_with(|| {
+                    groups.push(vec![]);
+                    groups.len() - 1
+                });
+                positions.push((group, groups[group].len()));
+                groups[group].push(scalar);
+            } else {
+                positions.push((0, 0));
+            }
+            extracted.push(value);
+        }
+        let mut casted = vec![new_null_array(&self.target, 1)];
+        for group in groups.into_iter().skip(1) {
+            let source = ScalarValue::iter_to_array(group)?;
+            let converted = spark_cast(

Review Comment:
   Fixed in 9ddf60b19. TRY now uses the existing Spark range check and 
saturating conversion for floating-point-to-integral casts. Both extraction 
modes return Long.MaxValue for FLOAT/DOUBLE 2^63. Native and stored-column SQL 
regressions cover the boundary, adjacent overflow values and nulls; the 
existing ordinary cast boundary tests also pass on Spark 4.1. Could a committer 
add `run-spark-4.1-tests` for the shared cast changes? My account cannot apply 
labels.



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