|
1 | 1 | import contextlib |
2 | 2 | import uuid |
3 | 3 | from datetime import datetime |
4 | | -from typing import Callable, ContextManager, Dict, Iterator, List, Optional, Union |
| 4 | +from typing import ( |
| 5 | + Callable, |
| 6 | + ContextManager, |
| 7 | + Dict, |
| 8 | + Iterator, |
| 9 | + List, |
| 10 | + Optional, |
| 11 | + Tuple, |
| 12 | + Union, |
| 13 | +) |
5 | 14 |
|
6 | 15 | import numpy as np |
7 | 16 | import pandas as pd |
8 | 17 | import pyarrow as pa |
| 18 | +from dateutil import parser |
9 | 19 | from pydantic import StrictStr |
10 | 20 | from pydantic.typing import Literal |
11 | 21 | from pytz import utc |
@@ -145,9 +155,21 @@ def query_generator() -> Iterator[str]: |
145 | 155 | entity_schema, expected_join_keys, entity_df_event_timestamp_col |
146 | 156 | ) |
147 | 157 |
|
| 158 | + entity_df_event_timestamp_range = _get_entity_df_event_timestamp_range( |
| 159 | + entity_df, |
| 160 | + entity_df_event_timestamp_col, |
| 161 | + redshift_client, |
| 162 | + config, |
| 163 | + table_name, |
| 164 | + ) |
| 165 | + |
148 | 166 | # Build a query context containing all information required to template the Redshift SQL query |
149 | 167 | query_context = offline_utils.get_feature_view_query_context( |
150 | | - feature_refs, feature_views, registry, project, |
| 168 | + feature_refs, |
| 169 | + feature_views, |
| 170 | + registry, |
| 171 | + project, |
| 172 | + entity_df_event_timestamp_range, |
151 | 173 | ) |
152 | 174 |
|
153 | 175 | # Generate the Redshift SQL query from the query context |
@@ -357,6 +379,48 @@ def _upload_entity_df_and_get_entity_schema( |
357 | 379 | raise InvalidEntityType(type(entity_df)) |
358 | 380 |
|
359 | 381 |
|
| 382 | +def _get_entity_df_event_timestamp_range( |
| 383 | + entity_df: Union[pd.DataFrame, str], |
| 384 | + entity_df_event_timestamp_col: str, |
| 385 | + redshift_client, |
| 386 | + config: RepoConfig, |
| 387 | + table_name: str, |
| 388 | +) -> Tuple[datetime, datetime]: |
| 389 | + if isinstance(entity_df, pd.DataFrame): |
| 390 | + entity_df_event_timestamp = entity_df.loc[ |
| 391 | + :, entity_df_event_timestamp_col |
| 392 | + ].infer_objects() |
| 393 | + if pd.api.types.is_string_dtype(entity_df_event_timestamp): |
| 394 | + entity_df_event_timestamp = pd.to_datetime( |
| 395 | + entity_df_event_timestamp, utc=True |
| 396 | + ) |
| 397 | + entity_df_event_timestamp_range = ( |
| 398 | + entity_df_event_timestamp.min(), |
| 399 | + entity_df_event_timestamp.max(), |
| 400 | + ) |
| 401 | + elif isinstance(entity_df, str): |
| 402 | + # If the entity_df is a string (SQL query), determine range |
| 403 | + # from table |
| 404 | + statement_id = aws_utils.execute_redshift_statement( |
| 405 | + redshift_client, |
| 406 | + config.offline_store.cluster_id, |
| 407 | + config.offline_store.database, |
| 408 | + config.offline_store.user, |
| 409 | + f"SELECT MIN({entity_df_event_timestamp_col}) AS min, MAX({entity_df_event_timestamp_col}) AS max FROM {table_name}", |
| 410 | + ) |
| 411 | + res = aws_utils.get_redshift_statement_result(redshift_client, statement_id)[ |
| 412 | + "Records" |
| 413 | + ][0] |
| 414 | + entity_df_event_timestamp_range = ( |
| 415 | + parser.parse(res[0]["stringValue"]), |
| 416 | + parser.parse(res[1]["stringValue"]), |
| 417 | + ) |
| 418 | + else: |
| 419 | + raise InvalidEntityType(type(entity_df)) |
| 420 | + |
| 421 | + return entity_df_event_timestamp_range |
| 422 | + |
| 423 | + |
360 | 424 | # This query is based on sdk/python/feast/infra/offline_stores/bigquery.py:MULTIPLE_FEATURE_VIEW_POINT_IN_TIME_JOIN |
361 | 425 | # There are couple of changes from BigQuery: |
362 | 426 | # 1. Use VARCHAR instead of STRING type |
@@ -428,9 +492,9 @@ def _upload_entity_df_and_get_entity_schema( |
428 | 492 | {{ feature }} as {% if full_feature_names %}{{ featureview.name }}__{{feature}}{% else %}{{ feature }}{% endif %}{% if loop.last %}{% else %}, {% endif %} |
429 | 493 | {% endfor %} |
430 | 494 | FROM {{ featureview.table_subquery }} |
431 | | - WHERE {{ featureview.event_timestamp_column }} <= (SELECT MAX(entity_timestamp) FROM entity_dataframe) |
| 495 | + WHERE {{ featureview.event_timestamp_column }} <= '{{ featureview.max_event_timestamp }}' |
432 | 496 | {% if featureview.ttl == 0 %}{% else %} |
433 | | - AND {{ featureview.event_timestamp_column }} >= (SELECT MIN(entity_timestamp) FROM entity_dataframe) - {{ featureview.ttl }} * interval '1' second |
| 497 | + AND {{ featureview.event_timestamp_column }} >= '{{ featureview.min_event_timestamp }}' |
434 | 498 | {% endif %} |
435 | 499 | ), |
436 | 500 |
|
|
0 commit comments