fix(sagemaker): read DirectInternetAccess instead of RootAccess on notebook instances (#12659)

This commit is contained in:
Jonathan Nguyen
2026-09-01 18:22:41 +02:00
committed by GitHub
parent 7c84822fa3
commit ae43d21efb
5 changed files with 151 additions and 8 deletions
@@ -0,0 +1 @@
`sagemaker_notebook_instance_without_direct_internet_access_configured` check logic to read the `DirectInternetAccess` setting instead of `RootAccess`, failing a notebook instance with direct internet access enabled even when root access is disabled
@@ -3,7 +3,21 @@ from prowler.providers.aws.services.sagemaker.sagemaker_client import sagemaker_
class sagemaker_notebook_instance_without_direct_internet_access_configured(Check): class sagemaker_notebook_instance_without_direct_internet_access_configured(Check):
def execute(self): """Ensure that SageMaker notebook instances have direct internet access disabled."""
def execute(self) -> list[Check_Report_AWS]:
"""Report whether each notebook instance has direct internet access disabled.
- PASS: DirectInternetAccess was read and is Disabled.
- FAIL: DirectInternetAccess was read and is Enabled, so the instance reaches the
internet directly rather than only through the VPC.
- MANUAL: DescribeNotebookInstance did not report DirectInternetAccess. Absent is not
Disabled, and the two cannot share a value here or an instance whose setting was
never read would be reported compliant.
Returns:
One report per notebook instance in the inventory.
"""
findings = [] findings = []
for notebook_instance in sagemaker_client.sagemaker_notebook_instances: for notebook_instance in sagemaker_client.sagemaker_notebook_instances:
report = Check_Report_AWS( report = Check_Report_AWS(
@@ -11,7 +25,13 @@ class sagemaker_notebook_instance_without_direct_internet_access_configured(Chec
) )
report.status = "PASS" report.status = "PASS"
report.status_extended = f"Sagemaker notebook instance {notebook_instance.name} has direct internet access disabled." report.status_extended = f"Sagemaker notebook instance {notebook_instance.name} has direct internet access disabled."
if notebook_instance.direct_internet_access: if notebook_instance.direct_internet_access is None:
# DescribeNotebookInstance did not report DirectInternetAccess, so the
# setting is unknown. Defaulting to PASS would report an unread instance
# as compliant.
report.status = "MANUAL"
report.status_extended = f"Sagemaker notebook instance {notebook_instance.name} did not report DirectInternetAccess, so it could not be determined; verify manually."
elif notebook_instance.direct_internet_access:
report.status = "FAIL" report.status = "FAIL"
report.status_extended = f"Sagemaker notebook instance {notebook_instance.name} has direct internet access enabled." report.status_extended = f"Sagemaker notebook instance {notebook_instance.name} has direct internet access enabled."
@@ -204,6 +204,39 @@ class SageMaker(AWSService):
) )
def _describe_notebook_instance(self, notebook_instance): def _describe_notebook_instance(self, notebook_instance):
"""Read one notebook instance's settings into the inventory.
Args:
notebook_instance: The NotebookInstance to populate, identified by
name and region.
DirectInternetAccess and RootAccess are unrelated settings, and this
method previously guarded on the presence of the first while reading
the value of the second. Three consequences followed, all of them here
rather than in the check that consumes this:
- An instance with DirectInternetAccess Enabled and RootAccess Disabled
was recorded as having no direct internet access, and reported PASS.
- An instance with DirectInternetAccess Disabled and RootAccess Enabled
was recorded as having it, and reported FAIL.
- With RootAccess absent, subscripting it raised KeyError. The
enclosing ``except Exception`` swallowed that, so the assignments
below it never ran and ``kms_key_id`` and ``lifecycle_config_name``
were left None for an instance that has them -- degrading checks that
read those fields and never touch this one.
The third is reachable rather than theoretical:
DescribeNotebookInstanceOutput declares 23 members and carries no
``required`` key at all at the pinned botocore, so both fields are
optional and a response with one and not the other is legal.
DirectInternetAccess is therefore assigned in BOTH states. Setting only
the True case would leave None meaning either Disabled or never-read,
and the check has to tell those apart: it reports MANUAL for never-read
rather than defaulting to PASS, which would assert compliance from an
absent answer. Collapsing the two states here would take that
distinction away from it.
"""
logger.info("SageMaker - describing notebook instances...") logger.info("SageMaker - describing notebook instances...")
try: try:
regional_client = self.regional_clients[notebook_instance.region] regional_client = self.regional_clients[notebook_instance.region]
@@ -223,11 +256,13 @@ class SageMaker(AWSService):
notebook_instance.root_access = True notebook_instance.root_access = True
if "SubnetId" in describe_notebook_instance: if "SubnetId" in describe_notebook_instance:
notebook_instance.subnet_id = describe_notebook_instance["SubnetId"] notebook_instance.subnet_id = describe_notebook_instance["SubnetId"]
if ( if "DirectInternetAccess" in describe_notebook_instance:
"DirectInternetAccess" in describe_notebook_instance # Assign both states, not just the enabled one. Left as None, "Disabled"
and describe_notebook_instance["RootAccess"] == "Enabled" # and "the field was never read" are the same value, and the check
): # defaults to PASS -- so an unreadable notebook instance reported clean.
notebook_instance.direct_internet_access = True notebook_instance.direct_internet_access = (
describe_notebook_instance["DirectInternetAccess"] == "Enabled"
)
if "KmsKeyId" in describe_notebook_instance: if "KmsKeyId" in describe_notebook_instance:
notebook_instance.kms_key_id = describe_notebook_instance["KmsKeyId"] notebook_instance.kms_key_id = describe_notebook_instance["KmsKeyId"]
if "NotebookInstanceLifecycleConfigName" in describe_notebook_instance: if "NotebookInstanceLifecycleConfigName" in describe_notebook_instance:
@@ -553,7 +588,8 @@ class NotebookInstance(BaseModel):
arn: str arn: str
root_access: bool = None root_access: bool = None
subnet_id: str = None subnet_id: str = None
direct_internet_access: bool = None # None when DescribeNotebookInstance did not report the field: unknown, not disabled.
direct_internet_access: Optional[bool] = None
kms_key_id: str = None kms_key_id: str = None
lifecycle_config_name: str = None lifecycle_config_name: str = None
# Decoded lifecycle scripts keyed by "<hook>[<index>]" (e.g. "OnStart[0]"), # Decoded lifecycle scripts keyed by "<hook>[<index>]" (e.g. "OnStart[0]"),
@@ -38,6 +38,48 @@ class Test_sagemaker_notebook_instance_without_direct_internet_access_configured
result = check.execute() result = check.execute()
assert len(result) == 0 assert len(result) == 0
def test_instance_direct_internet_unreported_is_manual(self):
"""An unreported DirectInternetAccess must not read as disabled.
The collector only ever assigned True, so "Disabled" and "the field was never
read" were both None and the check defaulted to PASS -- reporting an instance it
had not read as compliant. Now False means disabled and None means unknown.
"""
sagemaker_client = mock.MagicMock
sagemaker_client.sagemaker_notebook_instances = []
sagemaker_client.sagemaker_notebook_instances.append(
NotebookInstance(
name=test_notebook_instance,
arn=notebook_instance_arn,
region=AWS_REGION_EU_WEST_1,
direct_internet_access=None,
)
)
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_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 (
sagemaker_notebook_instance_without_direct_internet_access_configured,
)
check = (
sagemaker_notebook_instance_without_direct_internet_access_configured()
)
result = check.execute()
assert len(result) == 1
assert result[0].status == "MANUAL"
assert "did not report DirectInternetAccess" in result[0].status_extended
def test_instance_direct_internet_disabled(self): def test_instance_direct_internet_disabled(self):
sagemaker_client = mock.MagicMock sagemaker_client = mock.MagicMock
sagemaker_client.sagemaker_notebook_instances = [] sagemaker_client.sagemaker_notebook_instances = []
@@ -262,6 +262,50 @@ class Test_SageMaker_Service:
== lifecycle_config_name == lifecycle_config_name
) )
def test_describe_notebook_instance_direct_internet_independent_of_root_access(
self,
):
"""DirectInternetAccess and RootAccess are separate settings and must be read separately.
The shared fixture sets both to "Enabled", so a collector that reads RootAccess while
testing for the DirectInternetAccess key produces the right answer by coincidence. These
two cases separate the fields, which is the only way the confusion is visible.
"""
def only_direct_internet(self, operation_name, kwarg):
"""Serve a notebook instance with internet access on and root access off.
The combination a collector reading RootAccess records as having NO direct internet
access, which is the false PASS.
"""
if operation_name == "DescribeNotebookInstance":
return {"DirectInternetAccess": "Enabled", "RootAccess": "Disabled"}
return mock_make_api_call(self, operation_name, kwarg)
def only_root_access(self, operation_name, kwarg):
"""Serve a notebook instance with internet access off and root access on.
The mirror case: a collector reading RootAccess records direct internet access on an
instance that has none, which is the false FAIL.
"""
if operation_name == "DescribeNotebookInstance":
return {"DirectInternetAccess": "Disabled", "RootAccess": "Enabled"}
return mock_make_api_call(self, operation_name, kwarg)
aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1])
with patch(
"botocore.client.BaseClient._make_api_call", new=only_direct_internet
):
notebook = SageMaker(aws_provider).sagemaker_notebook_instances[0]
assert notebook.direct_internet_access
assert not notebook.root_access
with patch("botocore.client.BaseClient._make_api_call", new=only_root_access):
notebook = SageMaker(aws_provider).sagemaker_notebook_instances[0]
assert not notebook.direct_internet_access
assert notebook.root_access
# Test SageMaker describe notebook instance lifecycle config # Test SageMaker describe notebook instance lifecycle config
def test_describe_notebook_instance_lifecycle_config(self): def test_describe_notebook_instance_lifecycle_config(self):
aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1]) aws_provider = set_mocked_aws_provider([AWS_REGION_EU_WEST_1])