From 05c3a07e3716b74a3e02bd535b457bc7138a2bf4 Mon Sep 17 00:00:00 2001 From: Sifei Li Date: Fri, 14 Aug 2026 15:20:14 -0700 Subject: [PATCH] fix(train): Enforce Nova Restricted storage on all MTRL MPG branches Restricted managed storage was only applied when MultiTurnRLTrainer auto-created a fresh model package group. The explicit output_model_package_group branch and the continued-customization branch (group derived from the source ModelPackage) performed no Nova check, so a Nova run could silently land its output or intermediate checkpoints in a Standard-storage group, or split lineage across groups depending on the branch taken. Centralize the requirement in _nova_managed_configuration() and apply it on every resolution path for both the output and intermediate checkpoint groups. Existing groups that do not satisfy the requirement now fail fast at trainer construction with an actionable error, since groups cannot be converted after creation. Non-Nova (OSS) models are unaffected. Adds unit tests for the explicit, object, derived, and pre-existing default-named group branches, plus OSS pass-through. --- .../sagemaker/train/multi_turn_rl_trainer.py | 85 +++++++++++++-- .../unit/train/test_multi_turn_rl_trainer.py | 103 ++++++++++++++++++ 2 files changed, 176 insertions(+), 12 deletions(-) diff --git a/sagemaker-train/src/sagemaker/train/multi_turn_rl_trainer.py b/sagemaker-train/src/sagemaker/train/multi_turn_rl_trainer.py index 2ab58b1cbf..aa9f1bde66 100644 --- a/sagemaker-train/src/sagemaker/train/multi_turn_rl_trainer.py +++ b/sagemaker-train/src/sagemaker/train/multi_turn_rl_trainer.py @@ -643,23 +643,78 @@ def _validate_networking(vpc): "VPC config requires both non-empty 'security_group_ids' and 'subnets'." ) + def _nova_managed_configuration(self): + """Return the ManagedConfiguration required for the base model, or None. + + Nova (closed-source) models require every ModelPackageGroup used by the + job — output and intermediate checkpoint — to use Restricted managed + storage. Open-source models have no managed-storage requirement. + """ + if not _is_nova_model(self._model_name): + return None + from sagemaker.core.shapes import ManagedConfiguration + + return ManagedConfiguration(managed_storage_type="Restricted") + + def _validate_mpg_managed_storage(self, mpg, managed_configuration) -> None: + """Validate an existing ModelPackageGroup against a required ManagedConfiguration. + + Raises: + ValueError: If ``managed_configuration`` is required and the existing + group's managed storage type does not match. Existing groups + cannot be converted, so failing fast here (at trainer + construction) is preferable to a late job-submission rejection. + """ + if managed_configuration is None: + return + required = managed_configuration.managed_storage_type + existing = getattr(mpg, "managed_configuration", None) + existing_type = getattr(existing, "managed_storage_type", None) + if not isinstance(existing_type, str): + existing_type = None + if existing_type != required: + group_name = getattr(mpg, "model_package_group_name", None) or getattr( + mpg, "model_package_group_arn", "" + ) + raise ValueError( + f"ModelPackageGroup '{group_name}' uses " + f"'{existing_type or 'Standard'}' managed storage, but model " + f"'{self._model_name}' requires '{required}'. Existing groups cannot " + "be converted. Pass a ModelPackageGroup created with " + f"ManagedConfiguration(managed_storage_type='{required}'), or omit " + "the parameter to auto-create a compliant group." + ) + def _get_or_create_mpg(self, value, default_name: str, session, managed_configuration=None) -> str: """Resolve an existing ModelPackageGroup or auto-create one. If ``value`` is provided (object or string), validates it exists and returns its ARN. If ``value`` is None, creates a ModelPackageGroup with ``default_name`` (get-or-create). + When ``managed_configuration`` is provided, every resolution path validates + that the resulting group satisfies it (auto-created groups are created + with it). Returns: The ModelPackageGroup ARN. """ if value: if isinstance(value, ModelPackageGroup): + if managed_configuration is not None: + # A caller-constructed object may not carry managed_configuration; + # fetch the authoritative record before validating. + fetched = ModelPackageGroup.get( + model_package_group_name=value.model_package_group_name, + session=session.boto_session, + region=session.boto_session.region_name, + ) + self._validate_mpg_managed_storage(fetched, managed_configuration) return value.model_package_group_arn mpg = ModelPackageGroup.get( model_package_group_name=value, session=session.boto_session, region=session.boto_session.region_name, ) + self._validate_mpg_managed_storage(mpg, managed_configuration) return mpg.model_package_group_arn # Auto-create (get-or-create with deterministic name) @@ -670,6 +725,9 @@ def _get_or_create_mpg(self, value, default_name: str, session, managed_configur session=session.boto_session, region=session.boto_session.region_name, ) + self._validate_mpg_managed_storage(mpg, managed_configuration) + except ValueError: + raise except Exception: try: create_kwargs = { @@ -695,22 +753,27 @@ def _resolve_model_package_group(self, model, output_model_package_group, sessio 2. If ``model`` is a ModelPackage, derives the group from it. 3. Otherwise, auto-creates ``{model_name}-mtrl-mpg`` (get-or-create). + For Nova models, every branch validates (or creates) the group with + Restricted managed storage — including the explicit and + continued-customization (ModelPackage-derived) branches. + Returns: The ModelPackageGroup ARN. """ + managed_config = self._nova_managed_configuration() + if output_model_package_group: - return self._get_or_create_mpg(output_model_package_group, None, session) + return self._get_or_create_mpg( + output_model_package_group, None, session, managed_configuration=managed_config + ) # Derive from ModelPackage if isinstance(model, ModelPackage): group_name = model.model_package_group_name if group_name: - return self._get_or_create_mpg(group_name, None, session) - - managed_config = None - if _is_nova_model(self._model_name): - from sagemaker.core.shapes import ManagedConfiguration - managed_config = ManagedConfiguration(managed_storage_type="Restricted") + return self._get_or_create_mpg( + group_name, None, session, managed_configuration=managed_config + ) return self._get_or_create_mpg( None, f"{self._model_name}-mtrl-mpg", session, managed_configuration=managed_config @@ -719,17 +782,15 @@ def _resolve_model_package_group(self, model, output_model_package_group, sessio def _resolve_intermediate_checkpoint_mpg(self, intermediate_checkpoint_mpg, session) -> str: """Resolve or auto-create the intermediate checkpoint ModelPackageGroup. - If provided, validates it exists. Otherwise auto-creates + If provided, validates it exists (and, for Nova models, that it uses + Restricted managed storage). Otherwise auto-creates ``{model_name}-mtrl-checkpoint-mpg`` (get-or-create). Raises ValueError if the resolved ARN is the same as ``output_model_package_group``. Returns: The ModelPackageGroup ARN. """ - managed_config = None - if not intermediate_checkpoint_mpg and _is_nova_model(self._model_name): - from sagemaker.core.shapes import ManagedConfiguration - managed_config = ManagedConfiguration(managed_storage_type="Restricted") + managed_config = self._nova_managed_configuration() arn = self._get_or_create_mpg( intermediate_checkpoint_mpg, diff --git a/sagemaker-train/tests/unit/train/test_multi_turn_rl_trainer.py b/sagemaker-train/tests/unit/train/test_multi_turn_rl_trainer.py index 291c7cc79e..d0e4552f10 100644 --- a/sagemaker-train/tests/unit/train/test_multi_turn_rl_trainer.py +++ b/sagemaker-train/tests/unit/train/test_multi_turn_rl_trainer.py @@ -491,6 +491,109 @@ def test_nova_model_creates_restricted_mpg(self, mock_get, mock_create): call_kwargs = mock_create.call_args[1] assert call_kwargs["managed_configuration"].managed_storage_type == "Restricted" + @staticmethod + def _mock_mpg(storage_type=None, name="my-group", arn=MPG_ARN): + mpg = MagicMock() + mpg.model_package_group_name = name + mpg.model_package_group_arn = arn + if storage_type is None: + mpg.managed_configuration = None + else: + mpg.managed_configuration = MagicMock() + mpg.managed_configuration.managed_storage_type = storage_type + return mpg + + @patch("sagemaker.train.multi_turn_rl_trainer.ModelPackageGroup.get") + def test_nova_explicit_standard_group_raises(self, mock_get): + """Nova + explicit output group with Standard storage must fail fast.""" + mock_get.return_value = self._mock_mpg(storage_type=None) + + trainer = self._make_trainer() + trainer._model_name = "amazon-nova-pro" + with pytest.raises(ValueError, match="requires 'Restricted'"): + trainer._resolve_model_package_group("amazon-nova-pro", "my-group", self._mock_session()) + + @patch("sagemaker.train.multi_turn_rl_trainer.ModelPackageGroup.get") + def test_nova_explicit_restricted_group_passes(self, mock_get): + mock_get.return_value = self._mock_mpg(storage_type="Restricted") + + trainer = self._make_trainer() + trainer._model_name = "amazon-nova-pro" + result = trainer._resolve_model_package_group( + "amazon-nova-pro", "my-group", self._mock_session() + ) + assert result == MPG_ARN + + @patch("sagemaker.train.multi_turn_rl_trainer.ModelPackageGroup.get") + def test_nova_mpg_object_standard_group_raises(self, mock_get): + """Nova + explicit ModelPackageGroup object is validated via a fresh fetch.""" + from sagemaker.core.resources import ModelPackageGroup as MPG + + mock_get.return_value = self._mock_mpg(storage_type=None) + mpg_obj = MagicMock(spec=MPG) + mpg_obj.model_package_group_name = "my-group" + mpg_obj.model_package_group_arn = MPG_ARN + + trainer = self._make_trainer() + trainer._model_name = "amazon-nova-pro" + with pytest.raises(ValueError, match="requires 'Restricted'"): + trainer._resolve_model_package_group("amazon-nova-pro", mpg_obj, self._mock_session()) + + @patch("sagemaker.train.multi_turn_rl_trainer.ModelPackageGroup.get") + def test_nova_derived_standard_group_raises(self, mock_get): + """Nova continued-customization: group derived from source ModelPackage must be Restricted.""" + mock_get.return_value = self._mock_mpg(storage_type=None, name="derived-group") + mock_model = MagicMock(spec=ModelPackage) + mock_model.model_package_group_name = "derived-group" + + trainer = self._make_trainer() + trainer._model_name = "amazon-nova-pro" + with pytest.raises(ValueError, match="requires 'Restricted'"): + trainer._resolve_model_package_group(mock_model, None, self._mock_session()) + + @patch("sagemaker.train.multi_turn_rl_trainer.ModelPackageGroup.get") + def test_nova_derived_restricted_group_passes(self, mock_get): + mock_get.return_value = self._mock_mpg(storage_type="Restricted", name="derived-group") + mock_model = MagicMock(spec=ModelPackage) + mock_model.model_package_group_name = "derived-group" + + trainer = self._make_trainer() + trainer._model_name = "amazon-nova-pro" + result = trainer._resolve_model_package_group(mock_model, None, self._mock_session()) + assert result == MPG_ARN + + @patch("sagemaker.train.multi_turn_rl_trainer.ModelPackageGroup.get") + def test_nova_existing_default_named_standard_group_raises(self, mock_get): + """Nova auto-create path: a pre-existing default-named Standard group must not be reused.""" + mock_get.return_value = self._mock_mpg( + storage_type=None, name="amazon-nova-pro-mtrl-mpg" + ) + + trainer = self._make_trainer() + trainer._model_name = "amazon-nova-pro" + with pytest.raises(ValueError, match="requires 'Restricted'"): + trainer._resolve_model_package_group("amazon-nova-pro", None, self._mock_session()) + + @patch("sagemaker.train.multi_turn_rl_trainer.ModelPackageGroup.get") + def test_oss_model_standard_group_passes_unvalidated(self, mock_get): + """Non-Nova models keep the existing behavior: no storage validation.""" + mock_get.return_value = self._mock_mpg(storage_type=None) + + trainer = self._make_trainer() + result = trainer._resolve_model_package_group("test-model", "my-group", self._mock_session()) + assert result == MPG_ARN + + @patch("sagemaker.train.multi_turn_rl_trainer.ModelPackageGroup.get") + def test_nova_intermediate_explicit_standard_group_raises(self, mock_get): + """Nova + explicit intermediate checkpoint group with Standard storage must fail.""" + mock_get.return_value = self._mock_mpg(storage_type=None, name="my-ckpt-group") + + trainer = self._make_trainer() + trainer._model_name = "amazon-nova-pro" + trainer.output_model_package_group = "arn:other" + with pytest.raises(ValueError, match="requires 'Restricted'"): + trainer._resolve_intermediate_checkpoint_mpg("my-ckpt-group", self._mock_session()) + class TestAgentRuntimeIdPattern: def test_valid_runtime_id(self):