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
36 changes: 27 additions & 9 deletions sagemaker-mlops/src/sagemaker/mlops/feature_store/athena_query.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
"""Run Athena queries against Feature Store offline data and load the results."""

import logging
import os
import tempfile
from dataclasses import dataclass, field
Expand All @@ -18,6 +19,8 @@
from sagemaker.core.helper.session_helper import Session
from sagemaker.core.telemetry import Feature, _telemetry_emitter

logger = logging.getLogger(__name__)


@dataclass
class AthenaQuery:
Expand Down Expand Up @@ -91,6 +94,9 @@ def get_query_execution(self) -> Dict[str, Any]:
def as_dataframe(self, **kwargs) -> DataFrame:
"""Download the result of the current query and load it into a DataFrame.

The query result is downloaded to a temporary local CSV file, which is removed
after it has been loaded (or if downloading/loading fails).

Args:
**kwargs (object): key arguments used for the method pandas.read_csv to be able to
have a better tuning on data. For more info read:
Expand All @@ -106,12 +112,24 @@ def as_dataframe(self, **kwargs) -> DataFrame:
raise RuntimeError(f"Query {self._current_query_execution_id} failed.")

output_file = os.path.join(tempfile.gettempdir(), f"{self._current_query_execution_id}.csv")
download_athena_query_result(
session=self.sagemaker_session,
bucket=self._result_bucket,
prefix=self._result_file_prefix,
query_execution_id=self._current_query_execution_id,
filename=output_file,
)
kwargs.pop("delimiter", None)
return pd.read_csv(output_file, delimiter=",", **kwargs)
try:
download_athena_query_result(
session=self.sagemaker_session,
bucket=self._result_bucket,
prefix=self._result_file_prefix,
query_execution_id=self._current_query_execution_id,
filename=output_file,
)
kwargs.pop("delimiter", None)
return pd.read_csv(output_file, delimiter=",", **kwargs)
finally:
_remove_temp_file(output_file)


def _remove_temp_file(path: str) -> None:
"""Best-effort removal of a temporary file; never raises."""
try:
if os.path.exists(path):
os.remove(path)
except OSError as e:
logger.warning("Failed to remove temporary query result file %s: %s", path, e)
Original file line number Diff line number Diff line change
Expand Up @@ -114,3 +114,60 @@ def test_as_dataframe_raises_when_failed(self, mock_get, athena_query):

with pytest.raises(RuntimeError, match="failed"):
athena_query.as_dataframe()

@staticmethod
def _write_csv(**kwargs):
with open(kwargs["filename"], "w") as f:
f.write("col\n1\n2\n3\n")

@patch("sagemaker.mlops.feature_store.athena_query.get_query_execution")
@patch("sagemaker.mlops.feature_store.athena_query.download_athena_query_result")
def test_as_dataframe_removes_temp_file(self, mock_download, mock_get, athena_query, tmp_path):
athena_query._current_query_execution_id = "query-123"
athena_query._result_bucket = "bucket"
athena_query._result_file_prefix = "prefix"
mock_get.return_value = {"QueryExecution": {"Status": {"State": "SUCCEEDED"}}}
mock_download.side_effect = self._write_csv
expected_file = tmp_path / "query-123.csv"

with patch("tempfile.gettempdir", return_value=str(tmp_path)):
df = athena_query.as_dataframe()

assert df["col"].tolist() == [1, 2, 3]
assert mock_download.call_args[1]["filename"] == str(expected_file)
assert not expected_file.exists()

@patch("sagemaker.mlops.feature_store.athena_query.get_query_execution")
@patch("sagemaker.mlops.feature_store.athena_query.download_athena_query_result")
@patch("pandas.read_csv", side_effect=ValueError("bad csv"))
def test_as_dataframe_removes_temp_file_when_read_fails(
self, mock_read_csv, mock_download, mock_get, athena_query, tmp_path
):
athena_query._current_query_execution_id = "query-123"
athena_query._result_bucket = "bucket"
athena_query._result_file_prefix = "prefix"
mock_get.return_value = {"QueryExecution": {"Status": {"State": "SUCCEEDED"}}}
mock_download.side_effect = self._write_csv

with patch("tempfile.gettempdir", return_value=str(tmp_path)):
with pytest.raises(ValueError, match="bad csv"):
athena_query.as_dataframe()

assert not (tmp_path / "query-123.csv").exists()

@patch("sagemaker.mlops.feature_store.athena_query.get_query_execution")
@patch("sagemaker.mlops.feature_store.athena_query.download_athena_query_result")
def test_as_dataframe_cleanup_failure_does_not_raise(
self, mock_download, mock_get, athena_query, tmp_path
):
athena_query._current_query_execution_id = "query-123"
athena_query._result_bucket = "bucket"
athena_query._result_file_prefix = "prefix"
mock_get.return_value = {"QueryExecution": {"Status": {"State": "SUCCEEDED"}}}
mock_download.side_effect = self._write_csv

with patch("tempfile.gettempdir", return_value=str(tmp_path)):
with patch("os.remove", side_effect=PermissionError("file in use")):
df = athena_query.as_dataframe()

assert len(df) == 3
Loading