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
194 changes: 192 additions & 2 deletions datafusion/common/src/nested_struct.rs
Original file line number Diff line number Diff line change
Expand Up @@ -19,11 +19,12 @@ use crate::error::{_plan_err, Result};
use arrow::{
array::{
Array, ArrayRef, AsArray, DictionaryArray, FixedSizeListArray, GenericListArray,
GenericListViewArray, StructArray, downcast_integer, make_array, new_null_array,
GenericListViewArray, RecordBatch, StructArray, downcast_integer, make_array,
new_null_array,
},
buffer::NullBuffer,
compute::{CastOptions, can_cast_types, cast_with_options},
datatypes::{DataType, DataType::Struct, Field, FieldRef},
datatypes::{DataType, DataType::Struct, Field, FieldRef, SchemaRef},
};
use std::{collections::HashSet, sync::Arc};

Expand Down Expand Up @@ -1703,3 +1704,192 @@ mod tests {
));
}
}

/// Adapts a [`RecordBatch`] to a target [`SchemaRef`].
///
/// If `batch` already has the target schema, it is returned immediately.
///
/// If `batch` has a schema that is a legitimate subtype / stricter subset of
/// `target_schema` (as verified by [`arrow::datatypes::Schema::contains`]),
/// this function transforms the metadata/types of differing columns to match
/// `target_schema` without copying primitive buffer data.
///
/// If `batch` does not conform to `target_schema` under schema containment,
/// an error is returned.
pub fn adapt_batch_to_schema(
batch: RecordBatch,
target_schema: &SchemaRef,
) -> Result<RecordBatch> {
if Arc::ptr_eq(batch.schema_ref(), target_schema)
|| batch.schema().as_ref() == target_schema.as_ref()
{
return Ok(batch);
}

if !target_schema.contains(batch.schema().as_ref()) {
return _plan_err!(
"Batch schema does not conform to expected schema. Expected: {target_schema}, got: {}",
batch.schema()
);
}

let mut columns = Vec::with_capacity(batch.num_columns());
let mut needs_column_adaptation = false;
let cast_options = CastOptions::default();

for (target_field, col) in target_schema.fields().iter().zip(batch.columns()) {
if target_field.data_type() != col.data_type() {
needs_column_adaptation = true;
let adapted_col = cast_column(col, target_field.data_type(), &cast_options)?;
columns.push(adapted_col);
} else {
columns.push(Arc::clone(col));
}
}

if needs_column_adaptation {
Ok(RecordBatch::try_new(Arc::clone(target_schema), columns)?)
} else {
// Schema differs only in top-level metadata or field nullability, while
// column data types match exactly. Replace the schema on the batch.
Ok(RecordBatch::try_new(
Arc::clone(target_schema),
batch.columns().to_vec(),
)?)
}
}

#[cfg(test)]
mod adapt_schema_tests {
use super::*;
use arrow::array::{BooleanArray, Int32Array, StringArray};
use arrow::datatypes::{Field, Fields, Schema};

#[test]
fn test_adapt_batch_to_schema_identical() -> Result<()> {
let schema = Arc::new(Schema::new(vec![
Field::new("a", DataType::Int32, false),
Field::new("b", DataType::Utf8, true),
]));
let batch = RecordBatch::try_new(
Arc::clone(&schema),
vec![
Arc::new(Int32Array::from(vec![1, 2, 3])),
Arc::new(StringArray::from(vec![Some("x"), None, Some("z")])),
],
)?;

let adapted = adapt_batch_to_schema(batch.clone(), &schema)?;
assert!(Arc::ptr_eq(batch.schema_ref(), adapted.schema_ref()));
assert_eq!(batch, adapted);
Ok(())
}

#[test]
fn test_adapt_batch_to_schema_stricter_nested_struct() -> Result<()> {
let declared_schema = Arc::new(Schema::new(vec![
Field::new("id", DataType::Int32, false),
Field::new(
"nested",
Struct(Fields::from(vec![Field::new(
"val",
DataType::Boolean,
true,
)])),
false,
),
]));

let stricter_batch_schema = Arc::new(Schema::new(vec![
Field::new("id", DataType::Int32, false),
Field::new(
"nested",
Struct(Fields::from(vec![Field::new(
"val",
DataType::Boolean,
false,
)])),
false,
),
]));

let struct_col = Arc::new(StructArray::new(
Fields::from(vec![Field::new("val", DataType::Boolean, false)]),
vec![Arc::new(BooleanArray::from(vec![true, false, true]))],
None,
));

let batch = RecordBatch::try_new(
stricter_batch_schema,
vec![Arc::new(Int32Array::from(vec![1, 2, 3])), struct_col],
)?;

let adapted = adapt_batch_to_schema(batch, &declared_schema)?;
assert_eq!(adapted.schema().as_ref(), declared_schema.as_ref());
assert_eq!(adapted.num_rows(), 3);
assert_eq!(adapted.num_columns(), 2);
Ok(())
}

#[test]
fn test_adapt_batch_to_schema_top_level_nullability_only() -> Result<()> {
let declared_schema =
Arc::new(Schema::new(vec![Field::new("a", DataType::Int32, true)]));

let stricter_batch_schema =
Arc::new(Schema::new(vec![Field::new("a", DataType::Int32, false)]));

let batch = RecordBatch::try_new(
stricter_batch_schema,
vec![Arc::new(Int32Array::from(vec![1, 2, 3]))],
)?;

let adapted = adapt_batch_to_schema(batch, &declared_schema)?;
assert_eq!(adapted.schema().as_ref(), declared_schema.as_ref());
assert_eq!(adapted.column(0).len(), 3);
Ok(())
}

#[test]
fn test_adapt_batch_to_schema_incompatible_rejected() {
let declared_schema = Arc::new(Schema::new(vec![
Field::new("a", DataType::Int32, false), // target is non-nullable
]));

let batch_schema = Arc::new(Schema::new(vec![
Field::new("a", DataType::Int32, true), // batch is nullable (not contained!)
]));

let batch = RecordBatch::try_new(
batch_schema,
vec![Arc::new(Int32Array::from(vec![Some(1), None, Some(3)]))],
)
.unwrap();

let res = adapt_batch_to_schema(batch, &declared_schema);
assert!(res.is_err());
let err_msg = res.unwrap_err().to_string();
assert!(
err_msg.contains("does not conform to expected schema"),
"unexpected error message: {err_msg}"
);
}

#[test]
fn test_adapt_batch_to_schema_incompatible_type_rejected() {
let declared_schema =
Arc::new(Schema::new(vec![Field::new("a", DataType::Utf8, true)]));

let batch_schema =
Arc::new(Schema::new(vec![Field::new("a", DataType::Int32, true)]));

let batch = RecordBatch::try_new(
batch_schema,
vec![Arc::new(Int32Array::from(vec![1, 2, 3]))],
)
.unwrap();

let res = adapt_batch_to_schema(batch, &declared_schema);
assert!(res.is_err());
}
}
1 change: 1 addition & 0 deletions datafusion/core/tests/sql/aggregates/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1021,3 +1021,4 @@ pub fn split_fuzz_timestamp_data_into_batches(

pub mod basic;
pub mod dict_nulls;
mod nested_nullability;
Loading