diff --git a/pyiceberg/expressions/literals.py b/pyiceberg/expressions/literals.py index 61581d9b3c..208d78724f 100644 --- a/pyiceberg/expressions/literals.py +++ b/pyiceberg/expressions/literals.py @@ -75,6 +75,13 @@ def _parse_numeric_string(value: str) -> Decimal: return number +def _to_integral(value: Decimal, type_var: IcebergType) -> int: + """Convert a Decimal to an int, rejecting values with a fractional part instead of rounding them.""" + if value != value.to_integral_value(): + raise ValueError(f"Could not convert {value} into a {type_var}, value has a fractional part") + return int(value) + + class Literal(IcebergRootModel[L], Generic[L], ABC): # type: ignore """Literal which has a value and can be converted between types.""" @@ -527,24 +534,24 @@ def _(self, type_var: DecimalType) -> Literal[Decimal]: raise ValueError(f"Could not convert {self.value} into a {type_var}") @to.register(IntegerType) - def _(self, _: IntegerType) -> Literal[int]: + def _(self, type_var: IntegerType) -> Literal[int]: value_int = int(self.value.to_integral_value()) if value_int > IntegerType.max: return IntAboveMax() elif value_int < IntegerType.min: return IntBelowMin() else: - return LongLiteral(value_int) + return LongLiteral(_to_integral(self.value, type_var)) @to.register(LongType) - def _(self, _: LongType) -> Literal[int]: + def _(self, type_var: LongType) -> Literal[int]: value_int = int(self.value.to_integral_value()) if value_int > LongType.max: return LongAboveMax() elif value_int < LongType.min: return LongBelowMin() else: - return LongLiteral(value_int) + return LongLiteral(_to_integral(self.value, type_var)) @to.register(FloatType) def _(self, _: FloatType) -> Literal[float]: diff --git a/tests/catalog/test_catalog_behaviors.py b/tests/catalog/test_catalog_behaviors.py index b859e2d541..b7cff60e3d 100644 --- a/tests/catalog/test_catalog_behaviors.py +++ b/tests/catalog/test_catalog_behaviors.py @@ -1322,6 +1322,18 @@ def test_append_invalid_input_type_raises(catalog: Catalog) -> None: tbl.append("not an arrow object") +def test_scan_integer_column_with_decimal_literal(catalog: Catalog) -> None: + catalog.create_namespace("default") + identifier = f"default.scan_integer_decimal_literal_{catalog.name}" + tbl = catalog.create_table(identifier=identifier, schema=pa.schema([pa.field("x", pa.int32())])) + tbl.append(pa.table({"x": pa.array([1, 2, 3, 4], pa.int32())})) + + assert sorted(tbl.scan(row_filter="x > 2.0").to_arrow()["x"].to_pylist()) == [3, 4] + # Rounding 2.6 to 3 would turn x > 2.6 into x > 3 and silently drop x = 3 + with pytest.raises(ValueError, match="Could not convert 2.6 into a int"): + tbl.scan(row_filter="x > 2.6").to_arrow() + + 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 diff --git a/tests/expressions/test_literals.py b/tests/expressions/test_literals.py index 9251e79a7d..42e779bca8 100644 --- a/tests/expressions/test_literals.py +++ b/tests/expressions/test_literals.py @@ -919,6 +919,19 @@ def test_decimal_to_long_below_min() -> None: assert isinstance(DecimalLiteral(Decimal(LongType.min - 1)).to(LongType()), LongBelowMin) +@pytest.mark.parametrize("value", ["2.5", "2.6", "-2.5", "0.1"]) +@pytest.mark.parametrize("target_type", [IntegerType(), LongType()]) +def test_fractional_decimal_to_integral_type_raises(value: str, target_type: PrimitiveType) -> None: + # Rounding would change the predicate, e.g. x > 2.6 would become x > 3 and drop x = 3 + with pytest.raises(ValueError, match=f"Could not convert {value} into a {target_type}"): + _ = DecimalLiteral(Decimal(value)).to(target_type) + + +@pytest.mark.parametrize("target_type", [IntegerType(), LongType()]) +def test_integral_decimal_to_integral_type(target_type: PrimitiveType) -> None: + assert DecimalLiteral(Decimal("2.00")).to(target_type) == LongLiteral(2) + + def test_string_to_integer_type_invalid_value() -> None: with pytest.raises(ValueError) as e: _ = literal("abc").to(IntegerType())