mirror of
https://github.com/prowler-cloud/prowler.git
synced 2026-07-23 20:42:02 +00:00
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:
committed by
GitHub
parent
0974c5f333
commit
73c96f8346
+34
@@ -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": ""
|
||||
}
|
||||
+27
@@ -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] = []
|
||||
|
||||
+148
@@ -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"]
|
||||
+26
-4
@@ -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 (
|
||||
|
||||
+26
-4
@@ -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 (
|
||||
|
||||
+26
-4
@@ -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 (
|
||||
|
||||
+26
-4
@@ -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 (
|
||||
|
||||
+26
-4
@@ -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 (
|
||||
|
||||
+26
-4
@@ -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
|
||||
|
||||
+26
-4
@@ -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 (
|
||||
|
||||
+26
-4
@@ -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 (
|
||||
|
||||
+26
-4
@@ -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 (
|
||||
|
||||
+26
-4
@@ -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 (
|
||||
|
||||
Reference in New Issue
Block a user