feat(sagemaker): Ensure SageMaker Endpoint Production Variants have Initial Instance Count greater than one (#5045)

Co-authored-by: Sergio <sergio@prowler.com>
This commit is contained in:
Mario Rodriguez Lopez
2024-09-19 21:16:56 +02:00
committed by GitHub
parent 0974c5f333
commit 73c96f8346
16 changed files with 671 additions and 113 deletions
@@ -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-name> --endpoint-config-name <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": ""
}
@@ -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
@@ -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] = []
@@ -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"]
@@ -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 (
@@ -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 (
@@ -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 (
@@ -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 (
@@ -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 (
@@ -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 (
@@ -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
@@ -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 (
@@ -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 (
@@ -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 (
@@ -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 (