diff --git a/sagemaker-serve/src/sagemaker/serve/mode/local_container_mode.py b/sagemaker-serve/src/sagemaker/serve/mode/local_container_mode.py index 3633b850e6..acbbb3e85d 100644 --- a/sagemaker-serve/src/sagemaker/serve/mode/local_container_mode.py +++ b/sagemaker-serve/src/sagemaker/serve/mode/local_container_mode.py @@ -4,8 +4,9 @@ from pathlib import Path import logging import os +import re from datetime import datetime, timedelta -from typing import Dict, Type +from typing import Dict, Optional, Type import base64 import time import subprocess @@ -29,6 +30,15 @@ logger = logging.getLogger(__name__) +# Strict allowlist for a real ECR registry host: <12-digit-account>.dkr.ecr..amazonaws.com +# Also matches AWS China (.amazonaws.com.cn). This is intentionally an exact host match so that the +# "is this an ECR image?" classifier and the "which host do I docker login to?" extractor operate on +# the SAME value, preventing a crafted URI (e.g. attacker.com/x.dkr.ecr..amazonaws.com/repo) +# from being classified as ECR while login is pointed at an attacker-controlled host. +_ECR_HOST_RE = re.compile( + r"^[0-9]{12}\.dkr\.ecr\.[a-z0-9-]+\.amazonaws\.com(\.cn)?$" +) + _PING_HEALTH_CHECK_INTERVAL_SEC = 5 _PING_HEALTH_CHECK_FAIL_MSG = ( @@ -252,7 +262,10 @@ def _pull_image(self, image: str): ) decoded_token = base64.b64decode(encoded_token).decode("utf-8") username, password = decoded_token.split(":") - ecr_uri = image.split("/")[0] + # Reuse the same validated host the classifier accepted, so the ECR credential is + # only ever sent to a verified ECR endpoint (never to an attacker-controlled host + # embedded elsewhere in the image URI). + ecr_uri = self._ecr_registry_host(image) login_command = ["docker", "login", "-u", username, "-p", password, ecr_uri] result = subprocess.run(login_command, check=True, capture_output=True, text=True) @@ -277,7 +290,23 @@ def _pull_image(self, image: str): except docker.errors.APIError as e: raise RuntimeError(f"Failed to pull image '{image}': {e}") from e + def _ecr_registry_host(self, image: str) -> Optional[str]: + """Return the ECR registry host if ``image``'s registry is a valid ECR endpoint, else None. + + The registry host is the first "/"-delimited segment of the image URI -- i.e. the exact + value that would be passed to ``docker login``. It is validated against a strict ECR + endpoint pattern so that classification and the login target are derived from the same + value. + """ + host = image.split("/")[0] + return host if _ECR_HOST_RE.match(host) else None + def _is_ecr_image(self, image: str) -> bool: - """Check if image is from ECR.""" - return ".dkr.ecr." in image and ".amazonaws.com" in image + """Check if image is from ECR. + + Uses the registry host that would actually be used for ``docker login`` and validates it + against a strict ECR endpoint pattern, so the classifier and the login-host extractor can + never disagree. + """ + return self._ecr_registry_host(image) is not None diff --git a/sagemaker-serve/tests/unit/mode/test_local_container_mode_ecr.py b/sagemaker-serve/tests/unit/mode/test_local_container_mode_ecr.py new file mode 100644 index 0000000000..82dac80748 --- /dev/null +++ b/sagemaker-serve/tests/unit/mode/test_local_container_mode_ecr.py @@ -0,0 +1,125 @@ +# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"). You +# may not use this file except in compliance with the License. A copy of +# the License is located at +# +# http://aws.amazon.com/apache2.0/ +# +# or in the "license" file accompanying this file. This file is +# distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF +# ANY KIND, either express or implied. See the License for the specific +# language governing permissions and limitations under the License. +"""Tests for ECR host classification / login-target extraction in LocalContainerMode. + +These guard against the parser-confusion vulnerability where the "is this an ECR image?" +classifier and the "which host do I docker login to?" extractor disagreed, allowing a crafted +image URI to leak the ECR authorization token to an attacker-controlled host. +""" +from __future__ import absolute_import + +import unittest +from unittest.mock import MagicMock, patch + +from sagemaker.serve.mode.local_container_mode import LocalContainerMode + + +def _bare_mode(): + """Build a LocalContainerMode without running the heavy constructor. + + ``_is_ecr_image`` / ``_ecr_registry_host`` are pure w.r.t. their ``image`` argument, so we + only need a bound method, not a fully initialized instance. + """ + return LocalContainerMode.__new__(LocalContainerMode) + + +class TestEcrHostClassification(unittest.TestCase): + def setUp(self): + self.mode = _bare_mode() + + def test_valid_ecr_uri_is_recognized(self): + image = "123456789012.dkr.ecr.us-east-1.amazonaws.com/my-repo:latest" + self.assertTrue(self.mode._is_ecr_image(image)) + self.assertEqual( + self.mode._ecr_registry_host(image), + "123456789012.dkr.ecr.us-east-1.amazonaws.com", + ) + + def test_valid_ecr_uri_china_partition(self): + image = "123456789012.dkr.ecr.cn-north-1.amazonaws.com.cn/my-repo:latest" + self.assertTrue(self.mode._is_ecr_image(image)) + self.assertEqual( + self.mode._ecr_registry_host(image), + "123456789012.dkr.ecr.cn-north-1.amazonaws.com.cn", + ) + + def test_attacker_crafted_uri_is_not_ecr(self): + # The malicious host is first; the ECR-looking substrings appear only in a later segment. + image = "attacker.com/x.dkr.ecr.us-east-1.amazonaws.com/repo:tag" + self.assertFalse(self.mode._is_ecr_image(image)) + self.assertIsNone(self.mode._ecr_registry_host(image)) + + def test_ecr_substrings_in_repo_path_do_not_classify_as_ecr(self): + image = "evil.example.com/a.dkr.ecr.b.amazonaws.com:latest" + self.assertFalse(self.mode._is_ecr_image(image)) + self.assertIsNone(self.mode._ecr_registry_host(image)) + + def test_public_image_is_not_ecr(self): + for image in ( + "docker.io/library/nginx:latest", + "nginx:latest", + "ubuntu", + ): + self.assertFalse(self.mode._is_ecr_image(image)) + self.assertIsNone(self.mode._ecr_registry_host(image)) + + def test_bad_account_id_length_is_not_ecr(self): + # 11 digits instead of 12 must not match. + image = "12345678901.dkr.ecr.us-east-1.amazonaws.com/repo:tag" + self.assertFalse(self.mode._is_ecr_image(image)) + + +class TestPullImageLoginTarget(unittest.TestCase): + """Ensure docker login is only ever pointed at the validated ECR host.""" + + @patch("sagemaker.serve.mode.local_container_mode.subprocess.run") + @patch("sagemaker.serve.mode.local_container_mode._get_docker_client") + def test_login_uses_validated_host_for_valid_ecr(self, mock_client, mock_run): + import base64 + + mode = _bare_mode() + mode.client = MagicMock() + mock_client.return_value = mode.client + mode.ecr = MagicMock() + token = base64.b64encode(b"AWS:secret-token").decode("utf-8") + mode.ecr.get_authorization_token.return_value = { + "authorizationData": [{"authorizationToken": token}] + } + + image = "123456789012.dkr.ecr.us-east-1.amazonaws.com/my-repo:latest" + mode._pull_image(image) + + # docker login must target the validated ECR host, never any other segment. + self.assertTrue(mock_run.called) + login_args = mock_run.call_args[0][0] + self.assertEqual(login_args[0:2], ["docker", "login"]) + self.assertEqual(login_args[-1], "123456789012.dkr.ecr.us-east-1.amazonaws.com") + + @patch("sagemaker.serve.mode.local_container_mode.subprocess.run") + @patch("sagemaker.serve.mode.local_container_mode._get_docker_client") + def test_no_login_for_attacker_crafted_uri(self, mock_client, mock_run): + mode = _bare_mode() + mode.client = MagicMock() + mock_client.return_value = mode.client + mode.ecr = MagicMock() + + image = "attacker.com/x.dkr.ecr.us-east-1.amazonaws.com/repo:tag" + mode._pull_image(image) + + # The crafted URI is not ECR: no token fetched, no docker login, no credential leak. + mode.ecr.get_authorization_token.assert_not_called() + mock_run.assert_not_called() + + +if __name__ == "__main__": + unittest.main()