From 2a8ebec851cb3bea638681cb3123315c3f3eea52 Mon Sep 17 00:00:00 2001 From: Lucas Jia Date: Mon, 28 Sep 2026 10:41:38 -0700 Subject: [PATCH] fix: Clean up temp CSV in AthenaQuery.as_dataframe (v2) AthenaQuery.as_dataframe() downloaded the Athena query result to /.csv and never removed it, so every query left a file behind and repeated queries could fill the local disk. Wrap the download and read_csv in try/finally and remove the temporary file afterwards, including when the download or parsing fails. Cleanup is best-effort: an OSError is logged as a warning and does not affect the returned DataFrame. Fixes #5100 --- src/sagemaker/feature_store/feature_group.py | 31 ++++++++--- .../feature_store/test_feature_group.py | 53 +++++++++++++++++++ 2 files changed, 76 insertions(+), 8 deletions(-) diff --git a/src/sagemaker/feature_store/feature_group.py b/src/sagemaker/feature_store/feature_group.py index 082280d4d9..60e6ffb929 100644 --- a/src/sagemaker/feature_store/feature_group.py +++ b/src/sagemaker/feature_store/feature_group.py @@ -149,6 +149,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: @@ -168,15 +171,27 @@ def as_dataframe(self, **kwargs) -> DataFrame: output_filename = os.path.join( tempfile.gettempdir(), f"{self._current_query_execution_id}.csv" ) - self.sagemaker_session.download_athena_query_result( - bucket=self._result_bucket, - prefix=self._result_file_prefix, - query_execution_id=self._current_query_execution_id, - filename=output_filename, - ) + try: + self.sagemaker_session.download_athena_query_result( + bucket=self._result_bucket, + prefix=self._result_file_prefix, + query_execution_id=self._current_query_execution_id, + filename=output_filename, + ) + + kwargs.pop("delimiter", None) + return pd.read_csv(filepath_or_buffer=output_filename, delimiter=",", **kwargs) + finally: + _remove_temp_file(output_filename) + - kwargs.pop("delimiter", None) - return pd.read_csv(filepath_or_buffer=output_filename, delimiter=",", **kwargs) +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) @attr.s diff --git a/tests/unit/sagemaker/feature_store/test_feature_group.py b/tests/unit/sagemaker/feature_store/test_feature_group.py index a12d466ee1..390d21fe7a 100644 --- a/tests/unit/sagemaker/feature_store/test_feature_group.py +++ b/tests/unit/sagemaker/feature_store/test_feature_group.py @@ -1628,3 +1628,56 @@ def test_athena_query_as_dataframe_query_running(sagemaker_session_mock, query): with pytest.raises(RuntimeError) as error: query.as_dataframe() assert "Current query query_id is still being executed" in str(error) + + +def _write_query_result_csv(**kwargs): + with open(kwargs["filename"], "w") as f: + f.write("col\n1\n2\n3\n") + + +def _prepare_succeeded_query(sagemaker_session_mock, query): + sagemaker_session_mock.get_query_execution.return_value = { + "QueryExecution": {"Status": {"State": "SUCCEEDED"}} + } + sagemaker_session_mock.download_athena_query_result.side_effect = _write_query_result_csv + query._current_query_execution_id = "query_id" + query._result_bucket = "bucket" + query._result_file_prefix = "prefix" + + +def test_athena_query_as_dataframe_removes_temp_file(sagemaker_session_mock, query, tmp_path): + _prepare_succeeded_query(sagemaker_session_mock, query) + expected_file = tmp_path / "query_id.csv" + + with patch("tempfile.gettempdir", Mock(return_value=str(tmp_path))): + df = query.as_dataframe() + + download_kwargs = sagemaker_session_mock.download_athena_query_result.call_args[1] + assert df["col"].tolist() == [1, 2, 3] + assert download_kwargs["filename"] == str(expected_file) + assert not expected_file.exists() + + +@patch("pandas.read_csv", Mock(side_effect=ValueError("bad csv"))) +def test_athena_query_as_dataframe_removes_temp_file_when_read_fails( + sagemaker_session_mock, query, tmp_path +): + _prepare_succeeded_query(sagemaker_session_mock, query) + + with patch("tempfile.gettempdir", Mock(return_value=str(tmp_path))): + with pytest.raises(ValueError, match="bad csv"): + query.as_dataframe() + + assert not (tmp_path / "query_id.csv").exists() + + +def test_athena_query_as_dataframe_cleanup_failure_does_not_raise( + sagemaker_session_mock, query, tmp_path +): + _prepare_succeeded_query(sagemaker_session_mock, query) + + with patch("tempfile.gettempdir", Mock(return_value=str(tmp_path))): + with patch("os.remove", Mock(side_effect=PermissionError("file in use"))): + df = query.as_dataframe() + + assert len(df) == 3