Skip to content
Open
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
13 changes: 13 additions & 0 deletions sagemaker-core/src/sagemaker/core/jumpstart/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -1196,9 +1196,20 @@ def get_jumpstart_configs(
scope: enums.JumpStartScriptScope = enums.JumpStartScriptScope.INFERENCE,
model_type: enums.JumpStartModelType = enums.JumpStartModelType.OPEN_WEIGHTS,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
) -> Dict[str, JumpStartMetadataConfig]:
"""Returns metadata configs for the given model ID and region.

Args:
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
security vulnerabilities. (Default: False).
tolerate_deprecated_model (bool): True if deprecated models should be tolerated
(exception not raised). False if these models should raise an exception.
(Default: False).

Raises:
ValueError: If the script scope is not supported by JumpStart.
"""
Expand All @@ -1210,6 +1221,8 @@ def get_jumpstart_configs(
scope=scope,
model_type=model_type,
hub_arn=hub_arn,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
)

if scope == enums.JumpStartScriptScope.INFERENCE:
Expand Down
93 changes: 93 additions & 0 deletions sagemaker-core/tests/unit/test_jumpstart_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,7 @@
JumpStartBenchmarkStat,
DeploymentConfigMetadata,
)
from sagemaker.core.jumpstart.exceptions import VulnerableJumpStartModelError
from sagemaker.core.jumpstart.models import HubContentDocument
from sagemaker.core.helper.pipeline_variable import PipelineVariable

Expand Down Expand Up @@ -1356,6 +1357,98 @@ def test_get_jumpstart_configs_no_configs(self, mock_verify):
result = utils.get_jumpstart_configs("us-west-2", "test-model", "1.0.0")
assert result == {}

@patch("sagemaker.core.jumpstart.utils.verify_model_region_and_return_specs")
def test_get_jumpstart_configs_does_not_tolerate_by_default(self, mock_verify):
"""Test the model gate is left enabled when the caller asks for nothing"""
mock_specs = Mock()
mock_specs.inference_configs = None
mock_verify.return_value = mock_specs

utils.get_jumpstart_configs("us-west-2", "test-model", "1.0.0")

assert mock_verify.call_args.kwargs["tolerate_vulnerable_model"] is False
assert mock_verify.call_args.kwargs["tolerate_deprecated_model"] is False

@patch("sagemaker.core.jumpstart.utils.verify_model_region_and_return_specs")
def test_get_jumpstart_configs_forwards_tolerance(self, mock_verify):
"""Test tolerance reaches the spec lookup that runs the model gate"""
mock_specs = Mock()
mock_specs.inference_configs = None
mock_verify.return_value = mock_specs

utils.get_jumpstart_configs(
"us-west-2",
"test-model",
"1.0.0",
tolerate_vulnerable_model=True,
tolerate_deprecated_model=True,
)

assert mock_verify.call_args.kwargs["tolerate_vulnerable_model"] is True
assert mock_verify.call_args.kwargs["tolerate_deprecated_model"] is True

@patch("sagemaker.core.jumpstart.utils.verify_model_region_and_return_specs")
def test_get_jumpstart_configs_forwards_tolerance_for_training_scope(self, mock_verify):
"""Test tolerance reaches the spec lookup on the training scope too"""
mock_specs = Mock()
mock_specs.training_configs = None
mock_verify.return_value = mock_specs

utils.get_jumpstart_configs(
"us-west-2",
"test-model",
"1.0.0",
scope=enums.JumpStartScriptScope.TRAINING,
tolerate_vulnerable_model=True,
tolerate_deprecated_model=True,
)

assert mock_verify.call_args.kwargs["tolerate_vulnerable_model"] is True
assert mock_verify.call_args.kwargs["tolerate_deprecated_model"] is True

@patch("sagemaker.core.jumpstart.accessors.JumpStartModelsAccessor.get_model_specs")
def test_get_jumpstart_configs_vulnerable_model_raises_by_default(self, mock_get_specs):
"""Test a vulnerable model still trips the gate when tolerance is not requested"""
model_specs = Mock(spec=JumpStartModelSpecs)
model_specs.deprecated = False
model_specs.inference_vulnerable = True
model_specs.inference_vulnerabilities = ["CVE-2024-11393"]
mock_get_specs.return_value = model_specs

with pytest.raises(VulnerableJumpStartModelError):
utils.get_jumpstart_configs("us-west-2", "test-model", "1.0.0")

@patch("sagemaker.core.jumpstart.accessors.JumpStartModelsAccessor.get_model_specs")
def test_get_jumpstart_configs_tolerates_vulnerable_model(self, mock_get_specs):
"""Test a vulnerable model resolves configs instead of tripping the gate"""
model_specs = Mock(spec=JumpStartModelSpecs)
model_specs.deprecated = False
model_specs.inference_vulnerable = True
model_specs.inference_vulnerabilities = ["CVE-2024-11393"]
model_specs.inference_configs = None
mock_get_specs.return_value = model_specs

result = utils.get_jumpstart_configs(
"us-west-2", "test-model", "1.0.0", tolerate_vulnerable_model=True
)

assert result == {}

@patch("sagemaker.core.jumpstart.accessors.JumpStartModelsAccessor.get_model_specs")
def test_get_jumpstart_configs_tolerates_deprecated_model(self, mock_get_specs):
"""Test a deprecated model resolves configs instead of tripping the gate"""
model_specs = Mock(spec=JumpStartModelSpecs)
model_specs.deprecated = True
model_specs.inference_vulnerable = False
model_specs.inference_configs = None
mock_get_specs.return_value = model_specs

result = utils.get_jumpstart_configs(
"us-west-2", "test-model", "1.0.0", tolerate_deprecated_model=True
)

assert result == {}


class TestGetJumpstartUserAgentExtraSuffix:
"""Test cases for get_jumpstart_user_agent_extra_suffix function"""
Expand Down
2 changes: 2 additions & 0 deletions sagemaker-serve/src/sagemaker/serve/model_builder_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -2850,6 +2850,8 @@ def _ensure_metadata_configs(self) -> None:
model_id=model,
model_version=getattr(self, "model_version", None) or "*",
sagemaker_session=getattr(self, "sagemaker_session", None),
tolerate_vulnerable_model=getattr(self, "tolerate_vulnerable_model", None) or False,
tolerate_deprecated_model=getattr(self, "tolerate_deprecated_model", None) or False,
)

def _user_agent_decorator(self, func):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -611,6 +611,40 @@ def test_ensure_metadata_configs_not_string(self):
# Should remain None for non-string models
self.assertIsNone(utils._metadata_configs)

@patch("sagemaker.core.jumpstart.utils.get_jumpstart_configs")
def test_ensure_metadata_configs_forwards_tolerance(self, mock_get_configs):
"""Test tolerance flags reach the config lookup that runs the model gate."""
utils = _ModelBuilderUtils()
utils._metadata_configs = None
utils.model = "huggingface-llm-falcon-7b"
utils.region = "us-west-2"
utils.sagemaker_session = Mock()
utils.tolerate_vulnerable_model = True
utils.tolerate_deprecated_model = True

mock_get_configs.return_value = {}

utils._ensure_metadata_configs()

self.assertTrue(mock_get_configs.call_args.kwargs["tolerate_vulnerable_model"])
self.assertTrue(mock_get_configs.call_args.kwargs["tolerate_deprecated_model"])

@patch("sagemaker.core.jumpstart.utils.get_jumpstart_configs")
def test_ensure_metadata_configs_defaults_tolerance_to_false(self, mock_get_configs):
"""Test the model gate stays enabled when tolerance is not set."""
utils = _ModelBuilderUtils()
utils._metadata_configs = None
utils.model = "huggingface-llm-falcon-7b"
utils.region = "us-west-2"
utils.sagemaker_session = Mock()

mock_get_configs.return_value = {}

utils._ensure_metadata_configs()

self.assertFalse(mock_get_configs.call_args.kwargs["tolerate_vulnerable_model"])
self.assertFalse(mock_get_configs.call_args.kwargs["tolerate_deprecated_model"])


class TestGetServeSettings(unittest.TestCase):
"""Test _get_serve_setting method - skipped (requires proper session setup)."""
Expand Down
Loading