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]
