diff --git a/prowler/providers/aws/services/sagemaker/sagemaker_endpoint_config_prod_variant_instances/__init__.py b/prowler/providers/aws/services/sagemaker/sagemaker_endpoint_config_prod_variant_instances/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/prowler/providers/aws/services/sagemaker/sagemaker_endpoint_config_prod_variant_instances/sagemaker_endpoint_config_prod_variant_instances.metadata.json b/prowler/providers/aws/services/sagemaker/sagemaker_endpoint_config_prod_variant_instances/sagemaker_endpoint_config_prod_variant_instances.metadata.json new file mode 100644 index 0000000000..6a4622167b --- /dev/null +++ b/prowler/providers/aws/services/sagemaker/sagemaker_endpoint_config_prod_variant_instances/sagemaker_endpoint_config_prod_variant_instances.metadata.json @@ -0,0 +1,34 @@ +{ + "Provider": "aws", + "CheckID": "sagemaker_endpoint_config_prod_variant_instances", + "CheckTitle": "SageMaker endpoint production variants should have at least two initial instances", + "CheckType": [ + "Software and Configuration Checks/AWS Security Best Practices" + ], + "ServiceName": "sagemaker", + "SubServiceName": "", + "ResourceIdTemplate": "arn:aws:sagemaker:region:account-id:endpoint-config/resource-id", + "Severity": "medium", + "ResourceType": "Other", + "Description": "This control checks whether production variants of an Amazon SageMaker endpoint have an initial instance count greater than 1. A single instance creates a single point of failure and reduces availability.", + "Risk": "Having only one instance for a SageMaker endpoint production variant can lead to reduced availability, single points of failure, and slow recovery during incidents, especially if the instance becomes unavailable due to failure or security incidents.", + "RelatedUrl": "https://docs.aws.amazon.com/config/latest/developerguide/sagemaker-endpoint-config-prod-instance-count.html", + "Remediation": { + "Code": { + "CLI": "aws sagemaker update-endpoint --endpoint-name --endpoint-config-name ", + "NativeIaC": "", + "Other": "https://docs.aws.amazon.com/securityhub/latest/userguide/sagemaker-controls.html#sagemaker-4", + "Terraform": "" + }, + "Recommendation": { + "Text": "To increase the initial instance count, configure your SageMaker endpoint to use more than 1 instance in the production variant for high availability.", + "Url": "https://docs.aws.amazon.com/sagemaker/latest/dg/serverless-endpoints-create.html#serverless-endpoints-create-config" + } + }, + "Categories": [ + "redundancy" + ], + "DependsOn": [], + "RelatedTo": [], + "Notes": "" +} diff --git a/prowler/providers/aws/services/sagemaker/sagemaker_endpoint_config_prod_variant_instances/sagemaker_endpoint_config_prod_variant_instances.py b/prowler/providers/aws/services/sagemaker/sagemaker_endpoint_config_prod_variant_instances/sagemaker_endpoint_config_prod_variant_instances.py new file mode 100644 index 0000000000..43f07c6cf5 --- /dev/null +++ b/prowler/providers/aws/services/sagemaker/sagemaker_endpoint_config_prod_variant_instances/sagemaker_endpoint_config_prod_variant_instances.py @@ -0,0 +1,27 @@ +from prowler.lib.check.models import Check, Check_Report_AWS +from prowler.providers.aws.services.sagemaker.sagemaker_client import sagemaker_client + + +class sagemaker_endpoint_config_prod_variant_instances(Check): + def execute(self): + findings = [] + for endpoint_config in sagemaker_client.endpoint_configs.values(): + report = Check_Report_AWS(self.metadata()) + report.region = endpoint_config.region + report.resource_id = endpoint_config.name + report.resource_arn = endpoint_config.arn + report.resource_tags = endpoint_config.tags + report.status = "PASS" + report.status_extended = f"Sagemaker Endpoint Config {endpoint_config.name} has all production variants with more than one initial instance." + non_compliant_production_variants = [] + for production_variant in endpoint_config.production_variants: + if production_variant.initial_instance_count <= 1: + non_compliant_production_variants.append(production_variant.name) + + if non_compliant_production_variants: + report.status = "FAIL" + report.status_extended = f"Sagemaker Endpoint Config {endpoint_config.name}'s production variants {', '.join(non_compliant_production_variants)} with less than two initial instance." + + findings.append(report) + + return findings diff --git a/prowler/providers/aws/services/sagemaker/sagemaker_service.py b/prowler/providers/aws/services/sagemaker/sagemaker_service.py index 4954c18fa8..4591fa994d 100644 --- a/prowler/providers/aws/services/sagemaker/sagemaker_service.py +++ b/prowler/providers/aws/services/sagemaker/sagemaker_service.py @@ -16,12 +16,21 @@ class SageMaker(AWSService): self.sagemaker_notebook_instances = [] self.sagemaker_models = [] self.sagemaker_training_jobs = [] + self.endpoint_configs = {} self.__threading_call__(self._list_notebook_instances) self.__threading_call__(self._list_models) self.__threading_call__(self._list_training_jobs) - self._describe_model(self.regional_clients) - self._describe_notebook_instance(self.regional_clients) - self._describe_training_job(self.regional_clients) + self.__threading_call__(self._list_endpoint_configs) + self.__threading_call__(self._describe_model, self.sagemaker_models) + self.__threading_call__( + self._describe_notebook_instance, self.sagemaker_notebook_instances + ) + self.__threading_call__( + self._describe_training_job, self.sagemaker_training_jobs + ) + self.__threading_call__( + self._describe_endpoint_config, self.endpoint_configs.values() + ) self._list_tags_for_resource() def _list_notebook_instances(self, regional_client): @@ -96,93 +105,84 @@ class SageMaker(AWSService): f"{regional_client.region} -- {error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}" ) - def _describe_notebook_instance(self, regional_clients): + def _describe_notebook_instance(self, notebook_instance): logger.info("SageMaker - describing notebook instances...") try: - for notebook_instance in self.sagemaker_notebook_instances: - regional_client = regional_clients[notebook_instance.region] - try: - describe_notebook_instance = ( - regional_client.describe_notebook_instance( - NotebookInstanceName=notebook_instance.name - ) + regional_client = self.regional_clients[notebook_instance.region] + try: + describe_notebook_instance = regional_client.describe_notebook_instance( + NotebookInstanceName=notebook_instance.name + ) + except ClientError as error: + if error.response["Error"]["Code"] == "ValidationException": + logger.warning( + f"{regional_client.region} -- {error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}" ) - except ClientError as error: - if error.response["Error"]["Code"] == "ValidationException": - logger.warning( - f"{regional_client.region} -- {error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}" - ) - - continue - if ( - "RootAccess" in describe_notebook_instance - and describe_notebook_instance["RootAccess"] == "Enabled" - ): - notebook_instance.root_access = True - if "SubnetId" in describe_notebook_instance: - notebook_instance.subnet_id = describe_notebook_instance["SubnetId"] - if ( - "DirectInternetAccess" in describe_notebook_instance - and describe_notebook_instance["RootAccess"] == "Enabled" - ): - notebook_instance.direct_internet_access = True - if "KmsKeyId" in describe_notebook_instance: - notebook_instance.kms_key_id = describe_notebook_instance[ - "KmsKeyId" - ] + if ( + "RootAccess" in describe_notebook_instance + and describe_notebook_instance["RootAccess"] == "Enabled" + ): + notebook_instance.root_access = True + if "SubnetId" in describe_notebook_instance: + notebook_instance.subnet_id = describe_notebook_instance["SubnetId"] + if ( + "DirectInternetAccess" in describe_notebook_instance + and describe_notebook_instance["RootAccess"] == "Enabled" + ): + notebook_instance.direct_internet_access = True + if "KmsKeyId" in describe_notebook_instance: + notebook_instance.kms_key_id = describe_notebook_instance["KmsKeyId"] except Exception as error: logger.error( f"{regional_client.region} -- {error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}" ) - def _describe_model(self, regional_clients): + def _describe_model(self, model): logger.info("SageMaker - describing models...") try: - for model in self.sagemaker_models: - regional_client = regional_clients[model.region] - describe_model = regional_client.describe_model(ModelName=model.name) - if "EnableNetworkIsolation" in describe_model: - model.network_isolation = describe_model["EnableNetworkIsolation"] - if ( - "VpcConfig" in describe_model - and "Subnets" in describe_model["VpcConfig"] - ): - model.vpc_config_subnets = describe_model["VpcConfig"]["Subnets"] + regional_client = self.regional_clients[model.region] + describe_model = regional_client.describe_model(ModelName=model.name) + if "EnableNetworkIsolation" in describe_model: + model.network_isolation = describe_model["EnableNetworkIsolation"] + if ( + "VpcConfig" in describe_model + and "Subnets" in describe_model["VpcConfig"] + ): + model.vpc_config_subnets = describe_model["VpcConfig"]["Subnets"] except Exception as error: logger.error( f"{regional_client.region} -- {error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}" ) - def _describe_training_job(self, regional_clients): + def _describe_training_job(self, training_job): logger.info("SageMaker - describing training jobs...") try: - for training_job in self.sagemaker_training_jobs: - regional_client = regional_clients[training_job.region] - describe_training_job = regional_client.describe_training_job( - TrainingJobName=training_job.name - ) - if "EnableInterContainerTrafficEncryption" in describe_training_job: - training_job.container_traffic_encryption = describe_training_job[ - "EnableInterContainerTrafficEncryption" - ] - if ( - "ResourceConfig" in describe_training_job - and "VolumeKmsKeyId" in describe_training_job["ResourceConfig"] - ): - training_job.volume_kms_key_id = describe_training_job[ - "ResourceConfig" - ]["VolumeKmsKeyId"] - if "EnableNetworkIsolation" in describe_training_job: - training_job.network_isolation = describe_training_job[ - "EnableNetworkIsolation" - ] - if ( - "VpcConfig" in describe_training_job - and "Subnets" in describe_training_job["VpcConfig"] - ): - training_job.vpc_config_subnets = describe_training_job[ - "VpcConfig" - ]["Subnets"] + regional_client = self.regional_clients[training_job.region] + describe_training_job = regional_client.describe_training_job( + TrainingJobName=training_job.name + ) + if "EnableInterContainerTrafficEncryption" in describe_training_job: + training_job.container_traffic_encryption = describe_training_job[ + "EnableInterContainerTrafficEncryption" + ] + if ( + "ResourceConfig" in describe_training_job + and "VolumeKmsKeyId" in describe_training_job["ResourceConfig"] + ): + training_job.volume_kms_key_id = describe_training_job[ + "ResourceConfig" + ]["VolumeKmsKeyId"] + if "EnableNetworkIsolation" in describe_training_job: + training_job.network_isolation = describe_training_job[ + "EnableNetworkIsolation" + ] + if ( + "VpcConfig" in describe_training_job + and "Subnets" in describe_training_job["VpcConfig"] + ): + training_job.vpc_config_subnets = describe_training_job["VpcConfig"][ + "Subnets" + ] except Exception as error: logger.error( f"{regional_client.region} -- {error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}" @@ -217,6 +217,63 @@ class SageMaker(AWSService): logger.error( f"{regional_client.region} -- {error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}" ) + try: + for endpoint in self.endpoint_configs.values(): + regional_client = self.regional_clients[endpoint.region] + response = regional_client.list_tags(ResourceArn=endpoint.arn)["Tags"] + endpoint.tags = response + except Exception as error: + logger.error( + f"{regional_client.region} -- {error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}" + ) + + def _list_endpoint_configs(self, regional_client): + logger.info("SageMaker - listing endpoint configs...") + try: + list_endpoint_config_paginator = regional_client.get_paginator( + "list_endpoint_configs" + ) + for page in list_endpoint_config_paginator.paginate(): + for endpoint_config in page["EndpointConfigs"]: + if not self.audit_resources or ( + is_resource_filtered( + endpoint_config["EndpointConfigArn"], self.audit_resources + ) + ): + self.endpoint_configs[endpoint_config["EndpointConfigArn"]] = ( + EndpointConfig( + name=endpoint_config["EndpointConfigName"], + region=regional_client.region, + arn=endpoint_config["EndpointConfigArn"], + ) + ) + except Exception as error: + logger.error( + f"{regional_client.region} -- {error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}" + ) + + def _describe_endpoint_config(self, endpoint_config): + logger.info("SageMaker - describing endpoint configs...") + try: + regional_client = self.regional_clients[endpoint_config.region] + describe_endpoint_config = regional_client.describe_endpoint_config( + EndpointConfigName=endpoint_config.name + ) + production_variants = [] + for production_variant in describe_endpoint_config["ProductionVariants"]: + production_variants.append( + ProductionVariant( + name=production_variant["VariantName"], + initial_instance_count=production_variant[ + "InitialInstanceCount" + ], + ) + ) + endpoint_config.production_variants = production_variants + except Exception as error: + logger.error( + f"{regional_client.region} -- {error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}" + ) class NotebookInstance(BaseModel): @@ -248,3 +305,16 @@ class TrainingJob(BaseModel): network_isolation: bool = None vpc_config_subnets: list[str] = [] tags: Optional[list] = [] + + +class ProductionVariant(BaseModel): + name: str + initial_instance_count: int + + +class EndpointConfig(BaseModel): + name: str + region: str + arn: str + production_variants: list[ProductionVariant] = [] + tags: Optional[list] = [] diff --git a/tests/providers/aws/services/sagemaker/sagemaker_endpoint_config_prod_variant_instances/sagemaker_endpoint_config_prod_variant_instances_test.py b/tests/providers/aws/services/sagemaker/sagemaker_endpoint_config_prod_variant_instances/sagemaker_endpoint_config_prod_variant_instances_test.py new file mode 100644 index 0000000000..714726b50e --- /dev/null +++ b/tests/providers/aws/services/sagemaker/sagemaker_endpoint_config_prod_variant_instances/sagemaker_endpoint_config_prod_variant_instances_test.py @@ -0,0 +1,148 @@ +from unittest import mock + +from boto3 import client +from moto import mock_aws + +from tests.providers.aws.utils import ( + AWS_REGION_EU_WEST_1, + AWS_REGION_US_EAST_1, + set_mocked_aws_provider, +) + + +class Test_sagemaker_endpoint_config_prod_variant_instances: + @mock_aws + def test_no_endpoint_configs(self): + + from prowler.providers.aws.services.sagemaker.sagemaker_service import SageMaker + + aws_provider = set_mocked_aws_provider( + [AWS_REGION_EU_WEST_1, AWS_REGION_US_EAST_1] + ) + + with mock.patch( + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), mock.patch( + "prowler.providers.aws.services.sagemaker.sagemaker_endpoint_config_prod_variant_instances.sagemaker_endpoint_config_prod_variant_instances.sagemaker_client", + new=SageMaker(aws_provider), + ): + from prowler.providers.aws.services.sagemaker.sagemaker_endpoint_config_prod_variant_instances.sagemaker_endpoint_config_prod_variant_instances import ( + sagemaker_endpoint_config_prod_variant_instances, + ) + + check = sagemaker_endpoint_config_prod_variant_instances() + result = check.execute() + assert len(result) == 0 + + @mock_aws + def test_endpoint_config_non_compliant_prod_variant(self): + sagemaker_client = client("sagemaker", region_name=AWS_REGION_EU_WEST_1) + endpoint_config_name = "endpoint-config-test" + prod_variant_name = "Variant1" + prod_variant_name2 = "Variant2" + model_name = "mi-modelo-v1" + model_name2 = "mi-modelo-v2" + sagemaker_client.create_model(ModelName=model_name) + sagemaker_client.create_model(ModelName=model_name2) + endpoint_config = sagemaker_client.create_endpoint_config( + EndpointConfigName=endpoint_config_name, + ProductionVariants=[ + { + "VariantName": prod_variant_name, + "ModelName": "mi-modelo-v1", + "InitialInstanceCount": 1, + "InstanceType": "ml.m5.large", + "InitialVariantWeight": 0.6, + }, + { + "VariantName": prod_variant_name2, + "ModelName": "mi-modelo-v2", + "InitialInstanceCount": 2, + "InstanceType": "ml.m5.large", + "InitialVariantWeight": 0.4, + }, + ], + ) + + from prowler.providers.aws.services.sagemaker.sagemaker_service import SageMaker + + 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_endpoint_config_prod_variant_instances.sagemaker_endpoint_config_prod_variant_instances.sagemaker_client", + new=SageMaker(aws_provider), + ): + from prowler.providers.aws.services.sagemaker.sagemaker_endpoint_config_prod_variant_instances.sagemaker_endpoint_config_prod_variant_instances import ( + sagemaker_endpoint_config_prod_variant_instances, + ) + + check = sagemaker_endpoint_config_prod_variant_instances() + result = check.execute() + assert len(result) == 1 + assert result[0].status == "FAIL" + assert ( + result[0].status_extended + == f"Sagemaker Endpoint Config {endpoint_config_name}'s production variants {prod_variant_name} with less than two initial instance." + ) + assert result[0].resource_id == endpoint_config_name + assert result[0].resource_arn == endpoint_config["EndpointConfigArn"] + + @mock_aws + def test_endpoint_config_compliant_prod_variants(self): + sagemaker_client = client("sagemaker", region_name=AWS_REGION_EU_WEST_1) + endpoint_config_name = "endpoint-config-test" + prod_variant_name = "Variant1" + prod_variant_name2 = "Variant2" + model_name = "mi-modelo-v1" + model_name2 = "mi-modelo-v2" + sagemaker_client.create_model(ModelName=model_name) + sagemaker_client.create_model(ModelName=model_name2) + endpoint_config = sagemaker_client.create_endpoint_config( + EndpointConfigName=endpoint_config_name, + ProductionVariants=[ + { + "VariantName": prod_variant_name, + "ModelName": "mi-modelo-v1", + "InitialInstanceCount": 2, + "InstanceType": "ml.m5.large", + "InitialVariantWeight": 0.6, + }, + { + "VariantName": prod_variant_name2, + "ModelName": "mi-modelo-v2", + "InitialInstanceCount": 2, + "InstanceType": "ml.m5.large", + "InitialVariantWeight": 0.4, + }, + ], + ) + + from prowler.providers.aws.services.sagemaker.sagemaker_service import SageMaker + + 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_endpoint_config_prod_variant_instances.sagemaker_endpoint_config_prod_variant_instances.sagemaker_client", + new=SageMaker(aws_provider), + ): + from prowler.providers.aws.services.sagemaker.sagemaker_endpoint_config_prod_variant_instances.sagemaker_endpoint_config_prod_variant_instances import ( + sagemaker_endpoint_config_prod_variant_instances, + ) + + check = sagemaker_endpoint_config_prod_variant_instances() + result = check.execute() + assert len(result) == 1 + assert result[0].status == "PASS" + assert ( + result[0].status_extended + == f"Sagemaker Endpoint Config {endpoint_config_name} has all production variants with more than one initial instance." + ) + assert result[0].resource_id == endpoint_config_name + assert result[0].resource_arn == endpoint_config["EndpointConfigArn"] diff --git a/tests/providers/aws/services/sagemaker/sagemaker_models_network_isolation_enabled/sagemaker_models_network_isolation_enabled_test.py b/tests/providers/aws/services/sagemaker/sagemaker_models_network_isolation_enabled/sagemaker_models_network_isolation_enabled_test.py index bf51690923..b0a6db43be 100644 --- a/tests/providers/aws/services/sagemaker/sagemaker_models_network_isolation_enabled/sagemaker_models_network_isolation_enabled_test.py +++ b/tests/providers/aws/services/sagemaker/sagemaker_models_network_isolation_enabled/sagemaker_models_network_isolation_enabled_test.py @@ -2,7 +2,11 @@ from unittest import mock from uuid import uuid4 from prowler.providers.aws.services.sagemaker.sagemaker_service import Model -from tests.providers.aws.utils import AWS_ACCOUNT_NUMBER, AWS_REGION_EU_WEST_1 +from tests.providers.aws.utils import ( + AWS_ACCOUNT_NUMBER, + AWS_REGION_EU_WEST_1, + set_mocked_aws_provider, +) 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}" @@ -13,8 +17,14 @@ class Test_sagemaker_models_network_isolation_enabled: def test_no_models(self): sagemaker_client = mock.MagicMock sagemaker_client.sagemaker_models = [] + + aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1]) + with mock.patch( - "prowler.providers.aws.services.sagemaker.sagemaker_service.SageMaker", + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), mock.patch( + "prowler.providers.aws.services.sagemaker.sagemaker_models_network_isolation_enabled.sagemaker_models_network_isolation_enabled.sagemaker_client", sagemaker_client, ): from prowler.providers.aws.services.sagemaker.sagemaker_models_network_isolation_enabled.sagemaker_models_network_isolation_enabled import ( @@ -36,8 +46,14 @@ class Test_sagemaker_models_network_isolation_enabled: network_isolation=True, ) ) + + aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1]) + with mock.patch( - "prowler.providers.aws.services.sagemaker.sagemaker_service.SageMaker", + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), mock.patch( + "prowler.providers.aws.services.sagemaker.sagemaker_models_network_isolation_enabled.sagemaker_models_network_isolation_enabled.sagemaker_client", sagemaker_client, ): from prowler.providers.aws.services.sagemaker.sagemaker_models_network_isolation_enabled.sagemaker_models_network_isolation_enabled import ( @@ -66,8 +82,14 @@ class Test_sagemaker_models_network_isolation_enabled: network_isolation=False, ) ) + + aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1]) + with mock.patch( - "prowler.providers.aws.services.sagemaker.sagemaker_service.SageMaker", + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), mock.patch( + "prowler.providers.aws.services.sagemaker.sagemaker_models_network_isolation_enabled.sagemaker_models_network_isolation_enabled.sagemaker_client", sagemaker_client, ): from prowler.providers.aws.services.sagemaker.sagemaker_models_network_isolation_enabled.sagemaker_models_network_isolation_enabled import ( diff --git a/tests/providers/aws/services/sagemaker/sagemaker_models_vpc_settings_configured/sagemaker_models_vpc_settings_configured_test.py b/tests/providers/aws/services/sagemaker/sagemaker_models_vpc_settings_configured/sagemaker_models_vpc_settings_configured_test.py index 65ddfaea2e..42df229c0a 100644 --- a/tests/providers/aws/services/sagemaker/sagemaker_models_vpc_settings_configured/sagemaker_models_vpc_settings_configured_test.py +++ b/tests/providers/aws/services/sagemaker/sagemaker_models_vpc_settings_configured/sagemaker_models_vpc_settings_configured_test.py @@ -2,7 +2,11 @@ from unittest import mock from uuid import uuid4 from prowler.providers.aws.services.sagemaker.sagemaker_service import Model -from tests.providers.aws.utils import AWS_ACCOUNT_NUMBER, AWS_REGION_EU_WEST_1 +from tests.providers.aws.utils import ( + AWS_ACCOUNT_NUMBER, + AWS_REGION_EU_WEST_1, + set_mocked_aws_provider, +) 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}" @@ -13,8 +17,14 @@ class Test_sagemaker_models_vpc_settings_configured: def test_no_models(self): sagemaker_client = mock.MagicMock sagemaker_client.sagemaker_models = [] + + aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1]) + with mock.patch( - "prowler.providers.aws.services.sagemaker.sagemaker_service.SageMaker", + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), mock.patch( + "prowler.providers.aws.services.sagemaker.sagemaker_models_vpc_settings_configured.sagemaker_models_vpc_settings_configured.sagemaker_client", sagemaker_client, ): from prowler.providers.aws.services.sagemaker.sagemaker_models_vpc_settings_configured.sagemaker_models_vpc_settings_configured import ( @@ -36,8 +46,14 @@ class Test_sagemaker_models_vpc_settings_configured: vpc_config_subnets=[subnet_id], ) ) + + aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1]) + with mock.patch( - "prowler.providers.aws.services.sagemaker.sagemaker_service.SageMaker", + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), mock.patch( + "prowler.providers.aws.services.sagemaker.sagemaker_models_vpc_settings_configured.sagemaker_models_vpc_settings_configured.sagemaker_client", sagemaker_client, ): from prowler.providers.aws.services.sagemaker.sagemaker_models_vpc_settings_configured.sagemaker_models_vpc_settings_configured import ( @@ -65,8 +81,14 @@ class Test_sagemaker_models_vpc_settings_configured: region=AWS_REGION_EU_WEST_1, ) ) + + aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1]) + with mock.patch( - "prowler.providers.aws.services.sagemaker.sagemaker_service.SageMaker", + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), mock.patch( + "prowler.providers.aws.services.sagemaker.sagemaker_models_vpc_settings_configured.sagemaker_models_vpc_settings_configured.sagemaker_client", sagemaker_client, ): from prowler.providers.aws.services.sagemaker.sagemaker_models_vpc_settings_configured.sagemaker_models_vpc_settings_configured import ( diff --git a/tests/providers/aws/services/sagemaker/sagemaker_notebook_instance_encryption_enabled/sagemaker_notebook_instance_encryption_enabled_test.py b/tests/providers/aws/services/sagemaker/sagemaker_notebook_instance_encryption_enabled/sagemaker_notebook_instance_encryption_enabled_test.py index fe829c5ad9..5227206736 100644 --- a/tests/providers/aws/services/sagemaker/sagemaker_notebook_instance_encryption_enabled/sagemaker_notebook_instance_encryption_enabled_test.py +++ b/tests/providers/aws/services/sagemaker/sagemaker_notebook_instance_encryption_enabled/sagemaker_notebook_instance_encryption_enabled_test.py @@ -2,7 +2,11 @@ from unittest import mock from uuid import uuid4 from prowler.providers.aws.services.sagemaker.sagemaker_service import NotebookInstance -from tests.providers.aws.utils import AWS_ACCOUNT_NUMBER, AWS_REGION_EU_WEST_1 +from tests.providers.aws.utils import ( + AWS_ACCOUNT_NUMBER, + AWS_REGION_EU_WEST_1, + set_mocked_aws_provider, +) 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}" @@ -13,8 +17,14 @@ class Test_sagemaker_notebook_instance_encryption_enabled: def test_no_instances(self): sagemaker_client = mock.MagicMock sagemaker_client.sagemaker_notebook_instances = [] + + aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1]) + with mock.patch( - "prowler.providers.aws.services.sagemaker.sagemaker_service.SageMaker", + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), mock.patch( + "prowler.providers.aws.services.sagemaker.sagemaker_notebook_instance_encryption_enabled.sagemaker_notebook_instance_encryption_enabled.sagemaker_client", sagemaker_client, ): from prowler.providers.aws.services.sagemaker.sagemaker_notebook_instance_encryption_enabled.sagemaker_notebook_instance_encryption_enabled import ( @@ -36,8 +46,14 @@ class Test_sagemaker_notebook_instance_encryption_enabled: kms_key_id=kms_key, ) ) + + aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1]) + with mock.patch( - "prowler.providers.aws.services.sagemaker.sagemaker_service.SageMaker", + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), mock.patch( + "prowler.providers.aws.services.sagemaker.sagemaker_notebook_instance_encryption_enabled.sagemaker_notebook_instance_encryption_enabled.sagemaker_client", sagemaker_client, ): from prowler.providers.aws.services.sagemaker.sagemaker_notebook_instance_encryption_enabled.sagemaker_notebook_instance_encryption_enabled import ( @@ -65,8 +81,14 @@ class Test_sagemaker_notebook_instance_encryption_enabled: region=AWS_REGION_EU_WEST_1, ) ) + + aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1]) + with mock.patch( - "prowler.providers.aws.services.sagemaker.sagemaker_service.SageMaker", + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), mock.patch( + "prowler.providers.aws.services.sagemaker.sagemaker_notebook_instance_encryption_enabled.sagemaker_notebook_instance_encryption_enabled.sagemaker_client", sagemaker_client, ): from prowler.providers.aws.services.sagemaker.sagemaker_notebook_instance_encryption_enabled.sagemaker_notebook_instance_encryption_enabled import ( diff --git a/tests/providers/aws/services/sagemaker/sagemaker_notebook_instance_root_access_disabled/sagemaker_notebook_instance_root_access_disabled_test.py b/tests/providers/aws/services/sagemaker/sagemaker_notebook_instance_root_access_disabled/sagemaker_notebook_instance_root_access_disabled_test.py index be033d12e6..3b0a38382a 100644 --- a/tests/providers/aws/services/sagemaker/sagemaker_notebook_instance_root_access_disabled/sagemaker_notebook_instance_root_access_disabled_test.py +++ b/tests/providers/aws/services/sagemaker/sagemaker_notebook_instance_root_access_disabled/sagemaker_notebook_instance_root_access_disabled_test.py @@ -1,7 +1,11 @@ from unittest import mock from prowler.providers.aws.services.sagemaker.sagemaker_service import NotebookInstance -from tests.providers.aws.utils import AWS_ACCOUNT_NUMBER, AWS_REGION_EU_WEST_1 +from tests.providers.aws.utils import ( + AWS_ACCOUNT_NUMBER, + AWS_REGION_EU_WEST_1, + set_mocked_aws_provider, +) 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}" @@ -11,8 +15,14 @@ class Test_sagemaker_notebook_instance_root_access_disabled: def test_no_instances(self): sagemaker_client = mock.MagicMock sagemaker_client.sagemaker_notebook_instances = [] + + aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1]) + with mock.patch( - "prowler.providers.aws.services.sagemaker.sagemaker_service.SageMaker", + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), mock.patch( + "prowler.providers.aws.services.sagemaker.sagemaker_notebook_instance_root_access_disabled.sagemaker_notebook_instance_root_access_disabled.sagemaker_client", sagemaker_client, ): from prowler.providers.aws.services.sagemaker.sagemaker_notebook_instance_root_access_disabled.sagemaker_notebook_instance_root_access_disabled import ( @@ -34,8 +44,14 @@ class Test_sagemaker_notebook_instance_root_access_disabled: root_access=False, ) ) + + aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1]) + with mock.patch( - "prowler.providers.aws.services.sagemaker.sagemaker_service.SageMaker", + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), mock.patch( + "prowler.providers.aws.services.sagemaker.sagemaker_notebook_instance_root_access_disabled.sagemaker_notebook_instance_root_access_disabled.sagemaker_client", sagemaker_client, ): from prowler.providers.aws.services.sagemaker.sagemaker_notebook_instance_root_access_disabled.sagemaker_notebook_instance_root_access_disabled import ( @@ -64,8 +80,14 @@ class Test_sagemaker_notebook_instance_root_access_disabled: root_access=True, ) ) + + aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1]) + with mock.patch( - "prowler.providers.aws.services.sagemaker.sagemaker_service.SageMaker", + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), mock.patch( + "prowler.providers.aws.services.sagemaker.sagemaker_notebook_instance_root_access_disabled.sagemaker_notebook_instance_root_access_disabled.sagemaker_client", sagemaker_client, ): from prowler.providers.aws.services.sagemaker.sagemaker_notebook_instance_root_access_disabled.sagemaker_notebook_instance_root_access_disabled import ( diff --git a/tests/providers/aws/services/sagemaker/sagemaker_notebook_instance_vpc_settings_configured/sagemaker_notebook_instance_vpc_settings_configured_test.py b/tests/providers/aws/services/sagemaker/sagemaker_notebook_instance_vpc_settings_configured/sagemaker_notebook_instance_vpc_settings_configured_test.py index 4d0e70ff10..b613cd4c74 100644 --- a/tests/providers/aws/services/sagemaker/sagemaker_notebook_instance_vpc_settings_configured/sagemaker_notebook_instance_vpc_settings_configured_test.py +++ b/tests/providers/aws/services/sagemaker/sagemaker_notebook_instance_vpc_settings_configured/sagemaker_notebook_instance_vpc_settings_configured_test.py @@ -2,7 +2,11 @@ from unittest import mock from uuid import uuid4 from prowler.providers.aws.services.sagemaker.sagemaker_service import NotebookInstance -from tests.providers.aws.utils import AWS_ACCOUNT_NUMBER, AWS_REGION_EU_WEST_1 +from tests.providers.aws.utils import ( + AWS_ACCOUNT_NUMBER, + AWS_REGION_EU_WEST_1, + set_mocked_aws_provider, +) 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}" @@ -13,8 +17,14 @@ class Test_sagemaker_notebook_instance_vpc_settings_configured: def test_no_instances(self): sagemaker_client = mock.MagicMock sagemaker_client.sagemaker_notebook_instances = [] + + aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1]) + with mock.patch( - "prowler.providers.aws.services.sagemaker.sagemaker_service.SageMaker", + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), mock.patch( + "prowler.providers.aws.services.sagemaker.sagemaker_notebook_instance_vpc_settings_configured.sagemaker_notebook_instance_vpc_settings_configured.sagemaker_client", sagemaker_client, ): from prowler.providers.aws.services.sagemaker.sagemaker_notebook_instance_vpc_settings_configured.sagemaker_notebook_instance_vpc_settings_configured import ( @@ -36,8 +46,14 @@ class Test_sagemaker_notebook_instance_vpc_settings_configured: subnet_id=subnet_id, ) ) + + aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1]) + with mock.patch( - "prowler.providers.aws.services.sagemaker.sagemaker_service.SageMaker", + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), mock.patch( + "prowler.providers.aws.services.sagemaker.sagemaker_notebook_instance_vpc_settings_configured.sagemaker_notebook_instance_vpc_settings_configured.sagemaker_client", sagemaker_client, ): from prowler.providers.aws.services.sagemaker.sagemaker_notebook_instance_vpc_settings_configured.sagemaker_notebook_instance_vpc_settings_configured import ( @@ -66,8 +82,14 @@ class Test_sagemaker_notebook_instance_vpc_settings_configured: root_access=True, ) ) + + aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1]) + with mock.patch( - "prowler.providers.aws.services.sagemaker.sagemaker_service.SageMaker", + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), mock.patch( + "prowler.providers.aws.services.sagemaker.sagemaker_notebook_instance_vpc_settings_configured.sagemaker_notebook_instance_vpc_settings_configured.sagemaker_client", sagemaker_client, ): from prowler.providers.aws.services.sagemaker.sagemaker_notebook_instance_vpc_settings_configured.sagemaker_notebook_instance_vpc_settings_configured import ( diff --git a/tests/providers/aws/services/sagemaker/sagemaker_notebook_instance_without_direct_internet_access_configured/sagemaker_notebook_instance_without_direct_internet_access_configured_test.py b/tests/providers/aws/services/sagemaker/sagemaker_notebook_instance_without_direct_internet_access_configured/sagemaker_notebook_instance_without_direct_internet_access_configured_test.py index cbc08f5f74..8c80572443 100644 --- a/tests/providers/aws/services/sagemaker/sagemaker_notebook_instance_without_direct_internet_access_configured/sagemaker_notebook_instance_without_direct_internet_access_configured_test.py +++ b/tests/providers/aws/services/sagemaker/sagemaker_notebook_instance_without_direct_internet_access_configured/sagemaker_notebook_instance_without_direct_internet_access_configured_test.py @@ -1,7 +1,11 @@ from unittest import mock from prowler.providers.aws.services.sagemaker.sagemaker_service import NotebookInstance -from tests.providers.aws.utils import AWS_ACCOUNT_NUMBER, AWS_REGION_EU_WEST_1 +from tests.providers.aws.utils import ( + AWS_ACCOUNT_NUMBER, + AWS_REGION_EU_WEST_1, + set_mocked_aws_provider, +) 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}" @@ -11,8 +15,14 @@ class Test_sagemaker_notebook_instance_without_direct_internet_access_configured def test_no_instances(self): sagemaker_client = mock.MagicMock sagemaker_client.sagemaker_notebook_instances = [] + + aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1]) + with mock.patch( - "prowler.providers.aws.services.sagemaker.sagemaker_service.SageMaker", + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), mock.patch( + "prowler.providers.aws.services.sagemaker.sagemaker_notebook_instance_without_direct_internet_access_configured.sagemaker_notebook_instance_without_direct_internet_access_configured.sagemaker_client", sagemaker_client, ): from prowler.providers.aws.services.sagemaker.sagemaker_notebook_instance_without_direct_internet_access_configured.sagemaker_notebook_instance_without_direct_internet_access_configured import ( @@ -36,8 +46,14 @@ class Test_sagemaker_notebook_instance_without_direct_internet_access_configured direct_internet_access=False, ) ) + + aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1]) + with mock.patch( - "prowler.providers.aws.services.sagemaker.sagemaker_service.SageMaker", + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), mock.patch( + "prowler.providers.aws.services.sagemaker.sagemaker_notebook_instance_without_direct_internet_access_configured.sagemaker_notebook_instance_without_direct_internet_access_configured.sagemaker_client", sagemaker_client, ): from prowler.providers.aws.services.sagemaker.sagemaker_notebook_instance_without_direct_internet_access_configured.sagemaker_notebook_instance_without_direct_internet_access_configured import ( @@ -68,8 +84,14 @@ class Test_sagemaker_notebook_instance_without_direct_internet_access_configured direct_internet_access=True, ) ) + + aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1]) + with mock.patch( - "prowler.providers.aws.services.sagemaker.sagemaker_service.SageMaker", + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), mock.patch( + "prowler.providers.aws.services.sagemaker.sagemaker_notebook_instance_without_direct_internet_access_configured.sagemaker_notebook_instance_without_direct_internet_access_configured.sagemaker_client", sagemaker_client, ): from prowler.providers.aws.services.sagemaker.sagemaker_notebook_instance_without_direct_internet_access_configured.sagemaker_notebook_instance_without_direct_internet_access_configured import ( diff --git a/tests/providers/aws/services/sagemaker/sagemaker_service_test.py b/tests/providers/aws/services/sagemaker/sagemaker_service_test.py index d315128ae2..40fb047e5a 100644 --- a/tests/providers/aws/services/sagemaker/sagemaker_service_test.py +++ b/tests/providers/aws/services/sagemaker/sagemaker_service_test.py @@ -20,6 +20,9 @@ test_training_job = "test-training-job" test_arn_training_job = f"arn:aws:sagemaker:{AWS_REGION_EU_WEST_1}:{AWS_ACCOUNT_NUMBER}:training-job/{test_model}" subnet_id = "subnet-" + str(uuid4()) 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" make_api_call = botocore.client.BaseClient._make_api_call @@ -87,6 +90,29 @@ def mock_make_api_call(self, operation_name, kwarg): {"Key": "test", "Value": "test"}, ], } + if operation_name == "ListEndpointConfigs": + return { + "EndpointConfigs": [ + { + "EndpointConfigName": endpoint_config_name, + "EndpointConfigArn": endpoint_config_arn, + }, + ], + } + if operation_name == "DescribeEndpointConfig": + return { + "ProductionVariants": [ + { + "VariantName": prod_variant_name, + "InitialInstanceCount": 5, + }, + { + "VariantName": "Variant2", + "InitialInstanceCount": 2, + }, + ] + } + return make_api_call(self, operation_name, kwarg) @@ -186,3 +212,36 @@ class Test_SageMaker_Service: assert sagemaker.sagemaker_training_jobs[0].network_isolation assert sagemaker.sagemaker_training_jobs[0].volume_kms_key_id == kms_key_id assert sagemaker.sagemaker_training_jobs[0].vpc_config_subnets == [subnet_id] + + # Test SageMaker list endpoint configs + def test_list_endpoint_configs(self): + aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1]) + sagemaker = SageMaker(aws_provider) + assert len(sagemaker.endpoint_configs) == 1 + assert ( + sagemaker.endpoint_configs[endpoint_config_arn].name == endpoint_config_name + ) + assert ( + sagemaker.endpoint_configs[endpoint_config_arn].arn == endpoint_config_arn + ) + assert ( + sagemaker.endpoint_configs[endpoint_config_arn].region + == AWS_REGION_EU_WEST_1 + ) + assert sagemaker.sagemaker_notebook_instances[0].tags == [ + {"Key": "test", "Value": "test"}, + ] + + # Test SageMaker describe training jobs + def test_describe_endpoint_configs(self): + aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1]) + sagemaker = SageMaker(aws_provider) + assert len(sagemaker.endpoint_configs) == 1 + assert sagemaker.endpoint_configs[endpoint_config_arn].production_variants + for prod_variant in sagemaker.endpoint_configs[ + endpoint_config_arn + ].production_variants: + if prod_variant.name == prod_variant_name: + assert prod_variant.initial_instance_count == 5 + else: + assert prod_variant.initial_instance_count == 2 diff --git a/tests/providers/aws/services/sagemaker/sagemaker_training_jobs_intercontainer_encryption_enabled/sagemaker_training_jobs_intercontainer_encryption_enabled_test.py b/tests/providers/aws/services/sagemaker/sagemaker_training_jobs_intercontainer_encryption_enabled/sagemaker_training_jobs_intercontainer_encryption_enabled_test.py index 0e239540bf..80814733f9 100644 --- a/tests/providers/aws/services/sagemaker/sagemaker_training_jobs_intercontainer_encryption_enabled/sagemaker_training_jobs_intercontainer_encryption_enabled_test.py +++ b/tests/providers/aws/services/sagemaker/sagemaker_training_jobs_intercontainer_encryption_enabled/sagemaker_training_jobs_intercontainer_encryption_enabled_test.py @@ -1,7 +1,11 @@ from unittest import mock from prowler.providers.aws.services.sagemaker.sagemaker_service import TrainingJob -from tests.providers.aws.utils import AWS_ACCOUNT_NUMBER, AWS_REGION_EU_WEST_1 +from tests.providers.aws.utils import ( + AWS_ACCOUNT_NUMBER, + AWS_REGION_EU_WEST_1, + set_mocked_aws_provider, +) test_training_job = "test-training-job" training_job_arn = f"arn:aws:sagemaker:{AWS_REGION_EU_WEST_1}:{AWS_ACCOUNT_NUMBER}:training-job/{test_training_job}" @@ -11,8 +15,14 @@ class Test_sagemaker_training_jobs_intercontainer_encryption_enabled: def test_no_training_jobs(self): sagemaker_client = mock.MagicMock sagemaker_client.sagemaker_training_jobs = [] + + aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1]) + with mock.patch( - "prowler.providers.aws.services.sagemaker.sagemaker_service.SageMaker", + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), mock.patch( + "prowler.providers.aws.services.sagemaker.sagemaker_training_jobs_intercontainer_encryption_enabled.sagemaker_training_jobs_intercontainer_encryption_enabled.sagemaker_client", sagemaker_client, ): from prowler.providers.aws.services.sagemaker.sagemaker_training_jobs_intercontainer_encryption_enabled.sagemaker_training_jobs_intercontainer_encryption_enabled import ( @@ -34,8 +44,14 @@ class Test_sagemaker_training_jobs_intercontainer_encryption_enabled: container_traffic_encryption=True, ) ) + + aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1]) + with mock.patch( - "prowler.providers.aws.services.sagemaker.sagemaker_service.SageMaker", + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), mock.patch( + "prowler.providers.aws.services.sagemaker.sagemaker_training_jobs_intercontainer_encryption_enabled.sagemaker_training_jobs_intercontainer_encryption_enabled.sagemaker_client", sagemaker_client, ): from prowler.providers.aws.services.sagemaker.sagemaker_training_jobs_intercontainer_encryption_enabled.sagemaker_training_jobs_intercontainer_encryption_enabled import ( @@ -63,8 +79,14 @@ class Test_sagemaker_training_jobs_intercontainer_encryption_enabled: region=AWS_REGION_EU_WEST_1, ) ) + + aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1]) + with mock.patch( - "prowler.providers.aws.services.sagemaker.sagemaker_service.SageMaker", + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), mock.patch( + "prowler.providers.aws.services.sagemaker.sagemaker_training_jobs_intercontainer_encryption_enabled.sagemaker_training_jobs_intercontainer_encryption_enabled.sagemaker_client", sagemaker_client, ): from prowler.providers.aws.services.sagemaker.sagemaker_training_jobs_intercontainer_encryption_enabled.sagemaker_training_jobs_intercontainer_encryption_enabled import ( diff --git a/tests/providers/aws/services/sagemaker/sagemaker_training_jobs_network_isolation_enabled/sagemaker_training_jobs_network_isolation_enabled_test.py b/tests/providers/aws/services/sagemaker/sagemaker_training_jobs_network_isolation_enabled/sagemaker_training_jobs_network_isolation_enabled_test.py index f23aab5359..c0b9409347 100644 --- a/tests/providers/aws/services/sagemaker/sagemaker_training_jobs_network_isolation_enabled/sagemaker_training_jobs_network_isolation_enabled_test.py +++ b/tests/providers/aws/services/sagemaker/sagemaker_training_jobs_network_isolation_enabled/sagemaker_training_jobs_network_isolation_enabled_test.py @@ -2,7 +2,11 @@ from unittest import mock from uuid import uuid4 from prowler.providers.aws.services.sagemaker.sagemaker_service import TrainingJob -from tests.providers.aws.utils import AWS_ACCOUNT_NUMBER, AWS_REGION_EU_WEST_1 +from tests.providers.aws.utils import ( + AWS_ACCOUNT_NUMBER, + AWS_REGION_EU_WEST_1, + set_mocked_aws_provider, +) test_training_job = "test-training-job" training_job_arn = f"arn:aws:sagemaker:{AWS_REGION_EU_WEST_1}:{AWS_ACCOUNT_NUMBER}:training-job/{test_training_job}" @@ -13,8 +17,14 @@ class Test_sagemaker_training_jobs_network_isolation_enabled: def test_no_training_jobs(self): sagemaker_client = mock.MagicMock sagemaker_client.sagemaker_training_jobs = [] + + aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1]) + with mock.patch( - "prowler.providers.aws.services.sagemaker.sagemaker_service.SageMaker", + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), mock.patch( + "prowler.providers.aws.services.sagemaker.sagemaker_training_jobs_network_isolation_enabled.sagemaker_training_jobs_network_isolation_enabled.sagemaker_client", sagemaker_client, ): from prowler.providers.aws.services.sagemaker.sagemaker_training_jobs_network_isolation_enabled.sagemaker_training_jobs_network_isolation_enabled import ( @@ -36,8 +46,14 @@ class Test_sagemaker_training_jobs_network_isolation_enabled: network_isolation=True, ) ) + + aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1]) + with mock.patch( - "prowler.providers.aws.services.sagemaker.sagemaker_service.SageMaker", + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), mock.patch( + "prowler.providers.aws.services.sagemaker.sagemaker_training_jobs_network_isolation_enabled.sagemaker_training_jobs_network_isolation_enabled.sagemaker_client", sagemaker_client, ): from prowler.providers.aws.services.sagemaker.sagemaker_training_jobs_network_isolation_enabled.sagemaker_training_jobs_network_isolation_enabled import ( @@ -65,8 +81,14 @@ class Test_sagemaker_training_jobs_network_isolation_enabled: region=AWS_REGION_EU_WEST_1, ) ) + + aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1]) + with mock.patch( - "prowler.providers.aws.services.sagemaker.sagemaker_service.SageMaker", + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), mock.patch( + "prowler.providers.aws.services.sagemaker.sagemaker_training_jobs_network_isolation_enabled.sagemaker_training_jobs_network_isolation_enabled.sagemaker_client", sagemaker_client, ): from prowler.providers.aws.services.sagemaker.sagemaker_training_jobs_network_isolation_enabled.sagemaker_training_jobs_network_isolation_enabled import ( diff --git a/tests/providers/aws/services/sagemaker/sagemaker_training_jobs_volume_and_output_encryption_enabled/sagemaker_training_jobs_volume_and_output_encryption_enabled_test.py b/tests/providers/aws/services/sagemaker/sagemaker_training_jobs_volume_and_output_encryption_enabled/sagemaker_training_jobs_volume_and_output_encryption_enabled_test.py index 33fca077e0..365efa0960 100644 --- a/tests/providers/aws/services/sagemaker/sagemaker_training_jobs_volume_and_output_encryption_enabled/sagemaker_training_jobs_volume_and_output_encryption_enabled_test.py +++ b/tests/providers/aws/services/sagemaker/sagemaker_training_jobs_volume_and_output_encryption_enabled/sagemaker_training_jobs_volume_and_output_encryption_enabled_test.py @@ -2,7 +2,11 @@ from unittest import mock from uuid import uuid4 from prowler.providers.aws.services.sagemaker.sagemaker_service import TrainingJob -from tests.providers.aws.utils import AWS_ACCOUNT_NUMBER, AWS_REGION_EU_WEST_1 +from tests.providers.aws.utils import ( + AWS_ACCOUNT_NUMBER, + AWS_REGION_EU_WEST_1, + set_mocked_aws_provider, +) test_training_job = "test-training-job" training_job_arn = f"arn:aws:sagemaker:{AWS_REGION_EU_WEST_1}:{AWS_ACCOUNT_NUMBER}:training-job/{test_training_job}" @@ -13,8 +17,14 @@ class Test_sagemaker_training_jobs_volume_and_output_encryption_enabled: def test_no_training_jobs(self): sagemaker_client = mock.MagicMock sagemaker_client.sagemaker_training_jobs = [] + + aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1]) + with mock.patch( - "prowler.providers.aws.services.sagemaker.sagemaker_service.SageMaker", + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), mock.patch( + "prowler.providers.aws.services.sagemaker.sagemaker_training_jobs_volume_and_output_encryption_enabled.sagemaker_training_jobs_volume_and_output_encryption_enabled.sagemaker_client", sagemaker_client, ): from prowler.providers.aws.services.sagemaker.sagemaker_training_jobs_volume_and_output_encryption_enabled.sagemaker_training_jobs_volume_and_output_encryption_enabled import ( @@ -36,8 +46,14 @@ class Test_sagemaker_training_jobs_volume_and_output_encryption_enabled: volume_kms_key_id=kms_key_id, ) ) + + aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1]) + with mock.patch( - "prowler.providers.aws.services.sagemaker.sagemaker_service.SageMaker", + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), mock.patch( + "prowler.providers.aws.services.sagemaker.sagemaker_training_jobs_volume_and_output_encryption_enabled.sagemaker_training_jobs_volume_and_output_encryption_enabled.sagemaker_client", sagemaker_client, ): from prowler.providers.aws.services.sagemaker.sagemaker_training_jobs_volume_and_output_encryption_enabled.sagemaker_training_jobs_volume_and_output_encryption_enabled import ( @@ -65,8 +81,14 @@ class Test_sagemaker_training_jobs_volume_and_output_encryption_enabled: region=AWS_REGION_EU_WEST_1, ) ) + + aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1]) + with mock.patch( - "prowler.providers.aws.services.sagemaker.sagemaker_service.SageMaker", + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), mock.patch( + "prowler.providers.aws.services.sagemaker.sagemaker_training_jobs_volume_and_output_encryption_enabled.sagemaker_training_jobs_volume_and_output_encryption_enabled.sagemaker_client", sagemaker_client, ): from prowler.providers.aws.services.sagemaker.sagemaker_training_jobs_volume_and_output_encryption_enabled.sagemaker_training_jobs_volume_and_output_encryption_enabled import ( diff --git a/tests/providers/aws/services/sagemaker/sagemaker_training_jobs_vpc_settings_configured/sagemaker_training_jobs_vpc_settings_configured_test.py b/tests/providers/aws/services/sagemaker/sagemaker_training_jobs_vpc_settings_configured/sagemaker_training_jobs_vpc_settings_configured_test.py index 0e4c5a674b..c4a75a3f5b 100644 --- a/tests/providers/aws/services/sagemaker/sagemaker_training_jobs_vpc_settings_configured/sagemaker_training_jobs_vpc_settings_configured_test.py +++ b/tests/providers/aws/services/sagemaker/sagemaker_training_jobs_vpc_settings_configured/sagemaker_training_jobs_vpc_settings_configured_test.py @@ -2,7 +2,11 @@ from unittest import mock from uuid import uuid4 from prowler.providers.aws.services.sagemaker.sagemaker_service import TrainingJob -from tests.providers.aws.utils import AWS_ACCOUNT_NUMBER, AWS_REGION_EU_WEST_1 +from tests.providers.aws.utils import ( + AWS_ACCOUNT_NUMBER, + AWS_REGION_EU_WEST_1, + set_mocked_aws_provider, +) test_training_job = "test-training-job" training_job_arn = f"arn:aws:sagemaker:{AWS_REGION_EU_WEST_1}:{AWS_ACCOUNT_NUMBER}:training-job/{test_training_job}" @@ -13,8 +17,14 @@ class Test_sagemaker_training_jobs_vpc_settings_configured: def test_no_training_jobs(self): sagemaker_client = mock.MagicMock sagemaker_client.sagemaker_training_jobs = [] + + aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1]) + with mock.patch( - "prowler.providers.aws.services.sagemaker.sagemaker_service.SageMaker", + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), mock.patch( + "prowler.providers.aws.services.sagemaker.sagemaker_training_jobs_vpc_settings_configured.sagemaker_training_jobs_vpc_settings_configured.sagemaker_client", sagemaker_client, ): from prowler.providers.aws.services.sagemaker.sagemaker_training_jobs_vpc_settings_configured.sagemaker_training_jobs_vpc_settings_configured import ( @@ -36,8 +46,14 @@ class Test_sagemaker_training_jobs_vpc_settings_configured: vpc_config_subnets=[subnet_id], ) ) + + aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1]) + with mock.patch( - "prowler.providers.aws.services.sagemaker.sagemaker_service.SageMaker", + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), mock.patch( + "prowler.providers.aws.services.sagemaker.sagemaker_training_jobs_vpc_settings_configured.sagemaker_training_jobs_vpc_settings_configured.sagemaker_client", sagemaker_client, ): from prowler.providers.aws.services.sagemaker.sagemaker_training_jobs_vpc_settings_configured.sagemaker_training_jobs_vpc_settings_configured import ( @@ -65,8 +81,14 @@ class Test_sagemaker_training_jobs_vpc_settings_configured: region=AWS_REGION_EU_WEST_1, ) ) + + aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1]) + with mock.patch( - "prowler.providers.aws.services.sagemaker.sagemaker_service.SageMaker", + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), mock.patch( + "prowler.providers.aws.services.sagemaker.sagemaker_training_jobs_vpc_settings_configured.sagemaker_training_jobs_vpc_settings_configured.sagemaker_client", sagemaker_client, ): from prowler.providers.aws.services.sagemaker.sagemaker_training_jobs_vpc_settings_configured.sagemaker_training_jobs_vpc_settings_configured import (