mirror of
https://github.com/prowler-cloud/prowler.git
synced 2026-10-09 21:14:22 +00:00
feat(sagemaker): add sagemaker_domain_sso_configured check (#11094)
Co-authored-by: Daniel Barranquero <danielbo2001@gmail.com>
This commit is contained in:
co-authored by
Daniel Barranquero
parent
fb0ef391f2
commit
1f39b01fb2
+234
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user