diff --git a/crates/iceberg/src/arrow/nan_val_cnt_visitor.rs b/crates/iceberg/src/arrow/nan_val_cnt_visitor.rs index d01f3e9e56..18f5299358 100644 --- a/crates/iceberg/src/arrow/nan_val_cnt_visitor.rs +++ b/crates/iceberg/src/arrow/nan_val_cnt_visitor.rs @@ -19,17 +19,14 @@ use std::collections::HashMap; use std::collections::hash_map::Entry; -use std::sync::Arc; -use arrow_array::{ArrayRef, Float32Array, Float64Array, RecordBatch, StructArray}; -use arrow_schema::DataType; - -use crate::Result; -use crate::arrow::{ArrowArrayAccessor, FieldMatchMode}; -use crate::spec::{ - ListType, MapType, NestedFieldRef, PrimitiveType, Schema, SchemaRef, SchemaWithPartnerVisitor, - StructType, VariantType, visit_struct_with_partner, +use arrow_array::{ + Array, ArrayRef, Float32Array, Float64Array, ListArray, MapArray, RecordBatch, StructArray, }; +use arrow_schema::{DataType, FieldRef}; + +use crate::arrow::get_field_id_from_metadata; +use crate::{Error, ErrorKind, Result}; macro_rules! cast_and_update_cnt_map { ($t:ty, $col:ident, $self:ident, $field_id:ident) => { @@ -71,112 +68,67 @@ macro_rules! count_float_nans { pub struct NanValueCountVisitor { /// Stores field ID to NaN value count mapping pub nan_value_counts: HashMap, - match_mode: FieldMatchMode, } -impl SchemaWithPartnerVisitor for NanValueCountVisitor { - type T = (); - - fn schema( - &mut self, - _schema: &Schema, - _partner: &ArrayRef, - _value: Self::T, - ) -> Result { - Ok(()) - } - - fn field( - &mut self, - _field: &NestedFieldRef, - _partner: &ArrayRef, - _value: Self::T, - ) -> Result { - Ok(()) - } - - fn r#struct( - &mut self, - _struct: &StructType, - _partner: &ArrayRef, - _results: Vec, - ) -> Result { - Ok(()) - } - - fn list(&mut self, _list: &ListType, _list_arr: &ArrayRef, _value: Self::T) -> Result { - Ok(()) - } - - fn map( - &mut self, - _map: &MapType, - _partner: &ArrayRef, - _key_value: Self::T, - _value: Self::T, - ) -> Result { - Ok(()) - } - - fn primitive(&mut self, _p: &PrimitiveType, _col: &ArrayRef) -> Result { - Ok(()) - } - - fn variant(&mut self, _v: &VariantType, _col: &ArrayRef) -> Result { - Ok(()) - } - - fn after_struct_field(&mut self, field: &NestedFieldRef, partner: &ArrayRef) -> Result<()> { - let field_id = field.id; - count_float_nans!(partner, self, field_id); - Ok(()) - } - - fn after_list_element(&mut self, field: &NestedFieldRef, partner: &ArrayRef) -> Result<()> { - let field_id = field.id; - count_float_nans!(partner, self, field_id); - Ok(()) - } - - fn after_map_key(&mut self, field: &NestedFieldRef, partner: &ArrayRef) -> Result<()> { - let field_id = field.id; - count_float_nans!(partner, self, field_id); - Ok(()) - } +impl NanValueCountVisitor { + fn visit_field(&mut self, field: &FieldRef, array: &ArrayRef) -> Result<()> { + if matches!(array.data_type(), DataType::Float32 | DataType::Float64) { + let field_id = get_field_id_from_metadata(field)?; + count_float_nans!(array, self, field_id); + } - fn after_map_value(&mut self, field: &NestedFieldRef, partner: &ArrayRef) -> Result<()> { - let field_id = field.id; - count_float_nans!(partner, self, field_id); - Ok(()) + match field.data_type() { + DataType::Struct(fields) => { + let struct_array = + array + .as_any() + .downcast_ref::() + .ok_or_else(|| { + Error::new( + ErrorKind::DataInvalid, + "Expected struct array for NaN counts", + ) + })?; + for (field, column) in fields.iter().zip(struct_array.columns()) { + self.visit_field(field, column)?; + } + Ok(()) + } + DataType::List(element) => { + let list_array = array.as_any().downcast_ref::().ok_or_else(|| { + Error::new(ErrorKind::DataInvalid, "Expected list array for NaN counts") + })?; + self.visit_field(element, list_array.values()) + } + DataType::Map(entries, _) => { + let map_array = array.as_any().downcast_ref::().ok_or_else(|| { + Error::new(ErrorKind::DataInvalid, "Expected map array for NaN counts") + })?; + let DataType::Struct(fields) = entries.data_type() else { + return Err(Error::new( + ErrorKind::DataInvalid, + "Expected map entry struct for NaN counts", + )); + }; + self.visit_field(&fields[0], map_array.keys())?; + self.visit_field(&fields[1], map_array.values()) + } + _ => Ok(()), + } } -} -impl NanValueCountVisitor { /// Creates new instance of NanValueCountVisitor pub fn new() -> Self { - Self::new_with_match_mode(FieldMatchMode::Id) - } - - /// Creates new instance of NanValueCountVisitor with explicit match mode - pub fn new_with_match_mode(match_mode: FieldMatchMode) -> Self { Self { nan_value_counts: HashMap::new(), - match_mode, } } - /// Compute nan value counts in given schema and record batch - pub fn compute(&mut self, schema: SchemaRef, batch: RecordBatch) -> Result<()> { - let arrow_arr_partner_accessor = ArrowArrayAccessor::new_with_match_mode(self.match_mode); - - let struct_arr = Arc::new(StructArray::from(batch)) as ArrayRef; - visit_struct_with_partner( - schema.as_struct(), - &struct_arr, - self, - &arrow_arr_partner_accessor, - )?; - + /// Compute NaN counts from the validated, projected Arrow batch. + pub fn compute(&mut self, batch: &RecordBatch) -> Result<()> { + for (field, column) in batch.schema().fields().iter().zip(batch.columns()) { + self.visit_field(field, column)?; + } Ok(()) } } diff --git a/crates/iceberg/src/arrow/record_batch_partition_splitter.rs b/crates/iceberg/src/arrow/record_batch_partition_splitter.rs index 598bbf09f7..6ac3b74296 100644 --- a/crates/iceberg/src/arrow/record_batch_partition_splitter.rs +++ b/crates/iceberg/src/arrow/record_batch_partition_splitter.rs @@ -117,6 +117,7 @@ impl RecordBatchPartitionSplitter { } /// Split the record batch into multiple record batches based on the partition spec. + /// In pre-computed mode, the `_partition` routing column is removed from returned batches. pub fn split(&self, batch: &RecordBatch) -> Result> { let partition_structs = if let Some(calculator) = &self.calculator { // Compute partition values from source columns using calculator @@ -177,6 +178,23 @@ impl RecordBatchPartitionSplitter { .collect::>>()? }; + // The pre-computed partition column is routing metadata, not table data. Remove it + // before returning batches to the file writer, which validates against the table schema. + let batch = if self.calculator.is_none() { + let columns = batch + .schema() + .fields() + .iter() + .enumerate() + .filter_map(|(index, field)| { + (field.name() != PROJECTED_PARTITION_VALUE_COLUMN).then_some(index) + }) + .collect::>(); + batch.project(&columns)? + } else { + batch.clone() + }; + // Group the batch by row value. let mut group_ids = HashMap::new(); partition_structs @@ -207,7 +225,7 @@ impl RecordBatchPartitionSplitter { ); // filter the RecordBatch - partition_batches.push((partition_key, filter_record_batch(batch, &filter_array)?)); + partition_batches.push((partition_key, filter_record_batch(&batch, &filter_array)?)); } Ok(partition_batches) @@ -439,6 +457,12 @@ mod tests { }); assert_eq!(partitioned_batches.len(), 3); + assert!(partitioned_batches.iter().all(|(_, batch)| { + batch.num_columns() == 2 + && batch + .column_by_name(PROJECTED_PARTITION_VALUE_COLUMN) + .is_none() + })); // Helper to extract id and name values from a batch let extract_values = |batch: &RecordBatch| -> (Vec, Vec) { diff --git a/crates/iceberg/src/arrow/schema.rs b/crates/iceberg/src/arrow/schema.rs index 4ef7b29436..2c765debb2 100644 --- a/crates/iceberg/src/arrow/schema.rs +++ b/crates/iceberg/src/arrow/schema.rs @@ -326,7 +326,7 @@ pub fn arrow_type_to_type(ty: &DataType) -> Result { const ARROW_FIELD_DOC_KEY: &str = "doc"; -pub(super) fn get_field_id_from_metadata(field: &FieldRef) -> Result { +pub(crate) fn get_field_id_from_metadata(field: &FieldRef) -> Result { if let Some(value) = field.metadata().get(PARQUET_FIELD_ID_META_KEY) { return value.parse::().map_err(|e| { Error::new( @@ -793,6 +793,178 @@ pub fn schema_to_arrow_schema(schema: &Schema) -> Result { } } +fn parquet_write_arrow_field( + field: &NestedFieldRef, + arrow_field: &FieldRef, +) -> Result> { + let data_type = match field.field_type.as_ref() { + Type::Primitive(PrimitiveType::Unknown) => return Ok(None), + Type::Primitive(_) => arrow_field.data_type().clone(), + Type::Variant(_) => { + return Err(Error::new( + ErrorKind::FeatureUnsupported, + format!( + "Field {} has variant type, which is not yet implemented", + field.id + ), + )); + } + Type::Struct(struct_type) => { + let DataType::Struct(arrow_fields) = arrow_field.data_type() else { + return Err(Error::new( + ErrorKind::Unexpected, + format!( + "Expected Arrow struct for Iceberg field {}, got {}", + field.id, + arrow_field.data_type() + ), + )); + }; + if struct_type.fields().len() != arrow_fields.len() { + return Err(Error::new( + ErrorKind::Unexpected, + format!( + "Arrow and Iceberg struct field counts differ for field {}", + field.id + ), + )); + } + + let fields = struct_type + .fields() + .iter() + .zip(arrow_fields.iter()) + .filter_map(|(field, arrow_field)| { + parquet_write_arrow_field(field, arrow_field).transpose() + }) + .collect::>>()?; + if fields.is_empty() { + return Err(Error::new( + ErrorKind::FeatureUnsupported, + format!( + "Cannot write struct field {} with no Parquet physical fields", + field.id + ), + )); + } + DataType::Struct(fields.into()) + } + Type::List(list_type) => { + let DataType::List(element_field) = arrow_field.data_type() else { + return Err(Error::new( + ErrorKind::Unexpected, + format!( + "Expected Arrow list for Iceberg field {}, got {}", + field.id, + arrow_field.data_type() + ), + )); + }; + let element_field = parquet_write_arrow_field(&list_type.element_field, element_field)? + .ok_or_else(|| { + Error::new( + ErrorKind::FeatureUnsupported, + format!( + "Cannot write list element {} with no Parquet physical fields", + list_type.element_field.id + ), + ) + })?; + DataType::List(element_field) + } + Type::Map(map_type) => { + let DataType::Map(entries_field, ordered) = arrow_field.data_type() else { + return Err(Error::new( + ErrorKind::Unexpected, + format!( + "Expected Arrow map for Iceberg field {}, got {}", + field.id, + arrow_field.data_type() + ), + )); + }; + let DataType::Struct(entry_fields) = entries_field.data_type() else { + return Err(Error::new( + ErrorKind::Unexpected, + format!( + "Expected Arrow map entries struct for Iceberg field {}", + field.id + ), + )); + }; + if entry_fields.len() != 2 { + return Err(Error::new( + ErrorKind::Unexpected, + format!( + "Expected two Arrow map entry fields for Iceberg field {}", + field.id + ), + )); + } + + let key_field = parquet_write_arrow_field(&map_type.key_field, &entry_fields[0])? + .ok_or_else(|| { + Error::new( + ErrorKind::FeatureUnsupported, + format!( + "Cannot write map key {} with no Parquet physical fields", + map_type.key_field.id + ), + ) + })?; + let value_field = parquet_write_arrow_field(&map_type.value_field, &entry_fields[1])? + .ok_or_else(|| { + Error::new( + ErrorKind::FeatureUnsupported, + format!( + "Cannot write map value {} with no Parquet physical fields", + map_type.value_field.id + ), + ) + })?; + let entries_field = Arc::new( + entries_field + .as_ref() + .clone() + .with_data_type(DataType::Struct(vec![key_field, value_field].into())), + ); + DataType::Map(entries_field, *ordered) + } + }; + + Ok(Some(Arc::new( + arrow_field.as_ref().clone().with_data_type(data_type), + ))) +} + +/// Convert an Iceberg schema to the Arrow schema used for Parquet writes. +/// +/// Unknown fields are omitted because Iceberg has no Parquet physical mapping for them. Structs, +/// list elements, and map keys/values left with no physical fields cannot be omitted without +/// losing container semantics, so those schemas are rejected. +pub(crate) fn schema_to_arrow_schema_for_parquet_write(schema: &Schema) -> Result { + let arrow_schema = schema_to_arrow_schema(schema)?; + let fields = schema + .as_struct() + .fields() + .iter() + .zip(arrow_schema.fields().iter()) + .filter_map(|(field, arrow_field)| { + parquet_write_arrow_field(field, arrow_field).transpose() + }) + .collect::>>()?; + if fields.is_empty() { + return Err(Error::new( + ErrorKind::FeatureUnsupported, + "Cannot write a schema with no Parquet physical fields", + )); + } + Ok(ArrowSchema::new_with_metadata( + fields, + arrow_schema.metadata().clone(), + )) +} + /// Convert iceberg type to an arrow type. pub fn type_to_arrow_type(ty: &Type) -> Result { let mut converter = ToArrowSchemaConverter; @@ -2221,6 +2393,176 @@ mod tests { ); } + #[test] + fn test_parquet_arrow_schema_omits_unknown_struct_fields() { + let schema = Schema::builder() + .with_fields(vec![ + NestedField::optional(1, "unknown", PrimitiveType::Unknown.into()).into(), + NestedField::optional( + 2, + "struct", + Type::Struct(StructType::new(vec![ + NestedField::optional(3, "unknown", PrimitiveType::Unknown.into()).into(), + NestedField::optional(4, "known", PrimitiveType::Int.into()).into(), + ])), + ) + .into(), + ]) + .build() + .unwrap(); + + let arrow_schema = schema_to_arrow_schema_for_parquet_write(&schema).unwrap(); + assert_eq!(arrow_schema.fields().len(), 1); + assert_eq!(arrow_schema.field(0).name(), "struct"); + let DataType::Struct(fields) = arrow_schema.field(0).data_type() else { + panic!("expected struct field"); + }; + assert_eq!(fields.len(), 1); + assert_eq!(fields[0].name(), "known"); + } + + #[test] + fn test_parquet_arrow_schema_rejects_empty_struct() { + let schema = Schema::builder() + .with_fields(vec![ + NestedField::optional(1, "known", PrimitiveType::Int.into()).into(), + NestedField::optional( + 2, + "empty_struct", + Type::Struct(StructType::new(vec![ + NestedField::optional(3, "unknown", PrimitiveType::Unknown.into()).into(), + ])), + ) + .into(), + ]) + .build() + .unwrap(); + + let err = schema_to_arrow_schema_for_parquet_write(&schema) + .expect_err("conversion must fail when a struct has no physical fields"); + assert_eq!( + err.message(), + "Cannot write struct field 2 with no Parquet physical fields" + ); + } + + #[test] + fn test_parquet_arrow_schema_rejects_empty_schema() { + let schema = Schema::builder().build().unwrap(); + let err = schema_to_arrow_schema_for_parquet_write(&schema) + .expect_err("conversion must fail when a schema has no fields"); + assert_eq!( + err.message(), + "Cannot write a schema with no Parquet physical fields" + ); + } + + #[test] + fn test_parquet_arrow_schema_rejects_variant() { + let schema = Schema::builder() + .with_fields(vec![ + NestedField::optional(1, "variant", Type::Variant(VariantType)).into(), + ]) + .build() + .unwrap(); + let err = schema_to_arrow_schema_for_parquet_write(&schema) + .expect_err("variant Parquet writes are not supported"); + assert_eq!( + err.message(), + "Field 1 has variant type, which is not yet implemented" + ); + } + + #[test] + fn test_parquet_arrow_schema_rejects_all_unknown_fields() { + let schema = Schema::builder() + .with_fields(vec![ + NestedField::optional(1, "unknown", PrimitiveType::Unknown.into()).into(), + ]) + .build() + .unwrap(); + + let err = schema_to_arrow_schema_for_parquet_write(&schema) + .expect_err("conversion must fail when every field is unknown"); + assert_eq!( + err.message(), + "Cannot write a schema with no Parquet physical fields" + ); + } + + #[test] + fn test_parquet_arrow_schema_rejects_unknown_container_values() { + let list_schema = Schema::builder() + .with_fields(vec![ + NestedField::optional( + 1, + "list", + Type::List(ListType::new( + NestedField::list_element(2, PrimitiveType::Unknown.into(), false).into(), + )), + ) + .into(), + ]) + .build() + .unwrap(); + let err = schema_to_arrow_schema_for_parquet_write(&list_schema) + .expect_err("conversion must fail for an unknown list element"); + assert_eq!( + err.message(), + "Cannot write list element 2 with no Parquet physical fields" + ); + + let list_of_empty_struct_schema = Schema::builder() + .with_fields(vec![ + NestedField::optional( + 1, + "list", + Type::List(ListType::new( + NestedField::list_element( + 2, + Type::Struct(StructType::new(vec![ + NestedField::optional(3, "unknown", PrimitiveType::Unknown.into()) + .into(), + ])), + false, + ) + .into(), + )), + ) + .into(), + ]) + .build() + .unwrap(); + let err = schema_to_arrow_schema_for_parquet_write(&list_of_empty_struct_schema) + .expect_err("conversion must fail for a list of empty structs"); + assert_eq!( + err.message(), + "Cannot write struct field 2 with no Parquet physical fields" + ); + + let map_schema = Schema::builder() + .with_fields(vec![ + NestedField::optional( + 1, + "map", + Type::Map(MapType::new( + NestedField::map_key_element(2, PrimitiveType::String.into()).into(), + NestedField::map_value_element(3, PrimitiveType::Unknown.into(), false) + .into(), + )), + ) + .into(), + ]) + .build() + .unwrap(); + let err = schema_to_arrow_schema_for_parquet_write(&map_schema) + .expect_err("conversion must fail for an unknown map value"); + assert_eq!( + err.message(), + "Cannot write map value 3 with no Parquet physical fields" + ); + } + #[test] fn test_type_conversion() { // test primitive type diff --git a/crates/iceberg/src/writer/file_writer/parquet_writer.rs b/crates/iceberg/src/writer/file_writer/parquet_writer.rs index fbf333c7bf..36200b46b7 100644 --- a/crates/iceberg/src/writer/file_writer/parquet_writer.rs +++ b/crates/iceberg/src/writer/file_writer/parquet_writer.rs @@ -20,7 +20,10 @@ use std::collections::HashMap; use std::sync::Arc; -use arrow_schema::SchemaRef as ArrowSchemaRef; +use arrow_array::{ + Array, ArrayRef, ListArray, MapArray, RecordBatch, RecordBatchOptions, StructArray, +}; +use arrow_schema::{DataType, FieldRef, Fields, SchemaRef as ArrowSchemaRef}; use bytes::Bytes; use futures::future::BoxFuture; use itertools::Itertools; @@ -36,7 +39,8 @@ use parquet::file::statistics::Statistics; use super::{FileWriter, FileWriterBuilder}; use crate::arrow::{ ArrowFileReader, DEFAULT_MAP_FIELD_NAME, FieldMatchMode, NanValueCountVisitor, - get_parquet_stat_max_as_datum, get_parquet_stat_min_as_datum, + get_field_id_from_metadata, get_parquet_stat_max_as_datum, get_parquet_stat_min_as_datum, + schema_to_arrow_schema, schema_to_arrow_schema_for_parquet_write, }; use crate::compression::CompressionCodec; use crate::encryption::{EncryptionManager, StandardKeyMetadata}; @@ -50,7 +54,12 @@ use crate::transform::create_transform_function; use crate::writer::{CurrentFileStatus, DataFile}; use crate::{Error, ErrorKind, Result}; -/// ParquetWriterBuilder is used to builder a [`ParquetWriter`] +/// Builds a [`ParquetWriter`] for an Iceberg schema. +/// +/// Input batches use the logical Iceberg schema. Before writing, [`ParquetWriter`] projects each +/// batch to the physical Parquet schema and omits [`PrimitiveType::Unknown`] fields. Building the +/// writer fails when omission would leave no valid physical representation, such as an all-unknown +/// schema or an unknown list element or map key/value. #[derive(Clone, Debug)] pub struct ParquetWriterBuilder { props: WriterProperties, @@ -174,16 +183,334 @@ impl FileWriterBuilder for ParquetWriterBuilder { resolve_writer_properties(self.props.clone(), key_metadata.as_ref())?; Ok(ParquetWriter { schema: self.schema.clone(), + table_arrow_schema: Arc::new(schema_to_arrow_schema(&self.schema)?), + parquet_arrow_schema: Arc::new(schema_to_arrow_schema_for_parquet_write(&self.schema)?), + match_mode: self.match_mode, + cached_projection: None, inner_writer: None, writer_properties, current_row_num: 0, output_file, - nan_value_count_visitor: NanValueCountVisitor::new_with_match_mode(self.match_mode), + nan_value_count_visitor: NanValueCountVisitor::new(), key_metadata, }) } } +fn find_field_index( + fields: &Fields, + target: &FieldRef, + match_mode: FieldMatchMode, +) -> Result> { + match match_mode { + FieldMatchMode::Id => { + let target_id = get_field_id_from_metadata(target)?; + for (index, field) in fields.iter().enumerate() { + if get_field_id_from_metadata(field)? == target_id { + return Ok(Some(index)); + } + } + Ok(None) + } + FieldMatchMode::Name => Ok(fields + .iter() + .position(|field| field.name() == target.name())), + } +} + +/// Validate a source field against its full logical Iceberg field before projection. +/// Unknown fields must use Arrow Null. List elements and map keys/values are all checked because +/// removing one would change the shape of its container rather than omit a regular struct field. +fn validate_source_field_for_parquet_write( + source: &FieldRef, + target: &FieldRef, + match_mode: FieldMatchMode, +) -> Result<()> { + match (source.data_type(), target.data_type()) { + (DataType::Null, DataType::Null) => Ok(()), + (source_type, DataType::Null) => Err(Error::new( + ErrorKind::DataInvalid, + format!( + "Expected unknown field {} to use Arrow Null type, got {source_type}", + source.name() + ), + )), + (DataType::Struct(source_fields), DataType::Struct(target_fields)) => { + validate_source_fields_for_parquet_write(source_fields, target_fields, match_mode) + } + (DataType::List(source_element), DataType::List(target_element)) => { + // A list has one positional element. Its Arrow name may be `item` or `element`; + // only ID mode requires an explicit identity check. + if matches!(match_mode, FieldMatchMode::Id) + && get_field_id_from_metadata(source_element)? + != get_field_id_from_metadata(target_element)? + { + return Err(Error::new( + ErrorKind::DataInvalid, + format!( + "List element field {} does not match configured field {} for Parquet write", + source_element.name(), + target_element.name() + ), + )); + } + validate_source_field_for_parquet_write(source_element, target_element, match_mode) + } + (DataType::Map(source_entries, _), DataType::Map(target_entries, _)) => { + match (source_entries.data_type(), target_entries.data_type()) { + (DataType::Struct(source_fields), DataType::Struct(target_fields)) => { + validate_source_fields_for_parquet_write( + source_fields, + target_fields, + match_mode, + ) + } + _ => Ok(()), + } + } + _ => Ok(()), + } +} + +fn validate_source_fields_for_parquet_write( + source_fields: &Fields, + target_fields: &Fields, + match_mode: FieldMatchMode, +) -> Result<()> { + let mut matched_targets = vec![false; target_fields.len()]; + for source_field in source_fields { + let Some(target_index) = find_field_index(target_fields, source_field, match_mode)? else { + return Err(Error::new( + ErrorKind::DataInvalid, + format!( + "Field {} is not present in the configured Iceberg schema for Parquet write", + source_field.name() + ), + )); + }; + if matched_targets[target_index] { + return Err(Error::new( + ErrorKind::DataInvalid, + format!( + "Multiple source fields match configured field {} for Parquet write", + target_fields[target_index].name() + ), + )); + } + matched_targets[target_index] = true; + validate_source_field_for_parquet_write( + source_field, + &target_fields[target_index], + match_mode, + )?; + } + Ok(()) +} + +/// Array operations chosen once for a particular incoming Arrow schema. +enum ArrayProjection { + Identity, + Struct(Vec<(usize, ArrayProjection)>), + List(Box), + Map(Box), +} + +impl ArrayProjection { + fn new(source: &FieldRef, target: &FieldRef, match_mode: FieldMatchMode) -> Result { + if source.data_type() == target.data_type() { + return Ok(Self::Identity); + } + + match (source.data_type(), target.data_type()) { + (DataType::Struct(source_fields), DataType::Struct(target_fields)) => { + let columns = target_fields + .iter() + .map(|target_field| { + let index = find_field_index(source_fields, target_field, match_mode)? + .ok_or_else(|| { + Error::new( + ErrorKind::DataInvalid, + format!( + "Field {} is missing from struct array for Parquet write", + target_field.name() + ), + ) + })?; + Ok(( + index, + Self::new(&source_fields[index], target_field, match_mode)?, + )) + }) + .collect::>>()?; + Ok(Self::Struct(columns)) + } + (DataType::List(source_element), DataType::List(target_element)) => Ok(Self::List( + Box::new(Self::new(source_element, target_element, match_mode)?), + )), + (DataType::Map(source_entries, _), DataType::Map(target_entries, _)) => Ok(Self::Map( + Box::new(Self::new(source_entries, target_entries, match_mode)?), + )), + (source_type, target_type) => Err(Error::new( + ErrorKind::DataInvalid, + format!( + "Cannot project Arrow type {source_type} to {target_type} for Parquet field {}", + target.name() + ), + )), + } + } + + fn project(&self, array: &ArrayRef, target: &FieldRef) -> Result { + match self { + Self::Identity => Ok(array.clone()), + Self::Struct(columns) => { + let source = array + .as_any() + .downcast_ref::() + .ok_or_else(|| { + Error::new( + ErrorKind::DataInvalid, + "Expected struct array for Parquet write", + ) + })?; + let DataType::Struct(target_fields) = target.data_type() else { + unreachable!(); + }; + let projected = columns + .iter() + .zip(target_fields.iter()) + .map(|((index, projection), field)| { + projection.project(source.column(*index), field) + }) + .collect::>>()?; + Ok(Arc::new(StructArray::try_new_with_length( + target_fields.clone(), + projected, + source.nulls().cloned(), + source.len(), + )?)) + } + Self::List(element_projection) => { + let source = array.as_any().downcast_ref::().ok_or_else(|| { + Error::new( + ErrorKind::DataInvalid, + "Expected list array for Parquet write", + ) + })?; + let DataType::List(target_element) = target.data_type() else { + unreachable!(); + }; + let values = element_projection.project(source.values(), target_element)?; + Ok(Arc::new(ListArray::try_new( + target_element.clone(), + source.offsets().clone(), + values, + source.nulls().cloned(), + )?)) + } + Self::Map(entries_projection) => { + let source = array.as_any().downcast_ref::().ok_or_else(|| { + Error::new( + ErrorKind::DataInvalid, + "Expected map array for Parquet write", + ) + })?; + let DataType::Map(target_entries_field, ordered) = target.data_type() else { + unreachable!(); + }; + let source_entries: ArrayRef = Arc::new(source.entries().clone()); + let entries = entries_projection.project(&source_entries, target_entries_field)?; + let entries = entries + .as_any() + .downcast_ref::() + .ok_or_else(|| { + Error::new( + ErrorKind::Unexpected, + "Projected Parquet map entries are not a struct array", + ) + })? + .clone(); + Ok(Arc::new(MapArray::try_new( + target_entries_field.clone(), + source.offsets().clone(), + entries, + source.nulls().cloned(), + *ordered, + )?)) + } + } + } +} + +struct BatchProjection { + source_schema: ArrowSchemaRef, + columns: Vec<(usize, ArrayProjection)>, +} + +impl BatchProjection { + fn new( + source_schema: ArrowSchemaRef, + logical_schema: &ArrowSchemaRef, + target_schema: &ArrowSchemaRef, + match_mode: FieldMatchMode, + ) -> Result { + validate_source_fields_for_parquet_write( + source_schema.fields(), + logical_schema.fields(), + match_mode, + )?; + let columns = target_schema + .fields() + .iter() + .map(|target_field| { + let index = find_field_index(source_schema.fields(), target_field, match_mode)? + .ok_or_else(|| { + Error::new( + ErrorKind::DataInvalid, + format!( + "Field {} is missing from record batch for Parquet write", + target_field.name() + ), + ) + })?; + Ok(( + index, + ArrayProjection::new(&source_schema.fields()[index], target_field, match_mode)?, + )) + }) + .collect::>>()?; + Ok(Self { + source_schema, + columns, + }) + } + + fn project(&self, batch: &RecordBatch, target_schema: ArrowSchemaRef) -> Result { + let columns = self + .columns + .iter() + .zip(target_schema.fields().iter()) + .map(|((index, projection), field)| projection.project(batch.column(*index), field)) + .collect::>>()?; + let options = RecordBatchOptions::default().with_row_count(Some(batch.num_rows())); + RecordBatch::try_new_with_options(target_schema, columns, &options).map_err(Into::into) + } +} + +#[cfg(test)] +fn project_batch_for_parquet_write( + batch: &RecordBatch, + logical_schema: ArrowSchemaRef, + target_schema: ArrowSchemaRef, + match_mode: FieldMatchMode, +) -> Result { + if batch.schema_ref() == &target_schema { + return Ok(batch.clone()); + } + BatchProjection::new(batch.schema(), &logical_schema, &target_schema, match_mode)? + .project(batch, target_schema) +} + /// A mapping from Parquet column path names to internal field id struct IndexByParquetPathName { name_to_id: HashMap, @@ -309,9 +636,21 @@ impl SchemaVisitor for IndexByParquetPathName { } } -/// `ParquetWriter`` is used to write arrow data into parquet file on storage. +/// Writes logical Iceberg Arrow batches to a Parquet file. +/// +/// The input batches are matched to the configured Iceberg schema by field ID by default, or by +/// name when configured through [`ParquetWriterBuilder::with_match_mode`]. In ID mode, every input +/// field must contain valid Iceberg field-ID metadata. Fields with +/// [`PrimitiveType::Unknown`] are accepted in logical input batches and omitted from the physical +/// Parquet file. pub struct ParquetWriter { schema: SchemaRef, + /// Arrow schema converted from the full Iceberg table schema. + table_arrow_schema: ArrowSchemaRef, + /// Arrow schema projected to fields with a Parquet physical representation. + parquet_arrow_schema: ArrowSchemaRef, + match_mode: FieldMatchMode, + cached_projection: Option, output_file: OutputFile, inner_writer: Option>, writer_properties: WriterProperties, @@ -617,28 +956,44 @@ fn resolve_writer_properties( } impl FileWriter for ParquetWriter { - async fn write(&mut self, batch: &arrow_array::RecordBatch) -> Result<()> { + async fn write(&mut self, batch: &RecordBatch) -> Result<()> { // Skip empty batch if batch.num_rows() == 0 { return Ok(()); } - self.current_row_num += batch.num_rows(); - - let batch_c = batch.clone(); - self.nan_value_count_visitor - .compute(self.schema.clone(), batch_c)?; + let batch = if batch.schema_ref() == &self.parquet_arrow_schema { + batch.clone() + } else { + let source_schema = batch.schema(); + let same_schema = self.cached_projection.as_ref().is_some_and(|projection| { + Arc::ptr_eq(&projection.source_schema, &source_schema) + || projection.source_schema == source_schema + }); + if !same_schema { + self.cached_projection = Some(BatchProjection::new( + source_schema, + &self.table_arrow_schema, + &self.parquet_arrow_schema, + self.match_mode, + )?); + } + self.cached_projection + .as_ref() + .expect("projection initialized above") + .project(batch, self.parquet_arrow_schema.clone())? + }; + self.nan_value_count_visitor.compute(&batch)?; // Lazy initialize the writer let writer = if let Some(writer) = &mut self.inner_writer { writer } else { - let arrow_schema: ArrowSchemaRef = Arc::new(self.schema.as_ref().try_into()?); let inner_writer = self.output_file.writer().await?; let async_writer = AsyncFileWriter::new(inner_writer); let writer = AsyncArrowWriter::try_new( async_writer, - arrow_schema.clone(), + self.parquet_arrow_schema.clone(), Some(self.writer_properties.clone()), ) .map_err(|err| { @@ -649,7 +1004,7 @@ impl FileWriter for ParquetWriter { self.inner_writer.as_mut().unwrap() }; - writer.write(batch).await.map_err(|err| { + writer.write(&batch).await.map_err(|err| { Error::new( ErrorKind::Unexpected, "Failed to write using parquet writer.", @@ -657,6 +1012,8 @@ impl FileWriter for ParquetWriter { .with_source(err) })?; + self.current_row_num += batch.num_rows(); + Ok(()) } @@ -758,10 +1115,10 @@ mod tests { use anyhow::Result; use arrow_array::builder::{Float32Builder, Int32Builder, MapBuilder}; - use arrow_array::types::{Float32Type, Int64Type}; + use arrow_array::types::{Float32Type, Int32Type, Int64Type}; use arrow_array::{ Array, ArrayRef, BooleanArray, Decimal128Array, Float32Array, Float64Array, Int32Array, - Int64Array, ListArray, MapArray, RecordBatch, StructArray, + Int64Array, ListArray, MapArray, NullArray, RecordBatch, StructArray, }; use arrow_schema::{DataType, Field, Fields, SchemaRef as ArrowSchemaRef}; use arrow_select::concat::concat_batches; @@ -776,7 +1133,7 @@ mod tests { use super::*; use crate::Runtime; - use crate::arrow::{ArrowReaderBuilder, schema_to_arrow_schema}; + use crate::arrow::ArrowReaderBuilder; use crate::io::FileIO; use crate::scan::{FileScanTask, FileScanTaskStream}; use crate::spec::decimal_utils::{decimal_mantissa, decimal_new, decimal_scale}; @@ -848,6 +1205,381 @@ mod tests { .unwrap() } + #[test] + fn test_project_batch_for_parquet_omits_unknown_fields() { + let schema = Schema::builder() + .with_fields(vec![ + NestedField::optional(1, "unknown", PrimitiveType::Unknown.into()).into(), + NestedField::optional( + 2, + "struct", + Type::Struct(StructType::new(vec![ + NestedField::optional(3, "unknown", PrimitiveType::Unknown.into()).into(), + NestedField::optional(4, "known", PrimitiveType::Int.into()).into(), + ])), + ) + .into(), + ]) + .build() + .unwrap(); + let source_schema = Arc::new(schema_to_arrow_schema(&schema).unwrap()); + let DataType::Struct(struct_fields) = source_schema.field(1).data_type() else { + panic!("expected struct field"); + }; + let struct_array = Arc::new(StructArray::new( + struct_fields.clone(), + vec![ + Arc::new(NullArray::new(2)), + Arc::new(Int32Array::from(vec![Some(1), Some(2)])), + ], + None, + )); + let batch = RecordBatch::try_new(source_schema.clone(), vec![ + Arc::new(NullArray::new(2)), + struct_array, + ]) + .unwrap(); + let logical_schema = source_schema.clone(); + let target_schema = Arc::new(schema_to_arrow_schema_for_parquet_write(&schema).unwrap()); + + let projected = project_batch_for_parquet_write( + &batch, + logical_schema, + target_schema, + FieldMatchMode::Id, + ) + .unwrap(); + + assert_eq!(projected.num_rows(), 2); + assert_eq!(projected.num_columns(), 1); + let projected_struct = projected + .column(0) + .as_any() + .downcast_ref::() + .unwrap(); + assert_eq!(projected_struct.num_columns(), 1); + assert_eq!(projected_struct.fields()[0].name(), "known"); + + let mut nan_visitor = NanValueCountVisitor::new(); + nan_visitor.compute(&projected).unwrap(); + assert!(nan_visitor.nan_value_counts.is_empty()); + } + + #[test] + fn test_project_batch_for_parquet_honors_name_match_mode() { + let schema = Schema::builder() + .with_fields(vec![ + NestedField::optional(1, "known", PrimitiveType::Int.into()).into(), + ]) + .build() + .unwrap(); + let source_schema = Arc::new(arrow_schema::Schema::new(vec![ + Field::new("known", DataType::Int32, true).with_metadata(HashMap::from([( + PARQUET_FIELD_ID_META_KEY.to_string(), + "2".to_string(), + )])), + ])); + let batch = RecordBatch::try_new(source_schema, vec![Arc::new(Int32Array::from(vec![10]))]) + .unwrap(); + let logical_schema = Arc::new(schema_to_arrow_schema(&schema).unwrap()); + let target_schema = Arc::new(schema_to_arrow_schema_for_parquet_write(&schema).unwrap()); + + let projected = project_batch_for_parquet_write( + &batch, + logical_schema, + target_schema, + FieldMatchMode::Name, + ) + .unwrap(); + + assert_eq!(projected.num_columns(), 1); + let values = projected + .column(0) + .as_any() + .downcast_ref::() + .unwrap(); + assert_eq!(values.value(0), 10); + } + + #[test] + fn test_project_batch_for_parquet_rejects_unexpected_fields() { + let schema = Schema::builder() + .with_fields(vec![ + NestedField::optional(1, "known", PrimitiveType::Int.into()).into(), + NestedField::optional( + 2, + "struct", + Type::Struct(StructType::new(vec![ + NestedField::optional(3, "unknown", PrimitiveType::Unknown.into()).into(), + NestedField::optional(4, "known", PrimitiveType::Int.into()).into(), + ])), + ) + .into(), + ]) + .build() + .unwrap(); + let logical_schema = Arc::new(schema_to_arrow_schema(&schema).unwrap()); + let target_schema = Arc::new(schema_to_arrow_schema_for_parquet_write(&schema).unwrap()); + let DataType::Struct(logical_struct_fields) = logical_schema.field(1).data_type() else { + panic!("expected struct field"); + }; + let logical_struct = Arc::new(StructArray::new( + logical_struct_fields.clone(), + vec![ + Arc::new(NullArray::new(1)), + Arc::new(Int32Array::from(vec![10])), + ], + None, + )); + let unexpected_field = Field::new("unexpected", DataType::Int32, true).with_metadata( + HashMap::from([(PARQUET_FIELD_ID_META_KEY.to_string(), "5".to_string())]), + ); + + let top_level_batch = RecordBatch::try_new( + Arc::new(arrow_schema::Schema::new(vec![ + logical_schema.field(0).clone(), + logical_schema.field(1).clone(), + unexpected_field.clone(), + ])), + vec![ + Arc::new(Int32Array::from(vec![1])), + logical_struct.clone(), + Arc::new(Int32Array::from(vec![20])), + ], + ) + .unwrap(); + + let mut source_struct_fields: Vec = + logical_struct_fields.iter().cloned().collect(); + source_struct_fields.push(Arc::new(unexpected_field)); + let source_struct_fields: Fields = source_struct_fields.into(); + let source_struct = Arc::new(StructArray::new( + source_struct_fields.clone(), + vec![ + Arc::new(NullArray::new(1)), + Arc::new(Int32Array::from(vec![10])), + Arc::new(Int32Array::from(vec![20])), + ], + None, + )); + let source_struct_field = logical_schema + .field(1) + .clone() + .with_data_type(DataType::Struct(source_struct_fields)); + let nested_batch = RecordBatch::try_new( + Arc::new(arrow_schema::Schema::new(vec![ + logical_schema.field(0).clone(), + source_struct_field, + ])), + vec![Arc::new(Int32Array::from(vec![1])), source_struct], + ) + .unwrap(); + + for match_mode in [FieldMatchMode::Id, FieldMatchMode::Name] { + for batch in [&top_level_batch, &nested_batch] { + let error = project_batch_for_parquet_write( + batch, + logical_schema.clone(), + target_schema.clone(), + match_mode, + ) + .unwrap_err(); + assert!( + error + .to_string() + .contains("is not present in the configured Iceberg schema"), + "unexpected error: {error}" + ); + } + } + } + + #[test] + fn test_project_batch_for_parquet_validates_list_element_identity() { + let schema = Schema::builder() + .with_fields(vec![ + NestedField::optional( + 1, + "list", + Type::List(ListType::new( + NestedField::list_element(2, PrimitiveType::Int.into(), false).into(), + )), + ) + .into(), + ]) + .build() + .unwrap(); + let logical_schema = Arc::new(schema_to_arrow_schema(&schema).unwrap()); + let target_schema = Arc::new(schema_to_arrow_schema_for_parquet_write(&schema).unwrap()); + let DataType::List(target_element) = logical_schema.field(0).data_type() else { + panic!("expected list field"); + }; + let elements = [ + ( + Arc::new( + Field::new( + target_element.name(), + target_element.data_type().clone(), + target_element.is_nullable(), + ) + .with_metadata(HashMap::from([( + PARQUET_FIELD_ID_META_KEY.to_string(), + "3".to_string(), + )])), + ), + FieldMatchMode::Id, + true, + ), + ( + Arc::new(Field::new("item", DataType::Int32, true)), + FieldMatchMode::Name, + false, + ), + ]; + + for (source_element, match_mode, should_fail) in elements { + let list_parts = + ListArray::from_iter_primitive::([Some(vec![Some(10)])]) + .into_parts(); + let list_array = Arc::new(ListArray::new( + source_element.clone(), + list_parts.1, + list_parts.2, + list_parts.3, + )); + let source_field = logical_schema + .field(0) + .clone() + .with_data_type(DataType::List(source_element)); + let batch = RecordBatch::try_new( + Arc::new(arrow_schema::Schema::new(vec![source_field])), + vec![list_array], + ) + .unwrap(); + + let result = project_batch_for_parquet_write( + &batch, + logical_schema.clone(), + target_schema.clone(), + match_mode, + ); + if should_fail { + let error = result.expect_err("mismatched list element ID must fail"); + assert_eq!( + error.message(), + "List element field element does not match configured field element for Parquet write" + ); + } else { + let projected = result.expect("name mode accepts the Arrow item name"); + assert_eq!(projected.schema(), target_schema); + } + } + } + + #[test] + fn test_project_batch_for_parquet_rejects_invalid_field_id_metadata() { + let schema = Schema::builder() + .with_fields(vec![ + NestedField::optional(1, "known", PrimitiveType::Int.into()).into(), + ]) + .build() + .unwrap(); + let target_schema = Arc::new(schema_to_arrow_schema_for_parquet_write(&schema).unwrap()); + let logical_schema = Arc::new(schema_to_arrow_schema(&schema).unwrap()); + + let cases = [ + ( + Field::new("known", DataType::Int32, true), + "Field id not found in metadata", + ), + ( + Field::new("known", DataType::Int32, true).with_metadata(HashMap::from([( + PARQUET_FIELD_ID_META_KEY.to_string(), + "invalid".to_string(), + )])), + "Failed to parse field id", + ), + ]; + + for (field, expected_message) in cases { + let batch = + RecordBatch::try_new(Arc::new(arrow_schema::Schema::new(vec![field])), vec![ + Arc::new(Int32Array::from(vec![10])), + ]) + .unwrap(); + + let error = project_batch_for_parquet_write( + &batch, + logical_schema.clone(), + target_schema.clone(), + FieldMatchMode::Id, + ) + .unwrap_err(); + assert!( + error.to_string().contains(expected_message), + "unexpected error: {error}" + ); + } + } + + #[tokio::test] + async fn test_parquet_writer_reuses_projection_after_rejected_batch() -> Result<()> { + let schema = Arc::new( + Schema::builder() + .with_fields(vec![ + NestedField::optional(1, "known", PrimitiveType::Int.into()).into(), + NestedField::optional(2, "unknown", PrimitiveType::Unknown.into()).into(), + ]) + .build()?, + ); + let logical_schema = schema_to_arrow_schema(&schema)?; + let extra = Field::new("extra", DataType::Int32, true).with_metadata(HashMap::from([( + PARQUET_FIELD_ID_META_KEY.to_string(), + "3".to_string(), + )])); + let invalid = RecordBatch::try_new( + Arc::new(arrow_schema::Schema::new(vec![ + logical_schema.field(0).clone(), + extra, + ])), + vec![ + Arc::new(Int32Array::from(vec![1])), + Arc::new(Int32Array::from(vec![2])), + ], + )?; + let valid = RecordBatch::try_new(Arc::new(logical_schema), vec![ + Arc::new(Int32Array::from(vec![3])), + Arc::new(NullArray::new(1)), + ])?; + let temp_dir = TempDir::new()?; + let path = temp_dir.path().join("data.parquet"); + let output_file = FileIO::new_with_fs().new_output(path.to_str().unwrap())?; + let mut writer = ParquetWriterBuilder::new(WriterProperties::default(), schema) + .build(output_file) + .await?; + + let err = writer + .write(&invalid) + .await + .expect_err("extra field must fail"); + assert!(err.message().contains("Field extra is not present")); + assert_eq!(writer.current_row_num(), 0); + writer.write(&valid).await?; + writer.write(&valid).await?; + assert_eq!(writer.current_row_num(), 2); + let files = writer.close().await?; + assert_eq!(files.len(), 1); + let data_file = files + .into_iter() + .next() + .unwrap() + .partition(Struct::empty()) + .partition_spec_id(0) + .build()?; + assert_eq!(data_file.record_count(), 2); + Ok(()) + } + fn nested_schema_for_test() -> Schema { // Int, Struct(Int,Int), String, List(Int), Struct(Struct(Int)), Map(String, List(Int)) Schema::builder()