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
34 changes: 34 additions & 0 deletions pyiceberg/expressions/visitors.py
Original file line number Diff line number Diff line change
Expand Up @@ -1563,6 +1563,9 @@ def visit_is_null(self, term: BoundTerm) -> bool:
# if the column has any non-null values, the expression does not match
field_id = term.ref().field.field_id

if self._is_nested_column(field_id):
return ROWS_MIGHT_NOT_MATCH

if self._contains_nulls_only(field_id):
return ROWS_MUST_MATCH
else:
Expand All @@ -1573,6 +1576,9 @@ def visit_not_null(self, term: BoundTerm) -> bool:
# if the column has any non-null values, the expression does not match
field_id = term.ref().field.field_id

if self._is_nested_column(field_id):
return ROWS_MIGHT_NOT_MATCH

if (null_count := self.null_counts.get(field_id)) is not None and null_count == 0:
return ROWS_MUST_MATCH
else:
Expand Down Expand Up @@ -1602,6 +1608,9 @@ def visit_less_than(self, term: BoundTerm, literal: LiteralValue) -> bool:

field_id = term.ref().field.field_id

if self._is_nested_column(field_id):
return ROWS_MIGHT_NOT_MATCH

if self._can_contain_nulls(field_id) or self._can_contain_nans(field_id):
return ROWS_MIGHT_NOT_MATCH

Expand All @@ -1620,6 +1629,9 @@ def visit_less_than_or_equal(self, term: BoundTerm, literal: LiteralValue) -> bo

field_id = term.ref().field.field_id

if self._is_nested_column(field_id):
return ROWS_MIGHT_NOT_MATCH

if self._can_contain_nulls(field_id) or self._can_contain_nans(field_id):
return ROWS_MIGHT_NOT_MATCH

Expand All @@ -1638,6 +1650,9 @@ def visit_greater_than(self, term: BoundTerm, literal: LiteralValue) -> bool:

field_id = term.ref().field.field_id

if self._is_nested_column(field_id):
return ROWS_MIGHT_NOT_MATCH

if self._can_contain_nulls(field_id) or self._can_contain_nans(field_id):
return ROWS_MIGHT_NOT_MATCH

Expand All @@ -1660,6 +1675,9 @@ def visit_greater_than_or_equal(self, term: BoundTerm, literal: LiteralValue) ->
# Rows must match when: <-------X---Min----Max---------->
field_id = term.ref().field.field_id

if self._is_nested_column(field_id):
return ROWS_MIGHT_NOT_MATCH

if self._can_contain_nulls(field_id) or self._can_contain_nans(field_id):
return ROWS_MIGHT_NOT_MATCH

Expand All @@ -1682,6 +1700,9 @@ def visit_equal(self, term: BoundTerm, literal: LiteralValue) -> bool:
# Rows must match when Min == X == Max
field_id = term.ref().field.field_id

if self._is_nested_column(field_id):
return ROWS_MIGHT_NOT_MATCH

if self._can_contain_nulls(field_id) or self._can_contain_nans(field_id):
return ROWS_MIGHT_NOT_MATCH

Expand All @@ -1703,6 +1724,9 @@ def visit_not_equal(self, term: BoundTerm, literal: LiteralValue) -> bool:
# Rows must match when X < Min or Max < X because it is not in the range
field_id = term.ref().field.field_id

if self._is_nested_column(field_id):
return ROWS_MIGHT_NOT_MATCH

# If metrics prove the column contains only nulls or only NaNs, no row can have
# a value equal to the literal, so every row satisfies NotEqualTo. Partial
# null/NaN counts are not enough: a remaining non-null/non-NaN value may still
Expand Down Expand Up @@ -1736,6 +1760,9 @@ def visit_not_equal(self, term: BoundTerm, literal: LiteralValue) -> bool:
def visit_in(self, term: BoundTerm, literals: set[L]) -> bool:
field_id = term.ref().field.field_id

if self._is_nested_column(field_id):
return ROWS_MIGHT_NOT_MATCH

if self._can_contain_nulls(field_id) or self._can_contain_nans(field_id):
return ROWS_MIGHT_NOT_MATCH

Expand Down Expand Up @@ -1767,6 +1794,9 @@ def visit_in(self, term: BoundTerm, literals: set[L]) -> bool:
def visit_not_in(self, term: BoundTerm, literals: set[L]) -> bool:
field_id = term.ref().field.field_id

if self._is_nested_column(field_id):
return ROWS_MIGHT_NOT_MATCH

# If metrics prove the column contains only nulls or only NaNs, no row can have
# a value in the literal set, so every row satisfies NotIn. Partial null/NaN
# counts are not enough: a remaining non-null/non-NaN value may still be in the
Expand Down Expand Up @@ -1813,6 +1843,10 @@ def _get_field(self, field_id: int) -> NestedField:

return field

def _is_nested_column(self, field_id: int) -> bool:
# nested column metrics (e.g. null counts) are not reliable, so they cannot prove that all rows match
return self.struct.field(field_id) is None

def _can_contain_nulls(self, field_id: int) -> bool:
return (null_count := self.null_counts.get(field_id)) is not None and null_count > 0

Expand Down
45 changes: 45 additions & 0 deletions tests/expressions/test_evaluator.py
Original file line number Diff line number Diff line change
Expand Up @@ -64,6 +64,7 @@
NestedField,
PrimitiveType,
StringType,
StructType,
)

INT_MIN_VALUE = 30
Expand Down Expand Up @@ -1429,6 +1430,50 @@ def test_strict_missing_stats(strict_data_file_schema: Schema, strict_data_file_
assert not should_read, f"Should never match when stats are missing for expr: {expression}"


def test_strict_nested_column() -> None:
schema = Schema(
NestedField(
1,
"struct",
StructType(
NestedField(2, "nested_col_no_stats", IntegerType(), required=False),
NestedField(3, "nested_col_with_stats", IntegerType(), required=False),
),
required=False,
),
)

data_file = DataFile.from_args(
file_path="file_1.parquet",
file_format=FileFormat.PARQUET,
partition=Record(),
record_count=50,
value_counts={2: 50, 3: 50},
null_value_counts={2: 0, 3: 0},
nan_value_counts=None,
lower_bounds={3: INT_MIN},
upper_bounds={3: INT_MAX},
)

for name in ["struct.nested_col_no_stats", "struct.nested_col_with_stats"]:
expressions: list[BooleanExpression] = [
LessThan(name, INT_MAX_VALUE + 1),
LessThanOrEqual(name, INT_MAX_VALUE),
GreaterThan(name, INT_MIN_VALUE - 1),
GreaterThanOrEqual(name, INT_MIN_VALUE),
EqualTo(name, INT_MIN_VALUE),
NotEqualTo(name, INT_MAX_VALUE + 1),
In(name, {INT_MIN_VALUE, INT_MAX_VALUE}),
NotIn(name, {INT_MAX_VALUE + 1, INT_MAX_VALUE + 2}),
IsNull(name),
NotNull(name),
]

for expression in expressions:
should_read = _StrictMetricsEvaluator(schema, expression).eval(data_file)
assert should_read == ROWS_MIGHT_NOT_MATCH, f"Should not match: nested column metrics are not used for {expression}"


def test_strict_zero_record_file_stats(strict_data_file_schema: Schema) -> None:
zero_record_data_file = DataFile.from_args(
file_path="file_1.parquet", file_format=FileFormat.PARQUET, partition=Record(), record_count=0
Expand Down
Loading