diff --git a/vortex-python/python/vortex/polars_.py b/vortex-python/python/vortex/polars_.py index e976583fdd0..9a57cfa81f8 100644 --- a/vortex-python/python/vortex/polars_.py +++ b/vortex-python/python/vortex/polars_.py @@ -4,7 +4,7 @@ import json import operator from collections.abc import Callable -from typing import Any, cast +from typing import Any, Literal, cast import polars as pl @@ -108,6 +108,8 @@ def _polars_to_vortex(expr: dict[str, Any]) -> ve.Expr: return ve.literal(_dtype.date("days"), scalar["Date"]) elif "Binary" in scalar: return ve.literal(_dtype.binary(), bytes(scalar["Binary"])) + elif "Datetime" in scalar: + return _datetime_literal(scalar["Datetime"]) elif len(scalar) == 1 and next(iter(scalar)) in _LITERAL_TYPES: dtype, value = next(iter(scalar.items())) else: @@ -129,20 +131,7 @@ def _polars_to_vortex(expr: dict[str, Any]) -> ve.Expr: # Special-case date-times if literal_type == "DateTime": - (value, unit, tz) = expr[literal_type] - if unit == "Nanoseconds": - unit = "ns" - elif unit == "Microseconds": - unit = "us" - elif unit == "Milliseconds": - unit = "ms" - elif unit == "Seconds": - unit = "s" - else: - raise NotImplementedError(f"Unsupported Polars date time unit: {unit}") - - dtype = _dtype.timestamp(unit, tz=tz, nullable=value) - return ve.literal(dtype, value) + return _datetime_literal(expr[literal_type]) # Unwrap 'Dyn' scalars, whose type hasn't been established yet. # (post https://github.com/pola-rs/polars/pull/21849) @@ -164,6 +153,19 @@ def _polars_to_vortex(expr: dict[str, Any]) -> ve.Expr: return ve.fill_null(_inputs[0], _inputs[1]) return ve.zip_(ve.is_null(_inputs[0]), _inputs[1], _inputs[0]) fn = expr["function"] + if isinstance(fn, dict) and "ReplaceTimeZone" in fn.get("TemporalExpr", {}): + time_zone, non_existent = fn["TemporalExpr"]["ReplaceTimeZone"] + if isinstance(time_zone, dict): + time_zone = time_zone["inner"] + if non_existent not in ("Raise", "Null"): + raise NotImplementedError(f"Unsupported Polars nonexistent-time policy: {non_existent}") + return ve.replace_time_zone( + _inputs[0], + time_zone, + ambiguous=_inputs[1], + non_existent="raise" if non_existent == "Raise" else "null", + ) + if "Boolean" in fn: fn = fn["Boolean"] @@ -197,3 +199,19 @@ def _polars_to_vortex(expr: dict[str, Any]) -> ve.Expr: raise NotImplementedError(f"Unsupported Polars function: {fn}") raise NotImplementedError(f"Unsupported Polars expression: {expr}") + + +def _datetime_literal(data: list[Any]) -> ve.Expr: + value, unit, tz = data + units: dict[str, Literal["s", "ms", "us", "ns"]] = { + "Nanoseconds": "ns", + "Microseconds": "us", + "Milliseconds": "ms", + "Seconds": "s", + } + if unit not in units: + raise NotImplementedError(f"Unsupported Polars date time unit: {unit}") + if isinstance(tz, dict): + tz = tz["inner"] + dtype = _dtype.timestamp(units[unit], tz=tz, nullable=value is None) + return ve.literal(dtype, value) diff --git a/vortex-python/test/test_polars_.py b/vortex-python/test/test_polars_.py index c89713a5c25..dab406729cf 100644 --- a/vortex-python/test/test_polars_.py +++ b/vortex-python/test/test_polars_.py @@ -3,8 +3,9 @@ import math import os -from datetime import date, time +from datetime import UTC, date, datetime, time from decimal import Decimal +from zoneinfo import ZoneInfo import polars as pl import pyarrow as pa @@ -242,3 +243,129 @@ def test_polars_fill_null_string_literal(tmp_path): actual = vx.open(str(path)).to_polars().filter(expr).collect() assert_frame_equal(actual, expected_frame) assert actual["id"].to_list() == [1] + + +def test_datetime_predicate_pushdown(tmp_path): + table = pa.table( + { + "id": [0, 1, 2], + "value": pa.array( + [datetime(2026, 9, day, tzinfo=UTC) for day in [16, 17, 18]], + type=pa.timestamp("us", tz="UTC"), + ), + } + ) + path = tmp_path / "datetimes.vortex" + vx.io.write(vx.array(table), str(path)) + predicate = pl.col("value") >= datetime(2026, 9, 17, tzinfo=UTC) + expected = pl.DataFrame(table).lazy().filter(predicate).collect() + result = vx.open(str(path)).to_polars().filter(predicate).collect() + assert_frame_equal(result, expected) + assert result["id"].to_list() == [1, 2] + + +@pytest.mark.parametrize("unit, scale", [("ms", 1_000), ("us", 1_000_000), ("ns", 1_000_000_000)]) +@pytest.mark.parametrize("source_zone", [None, "UTC", "Europe/London"]) +@pytest.mark.parametrize("target_zone", [None, "UTC", "America/New_York"]) +def test_replace_time_zone_columns(tmp_path, unit, scale, source_zone, target_zone): + seconds = [1_705_320_000, 1_721_041_200 if source_zone == "Europe/London" else 1_721_044_800, None] + values = pa.array( + [None if value is None else value * scale for value in seconds], + type=pa.timestamp(unit, tz=source_zone), + ) + reference, frame = _time_zone_scan(tmp_path, values) + expression = pl.col("dt").dt.replace_time_zone(target_zone) + threshold = pl.lit( + datetime(2024, 7, 15, 12, tzinfo=None if target_zone is None else ZoneInfo(target_zone)), + dtype=pl.Datetime(unit, target_zone), + ) + expected_frame = reference.filter(expression >= threshold).collect() + actual = frame.filter(expression >= threshold).collect() + assert_frame_equal(actual, expected_frame) + assert actual["id"].to_list() == [1] + + +@pytest.mark.parametrize( + "ambiguous, fold, expected", + [ + ("earliest", 0, [0]), + ("earliest", 1, []), + ("latest", 0, []), + ("latest", 1, [0]), + ("null", 0, []), + ("null", 1, []), + ], +) +def test_replace_time_zone_ambiguous(tmp_path, ambiguous, fold, expected): + values = pa.array([datetime(2024, 11, 3, 1, 30), None], type=pa.timestamp("us")) + reference, frame = _time_zone_scan(tmp_path, values) + expression = pl.col("dt").dt.replace_time_zone("America/New_York", ambiguous=ambiguous) + threshold = datetime(2024, 11, 3, 1, 30, tzinfo=ZoneInfo("America/New_York"), fold=fold) + expected_frame = reference.filter(expression == threshold).collect() + actual = frame.filter(expression == threshold).collect() + assert_frame_equal(actual, expected_frame) + assert actual["id"].to_list() == expected + + +@pytest.mark.parametrize("fold, expected", [(0, [0]), (1, [1])]) +def test_replace_time_zone_policy_column(tmp_path, fold, expected): + values = pa.array([datetime(2024, 11, 3, 1, 30)] * 4, type=pa.timestamp("us")) + reference, frame = _time_zone_scan(tmp_path, values, ["earliest", "latest", "null", None]) + expression = pl.col("dt").dt.replace_time_zone("America/New_York", ambiguous=pl.col("policy")) + threshold = datetime(2024, 11, 3, 1, 30, tzinfo=ZoneInfo("America/New_York"), fold=fold) + expected_frame = reference.filter(expression == threshold).collect() + actual = frame.filter(expression == threshold).collect() + assert_frame_equal(actual, expected_frame) + assert actual["id"].to_list() == expected + + +def test_replace_time_zone_non_existent_null(tmp_path): + values = pa.array([datetime(2024, 3, 10, 2, 30), datetime(2024, 3, 10, 3, 30)], type=pa.timestamp("us")) + reference, frame = _time_zone_scan(tmp_path, values) + expression = pl.col("dt").dt.replace_time_zone("America/New_York", non_existent="null") + threshold = datetime(2024, 3, 10, 3, 30, tzinfo=ZoneInfo("America/New_York")) + expected_frame = reference.filter(expression == threshold).collect() + actual = frame.filter(expression == threshold).collect() + assert_frame_equal(actual, expected_frame) + assert actual["id"].to_list() == [1] + + +@pytest.mark.parametrize("value", [datetime(2024, 11, 3, 1, 30), datetime(2024, 3, 10, 2, 30)]) +def test_replace_time_zone_raises(tmp_path, value): + reference, frame = _time_zone_scan(tmp_path, pa.array([value], type=pa.timestamp("us"))) + expression = pl.col("dt").dt.replace_time_zone("America/New_York") + threshold = datetime(2024, 1, 1, tzinfo=ZoneInfo("America/New_York")) + with pytest.raises(pl.exceptions.ComputeError): + reference.filter(expression >= threshold).collect() + with pytest.raises((RuntimeError, pl.exceptions.ComputeError), match="ambiguous|gap|fold"): + frame.filter(expression >= threshold).collect() + + +def test_replace_time_zone_maps_to_native_expression(): + expression = pl.col("dt").dt.replace_time_zone("America/New_York", ambiguous=pl.col("policy"), non_existent="null") + expected = ve.replace_time_zone( + ve.column("dt"), "America/New_York", ambiguous=ve.column("policy"), non_existent="null" + ) + assert polars_to_vortex(expression).serialize() == expected.serialize() + + +def test_replace_time_zone_same_zone_during_fold(tmp_path): + # 2024-11-03 06:30 UTC is the second 01:30 in New York. + values = pa.array([1_730_615_400_000_000], type=pa.timestamp("us", tz="America/New_York")) + reference, frame = _time_zone_scan(tmp_path, values) + expression = pl.col("dt").dt.replace_time_zone("America/New_York") + threshold = datetime(2024, 11, 3, 1, 30, tzinfo=ZoneInfo("America/New_York"), fold=1) + expected_frame = reference.filter(expression == threshold).collect() + actual = frame.filter(expression == threshold).collect() + assert_frame_equal(actual, expected_frame) + assert actual["id"].to_list() == [0] + + +def _time_zone_scan(tmp_path, values, policy=None) -> tuple[pl.LazyFrame, pl.LazyFrame]: + columns = {"id": pa.array(range(len(values))), "dt": values} + if policy is not None: + columns["policy"] = pa.array(policy) + path = tmp_path / "timezones.vortex" + table = pa.table(columns) + vx.io.write(vx.array(table), str(path)) + return pl.DataFrame(table).lazy(), vx.open(str(path)).to_polars()