diff --git a/tests/test_lazy_roundtrip.py b/tests/test_lazy_roundtrip.py index d8148a5..e2a7db5 100644 --- a/tests/test_lazy_roundtrip.py +++ b/tests/test_lazy_roundtrip.py @@ -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) diff --git a/xarray_sql/lazyscan.py b/xarray_sql/lazyscan.py index 91de0a2..e8f36fc 100644 --- a/xarray_sql/lazyscan.py +++ b/xarray_sql/lazyscan.py @@ -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).""" @@ -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``. @@ -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 @@ -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])) 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())