Skip to content

Commit f908371

Browse files
committed
feat: Implement manual approval and eventbridge
1 parent 09c53f8 commit f908371

8 files changed

Lines changed: 129 additions & 22 deletions

File tree

modules/sagemaker/sagemaker-templates/README.md

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -109,6 +109,8 @@ As an example, if `sagemaker-templates-service-catalog` template configured to u
109109
- `model-package-group-name` - name of the model package group (required)
110110
- `model-bucket-name` - S3 bucket name for model artifacts (required)
111111
- `enable-network-isolation` - enable network isolation for endpoints (default: false)
112+
- `enable-manual-approval` - require manual approval before Pre-Prod and Prod deployments (default: true)
113+
- `enable-eventbridge-trigger` - automatically trigger pipeline when model is approved in Model Registry (default: true)
112114

113115
#### Hugging Face Import Models Template:
114116
- `hf-access-token-secret` - AWS Secret Manager secret containing Hugging Face access token (required)

modules/sagemaker/sagemaker-templates/settings.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -100,6 +100,8 @@ class ModelDeployProjectSettings(CdkBaseSettings):
100100
model_package_group_name: str
101101
model_bucket_name: str
102102
enable_network_isolation: bool = Field(default=False)
103+
enable_manual_approval: bool = Field(default=True)
104+
enable_eventbridge_trigger: bool = Field(default=True)
103105

104106

105107
class HfImportModelsProjectSettings(CdkBaseSettings):

modules/sagemaker/sagemaker-templates/stack.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -137,6 +137,8 @@ def __init__(
137137
model_package_group_name=model_deploy_project_settings.model_package_group_name,
138138
model_bucket_name=model_deploy_project_settings.model_bucket_name,
139139
enable_network_isolation=model_deploy_project_settings.enable_network_isolation,
140+
enable_manual_approval=model_deploy_project_settings.enable_manual_approval,
141+
enable_eventbridge_trigger=model_deploy_project_settings.enable_eventbridge_trigger,
140142
dev_vpc_id=dev_vpc_id,
141143
dev_subnet_ids=dev_subnet_ids,
142144
dev_security_group_ids=dev_security_group_ids,

modules/sagemaker/sagemaker-templates/templates/model_deploy/product_stack.py

Lines changed: 23 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -77,25 +77,27 @@ def __init__(
7777
model_package_group_name: str,
7878
model_bucket_name: str,
7979
enable_network_isolation: str,
80-
dev_vpc_id: str,
81-
dev_subnet_ids: List[str],
82-
dev_security_group_ids: List[str],
83-
pre_prod_account_id: str,
84-
pre_prod_region: str,
85-
pre_prod_vpc_id: str,
86-
pre_prod_subnet_ids: List[str],
87-
pre_prod_security_group_ids: List[str],
88-
prod_account_id: str,
89-
prod_region: str,
90-
prod_vpc_id: str,
91-
prod_subnet_ids: List[str],
92-
prod_security_group_ids: List[str],
93-
sagemaker_domain_id: str,
94-
sagemaker_domain_arn: str,
95-
repository_type: RepositoryType,
96-
access_token_secret_name: Optional[str],
97-
aws_codeconnection_arn: Optional[str],
98-
repository_owner: Optional[str],
80+
enable_manual_approval: bool = True,
81+
enable_eventbridge_trigger: bool = True,
82+
dev_vpc_id: str = "",
83+
dev_subnet_ids: List[str] = [],
84+
dev_security_group_ids: List[str] = [],
85+
pre_prod_account_id: str = "",
86+
pre_prod_region: str = "",
87+
pre_prod_vpc_id: str = "",
88+
pre_prod_subnet_ids: List[str] = [],
89+
pre_prod_security_group_ids: List[str] = [],
90+
prod_account_id: str = "",
91+
prod_region: str = "",
92+
prod_vpc_id: str = "",
93+
prod_subnet_ids: List[str] = [],
94+
prod_security_group_ids: List[str] = [],
95+
sagemaker_domain_id: str = "",
96+
sagemaker_domain_arn: str = "",
97+
repository_type: RepositoryType = RepositoryType.CODECOMMIT,
98+
access_token_secret_name: Optional[str] = None,
99+
aws_codeconnection_arn: Optional[str] = None,
100+
repository_owner: Optional[str] = None,
99101
**kwargs: Any,
100102
) -> None:
101103
super().__init__(scope, id)
@@ -146,6 +148,8 @@ def __init__(
146148
"PROD_SUBNET_IDS": codebuild.BuildEnvironmentVariable(value=json.dumps(prod_subnet_ids)),
147149
"PROD_SECURITY_GROUP_IDS": codebuild.BuildEnvironmentVariable(value=json.dumps(prod_security_group_ids)),
148150
"ENABLE_NETWORK_ISOLATION": codebuild.BuildEnvironmentVariable(value=enable_network_isolation),
151+
"ENABLE_MANUAL_APPROVAL": codebuild.BuildEnvironmentVariable(value=str(enable_manual_approval).lower()),
152+
"ENABLE_EVENTBRIDGE_TRIGGER": codebuild.BuildEnvironmentVariable(value=str(enable_eventbridge_trigger).lower()),
149153
}
150154
code_pipeline_deploy_project_name = "CodePipelineDeployProject"
151155

modules/sagemaker/sagemaker-templates/templates/model_deploy/seed_code/deploy_app/config/constants.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -51,3 +51,5 @@
5151
ECR_REPO_ARN = os.getenv("ECR_REPO_ARN", None)
5252

5353
ENABLE_NETWORK_ISOLATION = os.getenv("ENABLE_NETWORK_ISOLATION", "true").lower() == "true"
54+
ENABLE_MANUAL_APPROVAL = os.getenv("ENABLE_MANUAL_APPROVAL", "true").lower() == "true"
55+
ENABLE_EVENTBRIDGE_TRIGGER = os.getenv("ENABLE_EVENTBRIDGE_TRIGGER", "true").lower() == "true"

modules/sagemaker/sagemaker-templates/templates/model_deploy/seed_code/deploy_app/deploy_app/pipeline_stack.py

Lines changed: 33 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -4,8 +4,10 @@
44
import aws_cdk as cdk
55
import config.constants as constants
66
from aws_cdk import aws_codecommit as codecommit
7+
from aws_cdk import aws_events as events
8+
from aws_cdk import aws_events_targets as targets
79
from aws_cdk import aws_iam as iam
8-
from aws_cdk.pipelines import CodeBuildStep, CodePipeline, CodePipelineSource
10+
from aws_cdk.pipelines import CodeBuildStep, CodePipeline, CodePipelineSource, ManualApprovalStep
911
from constructs import Construct
1012

1113
from .deploy_endpoint_stack import DeployEndpointStack
@@ -36,6 +38,8 @@
3638
"PROD_SUBNET_IDS": json.dumps(constants.PROD_SUBNET_IDS),
3739
"PROD_SECURITY_GROUP_IDS": json.dumps(constants.PROD_SECURITY_GROUP_IDS),
3840
"ENABLE_NETWORK_ISOLATION": str(constants.ENABLE_NETWORK_ISOLATION),
41+
"ENABLE_MANUAL_APPROVAL": str(constants.ENABLE_MANUAL_APPROVAL),
42+
"ENABLE_EVENTBRIDGE_TRIGGER": str(constants.ENABLE_EVENTBRIDGE_TRIGGER),
3943
}
4044

4145

@@ -195,13 +199,39 @@ def __init__(self, scope: Construct, construct_id: str, **kwargs: Any) -> None:
195199
self,
196200
"preprod",
197201
env=cdk.Environment(account=constants.PRE_PROD_ACCOUNT_ID, region=constants.PRE_PROD_REGION),
198-
)
202+
),
203+
pre=[ManualApprovalStep("ApprovePreProd", comment="Approve deployment to Pre-Production")]
204+
if constants.ENABLE_MANUAL_APPROVAL
205+
else None,
199206
)
200207

201208
pipeline.add_stage(
202209
ProdStage(
203210
self,
204211
"prod",
205212
env=cdk.Environment(account=constants.PROD_ACCOUNT_ID, region=constants.PROD_REGION),
206-
)
213+
),
214+
pre=[ManualApprovalStep("ApproveProd", comment="Approve deployment to Production")]
215+
if constants.ENABLE_MANUAL_APPROVAL
216+
else None,
207217
)
218+
219+
# Build the pipeline to access the underlying CodePipeline construct
220+
pipeline.build_pipeline()
221+
222+
# Add EventBridge rule to trigger pipeline when model is approved in Model Registry
223+
if constants.ENABLE_EVENTBRIDGE_TRIGGER:
224+
events.Rule(
225+
self,
226+
"ModelApprovalEventRule",
227+
rule_name=f"{constants.PROJECT_NAME}-model-approval-trigger",
228+
event_pattern=events.EventPattern(
229+
source=["aws.sagemaker"],
230+
detail_type=["SageMaker Model Package State Change"],
231+
detail={
232+
"ModelPackageGroupName": [constants.MODEL_PACKAGE_GROUP_NAME],
233+
"ModelApprovalStatus": ["Approved"],
234+
},
235+
),
236+
targets=[targets.CodePipeline(pipeline.pipeline)],
237+
)
Lines changed: 63 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,63 @@
1+
# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved.
2+
# SPDX-License-Identifier: Apache-2.0
3+
4+
import pytest
5+
6+
from settings import ModelDeployProjectSettings
7+
8+
9+
class TestModelDeployProjectSettings:
10+
"""Tests for ModelDeployProjectSettings parameter handling."""
11+
12+
def test_default_values(self) -> None:
13+
"""Test that enable_manual_approval and enable_eventbridge_trigger default to True."""
14+
settings = ModelDeployProjectSettings(
15+
model_package_group_name="test-group",
16+
model_bucket_name="test-bucket",
17+
)
18+
assert settings.enable_manual_approval is True
19+
assert settings.enable_eventbridge_trigger is True
20+
assert settings.enable_network_isolation is False
21+
22+
def test_explicit_true_values(self) -> None:
23+
"""Test that explicit True values are accepted."""
24+
settings = ModelDeployProjectSettings(
25+
model_package_group_name="test-group",
26+
model_bucket_name="test-bucket",
27+
enable_manual_approval=True,
28+
enable_eventbridge_trigger=True,
29+
)
30+
assert settings.enable_manual_approval is True
31+
assert settings.enable_eventbridge_trigger is True
32+
33+
def test_explicit_false_values(self) -> None:
34+
"""Test that explicit False values are accepted."""
35+
settings = ModelDeployProjectSettings(
36+
model_package_group_name="test-group",
37+
model_bucket_name="test-bucket",
38+
enable_manual_approval=False,
39+
enable_eventbridge_trigger=False,
40+
)
41+
assert settings.enable_manual_approval is False
42+
assert settings.enable_eventbridge_trigger is False
43+
44+
def test_required_parameters(self) -> None:
45+
"""Test that required parameters raise validation errors when missing."""
46+
with pytest.raises(ValueError) as excinfo:
47+
ModelDeployProjectSettings()
48+
49+
assert "validation error" in str(excinfo.value).lower()
50+
51+
def test_model_package_group_name_required(self) -> None:
52+
"""Test that model_package_group_name is required."""
53+
with pytest.raises(ValueError):
54+
ModelDeployProjectSettings(
55+
model_bucket_name="test-bucket",
56+
)
57+
58+
def test_model_bucket_name_required(self) -> None:
59+
"""Test that model_bucket_name is required."""
60+
with pytest.raises(ValueError):
61+
ModelDeployProjectSettings(
62+
model_package_group_name="test-group",
63+
)

modules/sagemaker/sagemaker-templates/tests/test_stack.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -87,6 +87,8 @@ def stack(stack_defaults, project_template_type) -> cdk.Stack:
8787
model_package_group_name="test-model-package-group",
8888
model_bucket_name="test-model-bucket",
8989
enable_network_isolation="false",
90+
enable_manual_approval=True,
91+
enable_eventbridge_trigger=True,
9092
)
9193
elif project_template_type == ProjectTemplateType.HF_IMPORT_MODELS:
9294
hf_import_models_settings = HfImportModelsProjectSettings(

0 commit comments

Comments
 (0)