From bbebdd416c41a37232e35932616a8b259174c004 Mon Sep 17 00:00:00 2001 From: Lucas Jia Date: Thu, 30 Jul 2026 13:25:29 -0700 Subject: [PATCH] fix(serve): Validate ECR registry host before docker login in LocalContainerMode _is_ecr_image() classified an image as ECR by substring-matching ".dkr.ecr." and ".amazonaws.com" anywhere in the URI, while _pull_image() took the docker login target from image.split("/")[0]. The two used inconsistent values, so a crafted URI such as attacker.com/x.dkr.ecr..amazonaws.com/repo passed the classifier yet caused "docker login -u AWS -p attacker.com", leaking a valid, replayable ECR authorization token to an attacker-controlled host. Parse the registry host first and validate it against a strict ECR endpoint pattern, and reuse that same validated host for docker login, so the classifier and the login target can never disagree. Add unit tests covering valid ECR (incl. China partition), attacker-crafted, and public image URIs. --- .../serve/mode/local_container_mode.py | 37 +++++- .../mode/test_local_container_mode_ecr.py | 125 ++++++++++++++++++ 2 files changed, 158 insertions(+), 4 deletions(-) create mode 100644 sagemaker-serve/tests/unit/mode/test_local_container_mode_ecr.py 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()