Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 7 additions & 8 deletions docs/source/user-guide/latest/datatypes.md
Original file line number Diff line number Diff line change
Expand Up @@ -105,14 +105,13 @@ functions. Remaining work is tracked by

| Type | Status | Notes |
| ------------- | ------ | ------------------------------------------------------------------------------------------------------------------------------- |
| `VariantType` | ⚠️ | Spark 4.0+. Native Parquet scans support direct projection of top-level Variant columns. Non-null existence defaults fall back. |

Direct projection requires explicit configuration on every supported Spark version:
`spark.sql.variant.allowReadingShredded=true` (defaults to false in Spark 4.0) and
`spark.sql.variant.pushVariantIntoScan=false` (defaults to true in Spark 4.1+), with the default
Parquet timestamp inference settings. Support for Spark's whole-value pushdown rewrite is tracked
by [#5519](https://github.com/apache/datafusion-comet/issues/5519). Nested Variant columns, pushed-down
Variant field extraction, expressions, writes, shuffle and spill, Python operators, encrypted
| `VariantType` | ⚠️ | Spark 4.0+. Native Parquet scans support whole-value reads of top-level Variant columns. Non-null existence defaults fall back. |

Whole-value reads require `spark.sql.variant.allowReadingShredded=true` (defaults to false in
Spark 4.0) and the default Parquet timestamp inference settings. Spark's
`spark.sql.variant.pushVariantIntoScan` rewrite is supported when it requests only the whole value
of each Variant column. Nested Variant columns, pushed-down Variant field extraction,
expressions, writes, shuffle and spill, Python operators, encrypted
files, and Iceberg scans that read a Variant column fall back to Spark. Iceberg scans of tables
whose Variant columns the query does not read run natively. Spark also handles columnar-to-row
conversion of the native scan output and strict reads with `allowReadingShredded=false`. Broader
Expand Down
2 changes: 1 addition & 1 deletion native/core/src/execution/serde.rs
Original file line number Diff line number Diff line change
Expand Up @@ -181,7 +181,7 @@ pub fn to_arrow_datatype(dt_value: &DataType) -> ArrowDataType {
&info.field_datatypes[idx],
info.field_nullable[idx],
);
// Attach Spark field metadata (currently parquet.field.id) when present.
// Attach Spark field IDs and Variant request metadata when present.
// field_metadata is parallel to field_names; either empty or full length.
if let Some(meta) = info.field_metadata.get(idx) {
if !meta.metadata.is_empty() {
Expand Down
48 changes: 42 additions & 6 deletions native/core/src/parquet/cast_column.rs
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,7 @@ use self::variant::normalize_variant_array;
use arrow::{
array::{make_array, Array, ArrayRef, LargeListArray, ListArray, MapArray, StructArray},
compute::CastOptions,
datatypes::{DataType, FieldRef, Schema, TimeUnit},
datatypes::{DataType, Field, FieldRef, Schema, TimeUnit},
record_batch::RecordBatch,
};

Expand All @@ -36,6 +36,27 @@ use std::{
sync::Arc,
};

/// The Variant to reconstruct for a direct projection or Spark's full-value scan request.
pub(crate) fn variant_projection_field(field: &Field) -> Option<FieldRef> {
if field.has_valid_extension_type::<VariantType>() {
return Some(Arc::new(field.clone()));
}
let DataType::Struct(fields) = field.data_type() else {
return None;
};
if fields.len() != 1 || fields[0].name() != "0" {
return None;
}
let child = &fields[0];
if !child.has_valid_extension_type::<VariantType>() {
return None;
}
let metadata: serde_json::Value =
serde_json::from_str(child.metadata().get("__VARIANT_METADATA_KEY")?).ok()?;
(metadata["path"] == "$" && metadata["failOnError"] == true && metadata["timeZoneId"] == "UTC")
.then(|| Arc::clone(child))
}

/// Returns true if two DataTypes are structurally equivalent (same data layout)
/// but may differ in field names within nested types. With `use_field_id`, a struct
/// field that carries a Parquet field id must also find that id on the file field at
Expand Down Expand Up @@ -175,6 +196,8 @@ pub struct CometCastColumnExpr {
input_physical_field: FieldRef,
/// The field type required by query
target_field: FieldRef,
/// Derived once so request metadata is not parsed for each batch.
variant_field: Option<FieldRef>,
/// Options forwarded to [`cast_column`].
cast_options: CastOptions<'static>,
/// Spark parquet options for complex nested type conversions.
Expand Down Expand Up @@ -242,6 +265,7 @@ impl CometCastColumnExpr {
Ok(Self {
expr,
input_physical_field: physical_field,
variant_field: variant_projection_field(&target_field),
target_field,
cast_options: cast_options.unwrap_or(DEFAULT_CAST_OPTIONS),
parquet_options: None,
Expand Down Expand Up @@ -283,12 +307,24 @@ impl PhysicalExpr for CometCastColumnExpr {
fn evaluate(&self, batch: &RecordBatch) -> DataFusionResult<ColumnarValue> {
let value = self.expr.evaluate(batch)?;

if self.target_field.has_valid_extension_type::<VariantType>() {
if let Some(variant_field) = &self.variant_field {
return match value {
ColumnarValue::Array(array) => Ok(ColumnarValue::Array(normalize_variant_array(
&array,
&self.target_field,
)?)),
ColumnarValue::Array(array) => {
let normalized = normalize_variant_array(&array, variant_field)?;
if self.target_field.has_valid_extension_type::<VariantType>() {
Ok(ColumnarValue::Array(normalized))
} else {
let DataType::Struct(fields) = self.target_field.data_type() else {
unreachable!();
};
let nulls = normalized.nulls().cloned();
Ok(ColumnarValue::Array(Arc::new(StructArray::try_new(
fields.clone(),
vec![normalized],
nulls,
)?)))
}
}
ColumnarValue::Scalar(_) => Err(DataFusionError::Execution(
"Variant Parquet projection requires an array".to_string(),
)),
Expand Down
3 changes: 1 addition & 2 deletions native/core/src/parquet/parquet_exec.rs
Original file line number Diff line number Diff line change
Expand Up @@ -40,7 +40,6 @@ use datafusion::scalar::ScalarValue;
use datafusion_comet_spark_expr::jvm_udf::JvmScalarUdfExpr;
use datafusion_comet_spark_expr::EvalMode;
use datafusion_datasource::TableSchema;
use parquet::variant::VariantType;
use std::collections::HashMap;
use std::sync::Arc;

Expand Down Expand Up @@ -156,7 +155,7 @@ pub(crate) fn init_datasource_exec(
let projects_variant = required_schema
.fields()
.iter()
.any(|field| field.has_valid_extension_type::<VariantType>());
.any(|field| super::cast_column::variant_projection_field(field).is_some());
if projects_variant && encryption_enabled {
return Err(ExecutionError::GeneralError(
"Projected Variant with Parquet encryption requires Spark fallback".to_string(),
Expand Down
78 changes: 75 additions & 3 deletions native/core/src/parquet/parquet_exec/variant_tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -34,7 +34,10 @@ use parquet::{
},
file::{properties::WriterProperties, writer::SerializedFileWriter},
schema::types::{Type as ParquetType, TypePtr},
variant::{Variant, VariantArray, VariantBuilder, VariantDecimal4},
variant::{
shred_variant, Variant, VariantArray, VariantArrayBuilder, VariantBuilder,
VariantBuilderExt, VariantDecimal4, VariantType,
},
};
use std::{fs::File, path::PathBuf};
fn required_variant_schema() -> SchemaRef {
Expand Down Expand Up @@ -141,11 +144,16 @@ fn write_variant_typed_value<T: ParquetDataType>(typed_value: TypePtr, values: &
}

async fn scan_variant_file(filename: PathBuf) -> VariantArray {
let batch = scan_variant_batch(filename, required_variant_schema()).await;
VariantArray::try_new(batch.column(0).as_ref()).unwrap()
}

async fn scan_variant_batch(filename: PathBuf, required_schema: SchemaRef) -> RecordBatch {
let partitioned_file =
PartitionedFile::from_path(filename.to_string_lossy().into_owned()).unwrap();
let session_ctx = Arc::new(SessionContext::new());
let scan = init_datasource_exec(
required_variant_schema(),
required_schema,
None,
None,
ObjectStoreUrl::local_filesystem(),
Expand All @@ -168,7 +176,71 @@ async fn scan_variant_file(filename: PathBuf) -> VariantArray {
let mut stream = scan.execute(0, session_ctx.task_ctx()).unwrap();
let batch = stream.next().await.unwrap().unwrap();
assert!(stream.next().await.is_none());
VariantArray::try_new(batch.column(0).as_ref()).unwrap()
batch
}

#[tokio::test]
async fn full_value_variant_request_reads_canonical_and_shredded_parquet() {
let mut builder = VariantArrayBuilder::new(4);
builder
.new_object()
.with_field("a", 1_i8)
.with_field("extra", "text")
.finish();
builder.append_null();
builder.append_variant(Variant::Null);
builder.new_object().with_field("a", 2_i8).finish();
let canonical = builder.build();
let shredded = shred_variant(
&canonical,
&DataType::Struct(Fields::from(vec![Field::new("a", DataType::Int64, true)])),
)
.unwrap();
let mut child = required_variant_schema()
.field(0)
.clone()
.with_name("0")
.with_nullable(true);
let mut metadata = child.metadata().clone();
metadata.insert(
"__VARIANT_METADATA_KEY".to_string(),
r#"{"path":"$","failOnError":true,"timeZoneId":"UTC"}"#.to_string(),
);
child = child.with_metadata(metadata);
let required = Arc::new(Schema::new(vec![Field::new(
"v",
DataType::Struct(Fields::from(vec![child])),
true,
)]));

for input in [canonical.inner(), shredded.inner()] {
let schema = Arc::new(Schema::new(vec![Field::new(
"v",
input.data_type().clone(),
true,
)
.with_extension_type(VariantType)]));
let batch =
RecordBatch::try_new(Arc::clone(&schema), vec![Arc::new(input.clone())]).unwrap();
let file = tempfile::NamedTempFile::new().unwrap();
let mut writer = ArrowWriter::try_new(file.reopen().unwrap(), schema, None).unwrap();
writer.write(&batch).unwrap();
writer.close().unwrap();

let batch = scan_variant_batch(file.path().to_path_buf(), Arc::clone(&required)).await;
assert_eq!(batch.column(0).data_type(), required.field(0).data_type());
let wrapped = batch
.column(0)
.as_any()
.downcast_ref::<StructArray>()
.unwrap();
assert_eq!(wrapped.nulls(), canonical.inner().nulls());
let output = VariantArray::try_new(wrapped.column(0).as_ref()).unwrap();
assert!(output.is_null(1));
for index in [0, 2, 3] {
assert_eq!(output.value(index), canonical.value(index));
}
}
}

#[tokio::test]
Expand Down
12 changes: 4 additions & 8 deletions native/core/src/parquet/schema_adapter.rs
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@
// specific language governing permissions and limitations
// under the License.

use crate::parquet::cast_column::CometCastColumnExpr;
use crate::parquet::cast_column::{variant_projection_field, CometCastColumnExpr};
use crate::parquet::name_fold::{fold_name, fold_names, fold_schema_names};
use crate::parquet::parquet_support::{
duplicate_parquet_field_error, field_id, field_names_with_id, match_struct_fields,
Expand All @@ -37,7 +37,6 @@ use datafusion_physical_expr_adapter::{
replace_columns_with_literals, DefaultPhysicalExprAdapterFactory, PhysicalExprAdapter,
PhysicalExprAdapterFactory,
};
use parquet::variant::VariantType;
use std::collections::{HashMap, HashSet};
use std::fmt::{self, Display};
use std::hash::{Hash, Hasher};
Expand Down Expand Up @@ -1158,7 +1157,7 @@ impl SparkPhysicalExprAdapter {
let Ok(logical_field) = self.logical_file_schema.field_with_name(column.name()) else {
return Ok(expr);
};
if !logical_field.has_valid_extension_type::<VariantType>() {
if variant_projection_field(logical_field).is_none() {
return Ok(expr);
}
let Some(physical_field) = self.physical_file_schema.fields().get(column.index()) else {
Expand Down Expand Up @@ -1222,7 +1221,7 @@ impl SparkPhysicalExprAdapter {
Arc::clone(&e)
};

if logical_field.has_valid_extension_type::<VariantType>()
if variant_projection_field(logical_field).is_some()
|| logical_field.data_type() != physical_field.data_type()
{
// Apply the same Spark conversion rules as `replace_with_spark_cast`;
Expand Down Expand Up @@ -1303,10 +1302,7 @@ impl SparkPhysicalExprAdapter {
};
let physical_type = input_field.data_type();

if cast
.target_field()
.has_valid_extension_type::<VariantType>()
{
if variant_projection_field(cast.target_field()).is_some() {
let comet_cast: Arc<dyn PhysicalExpr> = Arc::new(
CometCastColumnExpr::try_new(
child,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -344,7 +344,8 @@ case class CometScanRule(session: SparkSession)
// preserve Spark's ENUM inference without losing Parquet decryption state.
// https://github.com/apache/datafusion-comet/issues/5477
if (encryptionEnabled(hadoopConf) &&
scanExec.requiredSchema.exists(field => isVariantType(field.dataType))) {
scanExec.requiredSchema.exists(field =>
isVariantType(field.dataType) || isWholeVariantStruct(field.dataType))) {
withFallbackReason(scanExec, "Native Parquet Variant scans do not support encryption")
return None
}
Expand Down Expand Up @@ -1048,7 +1049,8 @@ case class CometScanRule(session: SparkSession)
dt: DataType,
name: String,
reasons: ListBuffer[String]): Boolean =
isVariantType(dt) || typeChecker.isTypeSupported(dt, name, reasons)
isVariantType(dt) || isWholeVariantStruct(dt) ||
typeChecker.isTypeSupported(dt, name, reasons)
}
val schemaSupported =
requiredSchemaChecker.isSchemaSupported(scanExec.requiredSchema, fallbackReasons)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -769,14 +769,15 @@ object QueryPlanSerde extends Logging with CometExprShim with CometTypeShim {
if (includeFieldIds && ParquetUtils.hasFieldId(f)) Some(ParquetUtils.getFieldId(f))
else None
}
if (fieldIds.exists(_.isDefined)) {
// Emit one FieldMetadata entry per nested field, parallel to field_names. Entries
// for fields without an ID are empty so the slot index stays aligned.
fieldIds.foreach { idOpt =>
val variantMetadata = s.fields.map(variantRequestMetadata)
if (fieldIds.exists(_.isDefined) || variantMetadata.exists(_.isDefined)) {
// Keep metadata entries aligned with field_names, including empty slots.
fieldIds.zip(variantMetadata).foreach { case (idOpt, variantMeta) =>
val metaBuilder = Types.DataType.FieldMetadata.newBuilder()
idOpt.foreach { id =>
metaBuilder.putMetadata(CometParquetUtils.PARQUET_FIELD_ID_META_KEY, id.toString)
}
variantMeta.foreach { case (key, value) => metaBuilder.putMetadata(key, value) }
struct.addFieldMetadata(metaBuilder.build())
}
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -167,7 +167,8 @@ object CometNativeScan extends CometOperatorSerde[CometScanExec] with CometTypeS
withFallbackReason(scanExec, unsupportedDefaultReason)
}

if (scanExec.requiredSchema.exists(field => isVariantType(field.dataType))) {
if (scanExec.requiredSchema.exists(field =>
isVariantType(field.dataType) || isWholeVariantStruct(field.dataType))) {
// Spark's strict legacy reader owns malformed-layout errors (SPARK-47546).
// TODO: Remove this guard once the native reader implements Spark's strict Variant layout
// validation and malformed-input errors when allowReadingShredded=false.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,7 @@ import org.apache.comet.vector.CometVector

object Utils extends CometTypeShim with Logging {
private val VariantExtensionName = "arrow.parquet.variant"
private val VariantRequestMetadataKey = "__VARIANT_METADATA_KEY"

def getConfPath(confFileName: String): String = {
sys.env
Expand Down Expand Up @@ -89,7 +90,14 @@ object Utils extends CometTypeShim with Logging {
.getOrElse {
val fields = field.getChildren().asScala.map { child =>
val dt = fromArrowField(child)
StructField(child.getName, dt, child.isNullable)
val metadata = Option(child.getMetadata)
.flatMap(m => Option(m.get(VariantRequestMetadataKey)))
.map(json =>
new MetadataBuilder()
.putMetadata(VariantRequestMetadataKey, Metadata.fromJson(json))
.build())
.getOrElse(Metadata.empty)
StructField(child.getName, dt, child.isNullable, metadata)
}
StructType(fields.toSeq)
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -23,7 +23,7 @@ import java.nio.ByteBuffer
import java.nio.charset.{CharacterCodingException, CodingErrorAction, StandardCharsets}

import org.apache.spark.sql.catalyst.expressions.aggregate.Mode
import org.apache.spark.sql.types.{DataType, StructType}
import org.apache.spark.sql.types.{DataType, StructField, StructType}
import org.apache.spark.unsafe.types.UTF8String

trait CometTypeShim {
Expand All @@ -43,6 +43,10 @@ trait CometTypeShim {
// Spark 4 feature; Variant shredding doesn't exist in Spark 3.x.
def isVariantStruct(s: StructType): Boolean = false

def isWholeVariantStruct(dt: DataType): Boolean = false

def variantRequestMetadata(field: StructField): Option[(String, String)] = None

// Spark 4 feature; VariantType doesn't exist in Spark 3.x.
def isVariantType(dt: DataType): Boolean = false

Expand Down
Loading
Loading