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
26 changes: 26 additions & 0 deletions sagemaker-train/src/sagemaker/train/model_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand All @@ -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}
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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]]],
Expand Down
62 changes: 57 additions & 5 deletions sagemaker-train/tests/unit/train/test_model_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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"
Expand Down Expand Up @@ -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}"
Expand Down Expand Up @@ -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}"
Expand Down
Loading