Skip to content

Commit 9ca03d6

Browse files
author
zhilingc
committed
Ensure that batch retrieval tests clean up after themselves, reduce flakiness of file tests
1 parent 727ec0a commit 9ca03d6

1 file changed

Lines changed: 22 additions & 0 deletions

File tree

tests/e2e/bq-batch-retrieval.py

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -19,6 +19,7 @@
1919
from feast.feature_set import FeatureSet
2020
from feast.type_map import ValueType
2121
from google.cloud import storage, bigquery
22+
from google.cloud.storage import Blob
2223
from google.protobuf.duration_pb2 import Duration
2324
from pandavro import to_avro
2425

@@ -155,6 +156,7 @@ def test_batch_get_batch_features_with_file(client):
155156
client.ingest(file_fs1, features_1_df, timeout=480)
156157

157158
# Rename column (datetime -> event_timestamp)
159+
features_1_df['datetime'] + pd.Timedelta(seconds=1) # adds buffer to avoid rounding errors
158160
features_1_df = features_1_df.rename(columns={"datetime": "event_timestamp"})
159161

160162
to_avro(
@@ -169,6 +171,7 @@ def test_batch_get_batch_features_with_file(client):
169171
)
170172

171173
output = feature_retrieval_job.to_dataframe()
174+
clean_up_remote_files(feature_retrieval_job.get_avro_files())
172175
print(output.head())
173176

174177
assert output["entity_id"].to_list() == [
@@ -194,6 +197,7 @@ def test_batch_get_batch_features_with_gs_path(client, gcs_path):
194197
client.ingest(gcs_fs1, features_1_df, timeout=360)
195198

196199
# Rename column (datetime -> event_timestamp)
200+
features_1_df['datetime'] + pd.Timedelta(seconds=1) # adds buffer to avoid rounding errors
197201
features_1_df = features_1_df.rename(columns={"datetime": "event_timestamp"})
198202

199203
# Output file to local
@@ -220,6 +224,8 @@ def test_batch_get_batch_features_with_gs_path(client, gcs_path):
220224
)
221225

222226
output = feature_retrieval_job.to_dataframe()
227+
clean_up_remote_files(feature_retrieval_job.get_avro_files())
228+
blob.delete()
223229
print(output.head())
224230

225231
assert output["entity_id"].to_list() == [
@@ -256,6 +262,7 @@ def test_batch_order_by_creation_time(client):
256262
feature_refs=[f"{PROJECT_NAME}/feature_value3"],
257263
)
258264
output = feature_retrieval_job.to_dataframe()
265+
clean_up_remote_files(feature_retrieval_job.get_avro_files())
259266
print(output.head())
260267

261268
assert output["feature_value3"].to_list() == ["CORRECT"] * N_ROWS
@@ -291,6 +298,7 @@ def test_batch_additional_columns_in_entity_table(client):
291298
entity_rows=entity_df, feature_refs=[f"{PROJECT_NAME}/feature_value4"]
292299
)
293300
output = feature_retrieval_job.to_dataframe().sort_values(by=["entity_id"])
301+
clean_up_remote_files(feature_retrieval_job.get_avro_files())
294302
print(output.head(10))
295303

296304
assert np.allclose(
@@ -336,6 +344,7 @@ def test_batch_point_in_time_correctness_join(client):
336344
entity_rows=entity_df, feature_refs=[f"{PROJECT_NAME}/feature_value5"]
337345
)
338346
output = feature_retrieval_job.to_dataframe()
347+
clean_up_remote_files(feature_retrieval_job.get_avro_files())
339348
print(output.head())
340349

341350
assert output["feature_value5"].to_list() == ["CORRECT"] * N_EXAMPLES
@@ -384,6 +393,7 @@ def test_batch_multiple_featureset_joins(client):
384393
],
385394
)
386395
output = feature_retrieval_job.to_dataframe()
396+
clean_up_remote_files(feature_retrieval_job.get_avro_files())
387397
print(output.head())
388398

389399
assert output["entity_id"].to_list() == [
@@ -417,6 +427,7 @@ def test_batch_no_max_age(client):
417427
)
418428

419429
output = feature_retrieval_job.to_dataframe()
430+
clean_up_remote_files(feature_retrieval_job.get_avro_files())
420431
print(output.head())
421432

422433
assert output["entity_id"].to_list() == output["feature_value8"].to_list()
@@ -499,6 +510,7 @@ def test_update_featureset_apply_featureset_and_ingest_first_subset(
499510
)
500511

501512
output = feature_retrieval_job.to_dataframe().sort_values(by=["entity_id"])
513+
clean_up_remote_files(feature_retrieval_job.get_avro_files())
502514
print(output.head())
503515

504516
assert output["update_feature1"].to_list() == subset_df["update_feature1"].to_list()
@@ -552,6 +564,7 @@ def test_update_featureset_update_featureset_and_ingest_second_subset(
552564
)
553565

554566
output = feature_retrieval_job.to_dataframe().sort_values(by=["entity_id"])
567+
clean_up_remote_files(feature_retrieval_job.get_avro_files())
555568
print(output.head())
556569

557570
assert output["update_feature1"].to_list() == subset_df["update_feature1"].to_list()
@@ -587,6 +600,7 @@ def test_update_featureset_retrieve_valid_fields(client, update_featureset_dataf
587600
],
588601
)
589602
output = feature_retrieval_job.to_dataframe().sort_values(by=["entity_id"])
603+
clean_up_remote_files(feature_retrieval_job.get_avro_files())
590604
print(output.head(10))
591605
assert (
592606
output["update_feature1"].to_list()
@@ -623,3 +637,11 @@ def get_rows_ingested(
623637

624638
for row in rows:
625639
return row["count"]
640+
641+
642+
def clean_up_remote_files(files):
643+
for file_uri in files:
644+
if file_uri.scheme == "gs":
645+
blob = Blob.from_string(file_uri.geturl())
646+
blob.delete()
647+

0 commit comments

Comments
 (0)