diff --git a/pyiceberg/expressions/visitors.py b/pyiceberg/expressions/visitors.py index 5072d3de11..472b8a84bb 100644 --- a/pyiceberg/expressions/visitors.py +++ b/pyiceberg/expressions/visitors.py @@ -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: @@ -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: @@ -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 @@ -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 @@ -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 @@ -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 @@ -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 @@ -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 @@ -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 @@ -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 @@ -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 diff --git a/tests/expressions/test_evaluator.py b/tests/expressions/test_evaluator.py index bba4156e99..e6fa09454c 100644 --- a/tests/expressions/test_evaluator.py +++ b/tests/expressions/test_evaluator.py @@ -64,6 +64,7 @@ NestedField, PrimitiveType, StringType, + StructType, ) INT_MIN_VALUE = 30 @@ -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