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
23 changes: 23 additions & 0 deletions tests/test_lazy_roundtrip.py
Original file line number Diff line number Diff line change
Expand Up @@ -419,3 +419,26 @@ def test_timedelta_coordinates_returned_as_text(chunks):
)

xr.testing.assert_identical(out.compute(), ds)


@pytest.mark.parametrize("chunk", [1, 2, 3])
@pytest.mark.parametrize("via", ["polars", "spill"])
@pytest.mark.parametrize("kind", ["naive", "zoned", "duration"])
def test_chunked_windows_over_nanosecond_values(chunk, via, kind):
# Four values a nanosecond apart, in windows that leave a partial or
# single-step last window (sent as a value list, not a range).
pl = pytest.importorskip("polars")
steps = pd.to_timedelta([0, 1, 2, 3], unit="ns")
coord = steps if kind == "duration" else pd.Timestamp("2021-01-01") + steps
ds = xr.Dataset({"v": ("t", [1.0, 2.0, 3.0, 4.0])}, coords={"t": coord})
column = pa.array(coord.values)
if kind == "zoned":
column = column.cast(pa.timestamp("ns", tz="Asia/Tokyo"))
table = pa.table({"t": column, "v": ds.v.values})
result = pl.from_arrow(table).lazy() if via == "polars" else table

out = xql.to_dataset(
result, template=ds, chunks={"t": chunk}, spill=via == "spill"
)

xr.testing.assert_identical(out.compute(), ds)
77 changes: 56 additions & 21 deletions xarray_sql/lazyscan.py
Original file line number Diff line number Diff line change
Expand Up @@ -75,20 +75,6 @@ def _plain(value: Any) -> Any:
return value


def _zoned(value: Any, zone: str | None) -> Any:
"""A plain literal comparable with a column in time zone *zone*.

Window values come from numpy, which has no time zones: a
``datetime64`` is a UTC instant. Engines refuse to compare a naive
literal with a zone-aware column (or compare it as local time), so
the literal is labeled UTC and expressed in the column's zone.
"""
plain = _plain(value)
if zone and isinstance(plain, pd.Timestamp) and plain.tzinfo is None:
return plain.tz_localize("UTC").tz_convert(zone)
return plain


class LazyResultHandle(Protocol):
"""A re-executable query result (see module docstring)."""

Expand Down Expand Up @@ -285,6 +271,54 @@ def run() -> None:
self._run(run)


def _nanosecond(dtype: Any) -> str | None:
"""The numpy ``[ns]`` type matching a Polars time column, else ``None``."""
import polars as pl

if isinstance(dtype, pl.Datetime):
return "datetime64[ns]"
if isinstance(dtype, pl.Duration):
return "timedelta64[ns]"
return None


def _in_zone(expr: Any, dtype: Any) -> Any:
"""*expr*, a naive UTC time, in *dtype*'s zone if it has one."""
zone = getattr(dtype, "time_zone", None)
if zone:
return expr.dt.replace_time_zone("UTC").dt.convert_time_zone(zone)
return expr


def _polars_times(values: np.ndarray, dtype: Any) -> Any:
"""Window times or durations as a Polars Series of the column's type.

Built from the numpy values, so nanoseconds survive: Polars reads a
``pd.Timestamp``, ``pd.Timedelta``, or ``datetime`` literal as
microseconds, which never equals a value in a nanosecond column
(``is_in`` then matches nothing, or fails as a join on mismatched
key types). A naive window time is a UTC instant; it is expressed in
the column's zone, if any.
"""
import polars as pl

series = pl.Series(np.asarray(values, dtype=_nanosecond(dtype)))
return _in_zone(series, dtype).cast(dtype)


def _polars_value(value: Any, dtype: Any) -> Any:
"""One window bound as a Polars literal comparable with *dtype*."""
import polars as pl

unit = _nanosecond(dtype)
if unit and isinstance(value, (np.datetime64, np.timedelta64)):
# From the integer count, which is exact; a Python object is not.
nanos = int(value.astype(unit).astype("int64"))
literal = pl.lit(nanos, dtype=pl.Int64).cast(type(dtype)("ns"))
return _in_zone(literal, dtype).cast(dtype)
return _plain(value)


class PolarsHandle:
"""Handle over a ``polars.LazyFrame``.

Expand Down Expand Up @@ -317,10 +351,12 @@ def fetch(
schema = self._lf.collect_schema()
exprs = []
for dim, (kind, a, b) in specs.items():
zone = getattr(schema.get(dim), "time_zone", None)
dtype = schema.get(dim)
if kind == "range":
exprs.append(
pl.col(dim).is_between(_zoned(a, zone), _zoned(b, zone))
pl.col(dim).is_between(
_polars_value(a, dtype), _polars_value(b, dtype)
)
)
elif getattr(a, "dtype", None) is not None and a.dtype.kind == "f":
# Upstream Polars translates float ``is_in`` literals
Expand All @@ -331,14 +367,13 @@ def fetch(
# of values.
exprs.append(
pl.any_horizontal(
[
pl.col(dim).is_between(*(_zoned(v, zone),) * 2)
for v in a
]
[pl.col(dim).is_between(*(_plain(v),) * 2) for v in a]
)
)
elif _nanosecond(dtype):
exprs.append(pl.col(dim).is_in(_polars_times(a, dtype)))
else:
exprs.append(pl.col(dim).is_in([_zoned(v, zone) for v in a]))
exprs.append(pl.col(dim).is_in([_plain(v) for v in a]))

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This one is pre-existing but worth raising, np.timedelta64 is turned into pd.Timedelta, which Polars reads as microseconds, the same issue this PR fixes for pd.Timestamp. A pl.Duration("ns") column still falls through to this generic branch, so a timedelta64[ns] coordinate with a single-step window could hit the same silent NaN or join error.

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

馃 Confirmed, thanks! Six new duration cases failed exactly as you described. Fixed in 9da2806: pl.Duration columns now get literals of their own type from integer nanosecond counts, like pl.Datetime, for both range bounds and is_in value lists.

lf = self._lf.filter(*exprs) if exprs else self._lf
out = _collect_streaming(lf.select([pl.col(n) for n in columns]))
return cast(list[pa.RecordBatch], out.to_arrow().to_batches())
Expand Down
Loading