From ae3984e1cfd495254cc5e8ae7eeb1c4ceda4234b Mon Sep 17 00:00:00 2001 From: Vishakha Nerkar Date: Tue, 25 Aug 2026 11:22:27 -0700 Subject: [PATCH] fix(feature_store): register HubContent Dataset from DatasetBuilder CSV path DatasetBuilder defined _register_as_hub_content_dataset and the register_as_dataset flag, but _to_csv_from_feature_group never invoked the helper. Wire the call into the CSV extraction path, gated on register_as_dataset, passing the Athena QueryExecutionId. Add unit tests asserting the helper is invoked when the flag is set and skipped when it is not. --- .../mlops/feature_store/dataset_builder.py | 8 +- .../feature_store/test_dataset_builder.py | 74 +++++++++++++++++++ 2 files changed, 81 insertions(+), 1 deletion(-) diff --git a/sagemaker-mlops/src/sagemaker/mlops/feature_store/dataset_builder.py b/sagemaker-mlops/src/sagemaker/mlops/feature_store/dataset_builder.py index 12c79b380c..bdac896ba7 100644 --- a/sagemaker-mlops/src/sagemaker/mlops/feature_store/dataset_builder.py +++ b/sagemaker-mlops/src/sagemaker/mlops/feature_store/dataset_builder.py @@ -510,7 +510,13 @@ def _to_csv_from_feature_group(self) -> tuple[str, str]: query_string = self._construct_query_string(base_fg) result = self._run_query(query_string, base_fg.catalog, base_fg.database) - return self._extract_result(result) + csv_path, query = self._extract_result(result) + + if self._register_as_dataset: + query_execution_id = result.get("QueryExecution", {}).get("QueryExecutionId") + self._register_as_hub_content_dataset(csv_path, query_execution_id) + + return csv_path, query def _extract_result(self, query_result: dict) -> tuple[str, str]: execution = query_result.get("QueryExecution", {}) diff --git a/sagemaker-mlops/tests/unit/sagemaker/mlops/feature_store/test_dataset_builder.py b/sagemaker-mlops/tests/unit/sagemaker/mlops/feature_store/test_dataset_builder.py index 039251546e..4297ecb783 100644 --- a/sagemaker-mlops/tests/unit/sagemaker/mlops/feature_store/test_dataset_builder.py +++ b/sagemaker-mlops/tests/unit/sagemaker/mlops/feature_store/test_dataset_builder.py @@ -520,3 +520,77 @@ def test_register_skipped_when_no_fg_arns(self, mock_session): ) # Should NOT call DataSet.create since no FG ARNs mock_create.assert_not_called() + + def test_to_csv_from_feature_group_invokes_register_when_enabled( + self, mock_session, mock_feature_group + ): + """_to_csv_from_feature_group calls _register_as_hub_content_dataset when the + register_as_dataset flag is set. This guards the wiring: the helper existed but + was previously never invoked from the CSV extraction path.""" + builder = DatasetBuilder( + _sagemaker_session=mock_session, + _base=mock_feature_group, + _output_path="s3://bucket/output", + _register_as_dataset=True, + ) + + base_fg = MagicMock() + base_fg.event_time_identifier_feature.feature_type = FeatureTypeEnum.STRING + athena_result = { + "QueryExecution": { + "QueryExecutionId": "abc-123", + "ResultConfiguration": {"OutputLocation": "s3://bucket/output/result.csv"}, + "Query": "SELECT *", + } + } + + with patch( + "sagemaker.mlops.feature_store.dataset_builder.construct_feature_group_to_be_merged", + return_value=base_fg, + ), patch.object( + DatasetBuilder, "_construct_query_string", return_value="SELECT *" + ), patch.object( + DatasetBuilder, "_run_query", return_value=athena_result + ), patch.object( + DatasetBuilder, "_register_as_hub_content_dataset" + ) as mock_register: + csv_path, _ = builder._to_csv_from_feature_group() + + assert csv_path == "s3://bucket/output/result.csv" + mock_register.assert_called_once_with("s3://bucket/output/result.csv", "abc-123") + + def test_to_csv_from_feature_group_skips_register_when_disabled( + self, mock_session, mock_feature_group + ): + """_to_csv_from_feature_group does NOT register a Dataset when the flag is unset + (default behavior — no extra CreateHubContent call, no extra permissions).""" + builder = DatasetBuilder( + _sagemaker_session=mock_session, + _base=mock_feature_group, + _output_path="s3://bucket/output", + _register_as_dataset=False, + ) + + base_fg = MagicMock() + base_fg.event_time_identifier_feature.feature_type = FeatureTypeEnum.STRING + athena_result = { + "QueryExecution": { + "QueryExecutionId": "abc-123", + "ResultConfiguration": {"OutputLocation": "s3://bucket/output/result.csv"}, + "Query": "SELECT *", + } + } + + with patch( + "sagemaker.mlops.feature_store.dataset_builder.construct_feature_group_to_be_merged", + return_value=base_fg, + ), patch.object( + DatasetBuilder, "_construct_query_string", return_value="SELECT *" + ), patch.object( + DatasetBuilder, "_run_query", return_value=athena_result + ), patch.object( + DatasetBuilder, "_register_as_hub_content_dataset" + ) as mock_register: + builder._to_csv_from_feature_group() + + mock_register.assert_not_called()