Skip to content
Open
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
24 changes: 14 additions & 10 deletions pyiceberg/expressions/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -568,11 +568,6 @@ def __getnewargs__(self) -> tuple[BoundTerm]:


class BoundIsNull(BoundUnaryPredicate):
def __new__(cls, term: BoundTerm) -> BooleanExpression: # pylint: disable=W0221
if term.ref().field.required:
return AlwaysFalse()
return super().__new__(cls)

def __invert__(self) -> BoundNotNull:
"""Transform the Expression into its negated version."""
return BoundNotNull(self.term)
Expand All @@ -583,11 +578,6 @@ def as_unbound(self) -> type[IsNull]:


class BoundNotNull(BoundUnaryPredicate):
def __new__(cls, term: BoundTerm) -> BooleanExpression: # pylint: disable=W0221
if term.ref().field.required:
return AlwaysTrue()
return super().__new__(cls)

def __invert__(self) -> BoundIsNull:
"""Transform the Expression into its negated version."""
return BoundIsNull(self.term)
Expand All @@ -607,6 +597,13 @@ def __invert__(self) -> NotNull:
"""Transform the Expression into its negated version."""
return NotNull(self.term)

def bind(self, schema: Schema, case_sensitive: bool = True) -> BooleanExpression:
"""Bind the term, folding to AlwaysFalse() if the field and all its ancestors are required."""
bound_term = self.term.bind(schema, case_sensitive)
if schema.is_field_required_in_path(bound_term.ref().field.field_id):
return AlwaysFalse()
return BoundIsNull(bound_term)

@property
def as_bound(self) -> type[BoundIsNull]: # type: ignore
return BoundIsNull
Expand All @@ -622,6 +619,13 @@ def __invert__(self) -> IsNull:
"""Transform the Expression into its negated version."""
return IsNull(self.term)

def bind(self, schema: Schema, case_sensitive: bool = True) -> BooleanExpression:
"""Bind the term, folding to AlwaysTrue() if the field and all its ancestors are required."""
bound_term = self.term.bind(schema, case_sensitive)
if schema.is_field_required_in_path(bound_term.ref().field.field_id):
return AlwaysTrue()
return BoundNotNull(bound_term)

@property
def as_bound(self) -> type[BoundNotNull]: # type: ignore
return BoundNotNull
Expand Down
20 changes: 20 additions & 0 deletions pyiceberg/schema.py
Original file line number Diff line number Diff line change
Expand Up @@ -282,6 +282,26 @@ def accessor_for_field(self, field_id: int) -> Accessor:

return self._lazy_id_to_accessor[field_id]

def is_field_required_in_path(self, field_id: int) -> bool:
"""Check whether a field and every struct ancestor on its path to the root are required.

Args:
field_id (int): The ID of the field.

Returns:
bool: True if the field and all of its ancestors are required, False otherwise.
"""
if not self.find_field(field_id).required:
return False

parent_id = self._lazy_id_to_parent.get(field_id)
while parent_id is not None:
if not self.find_field(parent_id).required:
return False
parent_id = self._lazy_id_to_parent.get(parent_id)

return True

def identifier_field_names(self) -> set[str]:
"""Return the names of the identifier fields.

Expand Down
35 changes: 35 additions & 0 deletions tests/expressions/test_expressions.py
Original file line number Diff line number Diff line change
Expand Up @@ -818,6 +818,16 @@ def test_bound_is_not_null(term: BoundReference) -> None:
assert bound_not_null == eval(repr(bound_not_null))


def test_bound_is_null_direct_construction_does_not_fold_even_for_required_field() -> None:
# The required-field fold now lives in IsNull.bind()/NotNull.bind(), not here.
required_term = BoundReference(
field=NestedField(field_id=1, name="foo", field_type=StringType(), required=True),
accessor=Accessor(position=0),
)
assert isinstance(BoundIsNull(required_term), BoundIsNull)
assert isinstance(BoundNotNull(required_term), BoundNotNull)


def test_is_null() -> None:
ref = Reference("a")
is_null = IsNull(ref)
Expand Down Expand Up @@ -1292,6 +1302,31 @@ def test_nested_bind() -> None:
assert IsNull(Reference("foo.bar")).bind(schema) == bound


def test_is_null_required_field_under_optional_ancestor_does_not_bind_to_always_false() -> None:
# "bar" is required but "foo" is optional, so "bar" can still be missing.
schema = Schema(
NestedField(1, "foo", StructType(NestedField(2, "bar", StringType(), required=True)), required=False),
schema_id=1,
)
bound = IsNull(Reference("foo.bar")).bind(schema)
assert bound != AlwaysFalse()
assert isinstance(bound, BoundIsNull)

bound_not_null = NotNull(Reference("foo.bar")).bind(schema)
assert bound_not_null != AlwaysTrue()
assert isinstance(bound_not_null, BoundNotNull)


def test_is_null_required_field_under_required_ancestor_binds_to_always_false() -> None:
# Both "foo" and "bar" are required, so this still folds.
schema = Schema(
NestedField(1, "foo", StructType(NestedField(2, "bar", StringType(), required=True)), required=True),
schema_id=1,
)
assert IsNull(Reference("foo.bar")).bind(schema) == AlwaysFalse()
assert NotNull(Reference("foo.bar")).bind(schema) == AlwaysTrue()


def test_bind_dot_name() -> None:
schema = Schema(NestedField(1, "foo.bar", StringType()), schema_id=1)
bound = BoundIsNull(BoundReference(schema.find_field(1), schema.accessor_for_field(1)))
Expand Down
25 changes: 25 additions & 0 deletions tests/test_schema.py
Original file line number Diff line number Diff line change
Expand Up @@ -445,6 +445,31 @@ def __getitem__(self, pos: int) -> Any:
assert inner_accessor.get(container) == "name"


def test_is_field_required_in_path() -> None:
schema = Schema(
NestedField(1, "req_top", StringType(), required=True),
NestedField(2, "opt_top", StringType(), required=False),
NestedField(
3,
"opt_struct",
StructType(NestedField(4, "req_child", StringType(), required=True)),
required=False,
),
NestedField(
5,
"req_struct",
StructType(NestedField(6, "req_grandchild", StringType(), required=True)),
required=True,
),
schema_id=1,
)

assert schema.is_field_required_in_path(1) is True
assert schema.is_field_required_in_path(2) is False
assert schema.is_field_required_in_path(4) is False # required leaf, optional parent
assert schema.is_field_required_in_path(6) is True # required leaf and ancestors


def test_serialize_schema(table_schema_with_full_nested_fields: Schema) -> None:
actual = table_schema_with_full_nested_fields.model_dump_json()
expected = (
Expand Down
Loading