From 8b356c51843e2826c69c064c5f7735d426ab9f81 Mon Sep 17 00:00:00 2001 From: Mohamed Zeidan Date: Sun, 27 Sep 2026 16:52:00 -0700 Subject: [PATCH] fix: apply output KMS key to ModelTrainer source-code S3 uploads (#5956) V2 Estimator encrypted staged user code in S3 with output_kms_key; v3 ModelTrainer uploaded source code and driver files to S3 with no KMS encryption, so it could not be used in environments whose S3 policies enforce SSE-KMS. create_input_data_channel now passes S3 ExtraArgs (ServerSideEncryption=aws:kms, SSEKMSKeyId=...) on local uploads, using an explicit new kms_key argument when given and otherwise falling back to output_data_config.kms_key_id (mirroring V2, where the output KMS key also encrypted staged code). Behavior is unchanged when no KMS key is configured (extra_args stays None). --- .../src/sagemaker/train/model_trainer.py | 26 ++++++++ .../tests/unit/train/test_model_trainer.py | 62 +++++++++++++++++-- 2 files changed, 83 insertions(+), 5 deletions(-) diff --git a/sagemaker-train/src/sagemaker/train/model_trainer.py b/sagemaker-train/src/sagemaker/train/model_trainer.py index ce954036f6..b386cd5ffd 100644 --- a/sagemaker-train/src/sagemaker/train/model_trainer.py +++ b/sagemaker-train/src/sagemaker/train/model_trainer.py @@ -1034,6 +1034,7 @@ def create_input_data_channel( key_prefix: Optional[str] = None, ignore_patterns: Optional[List[str]] = None, instance_group_names: Optional[List[str]] = None, + kms_key: Optional[str] = None, ) -> Channel: """Create an input data channel for the training job. @@ -1058,9 +1059,16 @@ def create_input_data_channel( channel's data should be assigned to. Only applied when the channel is built from a URI/local-path data source (not a caller-supplied ``S3DataSource``/``FileSystemDataSource``, which the caller controls). + kms_key (Optional[str]): The Amazon Web Services KMS key id (or ARN/alias) to use for + server-side encryption when uploading local data to S3. If not specified, the + ``kms_key_id`` of this trainer's ``output_data_config`` is used (mirroring the V2 + ``Estimator`` behavior of encrypting staged user code with the output KMS key). + Only applied when ``data_source`` is a local file path. """ from sagemaker.core.helper.pipeline_variable import PipelineVariable + upload_extra_args = self._resolve_upload_extra_args(kms_key) + # Pass the field only when provided, so it stays unset (Unassigned) by default. instance_group_kwargs = ( {"instance_group_names": instance_group_names} @@ -1136,12 +1144,14 @@ def create_input_data_channel( path=copied_path, bucket=staging_bucket, key_prefix=effective_prefix, + extra_args=upload_extra_args, ) else: s3_uri = self.sagemaker_session.upload_data( path=data_source, bucket=staging_bucket, key_prefix=effective_prefix, + extra_args=upload_extra_args, ) channel = Channel( channel_name=channel_name, @@ -1170,6 +1180,22 @@ def create_input_data_channel( raise ValueError(f"Unsupported data_source type: {type(data_source)}") return channel + def _resolve_upload_extra_args(self, kms_key: Optional[str] = None) -> Optional[dict]: + """Build S3 ``ExtraArgs`` for encrypting local source/data uploads with KMS. + + Uses the explicit ``kms_key`` if given, otherwise falls back to the + ``kms_key_id`` on this trainer's ``output_data_config`` (V2 ``Estimator`` parity: + staged user code is encrypted with the output KMS key). Returns ``None`` when no + usable string key is configured, so behavior is unchanged by default (GH #5956). + A pipeline-variable key is ignored here because uploads happen at SDK compile time. + """ + resolved_key = kms_key + if resolved_key is None and self.output_data_config is not None: + resolved_key = getattr(self.output_data_config, "kms_key_id", None) + if isinstance(resolved_key, str) and resolved_key: + return {"ServerSideEncryption": "aws:kms", "SSEKMSKeyId": resolved_key} + return None + def _get_input_data_config( self, input_data_channels: Optional[List[Union[Channel, InputData]]], diff --git a/sagemaker-train/tests/unit/train/test_model_trainer.py b/sagemaker-train/tests/unit/train/test_model_trainer.py index 6022c60098..18146eddd9 100644 --- a/sagemaker-train/tests/unit/train/test_model_trainer.py +++ b/sagemaker-train/tests/unit/train/test_model_trainer.py @@ -474,6 +474,58 @@ def test_create_input_data_channel(mock_default_bucket, mock_upload_data, model_ assert channel.data_source.s3_data_source.s3_uri == expected_s3_uri +@patch("sagemaker.train.model_trainer.Session.upload_data") +@patch("sagemaker.train.model_trainer.Session.default_bucket") +def test_create_input_data_channel_encrypts_uploads_with_output_kms_key( + mock_default_bucket, mock_upload_data +): + """GH #5956: local source uploads are encrypted with output_data_config.kms_key_id.""" + mock_default_bucket.return_value = DEFAULT_BUCKET + mock_upload_data.return_value = f"s3://{DEFAULT_BUCKET}/code" + trainer = ModelTrainer( + training_image=DEFAULT_IMAGE, + role=DEFAULT_ROLE, + compute=DEFAULT_COMPUTE_CONFIG, + stopping_condition=DEFAULT_STOPPING_CONDITION, + output_data_config=OutputDataConfig( + s3_output_path=f"s3://{DEFAULT_BUCKET}/out", + kms_key_id="my-kms-key", + ), + ) + trainer.create_input_data_channel("code", DEFAULT_SOURCE_DIR) + assert mock_upload_data.call_args.kwargs["extra_args"] == { + "ServerSideEncryption": "aws:kms", + "SSEKMSKeyId": "my-kms-key", + } + + +@patch("sagemaker.train.model_trainer.Session.upload_data") +@patch("sagemaker.train.model_trainer.Session.default_bucket") +def test_create_input_data_channel_explicit_kms_key_overrides( + mock_default_bucket, mock_upload_data, model_trainer +): + """An explicit kms_key argument takes precedence over output_data_config.""" + mock_default_bucket.return_value = DEFAULT_BUCKET + mock_upload_data.return_value = f"s3://{DEFAULT_BUCKET}/code" + model_trainer.create_input_data_channel("code", DEFAULT_SOURCE_DIR, kms_key="explicit-key") + assert mock_upload_data.call_args.kwargs["extra_args"] == { + "ServerSideEncryption": "aws:kms", + "SSEKMSKeyId": "explicit-key", + } + + +@patch("sagemaker.train.model_trainer.Session.upload_data") +@patch("sagemaker.train.model_trainer.Session.default_bucket") +def test_create_input_data_channel_no_kms_by_default( + mock_default_bucket, mock_upload_data, model_trainer +): + """Default (no kms_key_id) leaves uploads unencrypted -> extra_args None (unchanged).""" + mock_default_bucket.return_value = DEFAULT_BUCKET + mock_upload_data.return_value = f"s3://{DEFAULT_BUCKET}/code" + model_trainer.create_input_data_channel("code", DEFAULT_SOURCE_DIR) + assert mock_upload_data.call_args.kwargs["extra_args"] is None + + def test_create_input_data_channel_with_instance_group_names(model_trainer): """instance_group_names is propagated onto the channel's S3DataSource.""" channel = model_trainer.create_input_data_channel( @@ -874,7 +926,7 @@ def test_remote_debug_config(mock_training_job, modules_session): @patch("sagemaker.train.model_trainer._get_unique_name") @patch("sagemaker.train.model_trainer.TrainingJob") def test_model_trainer_full_init(mock_training_job, mock_unique_name, modules_session): - def mock_upload_data(path, bucket, key_prefix): + def mock_upload_data(path, bucket, key_prefix, extra_args=None): return f"s3://{bucket}/{key_prefix}" modules_session.upload_data.side_effect = mock_upload_data @@ -1152,7 +1204,7 @@ def mock_upload_data(path, bucket, key_prefix): # def test_model_trainer_local_full_init( # mock_download_folder, mock_unique_name, mock_local_container, modules_session # ): -# def mock_upload_data(path, bucket, key_prefix): +# def mock_upload_data(path, bucket, key_prefix, extra_args=None): # return f"s3://{bucket}/{key_prefix}" # modules_session.upload_data.side_effect = mock_upload_data @@ -1406,7 +1458,7 @@ def test_hyperparameters_invalid(mock_exists, modules_session): @patch("sagemaker.train.model_trainer._get_unique_name") @patch("sagemaker.train.model_trainer.TrainingJob") def test_model_trainer_default_paths(mock_training_job, mock_unique_name, modules_session): - def mock_upload_data(path, bucket, key_prefix): + def mock_upload_data(path, bucket, key_prefix, extra_args=None): return f"s3://{bucket}/{key_prefix}" unique_name = "base-job-0123456789" @@ -1510,7 +1562,7 @@ def test_metric_definitions(mock_training_job, modules_session): @patch("sagemaker.train.model_trainer._get_unique_name") @patch("sagemaker.core.resources.TrainingJob") def test_nova_recipe(mock_training_job, mock_unique_name, modules_session): - def mock_upload_data(path, bucket, key_prefix): + def mock_upload_data(path, bucket, key_prefix, extra_args=None): if os.path.isfile(path): file_name = os.path.basename(path) return f"s3://{bucket}/{key_prefix}/{file_name}" @@ -1782,7 +1834,7 @@ def test_nova_recipe_model_package_config_only_mpg_from_recipe(modules_session): @patch("sagemaker.train.model_trainer._get_unique_name") @patch("sagemaker.train.model_trainer.TrainingJob") def test_llmft_recipe(mock_training_job, mock_unique_name, modules_session): - def mock_upload_data(path, bucket, key_prefix): + def mock_upload_data(path, bucket, key_prefix, extra_args=None): if os.path.isfile(path): file_name = os.path.basename(path) return f"s3://{bucket}/{key_prefix}/{file_name}"