Skip to content

Commit 3726653

Browse files
jackylee-chclaude
andcommitted
Move the NotEqualTo null fix into its own change
Pushing NotEqualTo down to Arrow drops rows where the column is null. That is a pre-existing bug, not one this change introduces: a single-literal NOT IN already folded to NotEqualTo before it. It changes the result of every `!=` row filter, so it belongs in its own change rather than here. Drop the null rows from the two scan tests so they no longer depend on it. `visit_not_in` needs no change, so a NOT IN that keeps two or more literals still keeps its nulls. Co-Authored-By: Claude Code <noreply@anthropic.com>
1 parent 75689b2 commit 3726653

2 files changed

Lines changed: 8 additions & 8 deletions

File tree

‎pyiceberg/io/pyarrow.py‎

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -919,8 +919,7 @@ def visit_equal(self, term: BoundTerm, literal: Literal[Any]) -> pc.Expression:
919919
return pc.field(self._get_field_name(term)) == _convert_scalar(literal.value, term.ref().field.field_type)
920920

921921
def visit_not_equal(self, term: BoundTerm, literal: Literal[Any]) -> pc.Expression:
922-
ref = pc.field(self._get_field_name(term))
923-
return ref.is_null(nan_is_null=False) | (ref != _convert_scalar(literal.value, term.ref().field.field_type))
922+
return pc.field(self._get_field_name(term)) != _convert_scalar(literal.value, term.ref().field.field_type)
924923

925924
def visit_greater_than_or_equal(self, term: BoundTerm, literal: Literal[Any]) -> pc.Expression:
926925
return pc.field(self._get_field_name(term)) >= _convert_scalar(literal.value, term.ref().field.field_type)

‎tests/io/test_pyarrow.py‎

Lines changed: 7 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -785,9 +785,10 @@ def test_expr_equal_to_pyarrow(bound_reference: BoundReference) -> None:
785785

786786

787787
def test_expr_not_equal_to_pyarrow(bound_reference: BoundReference) -> None:
788-
table = pa.table({"foo": [None, "hello", "world"]})
789-
expression = expression_to_pyarrow(BoundNotEqualTo(bound_reference, literal("hello")))
790-
assert table.filter(expression).column("foo").to_pylist() == [None, "world"]
788+
assert (
789+
repr(expression_to_pyarrow(BoundNotEqualTo(bound_reference, literal("hello"))))
790+
== '<pyarrow.compute.Expression (foo != "hello")>'
791+
)
791792

792793

793794
@pytest.mark.parametrize("boundary", [IntegerType.min, IntegerType.max])
@@ -796,7 +797,7 @@ def test_scan_in_out_of_range_literals(catalog: InMemoryCatalog, tmp_path: Path,
796797
schema = Schema(NestedField(1, "id", IntegerType(), required=False))
797798
catalog.create_namespace("default")
798799
table = catalog.create_table("default.out_of_range", schema=schema, location=str(tmp_path))
799-
values = [None, IntegerType.min, 1, 2, IntegerType.max]
800+
values = [IntegerType.min, 1, 2, IntegerType.max]
800801
table.append(pa.table({"id": pa.array(values, type=pa.int32())}))
801802
out_of_range = boundary - 1 if boundary == IntegerType.min else boundary + 1
802803
literals = [*valid_values, out_of_range, out_of_range * 2]
@@ -813,13 +814,13 @@ def test_scan_in_out_of_range_literals_after_type_promotion(catalog: InMemoryCat
813814
schema = Schema(NestedField(1, "id", IntegerType(), required=False))
814815
catalog.create_namespace("default")
815816
table = catalog.create_table("default.promoted_int", schema=schema, location=str(tmp_path))
816-
table.append(pa.table({"id": pa.array([None, 1, IntegerType.max], type=pa.int32())}))
817+
table.append(pa.table({"id": pa.array([1, IntegerType.max], type=pa.int32())}))
817818
with table.update_schema() as update:
818819
update.update_column("id", field_type=LongType())
819820
table.append(pa.table({"id": pa.array([2**40], type=pa.int64())}))
820821

821822
assert sorted(table.scan(row_filter=In("id", [1, 2**40])).to_arrow().column("id").to_pylist()) == [1, 2**40]
822-
assert table.scan(row_filter=NotIn("id", [1, 2**40])).to_arrow().column("id").to_pylist() == [None, IntegerType.max]
823+
assert table.scan(row_filter=NotIn("id", [1, 2**40])).to_arrow().column("id").to_pylist() == [IntegerType.max]
823824

824825

825826
def test_expr_greater_than_or_equal_equal_to_pyarrow(bound_reference: BoundReference) -> None:

0 commit comments

Comments
 (0)