feat(sagemaker): add sagemaker_models_registry_in_use check (#11196)

Co-authored-by: cascioli <simdon2015?gmail.com>
Co-authored-by: Claude Sonnet 4.6 <noreply@anthropic.com>
Co-authored-by: Daniel Barranquero <danielbo2001@gmail.com>
This commit is contained in:
Simone
2026-05-20 13:59:18 +02:00
committed by GitHub
co-authored by cascioli Claude Sonnet 4.6 Daniel Barranquero
parent cff1704d7b
commit 534dedb608
7 changed files with 358 additions and 0 deletions
+1
View File
@@ -10,6 +10,7 @@ All notable changes to the **Prowler SDK** are documented in this file.
- `cloudsql_instance_cmek_encryption_enabled` check for GCP provider [(#11023)](https://github.com/prowler-cloud/prowler/pull/11023)
- Google Workspace Groups service with 3 new checks [(#11186)](https://github.com/prowler-cloud/prowler/pull/11186)
- `ses_identity_dkim_enabled` check for AWS provider [(#10923)](https://github.com/prowler-cloud/prowler/pull/10923)
- `sagemaker_models_registry_in_use` check for AWS provider, verifying that at least one SageMaker Model Package Group has an approved model package to enforce ML governance workflows [(#11196)](https://github.com/prowler-cloud/prowler/pull/11196)
### 🔄 Changed
@@ -0,0 +1,42 @@
{
"Provider": "aws",
"CheckID": "sagemaker_models_registry_in_use",
"CheckTitle": "Amazon SageMaker Model Registry should have at least one approved model package",
"CheckType": [
"Software and Configuration Checks/AWS Security Best Practices"
],
"ServiceName": "sagemaker",
"SubServiceName": "",
"ResourceIdTemplate": "",
"Severity": "low",
"ResourceType": "Other",
"ResourceGroup": "ai_ml",
"Description": "**SageMaker Model Registry** is evaluated to verify that at least one Model Package Group exists and contains at least one model package with **ModelApprovalStatus = Approved**. This confirms that the ML governance workflow (register → review → approve → deploy) is actively in use.",
"Risk": "An empty Model Registry, or one with no approved packages, indicates that models are being deployed outside any review process. This breaks provenance and accountability for production ML workloads, making it impossible to enforce governance controls such as auditing, versioning, and approval workflows.",
"RelatedUrl": "",
"AdditionalURLs": [
"https://docs.aws.amazon.com/sagemaker/latest/dg/model-registry.html",
"https://docs.aws.amazon.com/sagemaker/latest/dg/model-registry-approve.html",
"https://docs.aws.amazon.com/sagemaker/latest/APIReference/API_ListModelPackageGroups.html",
"https://docs.aws.amazon.com/sagemaker/latest/APIReference/API_ListModelPackages.html"
],
"Remediation": {
"Code": {
"CLI": "aws sagemaker list-model-package-groups\naws sagemaker list-model-packages --model-package-group-name <group-name>\naws sagemaker update-model-package --model-package-arn <arn> --model-approval-status Approved",
"NativeIaC": "",
"Other": "1. In the AWS console, navigate to SageMaker > Models > Model Registry.\n2. Create a Model Package Group if none exists.\n3. Register a model version in the group.\n4. Review and approve at least one model package by setting its approval status to Approved.",
"Terraform": ""
},
"Recommendation": {
"Text": "Register all production models in the **SageMaker Model Registry** and enforce an approval workflow before deployment. Ensure at least one model package per group reaches **Approved** status. Use **IAM policies** to restrict who can approve model packages and integrate with **CI/CD pipelines** to automate registration.",
"Url": "https://hub.prowler.com/check/sagemaker_models_registry_in_use"
}
},
"Categories": [
"gen-ai",
"software-supply-chain"
],
"DependsOn": [],
"RelatedTo": [],
"Notes": ""
}
@@ -0,0 +1,28 @@
from prowler.lib.check.models import Check, Check_Report_AWS
from prowler.providers.aws.services.sagemaker.sagemaker_client import sagemaker_client
class sagemaker_models_registry_in_use(Check):
"""Ensure that SageMaker Model Registry has at least one approved model package."""
def execute(self) -> list[Check_Report_AWS]:
"""Execute the check logic.
Returns:
A list of reports indicating whether the SageMaker Model Registry
in each region contains at least one approved model package.
"""
findings = []
for registry in sagemaker_client.sagemaker_model_registries:
report = Check_Report_AWS(metadata=self.metadata(), resource=registry)
if not registry.has_groups:
report.status = "FAIL"
report.status_extended = f"SageMaker Model Registry in region {registry.region} has no Model Package Groups."
elif registry.has_approved_packages:
report.status = "PASS"
report.status_extended = f"SageMaker Model Registry in region {registry.region} has at least one approved model package."
else:
report.status = "FAIL"
report.status_extended = f"SageMaker Model Registry in region {registry.region} has Model Package Groups but no approved model packages."
findings.append(report)
return findings
@@ -17,6 +17,7 @@ class SageMaker(AWSService):
self.sagemaker_training_jobs = []
self.sagemaker_domains = []
self.endpoint_configs = {}
self.sagemaker_model_registries = []
# Retrieve resources concurrently
self.__threading_call__(self._list_notebook_instances)
@@ -24,6 +25,7 @@ class SageMaker(AWSService):
self.__threading_call__(self._list_training_jobs)
self.__threading_call__(self._list_endpoint_configs)
self.__threading_call__(self._list_domains)
self.__threading_call__(self._list_model_package_groups)
# Describe resources concurrently
self.__threading_call__(self._describe_model, self.sagemaker_models)
@@ -207,6 +209,71 @@ class SageMaker(AWSService):
f"{regional_client.region} -- {error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}"
)
def _list_model_package_groups(self, regional_client):
logger.info("SageMaker - listing model package groups...")
registry_arn = self.get_unknown_arn(
region=regional_client.region,
resource_type="model-registry",
)
has_groups = False
has_approved = False
try:
paginator = regional_client.get_paginator("list_model_package_groups")
for page in paginator.paginate():
for group in page["ModelPackageGroupSummaryList"]:
has_groups = True
if not has_approved:
group_name = group["ModelPackageGroupName"]
try:
pkg_paginator = regional_client.get_paginator(
"list_model_packages"
)
for pkg_page in pkg_paginator.paginate(
ModelPackageGroupName=group_name,
ModelApprovalStatus="Approved",
):
if pkg_page["ModelPackageSummaryList"]:
has_approved = True
break
except ClientError as pkg_error:
if pkg_error.response["Error"]["Code"] in (
"AccessDeniedException",
"UnrecognizedClientException",
):
raise
logger.error(
f"{regional_client.region} -- {pkg_error.__class__.__name__}[{pkg_error.__traceback__.tb_lineno}]: {pkg_error}"
)
except Exception as pkg_error:
logger.error(
f"{regional_client.region} -- {pkg_error.__class__.__name__}[{pkg_error.__traceback__.tb_lineno}]: {pkg_error}"
)
except ClientError as error:
if error.response["Error"]["Code"] in (
"AccessDeniedException",
"UnrecognizedClientException",
):
logger.warning(
f"{regional_client.region} -- {error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}"
)
return
logger.error(
f"{regional_client.region} -- {error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}"
)
except Exception as error:
logger.error(
f"{regional_client.region} -- {error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}"
)
self.sagemaker_model_registries.append(
ModelRegistry(
name="SageMaker Model Registry",
arn=registry_arn,
region=regional_client.region,
has_groups=has_groups,
has_approved_packages=has_approved,
)
)
def _list_tags_for_resource(self, resource):
"""
Lists tags for a specific SageMaker resource.
@@ -364,3 +431,13 @@ class EndpointConfig(BaseModel):
arn: str
production_variants: list[ProductionVariant] = []
tags: Optional[list] = []
class ModelRegistry(BaseModel):
"""Represents the SageMaker Model Registry state for a specific region."""
name: str
arn: str
region: str
has_groups: bool = False
has_approved_packages: bool = False
@@ -0,0 +1,156 @@
from unittest import mock
from prowler.providers.aws.services.sagemaker.sagemaker_service import ModelRegistry
from tests.providers.aws.utils import (
AWS_ACCOUNT_NUMBER,
AWS_REGION_EU_WEST_1,
set_mocked_aws_provider,
)
registry_arn = f"arn:aws:sagemaker:{AWS_REGION_EU_WEST_1}:{AWS_ACCOUNT_NUMBER}:model-registry/unknown"
class Test_sagemaker_models_registry_in_use:
def test_no_registries(self):
sagemaker_client = mock.MagicMock
sagemaker_client.sagemaker_model_registries = []
aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1])
with (
mock.patch(
"prowler.providers.common.provider.Provider.get_global_provider",
return_value=aws_provider,
),
mock.patch(
"prowler.providers.aws.services.sagemaker.sagemaker_models_registry_in_use.sagemaker_models_registry_in_use.sagemaker_client",
sagemaker_client,
),
):
from prowler.providers.aws.services.sagemaker.sagemaker_models_registry_in_use.sagemaker_models_registry_in_use import (
sagemaker_models_registry_in_use,
)
check = sagemaker_models_registry_in_use()
result = check.execute()
assert len(result) == 0
def test_registry_no_groups(self):
sagemaker_client = mock.MagicMock
sagemaker_client.sagemaker_model_registries = [
ModelRegistry(
name="SageMaker Model Registry",
arn=registry_arn,
region=AWS_REGION_EU_WEST_1,
has_groups=False,
has_approved_packages=False,
)
]
aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1])
with (
mock.patch(
"prowler.providers.common.provider.Provider.get_global_provider",
return_value=aws_provider,
),
mock.patch(
"prowler.providers.aws.services.sagemaker.sagemaker_models_registry_in_use.sagemaker_models_registry_in_use.sagemaker_client",
sagemaker_client,
),
):
from prowler.providers.aws.services.sagemaker.sagemaker_models_registry_in_use.sagemaker_models_registry_in_use import (
sagemaker_models_registry_in_use,
)
check = sagemaker_models_registry_in_use()
result = check.execute()
assert len(result) == 1
assert result[0].status == "FAIL"
assert (
result[0].status_extended
== f"SageMaker Model Registry in region {AWS_REGION_EU_WEST_1} has no Model Package Groups."
)
assert result[0].resource_id == "SageMaker Model Registry"
assert result[0].resource_arn == registry_arn
assert result[0].region == AWS_REGION_EU_WEST_1
def test_registry_groups_no_approved_packages(self):
sagemaker_client = mock.MagicMock
sagemaker_client.sagemaker_model_registries = [
ModelRegistry(
name="SageMaker Model Registry",
arn=registry_arn,
region=AWS_REGION_EU_WEST_1,
has_groups=True,
has_approved_packages=False,
)
]
aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1])
with (
mock.patch(
"prowler.providers.common.provider.Provider.get_global_provider",
return_value=aws_provider,
),
mock.patch(
"prowler.providers.aws.services.sagemaker.sagemaker_models_registry_in_use.sagemaker_models_registry_in_use.sagemaker_client",
sagemaker_client,
),
):
from prowler.providers.aws.services.sagemaker.sagemaker_models_registry_in_use.sagemaker_models_registry_in_use import (
sagemaker_models_registry_in_use,
)
check = sagemaker_models_registry_in_use()
result = check.execute()
assert len(result) == 1
assert result[0].status == "FAIL"
assert (
result[0].status_extended
== f"SageMaker Model Registry in region {AWS_REGION_EU_WEST_1} has Model Package Groups but no approved model packages."
)
assert result[0].resource_id == "SageMaker Model Registry"
assert result[0].resource_arn == registry_arn
assert result[0].region == AWS_REGION_EU_WEST_1
def test_registry_with_approved_packages(self):
sagemaker_client = mock.MagicMock
sagemaker_client.sagemaker_model_registries = [
ModelRegistry(
name="SageMaker Model Registry",
arn=registry_arn,
region=AWS_REGION_EU_WEST_1,
has_groups=True,
has_approved_packages=True,
)
]
aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1])
with (
mock.patch(
"prowler.providers.common.provider.Provider.get_global_provider",
return_value=aws_provider,
),
mock.patch(
"prowler.providers.aws.services.sagemaker.sagemaker_models_registry_in_use.sagemaker_models_registry_in_use.sagemaker_client",
sagemaker_client,
),
):
from prowler.providers.aws.services.sagemaker.sagemaker_models_registry_in_use.sagemaker_models_registry_in_use import (
sagemaker_models_registry_in_use,
)
check = sagemaker_models_registry_in_use()
result = check.execute()
assert len(result) == 1
assert result[0].status == "PASS"
assert (
result[0].status_extended
== f"SageMaker Model Registry in region {AWS_REGION_EU_WEST_1} has at least one approved model package."
)
assert result[0].resource_id == "SageMaker Model Registry"
assert result[0].resource_arn == registry_arn
assert result[0].region == AWS_REGION_EU_WEST_1
@@ -13,6 +13,11 @@ from tests.providers.aws.utils import (
set_mocked_aws_provider,
)
test_model_package_group_name = "test-model-package-group"
test_model_package_group_arn = f"arn:aws:sagemaker:{AWS_REGION_EU_WEST_1}:{AWS_ACCOUNT_NUMBER}:model-package-group/{test_model_package_group_name}"
test_model_package_name = "test-model-package"
test_model_package_arn = f"arn:aws:sagemaker:{AWS_REGION_EU_WEST_1}:{AWS_ACCOUNT_NUMBER}:model-package/{test_model_package_name}/1"
test_notebook_instance = "test-notebook-instance"
notebook_instance_arn = f"arn:aws:sagemaker:{AWS_REGION_EU_WEST_1}:{AWS_ACCOUNT_NUMBER}:notebook-instance/{test_notebook_instance}"
test_model = "test-model"
@@ -94,6 +99,25 @@ def mock_make_api_call(self, operation_name, kwarg):
"EnableNetworkIsolation": True,
"EnableInterContainerTrafficEncryption": True,
}
if operation_name == "ListModelPackageGroups":
return {
"ModelPackageGroupSummaryList": [
{
"ModelPackageGroupName": test_model_package_group_name,
"ModelPackageGroupArn": test_model_package_group_arn,
},
]
}
if operation_name == "ListModelPackages":
return {
"ModelPackageSummaryList": [
{
"ModelPackageName": test_model_package_name,
"ModelPackageArn": test_model_package_arn,
"ModelApprovalStatus": "Approved",
},
]
}
if operation_name == "ListTags":
return {
"Tags": [
@@ -379,3 +403,33 @@ class Test_SageMaker_Service:
if c[0][0] == sagemaker_service._list_tags_for_resource
]
assert len(tag_calls) == 5
# Test SageMaker list model package groups
def test_list_model_package_groups(self):
aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1])
sagemaker = SageMaker(aws_provider)
assert len(sagemaker.sagemaker_model_registries) == 1
registry = sagemaker.sagemaker_model_registries[0]
assert registry.region == AWS_REGION_EU_WEST_1
assert registry.has_groups is True
assert registry.has_approved_packages is True
def test_list_model_package_groups_access_denied(self):
aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1])
def mock_access_denied(self, operation_name, kwarg):
if operation_name == "ListModelPackageGroups":
raise botocore.exceptions.ClientError(
{
"Error": {
"Code": "AccessDeniedException",
"Message": "User is not authorized to perform sagemaker:ListModelPackageGroups",
}
},
"ListModelPackageGroups",
)
return make_api_call(self, operation_name, kwarg)
with patch("botocore.client.BaseClient._make_api_call", new=mock_access_denied):
sagemaker = SageMaker(aws_provider)
assert sagemaker.sagemaker_model_registries == []