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


##########
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:
   [P2] Match Spark's floating-point-to-BIGINT boundary semantics when reusing 
the TRY cast. For a stored Variant column containing 
`CAST(double('9223372036854775808') AS VARIANT)`, `try_variant_get(v, '$', 
'bigint')` returns `9223372036854775807` in Spark but NULL here. A FLOAT 
Variant containing `2^63` has the same mismatch. Spark's interpreted cast 
accepts this rounded boundary and saturates to `Long.MaxValue`; Comet's TRY 
path instead delegates to Arrow's stricter conversion. Strict extraction 
consequently fails a query that Spark accepts. Adapt this conversion to Spark's 
range check and saturation, or keep the affected output on fallback.
   
   Evidence: A disposable native test at the reviewed head passed 
`Variant::Double(9223372036854775808.0)` and the corresponding `Variant::Float` 
through `evaluate_array`. Both produced `Int64(NULL)`, and strict mode returned 
`InvalidVariantCast`. Spark 4.1.3 and 4.2.0 runtime probes returned 
`9223372036854775807` for both input types in both modes. Spark's 
`FloatExactNumeric.toLong` and `DoubleExactNumeric.toLong` implement this 
boundary behavior. Comet's shared cast bypasses its Spark-specific numeric 
conversion when `eval_mode == EvalMode::Try`.



##########
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:
   [P2] Preserve Spark's hexadecimal floating-point parsing before enabling 
these targets as `Compatible()`. With 
`spark.sql.variant.pushVariantIntoScan=false`, store `parse_json('"0x1.0p0"')` 
in a Parquet Variant column `v`, then select `try_variant_get(v, '$', 
'double')`. Spark returns `1.0`, while the newly admitted native expression 
returns NULL. FLOAT behaves identically. Strict `variant_get` also produces a 
native cast failure despite Spark accepting the value. The reused string cast 
relies on Rust parsing, which rejects Java's hexadecimal syntax. This changes 
previously correct fallback results. Match Spark's parser or retain fallback 
for affected targets.
   
   Evidence: At the reviewed head, a disposable test invoking 
`VariantGet::evaluate_array` with `Variant::String("0x1.0p0")` returned 
`Float32(NULL)` and `Float64(NULL)` in TRY mode and `InvalidVariantCast` in 
strict mode. Direct `VariantGet.variantGet` probes on Spark 4.1.3 and 4.2.0 
returned `1.0` for both targets and both modes. Spark's `Cast.scala` uses 
Java-compatible `toFloat`/`toDouble`; Comet's 
`conversion_funcs/string.rs::parse_string_to_float` uses Rust `parse`.



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