feat(sagemaker): add sagemaker_domain_sso_configured check (#11094)

Co-authored-by: Daniel Barranquero <danielbo2001@gmail.com>
This commit is contained in:
June
2026-05-14 11:42:30 +02:00
committed by GitHub
co-authored by Daniel Barranquero
parent fb0ef391f2
commit 1f39b01fb2
7 changed files with 414 additions and 2 deletions
@@ -0,0 +1,234 @@
from unittest import mock
from prowler.providers.aws.services.sagemaker.sagemaker_service import Domain
from tests.providers.aws.utils import (
AWS_ACCOUNT_NUMBER,
AWS_REGION_EU_WEST_1,
set_mocked_aws_provider,
)
test_domain_name = "test-domain"
test_domain_id = "d-testdomain123"
domain_arn = f"arn:aws:sagemaker:{AWS_REGION_EU_WEST_1}:{AWS_ACCOUNT_NUMBER}:domain/{test_domain_id}"
test_sso_instance_id = "app-test-instance-id"
test_sso_application_arn = (
f"arn:aws:sso::{AWS_ACCOUNT_NUMBER}:application/sagemaker/apl-test"
)
class Test_sagemaker_domain_sso_configured:
def test_no_domains(self):
sagemaker_client = mock.MagicMock
sagemaker_client.sagemaker_domains = []
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_domain_sso_configured.sagemaker_domain_sso_configured.sagemaker_client",
sagemaker_client,
),
):
from prowler.providers.aws.services.sagemaker.sagemaker_domain_sso_configured.sagemaker_domain_sso_configured import (
sagemaker_domain_sso_configured,
)
check = sagemaker_domain_sso_configured()
result = check.execute()
assert len(result) == 0
def test_domain_sso_configured_with_instance_id(self):
sagemaker_client = mock.MagicMock
sagemaker_client.sagemaker_domains = [
Domain(
domain_id=test_domain_id,
name=test_domain_name,
arn=domain_arn,
region=AWS_REGION_EU_WEST_1,
auth_mode="SSO",
single_sign_on_managed_application_instance_id=test_sso_instance_id,
)
]
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_domain_sso_configured.sagemaker_domain_sso_configured.sagemaker_client",
sagemaker_client,
),
):
from prowler.providers.aws.services.sagemaker.sagemaker_domain_sso_configured.sagemaker_domain_sso_configured import (
sagemaker_domain_sso_configured,
)
check = sagemaker_domain_sso_configured()
result = check.execute()
assert len(result) == 1
assert result[0].status == "PASS"
assert (
result[0].status_extended
== f"SageMaker domain {test_domain_name} is configured with SSO authentication and is associated with an IAM Identity Center instance."
)
assert result[0].resource_id == test_domain_name
assert result[0].resource_arn == domain_arn
def test_domain_sso_configured_with_application_arn(self):
sagemaker_client = mock.MagicMock
sagemaker_client.sagemaker_domains = [
Domain(
domain_id=test_domain_id,
name=test_domain_name,
arn=domain_arn,
region=AWS_REGION_EU_WEST_1,
auth_mode="SSO",
single_sign_on_application_arn=test_sso_application_arn,
)
]
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_domain_sso_configured.sagemaker_domain_sso_configured.sagemaker_client",
sagemaker_client,
),
):
from prowler.providers.aws.services.sagemaker.sagemaker_domain_sso_configured.sagemaker_domain_sso_configured import (
sagemaker_domain_sso_configured,
)
check = sagemaker_domain_sso_configured()
result = check.execute()
assert len(result) == 1
assert result[0].status == "PASS"
assert (
result[0].status_extended
== f"SageMaker domain {test_domain_name} is configured with SSO authentication and is associated with an IAM Identity Center instance."
)
def test_domain_sso_without_identity_center(self):
sagemaker_client = mock.MagicMock
sagemaker_client.sagemaker_domains = [
Domain(
domain_id=test_domain_id,
name=test_domain_name,
arn=domain_arn,
region=AWS_REGION_EU_WEST_1,
auth_mode="SSO",
)
]
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_domain_sso_configured.sagemaker_domain_sso_configured.sagemaker_client",
sagemaker_client,
),
):
from prowler.providers.aws.services.sagemaker.sagemaker_domain_sso_configured.sagemaker_domain_sso_configured import (
sagemaker_domain_sso_configured,
)
check = sagemaker_domain_sso_configured()
result = check.execute()
assert len(result) == 1
assert result[0].status == "FAIL"
assert (
result[0].status_extended
== f"SageMaker domain {test_domain_name} is configured with SSO authentication but is not associated with an IAM Identity Center instance."
)
assert result[0].resource_id == test_domain_name
assert result[0].resource_arn == domain_arn
def test_domain_iam_mode(self):
sagemaker_client = mock.MagicMock
sagemaker_client.sagemaker_domains = [
Domain(
domain_id=test_domain_id,
name=test_domain_name,
arn=domain_arn,
region=AWS_REGION_EU_WEST_1,
auth_mode="IAM",
)
]
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_domain_sso_configured.sagemaker_domain_sso_configured.sagemaker_client",
sagemaker_client,
),
):
from prowler.providers.aws.services.sagemaker.sagemaker_domain_sso_configured.sagemaker_domain_sso_configured import (
sagemaker_domain_sso_configured,
)
check = sagemaker_domain_sso_configured()
result = check.execute()
assert len(result) == 1
assert result[0].status == "FAIL"
assert (
result[0].status_extended
== f"SageMaker domain {test_domain_name} is not configured with SSO authentication, current mode is IAM."
)
assert result[0].resource_id == test_domain_name
assert result[0].resource_arn == domain_arn
def test_domain_auth_mode_unknown(self):
sagemaker_client = mock.MagicMock
sagemaker_client.sagemaker_domains = [
Domain(
domain_id=test_domain_id,
name=test_domain_name,
arn=domain_arn,
region=AWS_REGION_EU_WEST_1,
)
]
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_domain_sso_configured.sagemaker_domain_sso_configured.sagemaker_client",
sagemaker_client,
),
):
from prowler.providers.aws.services.sagemaker.sagemaker_domain_sso_configured.sagemaker_domain_sso_configured import (
sagemaker_domain_sso_configured,
)
check = sagemaker_domain_sso_configured()
result = check.execute()
assert len(result) == 1
assert result[0].status == "FAIL"
assert (
result[0].status_extended
== f"SageMaker domain {test_domain_name} is not configured with SSO authentication, current mode is unknown."
)
@@ -26,6 +26,13 @@ kms_key_id = str(uuid4())
endpoint_config_name = "endpoint-config-test"
endpoint_config_arn = f"arn:aws:sagemaker:{AWS_REGION_EU_WEST_1}:{AWS_ACCOUNT_NUMBER}:endpoint-config/{endpoint_config_name}"
prod_variant_name = "Variant1"
test_domain_name = "test-domain"
test_domain_id = "d-testdomain123"
test_domain_arn = f"arn:aws:sagemaker:{AWS_REGION_EU_WEST_1}:{AWS_ACCOUNT_NUMBER}:domain/{test_domain_id}"
test_sso_instance_id = "app-test-instance-id"
test_sso_application_arn = (
f"arn:aws:sso::{AWS_ACCOUNT_NUMBER}:application/sagemaker/apl-test"
)
make_api_call = botocore.client.BaseClient._make_api_call
@@ -115,6 +122,25 @@ def mock_make_api_call(self, operation_name, kwarg):
},
]
}
if operation_name == "ListDomains":
return {
"Domains": [
{
"DomainId": test_domain_id,
"DomainName": test_domain_name,
"DomainArn": test_domain_arn,
},
],
}
if operation_name == "DescribeDomain":
return {
"DomainId": test_domain_id,
"DomainName": test_domain_name,
"DomainArn": test_domain_arn,
"AuthMode": "SSO",
"SingleSignOnManagedApplicationInstanceId": test_sso_instance_id,
"SingleSignOnApplicationArn": test_sso_application_arn,
}
return make_api_call(self, operation_name, kwarg)
@@ -249,6 +275,33 @@ class Test_SageMaker_Service:
else:
assert prod_variant.initial_instance_count == 2
# Test SageMaker list domains
def test_list_domains(self):
aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1])
sagemaker = SageMaker(aws_provider)
assert len(sagemaker.sagemaker_domains) == 1
assert sagemaker.sagemaker_domains[0].domain_id == test_domain_id
assert sagemaker.sagemaker_domains[0].name == test_domain_name
assert sagemaker.sagemaker_domains[0].arn == test_domain_arn
assert sagemaker.sagemaker_domains[0].region == AWS_REGION_EU_WEST_1
# Test SageMaker describe domain
def test_describe_domain(self):
aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1])
sagemaker = SageMaker(aws_provider)
assert len(sagemaker.sagemaker_domains) == 1
assert sagemaker.sagemaker_domains[0].auth_mode == "SSO"
assert (
sagemaker.sagemaker_domains[
0
].single_sign_on_managed_application_instance_id
== test_sso_instance_id
)
assert (
sagemaker.sagemaker_domains[0].single_sign_on_application_arn
== test_sso_application_arn
)
# Test SageMaker _list_tags_for_resource
def test_list_tags_for_resource_calls_client(self):
"""Test that _list_tags_for_resource calls the correct AWS client and updates the resource."""
@@ -312,14 +365,17 @@ class Test_SageMaker_Service:
patch(
"prowler.providers.aws.services.sagemaker.sagemaker_service.SageMaker._list_endpoint_configs"
),
patch(
"prowler.providers.aws.services.sagemaker.sagemaker_service.SageMaker._list_domains"
),
):
sagemaker_service = SageMaker(audit_info)
# Check that __threading_call__ was called for _list_tags_for_resource
# (at least 4 calls expected, one for each resource type)
# (one for each resource type: models, notebooks, training jobs, endpoint configs, domains)
tag_calls = [
c
for c in mock_threading_call.call_args_list
if c[0][0] == sagemaker_service._list_tags_for_resource
]
assert len(tag_calls) == 4
assert len(tag_calls) == 5