diff --git a/sagemaker-core/src/sagemaker/core/shapes/shapes.py b/sagemaker-core/src/sagemaker/core/shapes/shapes.py index 5da1b83f50..ae67d598f4 100644 --- a/sagemaker-core/src/sagemaker/core/shapes/shapes.py +++ b/sagemaker-core/src/sagemaker/core/shapes/shapes.py @@ -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): diff --git a/sagemaker-core/src/sagemaker/core/tools/constants.py b/sagemaker-core/src/sagemaker/core/tools/constants.py index 0768664920..6c1917dc6a 100644 --- a/sagemaker-core/src/sagemaker/core/tools/constants.py +++ b/sagemaker-core/src/sagemaker/core/tools/constants.py @@ -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 diff --git a/sagemaker-core/tests/unit/generated/test_shapes.py b/sagemaker-core/tests/unit/generated/test_shapes.py index ffe5eef5e0..cfe118d651 100644 --- a/sagemaker-core/tests/unit/generated/test_shapes.py +++ b/sagemaker-core/tests/unit/generated/test_shapes.py @@ -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 @@ -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) diff --git a/sagemaker-core/tests/unit/tools/test_shapes_extractor.py b/sagemaker-core/tests/unit/tools/test_shapes_extractor.py index 4230042e0d..84cebacee7 100644 --- a/sagemaker-core/tests/unit/tools/test_shapes_extractor.py +++ b/sagemaker-core/tests/unit/tools/test_shapes_extractor.py @@ -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"