Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 11 additions & 5 deletions pyiceberg/io/pyarrow.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@
import importlib
import itertools
import logging
import math
import operator
import os
import re
Expand Down Expand Up @@ -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)
],
)
Expand All @@ -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.

Expand Down
25 changes: 24 additions & 1 deletion tests/catalog/test_catalog_behaviors.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
20 changes: 20 additions & 0 deletions tests/io/test_pyarrow.py
Original file line number Diff line number Diff line change
Expand Up @@ -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),
Expand Down
Loading