Skip to content

Commit de62865

Browse files
authored
Historical feature retrieval e2e test (#1067)
* Add e2e tests for historical feature retrieval Signed-off-by: Khor Shu Heng <khor.heng@gojek.com> * Seed the random number generator Signed-off-by: Khor Shu Heng <khor.heng@gojek.com> Co-authored-by: Khor Shu Heng <khor.heng@gojek.com>
1 parent 53d16a6 commit de62865

4 files changed

Lines changed: 183 additions & 51 deletions

File tree

tests/e2e/conftest.py

Lines changed: 55 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,11 @@
1+
import os
2+
from pathlib import Path
3+
4+
import pyspark
15
import pytest
26

7+
from feast import Client
8+
39

410
def pytest_addoption(parser):
511
parser.addoption("--core_url", action="store", default="localhost:6565")
@@ -34,3 +40,52 @@ def pytest_runtest_setup(item):
3440
previousfailed = getattr(item.parent, "_previousfailed", None)
3541
if previousfailed is not None:
3642
pytest.xfail("previous test failed (%s)" % previousfailed.name)
43+
44+
45+
@pytest.fixture(scope="session")
46+
def feast_version():
47+
return "0.8-SNAPSHOT"
48+
49+
50+
@pytest.fixture(scope="session")
51+
def ingestion_job_jar(pytestconfig, feast_version):
52+
default_path = (
53+
Path(__file__).parent.parent.parent
54+
/ "spark"
55+
/ "ingestion"
56+
/ "target"
57+
/ f"feast-ingestion-spark-{feast_version}.jar"
58+
)
59+
60+
return pytestconfig.getoption("ingestion_jar") or f"file://{default_path}"
61+
62+
63+
@pytest.fixture(scope="session")
64+
def feast_client(pytestconfig, ingestion_job_jar):
65+
redis_host, redis_port = pytestconfig.getoption("redis_url").split(":")
66+
67+
if pytestconfig.getoption("env") == "local":
68+
return Client(
69+
core_url=pytestconfig.getoption("core_url"),
70+
serving_url=pytestconfig.getoption("serving_url"),
71+
spark_launcher="standalone",
72+
spark_standalone_master="local",
73+
spark_home=os.getenv("SPARK_HOME") or os.path.dirname(pyspark.__file__),
74+
spark_ingestion_jar=ingestion_job_jar,
75+
redis_host=redis_host,
76+
redis_port=redis_port,
77+
)
78+
79+
if pytestconfig.getoption("env") == "gcloud":
80+
return Client(
81+
core_url=pytestconfig.getoption("core_url"),
82+
serving_url=pytestconfig.getoption("serving_url"),
83+
spark_launcher="dataproc",
84+
dataproc_cluster_name=pytestconfig.getoption("dataproc_cluster_name"),
85+
dataproc_project=pytestconfig.getoption("dataproc_project"),
86+
dataproc_region=pytestconfig.getoption("dataproc_region"),
87+
dataproc_staging_location=os.path.join(
88+
pytestconfig.getoption("staging_path"), "dataproc"
89+
),
90+
spark_ingestion_jar=ingestion_job_jar,
91+
)

tests/e2e/requirements.txt

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@ mock==2.0.0
22
numpy==1.16.4
33
pandas~=1.0.0
44
pandavro==1.5.*
5+
pyspark==2.4.2
56
pytest==6.0.0
67
pytest-benchmark==3.2.2
78
pytest-mock==1.10.4
Lines changed: 127 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,127 @@
1+
import os
2+
import tempfile
3+
import uuid
4+
from datetime import datetime, timedelta
5+
from urllib.parse import urlparse
6+
7+
import numpy as np
8+
import pandas as pd
9+
import pytest
10+
from google.protobuf.duration_pb2 import Duration
11+
from pandas._testing import assert_frame_equal
12+
13+
from feast import Client, Entity, Feature, FeatureTable, FileSource, ValueType
14+
from feast.data_format import ParquetFormat
15+
from feast.staging.storage_client import get_staging_client
16+
17+
np.random.seed(0)
18+
19+
20+
@pytest.fixture(scope="function")
21+
def staging_path(pytestconfig, tmp_path):
22+
if pytestconfig.getoption("env") == "local":
23+
return f"file://{tmp_path}"
24+
25+
staging_path = pytestconfig.getoption("staging_path")
26+
return os.path.join(staging_path, str(uuid.uuid4()))
27+
28+
29+
def test_historical_features(feast_client: Client, staging_path: str):
30+
customer_entity = Entity(
31+
name="customer_id", description="Customer", value_type=ValueType.INT64
32+
)
33+
feast_client.apply_entity(customer_entity)
34+
35+
max_age = Duration()
36+
max_age.FromSeconds(2 * 86400)
37+
38+
transactions_feature_table = FeatureTable(
39+
name="transactions",
40+
entities=["customer_id"],
41+
features=[
42+
Feature("daily_transactions", ValueType.DOUBLE),
43+
Feature("total_transactions", ValueType.DOUBLE),
44+
],
45+
batch_source=FileSource(
46+
"event_timestamp",
47+
"created_timestamp",
48+
ParquetFormat(),
49+
os.path.join(staging_path, "transactions"),
50+
),
51+
max_age=max_age,
52+
)
53+
54+
feast_client.apply_feature_table(transactions_feature_table)
55+
56+
retrieval_date = (
57+
datetime.utcnow()
58+
.replace(hour=0, minute=0, second=0, microsecond=0)
59+
.replace(tzinfo=None)
60+
)
61+
retrieval_outside_max_age_date = retrieval_date + timedelta(1)
62+
event_date = retrieval_date - timedelta(2)
63+
creation_date = retrieval_date - timedelta(1)
64+
65+
customers = [1001, 1002, 1003, 1004, 1005]
66+
daily_transactions = [np.random.rand() * 10 for _ in customers]
67+
total_transactions = [np.random.rand() * 100 for _ in customers]
68+
69+
transactions_df = pd.DataFrame(
70+
{
71+
"event_timestamp": [event_date for _ in customers],
72+
"created_timestamp": [creation_date for _ in customers],
73+
"customer_id": customers,
74+
"daily_transactions": daily_transactions,
75+
"total_transactions": total_transactions,
76+
}
77+
)
78+
79+
feast_client.ingest(transactions_feature_table, transactions_df)
80+
81+
feature_refs = ["transactions:daily_transactions"]
82+
83+
customer_df = pd.DataFrame(
84+
{
85+
"event_timestamp": [retrieval_date for _ in customers]
86+
+ [retrieval_outside_max_age_date for _ in customers],
87+
"customer_id": customers + customers,
88+
}
89+
)
90+
91+
with tempfile.TemporaryDirectory() as tempdir:
92+
df_export_path = os.path.join(tempdir, "customers.parquets")
93+
customer_df.to_parquet(df_export_path)
94+
scheme, _, remote_path, _, _, _ = urlparse(staging_path)
95+
staging_client = get_staging_client(scheme)
96+
staging_client.upload_file(df_export_path, None, remote_path)
97+
customer_source = FileSource(
98+
"event_timestamp",
99+
"event_timestamp",
100+
ParquetFormat(),
101+
os.path.join(staging_path, os.path.basename(df_export_path)),
102+
)
103+
104+
job = feast_client.get_historical_features(feature_refs, customer_source)
105+
output_dir = job.get_output_file_uri()
106+
107+
_, _, joined_df_destination_path, _, _, _ = urlparse(output_dir)
108+
joined_df = pd.read_parquet(joined_df_destination_path)
109+
110+
expected_joined_df = pd.DataFrame(
111+
{
112+
"event_timestamp": [retrieval_date for _ in customers]
113+
+ [retrieval_outside_max_age_date for _ in customers],
114+
"customer_id": customers + customers,
115+
"transactions__daily_transactions": daily_transactions
116+
+ [None] * len(customers),
117+
}
118+
)
119+
120+
assert_frame_equal(
121+
joined_df.sort_values(by=["customer_id", "event_timestamp"]).reset_index(
122+
drop=True
123+
),
124+
expected_joined_df.sort_values(
125+
by=["customer_id", "event_timestamp"]
126+
).reset_index(drop=True),
127+
)

tests/e2e/test_online_features.py

Lines changed: 0 additions & 51 deletions
Original file line numberDiff line numberDiff line change
@@ -4,12 +4,10 @@
44
import time
55
import uuid
66
from datetime import datetime, timedelta
7-
from pathlib import Path
87

98
import avro.schema
109
import numpy as np
1110
import pandas as pd
12-
import pyspark
1311
import pytest
1412
import pytz
1513
from avro.io import BinaryEncoder, DatumWriter
@@ -41,55 +39,6 @@ def generate_data():
4139
return df
4240

4341

44-
@pytest.fixture(scope="session")
45-
def feast_version():
46-
return "0.8-SNAPSHOT"
47-
48-
49-
@pytest.fixture(scope="session")
50-
def ingestion_job_jar(pytestconfig, feast_version):
51-
default_path = (
52-
Path(__file__).parent.parent.parent
53-
/ "spark"
54-
/ "ingestion"
55-
/ "target"
56-
/ f"feast-ingestion-spark-{feast_version}.jar"
57-
)
58-
59-
return pytestconfig.getoption("ingestion_jar") or f"file://{default_path}"
60-
61-
62-
@pytest.fixture(scope="session")
63-
def feast_client(pytestconfig, ingestion_job_jar):
64-
redis_host, redis_port = pytestconfig.getoption("redis_url").split(":")
65-
66-
if pytestconfig.getoption("env") == "local":
67-
return Client(
68-
core_url=pytestconfig.getoption("core_url"),
69-
serving_url=pytestconfig.getoption("serving_url"),
70-
spark_launcher="standalone",
71-
spark_standalone_master="local",
72-
spark_home=os.getenv("SPARK_HOME") or os.path.dirname(pyspark.__file__),
73-
spark_ingestion_jar=ingestion_job_jar,
74-
redis_host=redis_host,
75-
redis_port=redis_port,
76-
)
77-
78-
if pytestconfig.getoption("env") == "gcloud":
79-
return Client(
80-
core_url=pytestconfig.getoption("core_url"),
81-
serving_url=pytestconfig.getoption("serving_url"),
82-
spark_launcher="dataproc",
83-
dataproc_cluster_name=pytestconfig.getoption("dataproc_cluster_name"),
84-
dataproc_project=pytestconfig.getoption("dataproc_project"),
85-
dataproc_region=pytestconfig.getoption("dataproc_region"),
86-
dataproc_staging_location=os.path.join(
87-
pytestconfig.getoption("staging_path"), "dataproc"
88-
),
89-
spark_ingestion_jar=ingestion_job_jar,
90-
)
91-
92-
9342
@pytest.fixture(scope="function")
9443
def staging_path(pytestconfig, tmp_path):
9544
if pytestconfig.getoption("env") == "local":

0 commit comments

Comments
 (0)