Skip to content
Merged
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
31 changes: 23 additions & 8 deletions src/sagemaker/feature_store/feature_group.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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
Expand Down
53 changes: 53 additions & 0 deletions tests/unit/sagemaker/feature_store/test_feature_group.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Loading