Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion sagemaker-core/src/sagemaker/core/shapes/shapes.py
Original file line number Diff line number Diff line change
Expand Up @@ -10125,7 +10125,7 @@ class DataCaptureConfigSummary(Base):
capture_status: StrPipeVar
current_sampling_percentage: int
destination_s3_uri: StrPipeVar
kms_key_id: StrPipeVar
kms_key_id: Optional[StrPipeVar] = Unassigned()


class DebugRuleEvaluationStatus(Base):
Expand Down
3 changes: 3 additions & 0 deletions sagemaker-core/src/sagemaker/core/tools/constants.py
Original file line number Diff line number Diff line change
Expand Up @@ -126,6 +126,9 @@
"ModelPackageSecurityConfig": ["KmsKeyId"],
# S3Uri is optional when ModelDataSource references escrow-managed artifacts (RMP).
"S3ModelDataSource": ["S3Uri"],
# DescribeEndpoint omits DataCaptureConfig.KmsKeyId when data capture is enabled
# without a customer-managed KMS key (S3 default encryption is used instead).
"DataCaptureConfigSummary": ["KmsKeyId"],
}

# Members where the generated primitive type should be replaced with a PipelineVariable
Expand Down
56 changes: 55 additions & 1 deletion sagemaker-core/tests/unit/generated/test_shapes.py
Original file line number Diff line number Diff line change
@@ -1,10 +1,13 @@
import ast
import datetime
import unittest
from unittest.mock import MagicMock, patch

from pydantic import BaseModel, ValidationError

import os
from sagemaker.core.shapes import Base, AdditionalS3DataSource
from sagemaker.core.resources import Base as ResourceBase, Endpoint
from sagemaker.core.shapes import Base, AdditionalS3DataSource, DataCaptureConfigSummary
from sagemaker.core.utils.utils import Unassigned

# Use the installed package location
Expand Down Expand Up @@ -58,3 +61,54 @@ def _fetch_number_of_classes_in_file_not_inheriting_a_class(
if not any(base_class.id == base_class_name for base_class in node.bases):
count = count + 1
return count


class TestDataCaptureConfigSummaryOptionalKmsKeyId(unittest.TestCase):
"""DescribeEndpoint omits DataCaptureConfig.KmsKeyId when data capture is enabled
without a customer-managed KMS key (issue #5738)."""

_DESCRIBE_ENDPOINT_RESPONSE = {
"EndpointName": "my-endpoint",
"EndpointArn": "arn:aws:sagemaker:us-west-2:111122223333:endpoint/my-endpoint",
"EndpointConfigName": "my-endpoint-config",
"EndpointStatus": "InService",
"CreationTime": datetime.datetime(2026, 1, 1),
"LastModifiedTime": datetime.datetime(2026, 1, 1),
"DataCaptureConfig": {
"EnableCapture": True,
"CaptureStatus": "Started",
"CurrentSamplingPercentage": 100,
"DestinationS3Uri": "s3://my-bucket/data-capture",
},
}

def test_shape_validates_without_kms_key_id(self):
summary = DataCaptureConfigSummary(
enable_capture=True,
capture_status="Started",
current_sampling_percentage=100,
destination_s3_uri="s3://my-bucket/data-capture",
)
assert isinstance(summary.kms_key_id, Unassigned)

def test_shape_accepts_kms_key_id(self):
summary = DataCaptureConfigSummary(
enable_capture=True,
capture_status="Started",
current_sampling_percentage=100,
destination_s3_uri="s3://my-bucket/data-capture",
kms_key_id="my-kms-key",
)
assert summary.kms_key_id == "my-kms-key"

def test_endpoint_get_without_kms_key_id(self):
client = MagicMock()
client.describe_endpoint.return_value = self._DESCRIBE_ENDPOINT_RESPONSE
# Endpoint.get() resolves its client via resources.Base.get_sagemaker_client.
with patch.object(ResourceBase, "get_sagemaker_client", return_value=client):
endpoint = Endpoint.get("my-endpoint")

client.describe_endpoint.assert_called_once_with(EndpointName="my-endpoint")
assert endpoint.data_capture_config.enable_capture is True
assert endpoint.data_capture_config.destination_s3_uri == "s3://my-bucket/data-capture"
assert isinstance(endpoint.data_capture_config.kms_key_id, Unassigned)
57 changes: 57 additions & 0 deletions sagemaker-core/tests/unit/tools/test_shapes_extractor.py
Original file line number Diff line number Diff line change
Expand Up @@ -435,3 +435,60 @@ def test_override_is_targeted_not_blanket(self, extractor):
members = extractor.generate_shape_members("UnrelatedConfig")
assert "IntPipeVar" not in members["instance_count"]
assert "int" in members["instance_count"]


class TestDataCaptureConfigSummaryRequiredToOptionalOverride:
"""The service model marks DataCaptureConfigSummary.KmsKeyId as required, but
DescribeEndpoint omits it when data capture has no customer-managed KMS key
(issue #5738). REQUIRED_TO_OPTIONAL_OVERRIDES must keep codegen emitting it as
Optional, otherwise Endpoint.get() fails validation on such endpoints."""

_STRING = {"type": "string"}
_MODEL = {
"EnableCapture": {"type": "boolean"},
"CaptureStatus": _STRING,
"SamplingPercentage": {"type": "integer"},
"DestinationS3Uri": _STRING,
"KmsKeyId": _STRING,
"DataCaptureConfigSummary": {
"type": "structure",
"required": [
"EnableCapture",
"CaptureStatus",
"CurrentSamplingPercentage",
"DestinationS3Uri",
"KmsKeyId",
],
"members": {
"EnableCapture": {"shape": "EnableCapture"},
"CaptureStatus": {"shape": "CaptureStatus"},
"CurrentSamplingPercentage": {"shape": "SamplingPercentage"},
"DestinationS3Uri": {"shape": "DestinationS3Uri"},
"KmsKeyId": {"shape": "KmsKeyId"},
},
},
}

@pytest.fixture
def extractor(self, tmp_path):
# Constructing the extractor regenerates shape_dag.py; point that write
# at a temp file so the unit test leaves the checked-in file alone.
with (
patch("sagemaker.core.tools.shapes_extractor.reformat_file_with_black"),
patch(
"sagemaker.core.tools.shapes_extractor.SHAPE_DAG_FILE_PATH",
str(tmp_path / "shape_dag.py"),
),
):
return ShapesExtractor(combined_shapes=self._MODEL)

def test_kms_key_id_generated_as_optional(self, extractor):
members = extractor.generate_shape_members("DataCaptureConfigSummary")
assert members["kms_key_id"] == "Optional[StrPipeVar] = Unassigned()"

def test_other_members_stay_required(self, extractor):
members = extractor.generate_shape_members("DataCaptureConfigSummary")
assert members["enable_capture"] == "bool"
assert members["capture_status"] == "StrPipeVar"
assert members["current_sampling_percentage"] == "int"
assert members["destination_s3_uri"] == "StrPipeVar"
Loading