Skip to content
Merged
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
48 changes: 33 additions & 15 deletions vortex-python/python/vortex/polars_.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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:
Expand All @@ -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)
Expand All @@ -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"]

Expand Down Expand Up @@ -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)
129 changes: 128 additions & 1 deletion vortex-python/test/test_polars_.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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()
Loading