diff --git a/crates/datafusion/src/physical_plan/project.rs b/crates/datafusion/src/physical_plan/project.rs index 9e335c8..dae9b6d 100644 --- a/crates/datafusion/src/physical_plan/project.rs +++ b/crates/datafusion/src/physical_plan/project.rs @@ -48,7 +48,14 @@ use crate::to_datafusion_error; /// /// # Returns /// * `Ok(Arc)` - Extended plan with partition values column -/// * `Err` - If partition spec is not found or transformation fails +/// * `Err` - If `input`'s schema does not fit the table's schema, or the +/// partition spec is not found or transformation fails +/// +/// `input`'s schema fits the table's when it has the same columns, in the same +/// order, with the same types, ignoring field metadata. A column may be +/// non-nullable where the table's is optional, at any nesting depth, since its +/// values are valid in either. A nullable column where the table's is required +/// does not fit. pub fn project_with_partition( input: Arc, table: &Table, @@ -72,9 +79,11 @@ pub fn project_with_partition( let expected_schema_cleaned = strip_metadata_from_schema(&expected_arrow_schema) .map_err(to_datafusion_error)?; - if input_schema_cleaned != expected_schema_cleaned { + // `contains` rather than equality, so that a non-nullable input column fits + // an optional table column; see the function docs. + if !expected_schema_cleaned.contains(&input_schema_cleaned) { return plan_err!( - "Input schema does not match Iceberg table schema.\n\ + "Input schema is not compatible with Iceberg table schema.\n\ Expected schema: {expected_schema_cleaned}\n\ Input schema: {input_schema_cleaned}" ); @@ -252,14 +261,20 @@ impl std::hash::Hash for PartitionExpr { #[cfg(test)] mod tests { + use std::collections::HashMap; use std::collections::hash_map::DefaultHasher; use std::hash::{Hash, Hasher}; use datafusion::arrow::array::{ArrayRef, Int32Array, StructArray}; use datafusion::arrow::datatypes::{DataType, Field, Fields}; use datafusion::physical_plan::empty::EmptyExec; + use iceberg::TableIdent; + use iceberg::arrow::DEFAULT_MAP_FIELD_NAME; + use iceberg::io::FileIO; use iceberg::spec::{ - NestedField, PrimitiveType, Schema, StructType, Transform, Type, + FormatVersion, LIST_FIELD_NAME, ListType, MAP_KEY_FIELD_NAME, + MAP_VALUE_FIELD_NAME, MapType, NestedField, PrimitiveType, Schema, SortOrder, + StructType, TableMetadataBuilder, Transform, Type, }; use iceberg::test_utils::test_runtime; @@ -675,215 +690,194 @@ mod tests { assert_eq!(city_partition.value(1), "Los Angeles"); } - #[test] - fn test_schema_validation_matching_schemas() { - use iceberg::TableIdent; - use iceberg::io::FileIO; - use iceberg::spec::{FormatVersion, NestedField, PrimitiveType, Schema, Type}; - - let table_schema = Arc::new( - Schema::builder() - .with_fields(vec![ - NestedField::required(1, "id", Type::Primitive(PrimitiveType::Int)) - .into(), - NestedField::required( - 2, - "name", - Type::Primitive(PrimitiveType::String), - ) - .into(), - ]) - .build() - .unwrap(), - ); - + /// A table with `fields`, partitioned by identity on its `id` column. + fn table_partitioned_by_id(fields: Vec) -> Table { + let table_schema = Schema::builder() + .with_fields(fields.into_iter().map(Arc::new)) + .build() + .unwrap(); let partition_spec = PartitionSpec::builder(table_schema.clone()) .add_partition_field("id", "id_partition", Transform::Identity) .unwrap() .build() .unwrap(); - - let sort_order = iceberg::spec::SortOrder::builder() - .build(&table_schema) - .unwrap(); - - let table_metadata_builder = iceberg::spec::TableMetadataBuilder::new( - (*table_schema).clone(), + let sort_order = SortOrder::builder().build(&table_schema).unwrap(); + let metadata = TableMetadataBuilder::new( + table_schema, partition_spec, sort_order, "/test/table".to_string(), FormatVersion::V2, - std::collections::HashMap::new(), + HashMap::new(), ) - .unwrap(); - - let table_metadata = table_metadata_builder.build().unwrap(); - - // Create Arrow schema matching the table schema - let arrow_schema = Arc::new(ArrowSchema::new(vec![ - Field::new("id", DataType::Int32, false), - Field::new("name", DataType::Utf8, false), - ])); - - let input = Arc::new(EmptyExec::new(arrow_schema)); - - let table = Table::builder() - .metadata(table_metadata.metadata) + .unwrap() + .build() + .unwrap() + .metadata; + Table::builder() + .metadata(metadata) .identifier(TableIdent::from_strs(["test", "table"]).unwrap()) .file_io(FileIO::new_with_fs()) .metadata_location("/test/metadata.json") .runtime(test_runtime()) .build() - .unwrap(); - - let result = project_with_partition(input, &table); - assert!(result.is_ok(), "Schema validation should pass"); + .unwrap() } - #[test] - fn test_schema_validation_mismatched_schemas() { - use iceberg::TableIdent; - use iceberg::io::FileIO; - use iceberg::spec::{FormatVersion, NestedField, PrimitiveType, Schema, Type}; + /// Required `id: int` and `name: string` columns. + fn id_and_name_table() -> Table { + table_partitioned_by_id(vec![ + NestedField::required(1, "id", Type::Primitive(PrimitiveType::Int)), + NestedField::required(2, "name", Type::Primitive(PrimitiveType::String)), + ]) + } - let table_schema = Arc::new( - Schema::builder() - .with_fields(vec![ - NestedField::required(1, "id", Type::Primitive(PrimitiveType::Int)) - .into(), - NestedField::required( - 2, - "name", - Type::Primitive(PrimitiveType::String), - ) - .into(), - ]) - .build() - .unwrap(), - ); + fn input_of(fields: Vec) -> Arc { + Arc::new(EmptyExec::new(Arc::new(ArrowSchema::new(fields)))) + } - let partition_spec = PartitionSpec::builder(table_schema.clone()) - .add_partition_field("id", "id_partition", Transform::Identity) - .unwrap() - .build() - .unwrap(); + const INCOMPATIBLE: &str = "Input schema is not compatible with Iceberg table schema"; - let sort_order = iceberg::spec::SortOrder::builder() - .build(&table_schema) - .unwrap(); + #[test] + fn test_schema_validation_matching_schemas() { + let input = input_of(vec![ + Field::new("id", DataType::Int32, false), + Field::new("name", DataType::Utf8, false), + ]); + assert!(project_with_partition(input, &id_and_name_table()).is_ok()); + } - let table_metadata_builder = iceberg::spec::TableMetadataBuilder::new( - (*table_schema).clone(), - partition_spec, - sort_order, - "/test/table".to_string(), - FormatVersion::V2, - std::collections::HashMap::new(), - ) - .unwrap(); + #[test] + fn test_schema_validation_mismatched_schemas() { + let input = input_of(vec![ + Field::new("id", DataType::Int32, false), + Field::new("different_name", DataType::Utf8, false), + ]); + let err = project_with_partition(input, &id_and_name_table()) + .unwrap_err() + .to_string(); + assert!(err.contains(INCOMPATIBLE), "{err}"); + } - let table_metadata = table_metadata_builder.build().unwrap(); + #[test] + fn test_schema_validation_nullability() { + let id = |nullable| input_of(vec![Field::new("id", DataType::Int32, nullable)]); + let int = Type::Primitive(PrimitiveType::Int); + + // A non-nullable input fits an optional column, e.g. an INSERT from a + // NOT NULL source. + let optional = + table_partitioned_by_id(vec![NestedField::optional(1, "id", int.clone())]); + assert!(project_with_partition(id(false), &optional).is_ok()); + assert!(project_with_partition(id(true), &optional).is_ok()); + + // A nullable input could write nulls into a required column. + let required = table_partitioned_by_id(vec![NestedField::required(1, "id", int)]); + assert!(project_with_partition(id(false), &required).is_ok()); + let err = project_with_partition(id(true), &required) + .unwrap_err() + .to_string(); + assert!(err.contains(INCOMPATIBLE), "{err}"); + } - // Create Arrow schema with different field name (mismatched) - let arrow_schema = Arc::new(ArrowSchema::new(vec![ - Field::new("id", DataType::Int32, false), - Field::new("different_name", DataType::Utf8, false), // Wrong field name - ])); + /// Checks the nullability rule for a nested field: the input has a required + /// `id` and a column `c` of type `input_type(nullable)`, whose nested int is + /// `nullable`, and `nested_table(required)` has the matching table column. + fn assert_nested_nullability( + input_type: impl Fn(bool) -> DataType, + nested_table: impl Fn(bool) -> Type, + ) { + let input = |nullable| { + input_of(vec![ + Field::new("id", DataType::Int32, false), + Field::new("c", input_type(nullable), false), + ]) + }; + let table = |required| { + table_partitioned_by_id(vec![ + NestedField::required(1, "id", Type::Primitive(PrimitiveType::Int)), + NestedField::required(2, "c", nested_table(required)), + ]) + }; - let input = Arc::new(EmptyExec::new(arrow_schema)); + let optional = table(false); + assert!(project_with_partition(input(false), &optional).is_ok()); + assert!(project_with_partition(input(true), &optional).is_ok()); - let table = Table::builder() - .metadata(table_metadata.metadata) - .identifier(TableIdent::from_strs(["test", "table"]).unwrap()) - .file_io(FileIO::new_with_fs()) - .metadata_location("/test/metadata.json") - .runtime(test_runtime()) - .build() - .unwrap(); + let required = table(true); + assert!(project_with_partition(input(false), &required).is_ok()); + let err = project_with_partition(input(true), &required) + .unwrap_err() + .to_string(); + assert!(err.contains(INCOMPATIBLE), "{err}"); + } - let result = project_with_partition(input, &table); - assert!( - result.is_err(), - "Schema validation should fail for mismatched schemas" - ); - assert!( - result - .unwrap_err() - .to_string() - .contains("Input schema does not match Iceberg table schema") + #[test] + fn test_schema_validation_struct_nullability() { + let int = Type::Primitive(PrimitiveType::Int); + assert_nested_nullability( + |nullable| { + DataType::Struct(Fields::from(vec![Field::new( + "x", + DataType::Int32, + nullable, + )])) + }, + |required| { + let x = NestedField::new(3, "x", int.clone(), required); + Type::Struct(StructType::new(vec![Arc::new(x)])) + }, ); } #[test] - fn test_schema_validation_with_metadata_differences() { - use std::collections::HashMap; - - use iceberg::TableIdent; - use iceberg::io::FileIO; - use iceberg::spec::{FormatVersion, NestedField, PrimitiveType, Schema, Type}; - - let table_schema = Arc::new( - Schema::builder() - .with_fields(vec![ - NestedField::required(1, "id", Type::Primitive(PrimitiveType::Int)) - .into(), - NestedField::required( - 2, - "name", - Type::Primitive(PrimitiveType::String), - ) - .into(), - ]) - .build() - .unwrap(), + fn test_schema_validation_list_nullability() { + let int = Type::Primitive(PrimitiveType::Int); + assert_nested_nullability( + |nullable| { + DataType::List(Arc::new(Field::new( + LIST_FIELD_NAME, + DataType::Int32, + nullable, + ))) + }, + |required| { + let element = NestedField::list_element(3, int.clone(), required); + Type::List(ListType::new(Arc::new(element))) + }, ); + } - let partition_spec = PartitionSpec::builder(table_schema.clone()) - .add_partition_field("id", "id_partition", Transform::Identity) - .unwrap() - .build() - .unwrap(); - - let sort_order = iceberg::spec::SortOrder::builder() - .build(&table_schema) - .unwrap(); - - let table_metadata_builder = iceberg::spec::TableMetadataBuilder::new( - (*table_schema).clone(), - partition_spec, - sort_order, - "/test/table".to_string(), - FormatVersion::V2, - HashMap::new(), - ) - .unwrap(); - - let table_metadata = table_metadata_builder.build().unwrap(); - - // Create Arrow schema with metadata (should be ignored in comparison) - let mut metadata = HashMap::new(); - metadata.insert("extra".to_string(), "metadata".to_string()); + #[test] + fn test_schema_validation_map_nullability() { + let int = Type::Primitive(PrimitiveType::Int); + assert_nested_nullability( + |nullable| { + let entries = Fields::from(vec![ + Field::new(MAP_KEY_FIELD_NAME, DataType::Int32, false), + Field::new(MAP_VALUE_FIELD_NAME, DataType::Int32, nullable), + ]); + let entries = + Field::new(DEFAULT_MAP_FIELD_NAME, DataType::Struct(entries), false); + // Iceberg maps convert to unsorted Arrow maps. + DataType::Map(Arc::new(entries), false) + }, + |required| { + let key = NestedField::map_key_element(3, int.clone()); + let value = NestedField::map_value_element(4, int.clone(), required); + Type::Map(MapType::new(Arc::new(key), Arc::new(value))) + }, + ); + } - let arrow_schema = Arc::new(ArrowSchema::new(vec![ + #[test] + fn test_schema_validation_with_metadata_differences() { + // Field metadata is ignored in the comparison. + let metadata = HashMap::from([("extra".to_string(), "metadata".to_string())]); + let input = input_of(vec![ Field::new("id", DataType::Int32, false).with_metadata(metadata.clone()), Field::new("name", DataType::Utf8, false).with_metadata(metadata), - ])); - - let input = Arc::new(EmptyExec::new(arrow_schema)); - - let table = Table::builder() - .metadata(table_metadata.metadata) - .identifier(TableIdent::from_strs(["test", "table"]).unwrap()) - .file_io(FileIO::new_with_fs()) - .metadata_location("/test/metadata.json") - .runtime(test_runtime()) - .build() - .unwrap(); - - let result = project_with_partition(input, &table); - assert!( - result.is_ok(), - "Schema validation should pass even with metadata differences" - ); + ]); + assert!(project_with_partition(input, &id_and_name_table()).is_ok()); } } diff --git a/crates/datafusion/tests/integration_datafusion_test.rs b/crates/datafusion/tests/integration_datafusion_test.rs index dcea6a5..4c97516 100644 --- a/crates/datafusion/tests/integration_datafusion_test.rs +++ b/crates/datafusion/tests/integration_datafusion_test.rs @@ -22,8 +22,12 @@ use std::error::Error; use std::sync::Arc; use std::vec; -use datafusion::arrow::array::{Array, StringArray, UInt64Array}; +use datafusion::arrow::array::{ + Array, Int32Array, RecordBatch, StringArray, UInt64Array, +}; use datafusion::arrow::datatypes::{DataType, Field, Schema as ArrowSchema}; +use datafusion::arrow::util::pretty::pretty_format_batches; +use datafusion::datasource::MemTable; use datafusion::execution::context::SessionContext; use datafusion::parquet::arrow::PARQUET_FIELD_ID_META_KEY; use datafusion_iceberg::IcebergCatalogProvider; @@ -977,3 +981,72 @@ async fn test_insert_into_partitioned() -> Result<(), Box> { Ok(()) } + +/// An INSERT from a NOT NULL source into a partitioned table whose columns +/// are optional writes the rows. +#[tokio::test] +async fn test_insert_not_null_source_into_partitioned_table() -> Result<(), Box> +{ + let iceberg_catalog = get_iceberg_catalog().await; + let namespace = NamespaceIdent::new("test_not_null_source".to_string()); + set_test_namespace(&iceberg_catalog, &namespace).await?; + let schema = Schema::builder() + .with_schema_id(0) + .with_fields(vec![ + NestedField::optional(1, "id", Type::Primitive(PrimitiveType::Int)).into(), + NestedField::optional(2, "category", Type::Primitive(PrimitiveType::String)) + .into(), + ]) + .build()?; + let partition_spec = UnboundPartitionSpec::builder() + .with_spec_id(0) + .add_partition_field(2, "category", Transform::Identity)? + .build(); + let creation = TableCreation::builder() + .name("t".to_string()) + .location(temp_path()) + .schema(schema) + .partition_spec(partition_spec) + .properties(HashMap::new()) + .build(); + iceberg_catalog.create_table(&namespace, creation).await?; + + let ctx = SessionContext::new(); + let catalog = IcebergCatalogProvider::try_new(Arc::new(iceberg_catalog)).await?; + ctx.register_catalog("catalog", Arc::new(catalog)); + let source_schema = Arc::new(ArrowSchema::new(vec![ + Field::new("id", DataType::Int32, false), + Field::new("category", DataType::Utf8, false), + ])); + let batch = RecordBatch::try_new( + source_schema.clone(), + vec![ + Arc::new(Int32Array::from(vec![1, 2])), + Arc::new(StringArray::from(vec!["books", "games"])), + ], + )?; + let source = MemTable::try_new(source_schema, vec![vec![batch]])?; + ctx.register_table("source", Arc::new(source))?; + + ctx.sql("INSERT INTO catalog.test_not_null_source.t SELECT * FROM source") + .await? + .collect() + .await?; + + let batches = ctx + .sql("SELECT * FROM catalog.test_not_null_source.t ORDER BY id") + .await? + .collect() + .await?; + assert_eq!( + pretty_format_batches(&batches)?.to_string(), + "+----+----------+\n\ + | id | category |\n\ + +----+----------+\n\ + | 1 | books |\n\ + | 2 | games |\n\ + +----+----------+" + ); + + Ok(()) +}