Skip to content
Open
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
6 changes: 4 additions & 2 deletions native/core/src/execution/jni_api.rs
Original file line number Diff line number Diff line change
Expand Up @@ -686,6 +686,7 @@ fn prepare_output(
let schema_addrs = unsafe { schema_addrs.get_elements(env, ReleaseMode::NoCopyBack)? };
let schema_addrs = &*schema_addrs;

let output_schema = output_batch.schema();
let results = output_batch.columns();
let num_rows = output_batch.num_rows();

Expand All @@ -712,6 +713,7 @@ fn prepare_output(
let mut i = 0;
while i < results.len() {
let array_ref = results.get(i).ok_or(CometError::IndexOutOfBounds(i))?;
let field = output_schema.field(i);

if array_ref.offset() != 0 {
// https://github.com/apache/datafusion-comet/issues/2051
Expand All @@ -728,11 +730,11 @@ fn prepare_output(

new_array
.to_data()
.move_to_spark(array_addrs[i], schema_addrs[i])?;
.move_to_spark(field, array_addrs[i], schema_addrs[i])?;
} else {
array_ref
.to_data()
.move_to_spark(array_addrs[i], schema_addrs[i])?;
.move_to_spark(field, array_addrs[i], schema_addrs[i])?;
}
i += 1;
}
Expand Down
185 changes: 142 additions & 43 deletions native/core/src/execution/planner.rs
Original file line number Diff line number Diff line change
Expand Up @@ -43,7 +43,7 @@ use crate::execution::{
},
planner::expression_registry::ExpressionRegistry,
planner::operator_registry::OperatorRegistry,
serde::to_arrow_datatype,
serde::{to_arrow_datatype, to_arrow_field},
shuffle::{SchemaAlignExec, ShuffleWriterExec},
};
use crate::jvm_bridge::{jni_call, JVMClasses};
Expand Down Expand Up @@ -115,7 +115,7 @@ use crate::parquet::parquet_exec::init_datasource_exec;
use arrow::array::{
new_empty_array, Array, ArrayRef, BinaryBuilder, BooleanArray, Date32Array, Decimal128Array,
Float32Array, Float64Array, Int16Array, Int32Array, Int64Array, Int8Array, ListArray,
NullArray, StringBuilder, TimestampMicrosecondArray,
NullArray, RecordBatch, StringBuilder, TimestampMicrosecondArray,
};
use arrow::buffer::{BooleanBuffer, NullBuffer, OffsetBuffer};
use arrow::row::{OwnedRow, RowConverter, SortField};
Expand Down Expand Up @@ -923,6 +923,25 @@ impl PhysicalPlanner {
}
}

/// Decode a Spark existence default into the scalar consumed by the Parquet schema adapter.
/// Ordinary defaults are literals. A Variant default is transported as a constant
/// `CreateNamedStruct(value, metadata)` because Variant has Struct storage in Arrow.
fn create_default_value(
&self,
spark_expr: &Expr,
input_schema: SchemaRef,
) -> Result<ScalarValue, ExecutionError> {
let expr = self.create_expr(spark_expr, Arc::clone(&input_schema))?;
if let Some(literal) = expr.downcast_ref::<DataFusionLiteral>() {
return Ok(literal.value().clone());
}

let array = expr
.evaluate(&RecordBatch::new_empty(input_schema))?
.into_array_of_size(1)?;
Ok(ScalarValue::try_from_array(array.as_ref(), 0)?)
}

/// Create a DataFusion physical sort expression from Spark physical expression
fn create_sort_expr<'a>(
&'a self,
Expand Down Expand Up @@ -1580,44 +1599,38 @@ impl PhysicalPlanner {
.collect()
};

let default_values: Option<HashMap<Column, ScalarValue>> = if !common
.default_values
.is_empty()
{
// We have default values. Extract the two lists (same length) of values and
// indexes in the schema, and then create a HashMap to use in the SchemaMapper.
let default_values: Result<Vec<ScalarValue>, DataFusionError> = common
.default_values
.iter()
.map(|expr| {
let literal = self.create_expr(expr, Arc::clone(&required_schema))?;
let df_literal =
literal.downcast_ref::<DataFusionLiteral>().ok_or_else(|| {
GeneralError("Expected literal of default value.".to_string())
if common.default_values.len() != common.default_values_indexes.len() {
return Err(GeneralError(format!(
"NativeScan has {} default values but {} default indexes",
common.default_values.len(),
common.default_values_indexes.len()
)));
}
let default_values: Option<HashMap<Column, ScalarValue>> =
if common.default_values.is_empty() {
None
} else {
let defaults = common
.default_values
.iter()
.zip(&common.default_values_indexes)
.map(|(expr, offset)| {
let idx = usize::try_from(*offset).map_err(|_| {
GeneralError(format!("Invalid default value index {offset}"))
})?;
Ok(df_literal.value().clone())
})
.collect();
let default_values = default_values?;
let default_values_indexes: Vec<usize> = common
.default_values_indexes
.iter()
.map(|offset| *offset as usize)
.collect();
Some(
default_values_indexes
.into_iter()
.zip(default_values)
.map(|(idx, scalar_value)| {
let field = required_schema.field(idx);
let column = Column::new(field.name().as_str(), idx);
(column, scalar_value)
let field = required_schema.fields().get(idx).ok_or_else(|| {
GeneralError(format!(
"Default value index {idx} is outside schema with {} fields",
required_schema.fields().len()
))
})?;
let value =
self.create_default_value(expr, Arc::clone(&required_schema))?;
Ok((Column::new(field.name(), idx), value))
})
.collect(),
)
} else {
None
};
.collect::<Result<HashMap<_, _>, ExecutionError>>()?;
Some(defaults)
};

// Get one file from this partition (we know it's not empty due to early return above)
let one_file = partition_files
Expand Down Expand Up @@ -3772,15 +3785,17 @@ pub(crate) fn convert_spark_types_to_arrow_schema(
let arrow_fields = spark_types
.iter()
.map(|spark_type| {
let field = Field::new(
let field = to_arrow_field(
String::clone(&spark_type.name),
to_arrow_datatype(spark_type.data_type.as_ref().unwrap()),
spark_type.data_type.as_ref().unwrap(),
spark_type.nullable,
);
if spark_type.metadata.is_empty() {
field
} else {
field.with_metadata(spark_type.metadata.clone())
let mut metadata = spark_type.metadata.clone();
metadata.extend(field.metadata().clone());
field.with_metadata(metadata)
}
})
.collect_vec();
Expand Down Expand Up @@ -4598,7 +4613,8 @@ mod tests {
use std::{sync::Arc, task::Poll};

use arrow::array::{
Array, DictionaryArray, Int32Array, Int8Array, ListArray, RecordBatch, StringArray,
Array, BinaryArray, DictionaryArray, Int32Array, Int8Array, ListArray, RecordBatch,
StringArray,
};
use arrow::datatypes::{DataType, Field, FieldRef, Fields, Schema};
use datafusion::catalog::memory::DataSourceExec;
Expand All @@ -4611,7 +4627,10 @@ mod tests {
use datafusion::error::DataFusionError;
use datafusion::logical_expr::ScalarUDF;
use datafusion::physical_plan::ExecutionPlan;
use datafusion::{assert_batches_eq, physical_plan::common::collect, prelude::SessionContext};
use datafusion::{
assert_batches_eq, physical_plan::common::collect, prelude::SessionContext,
scalar::ScalarValue,
};
use datafusion_physical_expr_adapter::PhysicalExprAdapterFactory;
use tempfile::TempDir;
use tokio::sync::mpsc;
Expand All @@ -4621,6 +4640,7 @@ mod tests {
use crate::execution::operators::ExecutionError;
use crate::execution::planner::literal_to_array_ref;
use crate::execution::planner::parse_file_scan_tasks_from_common;
use crate::execution::serde::to_arrow_field;
use crate::parquet::parquet_support::SparkParquetOptions;
use crate::parquet::schema_adapter::SparkPhysicalExprAdapterFactory;
use datafusion_comet_proto::spark_expression::expr::ExprStruct;
Expand All @@ -4634,6 +4654,85 @@ mod tests {
};
use datafusion_comet_spark_expr::EvalMode;

#[test]
fn test_create_default_value_from_literal_and_variant_struct() {
fn literal(value: literal::Value, type_id: i32) -> Expr {
Expr {
expr_struct: Some(ExprStruct::Literal(spark_expression::Literal {
value: Some(value),
datatype: Some(spark_expression::DataType {
type_id,
type_info: None,
}),
is_null: false,
})),
query_context: None,
expr_id: None,
}
}

let planner = PhysicalPlanner::default();
let int_field = Field::new("n", DataType::Int32, true);
let int_default = literal(literal::Value::IntVal(7), 3);
assert_eq!(
planner
.create_default_value(&int_default, Arc::new(Schema::new(vec![int_field.clone()])),)
.unwrap(),
ScalarValue::Int32(Some(7))
);

let variant_field = to_arrow_field(
"v",
&spark_expression::DataType {
type_id: 21,
type_info: None,
},
true,
);
let variant_default = Expr {
expr_struct: Some(ExprStruct::CreateNamedStruct(
spark_expression::CreateNamedStruct {
values: vec![
literal(literal::Value::BytesVal(vec![1, 2]), 8),
literal(literal::Value::BytesVal(vec![3, 4]), 8),
],
names: vec!["value".to_string(), "metadata".to_string()],
},
)),
query_context: None,
expr_id: None,
};
let value = planner
.create_default_value(
&variant_default,
Arc::new(Schema::new(vec![variant_field.clone()])),
)
.unwrap();
let ScalarValue::Struct(value) = value else {
panic!("expected Variant default to use Struct storage")
};
assert_eq!(value.fields()[0].name(), "value");
assert_eq!(value.fields()[1].name(), "metadata");
assert_eq!(
value
.column(0)
.as_any()
.downcast_ref::<BinaryArray>()
.unwrap()
.value(0),
&[1, 2]
);
assert_eq!(
value
.column(1)
.as_any()
.downcast_ref::<BinaryArray>()
.unwrap()
.value(0),
&[3, 4]
);
}

#[test]
fn test_unpack_dictionary_primitive() {
let op_scan = Operator {
Expand Down
Loading
Loading