Skip to content

Commit 317fdd3

Browse files
author
Pradithya Aria
committed
Use datetime type for time range filter
1 parent 22b3d44 commit 317fdd3

2 files changed

Lines changed: 32 additions & 23 deletions

File tree

sdk/python/feast/sdk/client.py

Lines changed: 15 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -18,7 +18,6 @@
1818

1919
import os
2020
from datetime import datetime
21-
import dateutil.parser
2221

2322
import grpc
2423
import pandas as pd
@@ -220,17 +219,27 @@ def get_serving_data(self, feature_set, entity_keys, ts_range=None):
220219
representing the data wanted
221220
entity_keys (:obj: `list` of :obj: `str): list of entity keys
222221
ts_range (:obj: `list` of str, optional): size 2 list of start
223-
timestamp and end timestamp, in ISO 8601 format. It will
222+
and end time, in datetime type. It will
224223
filter out any feature value having event timestamp outside
225224
of the ts_range.
226225
227226
Returns:
228227
pandas.DataFrame: DataFrame of results
229228
"""
229+
start = None
230+
end = None
231+
if ts_range is not None:
232+
if len(ts_range) != 2:
233+
raise ValueError("ts_range must have len 2")
234+
start = ts_range[0]
235+
end = ts_range[1]
236+
if type(start) is not datetime or type(end) is not datetime:
237+
raise TypeError("start and end must be datetime type")
238+
230239
request = self._build_serving_request(feature_set, entity_keys)
231240
self._connect_serving()
232241
return self._response_to_df(feature_set, self._serving_service_stub
233-
.QueryFeatures(request), ts_range)
242+
.QueryFeatures(request), start, end)
234243

235244
def download_dataset(self, dataset_info, dest, staging_location,
236245
file_type=FileType.CSV):
@@ -299,22 +308,15 @@ def _build_serving_request(self, feature_set, entity_keys):
299308
entityId=entity_keys,
300309
featureId=feature_set.features)
301310

302-
def _response_to_df(self, feature_set, response, ts_range = None):
303-
start = None
304-
end = None
305-
if ts_range is not None:
306-
if len(ts_range) != 2:
307-
raise ValueError("ts_range must have len 2")
308-
start = dateutil.parser.parse(ts_range[0])
309-
end = dateutil.parser.parse(ts_range[1])
310-
311+
def _response_to_df(self, feature_set, response, start=None, end=None):
312+
is_filter_time = start is not None and end is not None
311313
df = pd.DataFrame(columns=[feature_set.entity] + feature_set.features)
312314
for entity_id in response.entities:
313315
feature_map = response.entities[entity_id].features
314316
row = {response.entityName: entity_id}
315317
for feature_id in feature_map:
316318
v = feature_map[feature_id].value
317-
if ts_range is not None and not _is_granularity_none(
319+
if is_filter_time and not _is_granularity_none(
318320
feature_id):
319321
ts = feature_map[feature_id].timestamp.ToDatetime()
320322
if ts < start or ts > end:

sdk/python/tests/sdk/test_client.py

Lines changed: 17 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -360,13 +360,11 @@ def test_serving_response_to_df_with_time_filter(self, client):
360360
'entity.feat1': [np.NaN, 3],
361361
'entity.feat2': [np.NaN, np.NaN]}) \
362362
.reset_index(drop=True)
363-
start = datetime.utcfromtimestamp(2).isoformat()
364-
end = datetime.utcfromtimestamp(5).isoformat()
365-
366-
ts_range = [start, end]
363+
start = datetime.utcfromtimestamp(2)
364+
end = datetime.utcfromtimestamp(5)
367365
df = client._response_to_df(FeatureSet("entity", ["entity.feat1",
368366
"entity.feat2"]),
369-
response, ts_range) \
367+
response, start, end) \
370368
.sort_values(['entity']) \
371369
.reset_index(drop=True)[expected_df.columns]
372370
print(df)
@@ -386,13 +384,11 @@ def test_serving_response_to_df_with_time_filter_granularity_none(self,
386384
'entity.none.feat1': [1, 3],
387385
'entity.none.feat2': [np.NaN, np.NaN]}) \
388386
.reset_index(drop=True)
389-
start = datetime.utcfromtimestamp(2).isoformat()
390-
end = datetime.utcfromtimestamp(5).isoformat()
391-
392-
ts_range = [start, end]
387+
start = datetime.utcfromtimestamp(2)
388+
end = datetime.utcfromtimestamp(5)
393389
df = client._response_to_df(FeatureSet("entity", ["entity.none.feat1",
394390
"entity.none.feat2"]),
395-
response, ts_range) \
391+
response, start, end) \
396392
.sort_values(['entity']) \
397393
.reset_index(drop=True)[expected_df.columns]
398394
print(df)
@@ -401,6 +397,17 @@ def test_serving_response_to_df_with_time_filter_granularity_none(self,
401397
check_column_type=False,
402398
check_index_type=False)
403399

400+
def test_serving_invalid_type(self, client):
401+
start = "2018-01-01T01:01:01"
402+
end = "2018-01-01T01:01:01"
403+
ts_range = [start, end]
404+
with pytest.raises(TypeError, match="start and end must be datetime "
405+
"type"):
406+
client.get_serving_data(FeatureSet("entity", ["entity.none.feat1",
407+
"entity.none.feat2"]),
408+
["1234", "5678"],
409+
ts_range)
410+
404411
def test_download_dataset_as_file(self, client, mocker):
405412
destination = "/tmp/dest_file"
406413

0 commit comments

Comments
 (0)