Visitar URL original
fix: Normalize Dask timestamps to UTC without a row-wise apply by Daksha1611 · Pull Request #6818 · feast-dev/feast · GitHub
Skip to content
Open
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
28 changes: 19 additions & 9 deletions sdk/python/feast/infra/offline_stores/dask.py
Original file line number Diff line number Diff line change
Expand Up @@ -1140,6 +1140,22 @@ def _merge(
return df_to_join


def _to_utc(series: dd.Series) -> dd.Series:
"""Return *series* as tz-aware UTC, correctly for an empty partition.

The row-wise ``apply`` this replaces declared ``meta="datetime64[ns, UTC]"``
but produced whatever the lambda returned. With zero rows the lambda never
runs, so the computed partition stayed tz-naive while meta claimed
otherwise, and the later tz-aware comparison in ``_filter_ttl`` raised
``TypeError: Invalid comparison``. A column carrying a non-UTC tz hit the
same mismatch, since the lambda passed those values through untouched.

``to_datetime(..., utc=True)`` localizes tz-naive values, converts tz-aware
ones, and yields the right dtype for an empty frame.
"""
return dd.to_datetime(series, utc=True)


def _normalize_timestamp(
df_to_join: dd.DataFrame,
timestamp_field: str,
Expand All @@ -1159,10 +1175,7 @@ def _normalize_timestamp(
df_to_join = df_to_join.drop(columns=dups)

# Make sure all timestamp fields are tz-aware. We default tz-naive fields to UTC
df_to_join[timestamp_field] = df_to_join[timestamp_field].apply(
lambda x: x if x.tzinfo else x.replace(tzinfo=timezone.utc),
meta=(timestamp_field, "datetime64[ns, UTC]"),
)
df_to_join[timestamp_field] = _to_utc(df_to_join[timestamp_field])

# TODO: need to figure out why the value of created_timestamp_column_type.tz is pytz.UTC
if created_timestamp_column and (
Expand All @@ -1174,11 +1187,8 @@ def _normalize_timestamp(
df_to_join, dups = _df_column_uniquify(df_to_join)
df_to_join = df_to_join.drop(columns=dups)

df_to_join[created_timestamp_column] = df_to_join[
created_timestamp_column
].apply(
lambda x: x if x.tzinfo else x.replace(tzinfo=timezone.utc),
meta=(timestamp_field, "datetime64[ns, UTC]"),
df_to_join[created_timestamp_column] = _to_utc(
df_to_join[created_timestamp_column]
)

return df_to_join.persist()
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,58 @@
"""Timestamp normalization must produce tz-aware UTC even for an empty frame.

A zero-row ``entity_df`` is a normal degenerate case in batch scoring, when the
upstream query matched nothing for that run.
"""

from datetime import datetime, timedelta
from unittest.mock import MagicMock

import dask.dataframe as dd
import pandas as pd
import pytest

from feast.infra.offline_stores.dask import _filter_ttl, _normalize_timestamp

ENTITY_TS = "entity_timestamp"
EVENT_TS = "event_timestamp"


def _frame(n: int, tz: str | None = None) -> dd.DataFrame:
stamps = pd.to_datetime([datetime(2026, 2, 1)] * n)
if tz is not None:
stamps = stamps.tz_localize(tz)
return dd.from_pandas(
pd.DataFrame(
{
EVENT_TS: stamps,
ENTITY_TS: pd.to_datetime([datetime(2026, 2, 1)] * n, utc=True),
"conv_rate": [0.5] * n,
}
),
npartitions=1,
)


@pytest.mark.parametrize("rows", [0, 1])
def test_normalize_timestamp_is_utc_aware_regardless_of_row_count(rows):
normalized = _normalize_timestamp(_frame(rows), EVENT_TS).compute()
assert isinstance(normalized[EVENT_TS].dtype, pd.DatetimeTZDtype)
assert str(normalized[EVENT_TS].dtype.tz) == "UTC"


@pytest.mark.parametrize("rows", [0, 1])
def test_filter_ttl_on_empty_frame_does_not_raise(rows):
"""The tz-naive/tz-aware mismatch used to surface here as a TypeError."""
feature_view = MagicMock()
feature_view.ttl = timedelta(days=3650)

normalized = _normalize_timestamp(_frame(rows), EVENT_TS)
result = _filter_ttl(normalized, feature_view, ENTITY_TS, EVENT_TS).compute()

assert len(result) == rows


def test_non_utc_timezone_is_converted_not_passed_through():
"""A non-UTC column also diverged from the declared UTC meta."""
normalized = _normalize_timestamp(_frame(1, tz="US/Eastern"), EVENT_TS).compute()
assert str(normalized[EVENT_TS].dtype.tz) == "UTC"