diff --git a/sagemaker-mlops/src/sagemaker/mlops/feature_store/athena_query.py b/sagemaker-mlops/src/sagemaker/mlops/feature_store/athena_query.py index 6131442b94..b70c4562f9 100644 --- a/sagemaker-mlops/src/sagemaker/mlops/feature_store/athena_query.py +++ b/sagemaker-mlops/src/sagemaker/mlops/feature_store/athena_query.py @@ -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 @@ -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: @@ -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: @@ -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) diff --git a/sagemaker-mlops/tests/unit/sagemaker/mlops/feature_store/test_athena_query.py b/sagemaker-mlops/tests/unit/sagemaker/mlops/feature_store/test_athena_query.py index e99eb6c7fc..9bb3340f50 100644 --- a/sagemaker-mlops/tests/unit/sagemaker/mlops/feature_store/test_athena_query.py +++ b/sagemaker-mlops/tests/unit/sagemaker/mlops/feature_store/test_athena_query.py @@ -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