diff --git a/sagemaker-serve/src/sagemaker/serve/model_builder.py b/sagemaker-serve/src/sagemaker/serve/model_builder.py index 948b87bcaa..bc151e38b5 100644 --- a/sagemaker-serve/src/sagemaker/serve/model_builder.py +++ b/sagemaker-serve/src/sagemaker/serve/model_builder.py @@ -3141,13 +3141,16 @@ def _create_sagemaker_model(self): sagemaker_session=self.sagemaker_session, ) - # Nova 1P images require network isolation on the Model resource + # Nova 1P images require network isolation on the Model resource. + # A pipeline-variable image is only resolved at execution time, so it + # cannot be inspected here; set enable_network_isolation explicitly then. enable_network_isolation = self._enable_network_isolation resolved_image_uri = ( container_def["Image"] if isinstance(container_def, dict) else container_def[0]["Image"] ) if ( not enable_network_isolation + and isinstance(resolved_image_uri, str) and "nova-" in resolved_image_uri and is_1p_image_uri(resolved_image_uri) ): diff --git a/sagemaker-serve/src/sagemaker/serve/validations/check_image_uri.py b/sagemaker-serve/src/sagemaker/serve/validations/check_image_uri.py index a71e0b2db4..017f6c691e 100644 --- a/sagemaker-serve/src/sagemaker/serve/validations/check_image_uri.py +++ b/sagemaker-serve/src/sagemaker/serve/validations/check_image_uri.py @@ -4,6 +4,9 @@ import logging import re +from typing import Optional + +from sagemaker.core.helper.pipeline_variable import StrPipeVar logger = logging.getLogger(__name__) @@ -318,8 +321,15 @@ } -def is_1p_image_uri(image_uri: str) -> bool: - """Shows if the given image_uri is owned by a 1st party account""" +def is_1p_image_uri(image_uri: Optional[StrPipeVar]) -> bool: + """Shows if the given image_uri is owned by a 1st party account. + + A pipeline variable (e.g. ``ParameterString``) is only resolved at pipeline + execution time, so its account cannot be inspected at build time. Such values, + and any other non-string value, are treated as not 1st party. + """ + if not isinstance(image_uri, str): + return False image_uri_account = image_uri[0:12] return image_uri_account in all_accounts diff --git a/sagemaker-serve/tests/unit/test_model_builder_pipeline_variable_image_uri.py b/sagemaker-serve/tests/unit/test_model_builder_pipeline_variable_image_uri.py new file mode 100644 index 0000000000..44730c8a9f --- /dev/null +++ b/sagemaker-serve/tests/unit/test_model_builder_pipeline_variable_image_uri.py @@ -0,0 +1,124 @@ +"""Unit tests for ModelBuilder with a pipeline-variable image_uri (issue #5760).""" + +import unittest +from unittest.mock import Mock, patch + +import boto3 + +from sagemaker.core.workflow.parameters import ParameterString +from sagemaker.core.workflow.pipeline_context import PipelineSession, _ModelStepArguments +from sagemaker.serve.model_builder import ModelBuilder + +ROLE_ARN = "arn:aws:iam::123456789012:role/TestRole" +TELEMETRY_MODULE = "sagemaker.core.telemetry.telemetry_logging" +NOVA_1P_IMAGE = "708977205387.dkr.ecr.us-east-1.amazonaws.com/nova-inference:latest" + + +def _mock_session(): + session = Mock() + session.boto_region_name = "us-west-2" + session.default_bucket.return_value = "test-bucket" + session.default_bucket_prefix = None + session.config = {} + session.sagemaker_config = {} + return session + + +class TestBuildValidationsWithPipelineVariableImageUri(unittest.TestCase): + """_build_validations must not slice a pipeline-variable image_uri.""" + + def test_image_only_pipeline_variable_is_passthrough(self): + builder = ModelBuilder( + image_uri=ParameterString(name="ServingImageUri"), + role_arn=ROLE_ARN, + sagemaker_session=_mock_session(), + ) + + builder._build_validations() + + self.assertTrue(builder._passthrough) + + def test_pipeline_variable_with_model_requires_model_server(self): + builder = ModelBuilder( + model=Mock(), + image_uri=ParameterString(name="ServingImageUri"), + role_arn=ROLE_ARN, + sagemaker_session=_mock_session(), + ) + + with self.assertRaises(ValueError) as context: + builder._build_validations() + + self.assertIn("Model_server must be set", str(context.exception)) + + +@patch("sagemaker.serve.model_builder.resolve_nested_dict_value_from_config") +@patch("sagemaker.serve.model_builder.resolve_value_from_config") +@patch.object(ModelBuilder, "_init_sagemaker_session_if_does_not_exist") +@patch.object(ModelBuilder, "_prepare_container_def") +class TestCreateSageMakerModelNetworkIsolation(unittest.TestCase): + """The Nova network-isolation check must tolerate a pipeline-variable image.""" + + def _run(self, image_uri, mock_prepare, mock_resolve, mock_resolve_nested): + mock_prepare.return_value = {"Image": image_uri, "Environment": {}} + mock_resolve.side_effect = lambda value, *args, **kwargs: value + mock_resolve_nested.side_effect = lambda value, *args, **kwargs: value + + session = Mock(spec=PipelineSession) + builder = ModelBuilder( + image_uri=image_uri, + role_arn=ROLE_ARN, + sagemaker_session=_mock_session(), + ) + builder.sagemaker_session = session + builder.model_name = "test-model" + + builder._create_sagemaker_model() + + session.create_model.assert_called_once() + return session.create_model.call_args.kwargs + + def test_pipeline_variable_image_does_not_raise( + self, mock_prepare, mock_init, mock_resolve, mock_resolve_nested + ): + image_uri = ParameterString(name="ServingImageUri") + + kwargs = self._run(image_uri, mock_prepare, mock_resolve, mock_resolve_nested) + + self.assertIs(kwargs["container_defs"]["Image"], image_uri) + self.assertFalse(kwargs["enable_network_isolation"]) + + def test_nova_1p_string_image_still_enables_network_isolation( + self, mock_prepare, mock_init, mock_resolve, mock_resolve_nested + ): + kwargs = self._run(NOVA_1P_IMAGE, mock_prepare, mock_resolve, mock_resolve_nested) + + self.assertTrue(kwargs["enable_network_isolation"]) + + +@patch(f"{TELEMETRY_MODULE}.resolve_value_from_config", return_value=False) +@patch(f"{TELEMETRY_MODULE}._send_telemetry_request") +@patch("sagemaker.serve.model_builder.resolve_and_validate_role", return_value=ROLE_ARN) +@patch.object(PipelineSession, "default_bucket", return_value="test-bucket") +class TestBuildWithPipelineVariableImageUri(unittest.TestCase): + """End-to-end build() under a PipelineSession, as reported in #5760.""" + + def test_build_keeps_pipeline_variable_in_create_model_request(self, *_mocks): + image_uri = ParameterString(name="ServingImageUri") + session = PipelineSession( + boto_session=boto3.Session( + region_name="us-west-2", + aws_access_key_id="testing", + aws_secret_access_key="testing", + ) + ) + builder = ModelBuilder(image_uri=image_uri, role_arn=ROLE_ARN, sagemaker_session=session) + + step_args = builder.build() + + self.assertIsInstance(step_args, _ModelStepArguments) + self.assertIs(step_args.create_model_request["PrimaryContainer"]["Image"], image_uri) + + +if __name__ == "__main__": + unittest.main() diff --git a/sagemaker-serve/tests/unit/validations/test_check_image_uri.py b/sagemaker-serve/tests/unit/validations/test_check_image_uri.py index e7871f5573..4b95f3dab8 100644 --- a/sagemaker-serve/tests/unit/validations/test_check_image_uri.py +++ b/sagemaker-serve/tests/unit/validations/test_check_image_uri.py @@ -1,4 +1,6 @@ import unittest +from sagemaker.core.workflow.functions import Join +from sagemaker.core.workflow.parameters import ParameterString from sagemaker.serve.validations.check_image_uri import ( is_1p_image_uri, all_accounts, @@ -73,6 +75,23 @@ def test_all_accounts_contains_known_accounts(self): self.assertIn("763104351884", all_accounts) self.assertIn("246618743249", all_accounts) + def test_is_1p_image_uri_pipeline_variable_returns_false(self): + # A pipeline variable is only resolved at execution time, so it must not be + # sliced (TypeError, #5760); it is treated as not 1P, even if its default + # value would be a 1P image. + image_uri = ParameterString( + name="ServingImageUri", + default_value="763104351884.dkr.ecr.us-east-1.amazonaws.com/pytorch:latest", + ) + self.assertFalse(is_1p_image_uri(image_uri)) + + def test_is_1p_image_uri_pipeline_function_returns_false(self): + image_uri = Join(on="", values=["763104351884", ParameterString(name="Suffix")]) + self.assertFalse(is_1p_image_uri(image_uri)) + + def test_is_1p_image_uri_none_returns_false(self): + self.assertFalse(is_1p_image_uri(None)) + if __name__ == "__main__": unittest.main()