diff --git a/pyiceberg/io/pyarrow.py b/pyiceberg/io/pyarrow.py index 1e59107da3..1179f763a7 100644 --- a/pyiceberg/io/pyarrow.py +++ b/pyiceberg/io/pyarrow.py @@ -31,6 +31,7 @@ import importlib import itertools import logging +import math import operator import os import re @@ -3097,11 +3098,7 @@ def _determine_partitions(spec: PartitionSpec, schema: Schema, arrow_table: pa.T functools.reduce( operator.and_, [ - ( - pc.field(partition_field_name) == unique_partition[partition_field_name] - if unique_partition[partition_field_name] is not None - else pc.field(partition_field_name).is_null() - ) + _partition_value_filter(partition_field_name, unique_partition[partition_field_name]) for field, partition_field_name in zip(spec.fields, partition_fields, strict=True) ], ) @@ -3116,6 +3113,15 @@ def _determine_partitions(spec: PartitionSpec, schema: Schema, arrow_table: pa.T ) +def _partition_value_filter(partition_field_name: str, value: Any) -> pc.Expression: + """Build a filter that selects the rows of one partition value, where null and NaN never compare equal.""" + if value is None: + return pc.field(partition_field_name).is_null() + elif isinstance(value, float) and math.isnan(value): + return pc.is_nan(pc.field(partition_field_name)) + return pc.field(partition_field_name) == value + + def _get_field_from_arrow_table(arrow_table: pa.Table, field_path: str) -> pa.Array: """Get a field from an Arrow table, supporting both literal field names and nested field paths. diff --git a/tests/catalog/test_catalog_behaviors.py b/tests/catalog/test_catalog_behaviors.py index b859e2d541..4310ed1eb6 100644 --- a/tests/catalog/test_catalog_behaviors.py +++ b/tests/catalog/test_catalog_behaviors.py @@ -48,7 +48,7 @@ from pyiceberg.table.update import AddSchemaUpdate, SetCurrentSchemaUpdate from pyiceberg.transforms import IdentityTransform from pyiceberg.typedef import Identifier -from pyiceberg.types import BooleanType, IntegerType, LongType, NestedField, StringType +from pyiceberg.types import BooleanType, DoubleType, IntegerType, LongType, NestedField, StringType # Name parsing tests @@ -1322,6 +1322,29 @@ def test_append_invalid_input_type_raises(catalog: Catalog) -> None: tbl.append("not an arrow object") +def test_append_nan_to_identity_partitioned_table(catalog: Catalog) -> None: + catalog.create_namespace("default") + identifier = f"default.append_nan_identity_partition_{catalog.name}" + iceberg_schema = Schema( + NestedField(1, "id", IntegerType(), required=False), + NestedField(2, "value", DoubleType(), required=False), + ) + partition_spec = PartitionSpec( + PartitionField(source_id=2, field_id=1000, transform=IdentityTransform(), name="value"), + ) + tbl = catalog.create_table(identifier=identifier, schema=iceberg_schema, partition_spec=partition_spec) + + tbl.append( + pa.Table.from_pydict( + {"id": [1, 2, 3, 4], "value": [1.0, float("nan"), None, float("nan")]}, + schema=schema_to_pyarrow(iceberg_schema), + ) + ) + + assert sorted(tbl.scan().to_arrow()["id"].to_pylist()) == [1, 2, 3, 4] + assert sorted(tbl.scan(row_filter="value is nan").to_arrow()["id"].to_pylist()) == [2, 4] + + def test_record_batch_reader_consumed_exactly_once(catalog: Catalog) -> None: """The streaming path must consume the underlying generator exactly once. A regression that drained the reader twice (e.g. an extra .schema access diff --git a/tests/io/test_pyarrow.py b/tests/io/test_pyarrow.py index 892d8e54eb..386ea076db 100644 --- a/tests/io/test_pyarrow.py +++ b/tests/io/test_pyarrow.py @@ -2830,6 +2830,26 @@ def test_partition_for_deep_nested_field() -> None: assert partition_values == {"data-1", "data-2"} +def test_determine_partitions_identity_nan() -> None: + schema = Schema( + NestedField(id=1, name="id", field_type=IntegerType(), required=False), + NestedField(id=2, name="value", field_type=DoubleType(), required=False), + ) + spec = PartitionSpec(PartitionField(source_id=2, field_id=1000, transform=IdentityTransform(), name="value")) + arrow_table = pa.Table.from_pydict( + {"id": [1, 2, 3, 4], "value": [1.0, float("nan"), None, float("nan")]}, + schema=schema.as_arrow(), + ) + + partitions = list(_determine_partitions(spec, schema, arrow_table)) + + # NaN never compares equal to itself, so its rows must be selected with is_nan + rows_by_partition = { + repr(p.partition_key.partition[0]): sorted(p.arrow_table_partition["id"].to_pylist()) for p in partitions + } + assert rows_by_partition == {"1.0": [1], "nan": [2, 4], "None": [3]} + + def test_inspect_partition_for_nested_field(catalog: InMemoryCatalog) -> None: schema = Schema( NestedField(id=1, name="foo", field_type=StringType(), required=True),