diff --git a/crates/catalog/glue/src/schema.rs b/crates/catalog/glue/src/schema.rs index 59d51f3dc1..747435611b 100644 --- a/crates/catalog/glue/src/schema.rs +++ b/crates/catalog/glue/src/schema.rs @@ -149,6 +149,12 @@ impl SchemaVisitor for GlueSchemaBuilder { fn primitive(&mut self, p: &PrimitiveType) -> Result { let glue_type = match p { + PrimitiveType::Unknown => { + return Err(Error::new( + ErrorKind::FeatureUnsupported, + format!("Conversion from {p:?} is not supported"), + )); + } PrimitiveType::Boolean => "boolean".to_string(), PrimitiveType::Int => "int".to_string(), PrimitiveType::Long => "bigint".to_string(), diff --git a/crates/catalog/hms/src/schema.rs b/crates/catalog/hms/src/schema.rs index e21a75f976..f04ee11c17 100644 --- a/crates/catalog/hms/src/schema.rs +++ b/crates/catalog/hms/src/schema.rs @@ -100,6 +100,12 @@ impl SchemaVisitor for HiveSchemaBuilder { fn primitive(&mut self, p: &PrimitiveType) -> Result { let hive_type = match p { + PrimitiveType::Unknown => { + return Err(Error::new( + ErrorKind::FeatureUnsupported, + format!("Conversion from {p:?} is not supported"), + )); + } PrimitiveType::Boolean => "boolean".to_string(), PrimitiveType::Int => "int".to_string(), PrimitiveType::Long => "bigint".to_string(), diff --git a/crates/iceberg/public-api.txt b/crates/iceberg/public-api.txt index 3d9abf2275..18e03cb257 100644 --- a/crates/iceberg/public-api.txt +++ b/crates/iceberg/public-api.txt @@ -1655,6 +1655,7 @@ pub iceberg::spec::PrimitiveType::Timestamp pub iceberg::spec::PrimitiveType::TimestampNs pub iceberg::spec::PrimitiveType::Timestamptz pub iceberg::spec::PrimitiveType::TimestamptzNs +pub iceberg::spec::PrimitiveType::Unknown pub iceberg::spec::PrimitiveType::Uuid impl iceberg::spec::PrimitiveType pub fn iceberg::spec::PrimitiveType::compatible(&self, literal: &iceberg::spec::PrimitiveLiteral) -> bool diff --git a/crates/iceberg/src/arrow/nan_val_cnt_visitor.rs b/crates/iceberg/src/arrow/nan_val_cnt_visitor.rs index d01f3e9e56..6897f02c06 100644 --- a/crates/iceberg/src/arrow/nan_val_cnt_visitor.rs +++ b/crates/iceberg/src/arrow/nan_val_cnt_visitor.rs @@ -21,15 +21,17 @@ use std::collections::HashMap; use std::collections::hash_map::Entry; use std::sync::Arc; -use arrow_array::{ArrayRef, Float32Array, Float64Array, RecordBatch, StructArray}; +use arrow_array::{ + ArrayRef, Float32Array, Float64Array, ListArray, MapArray, RecordBatch, StructArray, +}; use arrow_schema::DataType; -use crate::Result; -use crate::arrow::{ArrowArrayAccessor, FieldMatchMode}; +use crate::arrow::FieldMatchMode; use crate::spec::{ ListType, MapType, NestedFieldRef, PrimitiveType, Schema, SchemaRef, SchemaWithPartnerVisitor, - StructType, VariantType, visit_struct_with_partner, + StructType, Type, VariantType, }; +use crate::{Error, ErrorKind, Result}; macro_rules! cast_and_update_cnt_map { ($t:ty, $col:ident, $self:ident, $field_id:ident) => { @@ -152,6 +154,78 @@ impl SchemaWithPartnerVisitor for NanValueCountVisitor { } impl NanValueCountVisitor { + fn visit_field(&mut self, field: &NestedFieldRef, array: &ArrayRef) -> Result<()> { + let field_id = field.id; + count_float_nans!(array, self, field_id); + + match field.field_type.as_ref() { + Type::Primitive(_) | Type::Variant(_) => Ok(()), + Type::Struct(struct_type) => self.visit_struct(struct_type, array), + Type::List(list_type) => { + let list_array = array.as_any().downcast_ref::().ok_or_else(|| { + Error::new( + ErrorKind::DataInvalid, + format!( + "Expected list array for field {}, got {}", + field.id, + array.data_type() + ), + ) + })?; + self.visit_field(&list_type.element_field, list_array.values()) + } + Type::Map(map_type) => { + let map_array = array.as_any().downcast_ref::().ok_or_else(|| { + Error::new( + ErrorKind::DataInvalid, + format!( + "Expected map array for field {}, got {}", + field.id, + array.data_type() + ), + ) + })?; + self.visit_field(&map_type.key_field, map_array.keys())?; + self.visit_field(&map_type.value_field, map_array.values()) + } + } + } + + fn visit_struct(&mut self, struct_type: &StructType, array: &ArrayRef) -> Result<()> { + let struct_array = array + .as_any() + .downcast_ref::() + .ok_or_else(|| { + Error::new( + ErrorKind::DataInvalid, + format!("Expected struct array, got {}", array.data_type()), + ) + })?; + + for field in struct_type.fields() { + if matches!( + field.field_type.as_ref(), + Type::Primitive(PrimitiveType::Unknown) + ) { + continue; + } + + let field_position = struct_array + .fields() + .iter() + .position(|arrow_field| self.match_mode.match_field(arrow_field, field)) + .ok_or_else(|| { + Error::new( + ErrorKind::DataInvalid, + format!("Field id {} not found in struct array", field.id), + ) + })?; + self.visit_field(field, struct_array.column(field_position))?; + } + + Ok(()) + } + /// Creates new instance of NanValueCountVisitor pub fn new() -> Self { Self::new_with_match_mode(FieldMatchMode::Id) @@ -167,17 +241,8 @@ impl NanValueCountVisitor { /// 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, - )?; - - Ok(()) + self.visit_struct(schema.as_struct(), &struct_arr) } } diff --git a/crates/iceberg/src/arrow/schema.rs b/crates/iceberg/src/arrow/schema.rs index dd26824c3a..004ce6296d 100644 --- a/crates/iceberg/src/arrow/schema.rs +++ b/crates/iceberg/src/arrow/schema.rs @@ -191,6 +191,7 @@ fn visit_type(r#type: &DataType, visitor: &mut V) -> Resu | DataType::Utf8 | DataType::LargeUtf8 | DataType::Utf8View + | DataType::Null | DataType::Binary | DataType::LargeBinary | DataType::BinaryView @@ -509,6 +510,7 @@ impl ArrowSchemaVisitor for ArrowSchemaConverter { fn primitive(&mut self, p: &DataType) -> Result { match p { + DataType::Null => Ok(Type::Primitive(PrimitiveType::Unknown)), DataType::Boolean => Ok(Type::Primitive(PrimitiveType::Boolean)), DataType::Int8 | DataType::Int16 | DataType::Int32 => { Ok(Type::Primitive(PrimitiveType::Int)) @@ -697,6 +699,7 @@ impl SchemaVisitor for ToArrowSchemaConverter { fn primitive(&mut self, p: &PrimitiveType) -> Result { match p { + PrimitiveType::Unknown => Ok(ArrowSchemaOrFieldOrType::Type(DataType::Null)), PrimitiveType::Boolean => Ok(ArrowSchemaOrFieldOrType::Type(DataType::Boolean)), PrimitiveType::Int => Ok(ArrowSchemaOrFieldOrType::Type(DataType::Int32)), PrimitiveType::Long => Ok(ArrowSchemaOrFieldOrType::Type(DataType::Int64)), @@ -790,6 +793,164 @@ pub fn schema_to_arrow_schema(schema: &Schema) -> Result { } } +fn parquet_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::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_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_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_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_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) + } + Type::Primitive(_) | Type::Variant(_) => arrow_field.data_type().clone(), + }; + + 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(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_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; @@ -1209,6 +1370,7 @@ pub(crate) fn primitive_type_to_arrow_type_with_ree(primitive_type: &PrimitiveTy }; match primitive_type { + PrimitiveType::Unknown => make_ree(DataType::Null), PrimitiveType::Boolean => make_ree(DataType::Boolean), PrimitiveType::Int => make_ree(DataType::Int32), PrimitiveType::Long => make_ree(DataType::Int64), @@ -2217,6 +2379,149 @@ 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(&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(); + + assert!( + schema_to_arrow_schema_for_parquet(&schema) + .unwrap_err() + .message() + .contains("struct field 2") + ); + } + + #[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(); + + assert!( + schema_to_arrow_schema_for_parquet(&schema) + .unwrap_err() + .message() + .contains("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::optional(2, "element", PrimitiveType::Unknown.into()).into(), + )), + ) + .into(), + ]) + .build() + .unwrap(); + assert!( + schema_to_arrow_schema_for_parquet(&list_schema) + .unwrap_err() + .message() + .contains("list element") + ); + + let list_of_empty_struct_schema = Schema::builder() + .with_fields(vec![ + NestedField::optional( + 1, + "list", + Type::List(ListType::new( + NestedField::optional( + 2, + "element", + Type::Struct(StructType::new(vec![ + NestedField::optional(3, "unknown", PrimitiveType::Unknown.into()) + .into(), + ])), + ) + .into(), + )), + ) + .into(), + ]) + .build() + .unwrap(); + assert!( + schema_to_arrow_schema_for_parquet(&list_of_empty_struct_schema) + .unwrap_err() + .message() + .contains("struct field 2") + ); + + 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(); + assert!( + schema_to_arrow_schema_for_parquet(&map_schema) + .unwrap_err() + .message() + .contains("map value") + ); + } + #[test] fn test_type_conversion() { // test primitive type @@ -2227,6 +2532,13 @@ mod tests { assert_eq!(iceberg_type, arrow_type_to_type(&arrow_type).unwrap()); } + { + let arrow_type = DataType::Null; + let iceberg_type = Type::Primitive(PrimitiveType::Unknown); + assert_eq!(arrow_type, type_to_arrow_type(&iceberg_type).unwrap()); + assert_eq!(iceberg_type, arrow_type_to_type(&arrow_type).unwrap()); + } + // test struct type { // no metadata will cause error diff --git a/crates/iceberg/src/arrow/value.rs b/crates/iceberg/src/arrow/value.rs index d22e565a4a..0de3b5650a 100644 --- a/crates/iceberg/src/arrow/value.rs +++ b/crates/iceberg/src/arrow/value.rs @@ -206,6 +206,7 @@ impl SchemaWithPartnerVisitor for ArrowArrayToIcebergStructConverter { fn primitive(&mut self, p: &PrimitiveType, partner: &ArrayRef) -> Result>> { match p { + PrimitiveType::Unknown => Ok(vec![None; partner.len()]), PrimitiveType::Boolean => { let array = partner .as_any() @@ -636,6 +637,7 @@ pub(crate) fn create_primitive_array_single_element( prim_lit: &Option, ) -> Result { match (data_type, prim_lit) { + (DataType::Null, None) => Ok(Arc::new(arrow_array::NullArray::new(1))), (DataType::Boolean, Some(PrimitiveLiteral::Boolean(v))) => { Ok(Arc::new(BooleanArray::from(vec![*v]))) } @@ -943,11 +945,10 @@ pub(crate) fn create_primitive_array_repeated( Some(NullBuffer::new_null(num_rows)), )) } - (DataType::Null, _) => Arc::new(arrow_array::NullArray::new(num_rows)), + (DataType::Null, None) => Arc::new(arrow_array::NullArray::new(num_rows)), // --- Catch-all null arm: use arrow-rs new_null_array for any remaining DataType --- (dt, None) => new_null_array(dt, num_rows), - (dt, _) => { return Err(Error::new( ErrorKind::Unexpected, @@ -1866,6 +1867,26 @@ mod test { } } + #[test] + fn test_create_null_array_rejects_non_null_literal() { + let literal = Some(PrimitiveLiteral::Int(1)); + + assert!(create_primitive_array_single_element(&DataType::Null, &literal).is_err()); + assert!(create_primitive_array_repeated(&DataType::Null, &literal, 2).is_err()); + assert_eq!( + create_primitive_array_single_element(&DataType::Null, &None) + .unwrap() + .len(), + 1 + ); + assert_eq!( + create_primitive_array_repeated(&DataType::Null, &None, 2) + .unwrap() + .len(), + 2 + ); + } + #[test] fn test_create_decimal_array_repeated_respects_precision() { // Ensure repeated arrays also respect target precision, not Arrow's default. diff --git a/crates/iceberg/src/avro/schema.rs b/crates/iceberg/src/avro/schema.rs index 528b89b7c1..fc7647fbb7 100644 --- a/crates/iceberg/src/avro/schema.rs +++ b/crates/iceberg/src/avro/schema.rs @@ -74,7 +74,7 @@ impl SchemaVisitor for SchemaToAvroSchema { record.name = Name::from(format!("r{}", field.id).as_str()); } - if !field.required { + if !field.required && !matches!(field_schema, AvroSchema::Null) { field_schema = avro_optional(field_schema)?; } @@ -126,7 +126,7 @@ impl SchemaVisitor for SchemaToAvroSchema { record.name = Name::from(format!("r{}", list.element_field.id).as_str()); } - if !list.element_field.required { + if !list.element_field.required && !matches!(field_schema, AvroSchema::Null) { field_schema = avro_optional(field_schema)?; } @@ -147,7 +147,7 @@ impl SchemaVisitor for SchemaToAvroSchema { ) -> Result { let key_field_schema = key_value.unwrap_left(); let mut value_field_schema = value.unwrap_left(); - if !map.value_field.required { + if !map.value_field.required && !matches!(value_field_schema, AvroSchema::Null) { value_field_schema = avro_optional(value_field_schema)?; } @@ -222,6 +222,7 @@ impl SchemaVisitor for SchemaToAvroSchema { fn primitive(&mut self, p: &PrimitiveType) -> Result { let avro_schema = match p { + PrimitiveType::Unknown => AvroSchema::Null, PrimitiveType::Boolean => AvroSchema::Boolean, PrimitiveType::Int => AvroSchema::Int, PrimitiveType::Long => AvroSchema::Long, @@ -311,6 +312,10 @@ pub(crate) fn avro_decimal_schema(precision: usize, scale: usize) -> Result Result { + if matches!(avro_schema, AvroSchema::Null) { + return Ok(AvroSchema::Null); + } + Ok(AvroSchema::Union(UnionSchema::new(vec![ AvroSchema::Null, avro_schema, @@ -447,10 +452,11 @@ impl AvroSchemaVisitor for AvroSchemaToSchema { let field_id = Self::get_element_id_from_attributes(&avro_field.custom_attributes, FIELD_ID_PROP)?; - let optional = is_avro_optional(&avro_field.schema); + let optional = is_avro_optional(&avro_field.schema) + || matches!(&avro_field.schema, AvroSchema::Null); - let mut field = - NestedField::new(field_id, &avro_field.name, field_type.unwrap(), !optional); + let field_type = field_type.unwrap_or(Type::Primitive(PrimitiveType::Unknown)); + let mut field = NestedField::new(field_id, &avro_field.name, field_type, !optional); if let Some(doc) = &avro_field.doc { field = field.with_doc(doc); @@ -482,7 +488,9 @@ impl AvroSchemaVisitor for AvroSchemaToSchema { } if options.len() == 1 { - Ok(Some(options.remove(0).unwrap())) + Ok(options + .remove(0) + .or(Some(Type::Primitive(PrimitiveType::Unknown)))) } else { Ok(Some(options.remove(1).unwrap())) } @@ -490,10 +498,11 @@ impl AvroSchemaVisitor for AvroSchemaToSchema { fn array(&mut self, array: &ArraySchema, item: Option) -> Result { let element_field_id = Self::get_element_id_from_attributes(&array.attributes, ELEMENT_ID)?; + let item = item.unwrap_or(Type::Primitive(PrimitiveType::Unknown)); let element_field = NestedField::list_element( element_field_id, - item.unwrap(), - !is_avro_optional(&array.items), + item, + !is_avro_optional(&array.items) && !matches!(array.items.as_ref(), AvroSchema::Null), ) .into(); Ok(Some(Type::List(ListType { element_field }))) @@ -504,10 +513,11 @@ impl AvroSchemaVisitor for AvroSchemaToSchema { let key_field = NestedField::map_key_element(key_field_id, Type::Primitive(PrimitiveType::String)); let value_field_id = Self::get_element_id_from_attributes(&map.attributes, VALUE_ID)?; + let value = value.unwrap_or(Type::Primitive(PrimitiveType::Unknown)); let value_field = NestedField::map_value_element( value_field_id, - value.unwrap(), - !is_avro_optional(&map.types), + value, + !is_avro_optional(&map.types) && !matches!(map.types.as_ref(), AvroSchema::Null), ); Ok(Some(Type::Map(MapType { key_field: key_field.into(), @@ -557,12 +567,7 @@ impl AvroSchemaVisitor for AvroSchemaToSchema { "Can't convert avro map schema, missing key schema.", ) })?; - let value = value.ok_or_else(|| { - Error::new( - ErrorKind::DataInvalid, - "Can't convert avro map schema, missing value schema.", - ) - })?; + let value = value.unwrap_or(Type::Primitive(PrimitiveType::Unknown)); let key_id = Self::get_element_id_from_attributes( &array.fields[0].custom_attributes, FIELD_ID_PROP, @@ -575,7 +580,8 @@ impl AvroSchemaVisitor for AvroSchemaToSchema { let value_field = NestedField::map_value_element( value_id, value, - !is_avro_optional(&array.fields[1].schema), + !is_avro_optional(&array.fields[1].schema) + && !matches!(&array.fields[1].schema, AvroSchema::Null), ); Ok(Some(Type::Map(MapType { key_field: key_field.into(), @@ -659,6 +665,25 @@ mod tests { assert_eq!(iceberg_schema, converted_avro_converted_iceberg_schema); } + #[test] + fn test_unknown_type_schema_conversion() { + let schema = Schema::builder() + .with_fields(vec![ + NestedField::optional(1, "empty", PrimitiveType::Unknown.into()).into(), + ]) + .build() + .unwrap(); + + let avro_schema = schema_to_avro_schema("table", &schema).unwrap(); + let AvroSchema::Record(record) = &avro_schema else { + panic!("expected avro record schema"); + }; + assert!(matches!(record.fields[0].schema, AvroSchema::Null)); + assert_eq!(record.fields[0].default, Some(Value::Null)); + + assert_eq!(schema, avro_schema_to_schema(&avro_schema).unwrap()); + } + #[test] fn test_manifest_file_v1_schema() { let fields = vec![ diff --git a/crates/iceberg/src/spec/datatypes.rs b/crates/iceberg/src/spec/datatypes.rs index 79c48c1318..bf6d590a53 100644 --- a/crates/iceberg/src/spec/datatypes.rs +++ b/crates/iceberg/src/spec/datatypes.rs @@ -19,7 +19,6 @@ * Data Types */ use std::collections::HashMap; -use std::convert::identity; use std::fmt; use std::ops::Index; use std::sync::{Arc, OnceLock}; @@ -134,8 +133,8 @@ impl Type { /// Minimum [`FormatVersion`] required to support this type, **without** taking /// nested field types into account. /// - /// `TimestampNs` / `TimestamptzNs` / `Variant` require [`FormatVersion::V3`]; every - /// other type is valid from [`FormatVersion::V1`]. Mirrors Java's + /// `Unknown` / `TimestampNs` / `TimestamptzNs` / `Variant` require + /// [`FormatVersion::V3`]; every other type is valid from [`FormatVersion::V1`]. Mirrors Java's /// `Schema.MIN_FORMAT_VERSIONS` (a shallow lookup keyed by type id), so it /// intentionally does not recurse: callers needing the floor for a whole schema /// iterate its flattened fields (see [`Schema::calc_min_compatible_format`]). @@ -143,7 +142,9 @@ impl Type { /// [`Schema::calc_min_compatible_format`]: crate::spec::Schema::calc_min_compatible_format pub(crate) fn min_format_version(&self) -> FormatVersion { match self { - Type::Primitive(PrimitiveType::TimestampNs | PrimitiveType::TimestamptzNs) + Type::Primitive( + PrimitiveType::Unknown | PrimitiveType::TimestampNs | PrimitiveType::TimestamptzNs, + ) | Type::Variant(_) => FormatVersion::V3, _ => FormatVersion::V1, } @@ -274,6 +275,8 @@ pub enum PrimitiveType { Fixed(u64), /// Arbitrary-length byte array. Binary, + /// Default / null column type used when a more specific type is not known. + Unknown, } impl PrimitiveType { @@ -391,6 +394,7 @@ where S: Serializer { impl fmt::Display for PrimitiveType { fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result { match self { + PrimitiveType::Unknown => write!(f, "unknown"), PrimitiveType::Boolean => write!(f, "boolean"), PrimitiveType::Int => write!(f, "int"), PrimitiveType::Long => write!(f, "long"), @@ -554,7 +558,7 @@ impl fmt::Display for StructType { } #[derive(Debug, PartialEq, Serialize, Deserialize, Eq, Clone)] -#[serde(from = "SerdeNestedField", into = "SerdeNestedField")] +#[serde(try_from = "SerdeNestedField", into = "SerdeNestedField")] /// A struct is a tuple of typed values. Each field in the tuple is named and has an integer id that is unique in the table schema. /// Each field can be either optional or required, meaning that values can (or cannot) be null. Fields may be any type. /// Fields may have an optional comment or doc string. Fields can have default values. @@ -591,25 +595,72 @@ struct SerdeNestedField { pub write_default: Option, } -impl From for NestedField { - fn from(value: SerdeNestedField) -> Self { - NestedField { +impl TryFrom for NestedField { + type Error = crate::Error; + + fn try_from(value: SerdeNestedField) -> Result { + fn validate_unknown_values(default: &JsonValue, field_type: &Type) -> Result<()> { + if default.is_null() { + return Ok(()); + } + + match (field_type, default) { + (Type::Primitive(PrimitiveType::Unknown), _) => { + ensure_data_valid!(false, "Unknown type only supports null default values",); + } + (Type::Struct(struct_type), JsonValue::Object(object)) => { + for field in struct_type.fields() { + if let Some(value) = object.get(&field.id.to_string()) { + validate_unknown_values(value, &field.field_type)?; + } + } + } + (Type::List(list_type), JsonValue::Array(values)) => { + for value in values { + validate_unknown_values(value, &list_type.element_field.field_type)?; + } + } + (Type::Map(map_type), JsonValue::Object(object)) => { + if let Some(JsonValue::Array(keys)) = object.get("keys") { + for key in keys { + validate_unknown_values(key, &map_type.key_field.field_type)?; + } + } + if let Some(JsonValue::Array(values)) = object.get("values") { + for value in values { + validate_unknown_values(value, &map_type.value_field.field_type)?; + } + } + } + _ => {} + } + + Ok(()) + } + + fn parse_default(default: Option, field_type: &Type) -> Result> { + let Some(default) = default else { + return Ok(None); + }; + validate_unknown_values(&default, field_type)?; + match Literal::try_from_json(default, field_type) { + Ok(default) => Ok(default), + Err(_) => Ok(None), + } + } + + let initial_default = parse_default(value.initial_default, &value.field_type)?; + let write_default = parse_default(value.write_default, &value.field_type)?; + + Ok(NestedField { id: value.id, name: value.name, required: value.required, - initial_default: value.initial_default.and_then(|x| { - Literal::try_from_json(x, &value.field_type) - .ok() - .and_then(identity) - }), - write_default: value.write_default.and_then(|x| { - Literal::try_from_json(x, &value.field_type) - .ok() - .and_then(identity) - }), + initial_default, + write_default, field_type: value.field_type, doc: value.doc, - } + }) } } @@ -936,6 +987,7 @@ mod tests { { "type": "struct", "fields": [ + {"id": 17, "name": "unknown_field", "required": false, "type": "unknown"}, {"id": 1, "name": "bool_field", "required": true, "type": "boolean"}, {"id": 2, "name": "int_field", "required": true, "type": "int"}, {"id": 3, "name": "long_field", "required": true, "type": "long"}, @@ -960,6 +1012,12 @@ mod tests { record, Type::Struct(StructType { fields: vec![ + NestedField::optional( + 17, + "unknown_field", + Type::Primitive(PrimitiveType::Unknown), + ) + .into(), NestedField::required(1, "bool_field", Type::Primitive(PrimitiveType::Boolean)) .into(), NestedField::required(2, "int_field", Type::Primitive(PrimitiveType::Int)) @@ -1341,6 +1399,8 @@ mod tests { for (ty, literal) in pairs { assert!(ty.compatible(&literal)); } + + assert!(!PrimitiveType::Unknown.compatible(&PrimitiveLiteral::Int(1))); } #[test] diff --git a/crates/iceberg/src/spec/schema/mod.rs b/crates/iceberg/src/spec/schema/mod.rs index 652f98b649..ca51f0e6c3 100644 --- a/crates/iceberg/src/spec/schema/mod.rs +++ b/crates/iceberg/src/spec/schema/mod.rs @@ -39,11 +39,11 @@ pub use self::prune_columns::prune_columns; use super::NestedField; use crate::error::Result; use crate::expr::accessor::StructAccessor; -use crate::spec::FormatVersion; use crate::spec::datatypes::{ LIST_FIELD_NAME, ListType, MAP_KEY_FIELD_NAME, MAP_VALUE_FIELD_NAME, MapType, NestedFieldRef, PrimitiveType, StructType, Type, }; +use crate::spec::{FormatVersion, Literal}; use crate::{Error, ErrorKind, ensure_data_valid}; /// Type alias for schema id. @@ -133,6 +133,10 @@ impl SchemaBuilder { /// Builds the schema. pub fn build(self) -> Result { + for field in &self.fields { + Self::validate_unknown_type_field(field)?; + } + let field_id_to_accessor = self.build_accessors(); let r#struct = StructType::new(self.fields); @@ -190,6 +194,77 @@ impl SchemaBuilder { Ok(schema) } + fn validate_unknown_type_field(field: &NestedFieldRef) -> Result<()> { + ensure_data_valid!( + !field + .initial_default + .iter() + .chain(field.write_default.iter()) + .any(|default| { + Self::default_contains_non_null_unknown(default, &field.field_type) + }), + "Field {} cannot have non-null defaults because unknown type requires null defaults", + field.name + ); + + match field.field_type.as_ref() { + Type::Primitive(PrimitiveType::Unknown) => { + ensure_data_valid!( + !field.required, + "Field {} cannot be required because unknown type must be optional", + field.name + ); + } + Type::Struct(struct_type) => { + for nested_field in struct_type.fields() { + Self::validate_unknown_type_field(nested_field)?; + } + } + Type::List(list_type) => { + Self::validate_unknown_type_field(&list_type.element_field)?; + } + Type::Map(map_type) => { + Self::validate_unknown_type_field(&map_type.key_field)?; + Self::validate_unknown_type_field(&map_type.value_field)?; + } + Type::Primitive(_) | Type::Variant(_) => {} + } + + Ok(()) + } + + fn default_contains_non_null_unknown(default: &Literal, field_type: &Type) -> bool { + match (default, field_type) { + (_, Type::Primitive(PrimitiveType::Unknown)) => true, + (Literal::Struct(value), Type::Struct(struct_type)) => value + .iter() + .zip(struct_type.fields()) + .any(|(value, field)| { + value.is_some_and(|value| { + Self::default_contains_non_null_unknown(value, &field.field_type) + }) + }), + (Literal::List(values), Type::List(list_type)) => values.iter().any(|value| { + value.as_ref().is_some_and(|value| { + Self::default_contains_non_null_unknown( + value, + &list_type.element_field.field_type, + ) + }) + }), + (Literal::Map(map), Type::Map(map_type)) => map.iter().any(|(key, value)| { + Self::default_contains_non_null_unknown(key, &map_type.key_field.field_type) + || value.as_ref().is_some_and(|value| { + Self::default_contains_non_null_unknown( + value, + &map_type.value_field.field_type, + ) + }) + }), + _ => false, + } + } + fn build_accessors(&self) -> HashMap> { let mut map = HashMap::new(); @@ -647,6 +722,17 @@ mod tests { ]); assert_eq!(variant.calc_min_compatible_format(), FormatVersion::V3); + // Unknown is a v3-only primitive type. + let unknown = schema_with(vec![ + NestedField::optional(1, "u", Primitive(PrimitiveType::Unknown)).into(), + ]); + assert_eq!(unknown.calc_min_compatible_format(), FormatVersion::V3); + assert!( + unknown + .check_format_compatibility(FormatVersion::V2) + .is_err() + ); + // A v3-only type nested inside a list inside a struct → V3 (flattened fields). let nested = schema_with(vec![ NestedField::required( @@ -1475,4 +1561,245 @@ table { .is_err() ); } + + #[test] + fn test_unknown_type_deserialization_rejects_non_null_default() { + let field_json = serde_json::json!({ + "id": 1, + "name": "empty", + "required": false, + "type": "unknown", + "initial-default": 1 + }); + + let error = serde_json::from_value::(field_json.clone()).unwrap_err(); + assert!( + error + .to_string() + .contains("Unknown type only supports null default values"), + "unexpected error: {error}" + ); + + let schema_json = serde_json::json!({ + "type": "struct", + "schema-id": 1, + "fields": [field_json] + }); + assert!(serde_json::from_value::(schema_json).is_err()); + } + + #[test] + fn test_unknown_type_deserialization_accepts_null_defaults() { + let schema_json = serde_json::json!({ + "type": "struct", + "schema-id": 1, + "fields": [ + { + "id": 1, + "name": "empty", + "required": false, + "type": "unknown", + "initial-default": null, + "write-default": null + } + ] + }); + + serde_json::from_value::(schema_json).unwrap(); + } + + #[test] + fn test_unknown_type_must_be_optional_with_null_defaults() { + assert!( + Schema::builder() + .with_schema_id(1) + .with_fields(vec![ + NestedField::optional(1, "empty", Primitive(PrimitiveType::Unknown)).into() + ]) + .build() + .is_ok() + ); + + let required_error = Schema::builder() + .with_schema_id(1) + .with_fields(vec![ + NestedField::required(1, "empty", Primitive(PrimitiveType::Unknown)).into(), + ]) + .build() + .unwrap_err(); + assert!( + required_error + .message() + .contains("unknown type must be optional") + ); + + let default_error = Schema::builder() + .with_schema_id(1) + .with_fields(vec![ + NestedField::optional(1, "empty", Primitive(PrimitiveType::Unknown)) + .with_initial_default(Literal::int(1)) + .into(), + ]) + .build() + .unwrap_err(); + assert!( + default_error + .message() + .contains("unknown type requires null defaults") + ); + } + + #[test] + fn test_unknown_type_rejects_non_null_container_defaults() { + let cases = [ + ( + "struct", + Struct(StructType::new(vec![ + NestedField::optional(2, "empty", Primitive(PrimitiveType::Unknown)).into(), + ])), + Literal::Struct(crate::spec::Struct::from_iter([Some(Literal::int(1))])), + ), + ( + "list", + List(ListType::new( + NestedField::list_element(2, Primitive(PrimitiveType::Unknown), false).into(), + )), + Literal::List(vec![Some(Literal::int(1))]), + ), + ( + "map", + Map(MapType::optional( + 2, + Primitive(PrimitiveType::String), + 3, + Primitive(PrimitiveType::Unknown), + )), + Literal::Map(MapValue::from([( + Literal::string("key"), + Some(Literal::int(1)), + )])), + ), + ]; + + for (name, field_type, default) in cases { + let error = Schema::builder() + .with_schema_id(1) + .with_fields(vec![ + NestedField::optional(1, name, field_type) + .with_initial_default(default) + .into(), + ]) + .build() + .unwrap_err(); + assert!( + error + .message() + .contains("unknown type requires null defaults"), + "unexpected error for {name}: {error}" + ); + } + } + + #[test] + fn test_unknown_type_accepts_null_container_defaults() { + let cases = [ + ( + "struct", + Struct(StructType::new(vec![ + NestedField::optional(2, "empty", Primitive(PrimitiveType::Unknown)).into(), + ])), + Literal::Struct(crate::spec::Struct::from_iter([None])), + ), + ( + "list", + List(ListType::new( + NestedField::list_element(2, Primitive(PrimitiveType::Unknown), false).into(), + )), + Literal::List(vec![None]), + ), + ( + "map", + Map(MapType::optional( + 2, + Primitive(PrimitiveType::String), + 3, + Primitive(PrimitiveType::Unknown), + )), + Literal::Map(MapValue::from([(Literal::string("key"), None)])), + ), + ]; + + for (name, field_type, default) in cases { + let schema = Schema::builder() + .with_schema_id(1) + .with_fields(vec![ + NestedField::optional(1, name, field_type) + .with_write_default(default) + .into(), + ]) + .build() + .unwrap(); + serde_json::to_value(schema).unwrap(); + } + } + + #[test] + fn test_unknown_type_deserialization_rejects_non_null_container_defaults() { + let cases = [ + ( + "struct", + serde_json::json!({ + "type": "struct", + "fields": [{ + "id": 2, + "name": "empty", + "required": false, + "type": "unknown" + }] + }), + serde_json::json!({"2": 1}), + ), + ( + "list", + serde_json::json!({ + "type": "list", + "element-id": 2, + "element-required": false, + "element": "unknown" + }), + serde_json::json!([1]), + ), + ( + "map", + serde_json::json!({ + "type": "map", + "key-id": 2, + "key": "string", + "value-id": 3, + "value-required": false, + "value": "unknown" + }), + serde_json::json!({"keys": ["key"], "values": [1]}), + ), + ]; + + for (name, field_type, default) in cases { + let schema_json = serde_json::json!({ + "type": "struct", + "schema-id": 1, + "fields": [{ + "id": 1, + "name": name, + "required": false, + "type": field_type, + "initial-default": default + }] + }); + + assert!( + serde_json::from_value::(schema_json).is_err(), + "non-null unknown default in {name} should be rejected" + ); + } + } } diff --git a/crates/iceberg/src/spec/values/datum.rs b/crates/iceberg/src/spec/values/datum.rs index f170a09df5..593544b94b 100644 --- a/crates/iceberg/src/spec/values/datum.rs +++ b/crates/iceberg/src/spec/values/datum.rs @@ -368,6 +368,12 @@ impl Datum { /// See [this spec](https://iceberg.apache.org/spec/#binary-single-value-serialization) for reference. pub fn try_from_bytes(bytes: &[u8], data_type: PrimitiveType) -> Result { let literal = match data_type { + PrimitiveType::Unknown => { + return Err(Error::new( + ErrorKind::FeatureUnsupported, + "Cannot create datum for unknown type from bytes", + )); + } PrimitiveType::Boolean => { if bytes.len() == 1 && bytes[0] == 0u8 { PrimitiveLiteral::Boolean(false) diff --git a/crates/iceberg/src/spec/values/map.rs b/crates/iceberg/src/spec/values/map.rs index e0f75205f0..81adb8da23 100644 --- a/crates/iceberg/src/spec/values/map.rs +++ b/crates/iceberg/src/spec/values/map.rs @@ -87,6 +87,10 @@ impl Map { self.index.get(key).map(|index| &self.pair[*index].1) } + pub(crate) fn iter(&self) -> impl Iterator)> { + self.pair.iter().map(|(key, value)| (key, value)) + } + /// The order of map is matter, so this method used to compare two maps has same key-value pairs without considering the order. pub fn has_same_content(&self, other: &Map) -> bool { if self.len() != other.len() { diff --git a/crates/iceberg/src/writer/file_writer/parquet_writer.rs b/crates/iceberg/src/writer/file_writer/parquet_writer.rs index 95654c89a7..a6274f2761 100644 --- a/crates/iceberg/src/writer/file_writer/parquet_writer.rs +++ b/crates/iceberg/src/writer/file_writer/parquet_writer.rs @@ -20,13 +20,16 @@ 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; -use parquet::arrow::AsyncArrowWriter; use parquet::arrow::async_reader::AsyncFileReader; use parquet::arrow::async_writer::AsyncFileWriter as ArrowAsyncFileWriter; +use parquet::arrow::{AsyncArrowWriter, PARQUET_FIELD_ID_META_KEY}; use parquet::basic::{BrotliLevel, Compression, GzipLevel, ZstdLevel}; use parquet::encryption::encrypt::FileEncryptionProperties; use parquet::file::metadata::ParquetMetaData; @@ -37,6 +40,7 @@ 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, + schema_to_arrow_schema_for_parquet, }; use crate::compression::CompressionCodec; use crate::encryption::{EncryptionManager, StandardKeyMetadata}; @@ -167,6 +171,8 @@ impl FileWriterBuilder for ParquetWriterBuilder { resolve_writer_properties(self.props.clone(), key_metadata.as_ref())?; Ok(ParquetWriter { schema: self.schema.clone(), + arrow_schema: Arc::new(schema_to_arrow_schema_for_parquet(&self.schema)?), + match_mode: self.match_mode, inner_writer: None, writer_properties, current_row_num: 0, @@ -177,6 +183,172 @@ impl FileWriterBuilder for ParquetWriterBuilder { } } +fn field_id(field: &FieldRef) -> Option<&str> { + field + .metadata() + .get(PARQUET_FIELD_ID_META_KEY) + .map(String::as_str) +} + +fn find_field_index( + fields: &Fields, + target: &FieldRef, + match_mode: FieldMatchMode, +) -> Option { + match match_mode { + FieldMatchMode::Id => field_id(target).and_then(|target_id| { + fields + .iter() + .position(|field| field_id(field) == Some(target_id)) + }), + FieldMatchMode::Name => fields + .iter() + .position(|field| field.name() == target.name()), + } +} + +fn project_array_for_parquet( + array: &ArrayRef, + target: &FieldRef, + match_mode: FieldMatchMode, +) -> Result { + if array.data_type() == target.data_type() { + return Ok(array.clone()); + } + + match target.data_type() { + DataType::Struct(target_fields) => { + let source = array + .as_any() + .downcast_ref::() + .ok_or_else(|| { + Error::new( + ErrorKind::DataInvalid, + format!( + "Expected struct array for Parquet field {}, got {}", + target.name(), + array.data_type() + ), + ) + })?; + 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() + ), + ) + })?; + project_array_for_parquet(source.column(index), target_field, match_mode) + }) + .collect::>>()?; + Ok(Arc::new(StructArray::try_new_with_length( + target_fields.clone(), + columns, + source.nulls().cloned(), + source.len(), + )?)) + } + DataType::List(target_element) => { + let source = array.as_any().downcast_ref::().ok_or_else(|| { + Error::new( + ErrorKind::DataInvalid, + format!( + "Expected list array for Parquet field {}, got {}", + target.name(), + array.data_type() + ), + ) + })?; + let values = project_array_for_parquet(source.values(), target_element, match_mode)?; + Ok(Arc::new(ListArray::try_new( + target_element.clone(), + source.offsets().clone(), + values, + source.nulls().cloned(), + )?)) + } + DataType::Map(target_entries, ordered) => { + let source = array.as_any().downcast_ref::().ok_or_else(|| { + Error::new( + ErrorKind::DataInvalid, + format!( + "Expected map array for Parquet field {}, got {}", + target.name(), + array.data_type() + ), + ) + })?; + let source_entries: ArrayRef = Arc::new(source.entries().clone()); + let entries = project_array_for_parquet(&source_entries, target_entries, match_mode)?; + 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.clone(), + source.offsets().clone(), + entries, + source.nulls().cloned(), + *ordered, + )?)) + } + _ => Err(Error::new( + ErrorKind::DataInvalid, + format!( + "Cannot project Arrow type {} to {} for Parquet field {}", + array.data_type(), + target.data_type(), + target.name() + ), + )), + } +} + +fn project_batch_for_parquet( + batch: &RecordBatch, + target_schema: ArrowSchemaRef, + match_mode: FieldMatchMode, +) -> Result { + if batch.schema_ref() == &target_schema { + return Ok(batch.clone()); + } + + let source_schema = batch.schema(); + 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() + ), + ) + })?; + project_array_for_parquet(batch.column(index), target_field, match_mode) + }) + .collect::>>()?; + let options = RecordBatchOptions::default() + .with_match_field_names(false) + .with_row_count(Some(batch.num_rows())); + RecordBatch::try_new_with_options(target_schema, columns, &options).map_err(Into::into) +} + /// A mapping from Parquet column path names to internal field id struct IndexByParquetPathName { name_to_id: HashMap, @@ -305,6 +477,8 @@ impl SchemaVisitor for IndexByParquetPathName { /// `ParquetWriter`` is used to write arrow data into parquet file on storage. pub struct ParquetWriter { schema: SchemaRef, + arrow_schema: ArrowSchemaRef, + match_mode: FieldMatchMode, output_file: OutputFile, inner_writer: Option>, writer_properties: WriterProperties, @@ -610,7 +784,7 @@ 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(()); @@ -618,20 +792,19 @@ impl FileWriter for ParquetWriter { self.current_row_num += batch.num_rows(); - let batch_c = batch.clone(); + let batch = project_batch_for_parquet(batch, self.arrow_schema.clone(), self.match_mode)?; self.nan_value_count_visitor - .compute(self.schema.clone(), batch_c)?; + .compute(self.schema.clone(), batch.clone())?; // 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.arrow_schema.clone(), Some(self.writer_properties.clone()), ) .map_err(|err| { @@ -642,7 +815,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.", @@ -754,7 +927,7 @@ mod tests { use arrow_array::types::{Float32Type, 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; @@ -841,6 +1014,96 @@ 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, vec![ + Arc::new(NullArray::new(2)), + struct_array, + ]) + .unwrap(); + let target_schema = Arc::new(schema_to_arrow_schema_for_parquet(&schema).unwrap()); + + let projected = + project_batch_for_parquet(&batch, 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(Arc::new(schema), 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(), + )])), + Field::new("wrong", DataType::Int32, true).with_metadata(HashMap::from([( + PARQUET_FIELD_ID_META_KEY.to_string(), + "1".to_string(), + )])), + ])); + let batch = RecordBatch::try_new(source_schema, vec![ + Arc::new(Int32Array::from(vec![10])), + Arc::new(Int32Array::from(vec![20])), + ]) + .unwrap(); + let target_schema = Arc::new(schema_to_arrow_schema_for_parquet(&schema).unwrap()); + + let projected = + project_batch_for_parquet(&batch, target_schema, FieldMatchMode::Name).unwrap(); + + let values = projected + .column(0) + .as_any() + .downcast_ref::() + .unwrap(); + assert_eq!(values.value(0), 10); + } + fn nested_schema_for_test() -> Schema { // Int, Struct(Int,Int), String, List(Int), Struct(Struct(Int)), Map(String, List(Int)) Schema::builder()