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]
