Skip to content

Commit afe794a

Browse files
authored
Merge branch 'main' into feat/logging-config
2 parents 6e61ced + cc14571 commit afe794a

4 files changed

Lines changed: 162 additions & 38 deletions

File tree

CHANGELOG.md

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,8 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
1414

1515
### **Changed**
1616

17+
- consolidated redundant `DevStage`/`PreProdStage`/`ProdStage` classes into a single `DeployStage` in `sagemaker-templates` model deploy seed code, fixing redundant CF stack names (e.g. `dev-dev-endpoint``dev-{project}-endpoint`) and adding project uniqueness to prevent cross-project collisions
18+
1719
## v3.2.3
1820

1921
### **Added**

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

Lines changed: 5 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -80,6 +80,8 @@ def __init__(
8080
self,
8181
scope: constructs.Construct,
8282
id: str,
83+
*,
84+
stage_name: str,
8385
vpc_id: str,
8486
subnet_ids: List[str],
8587
security_group_ids: List[str],
@@ -164,7 +166,7 @@ def __init__(
164166
latest_approved_model_package = get_approved_package()
165167

166168
# Sagemaker Model
167-
model_name = f"-{id}-{timestamp}"
169+
model_name = f"-{stage_name}-{timestamp}"
168170
model_name = MODEL_PACKAGE_GROUP_NAME[: MAX_NAME_LENGTH - len(model_name)] + model_name
169171

170172
vpc_config = None
@@ -192,7 +194,7 @@ def __init__(
192194
)
193195

194196
# Sagemaker Endpoint Config
195-
endpoint_config_name = f"-{id}-ec-{timestamp}"
197+
endpoint_config_name = f"-{stage_name}-ec-{timestamp}"
196198
endpoint_config_name = (
197199
MODEL_PACKAGE_GROUP_NAME[: MAX_NAME_LENGTH - len(endpoint_config_name)] + endpoint_config_name
198200
)
@@ -261,7 +263,7 @@ def __init__(
261263
endpoint_config.add_dependency(model)
262264

263265
# Sagemaker Endpoint
264-
endpoint_name = f"-{id}-endpoint"
266+
endpoint_name = f"-{stage_name}-ep"
265267
endpoint_name = MODEL_PACKAGE_GROUP_NAME[: MAX_NAME_LENGTH - len(endpoint_name)] + endpoint_name
266268

267269
endpoint = sagemaker.CfnEndpoint(

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

Lines changed: 32 additions & 35 deletions
Original file line numberDiff line numberDiff line change
@@ -44,42 +44,27 @@
4444
}
4545

4646

47-
class DevStage(cdk.Stage):
48-
def __init__(self, scope: Construct, construct_id: str, **kwargs: Any) -> None:
49-
super().__init__(scope, construct_id, **kwargs)
50-
51-
DeployEndpointStack(
52-
self,
53-
"dev-endpoint",
54-
vpc_id=constants.DEV_VPC_ID,
55-
subnet_ids=constants.DEV_SUBNET_IDS,
56-
security_group_ids=constants.DEV_SECURITY_GROUP_IDS,
57-
)
58-
59-
60-
class PreProdStage(cdk.Stage):
61-
def __init__(self, scope: Construct, construct_id: str, **kwargs: Any) -> None:
62-
super().__init__(scope, construct_id, **kwargs)
63-
64-
DeployEndpointStack(
65-
self,
66-
"preprod-endpoint",
67-
vpc_id=constants.PRE_PROD_VPC_ID,
68-
subnet_ids=constants.PRE_PROD_SUBNET_IDS,
69-
security_group_ids=constants.PRE_PROD_SECURITY_GROUP_IDS,
70-
)
71-
72-
73-
class ProdStage(cdk.Stage):
74-
def __init__(self, scope: Construct, construct_id: str, **kwargs: Any) -> None:
47+
class DeployStage(cdk.Stage):
48+
def __init__(
49+
self,
50+
scope: Construct,
51+
construct_id: str,
52+
*,
53+
stage_name: str,
54+
vpc_id: str,
55+
subnet_ids: list[str],
56+
security_group_ids: list[str],
57+
**kwargs: Any,
58+
) -> None:
7559
super().__init__(scope, construct_id, **kwargs)
7660

7761
DeployEndpointStack(
7862
self,
79-
"prod-endpoint",
80-
vpc_id=constants.PROD_VPC_ID,
81-
subnet_ids=constants.PROD_SUBNET_IDS,
82-
security_group_ids=constants.PROD_SECURITY_GROUP_IDS,
63+
f"{constants.PROJECT_NAME}-endpoint",
64+
stage_name=stage_name,
65+
vpc_id=vpc_id,
66+
subnet_ids=subnet_ids,
67+
security_group_ids=security_group_ids,
8368
)
8469

8570

@@ -188,17 +173,25 @@ def __init__(self, scope: Construct, construct_id: str, **kwargs: Any) -> None:
188173
)
189174

190175
pipeline.add_stage(
191-
DevStage(
176+
DeployStage(
192177
self,
193178
"dev",
179+
stage_name="dev",
180+
vpc_id=constants.DEV_VPC_ID,
181+
subnet_ids=constants.DEV_SUBNET_IDS,
182+
security_group_ids=constants.DEV_SECURITY_GROUP_IDS,
194183
env=cdk.Environment(account=constants.DEV_ACCOUNT_ID, region=constants.DEV_REGION),
195184
)
196185
)
197186

198187
pipeline.add_stage(
199-
PreProdStage(
188+
DeployStage(
200189
self,
201190
"preprod",
191+
stage_name="preprod",
192+
vpc_id=constants.PRE_PROD_VPC_ID,
193+
subnet_ids=constants.PRE_PROD_SUBNET_IDS,
194+
security_group_ids=constants.PRE_PROD_SECURITY_GROUP_IDS,
202195
env=cdk.Environment(account=constants.PRE_PROD_ACCOUNT_ID, region=constants.PRE_PROD_REGION),
203196
),
204197
pre=[ManualApprovalStep("ApprovePreProd", comment="Approve deployment to Pre-Production")]
@@ -207,9 +200,13 @@ def __init__(self, scope: Construct, construct_id: str, **kwargs: Any) -> None:
207200
)
208201

209202
pipeline.add_stage(
210-
ProdStage(
203+
DeployStage(
211204
self,
212205
"prod",
206+
stage_name="prod",
207+
vpc_id=constants.PROD_VPC_ID,
208+
subnet_ids=constants.PROD_SUBNET_IDS,
209+
security_group_ids=constants.PROD_SECURITY_GROUP_IDS,
213210
env=cdk.Environment(account=constants.PROD_ACCOUNT_ID, region=constants.PROD_REGION),
214211
),
215212
pre=[ManualApprovalStep("ApproveProd", comment="Approve deployment to Production")]
Lines changed: 123 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,123 @@
1+
"""Verify that cdk synth produces correct SageMaker resource names.
2+
3+
This test mocks the SageMaker API call (get_approved_package) and all
4+
required environment variables, then synthesizes the DeployEndpointStack
5+
and asserts that stage_name drives unique SageMaker resource names.
6+
"""
7+
8+
import json
9+
import os
10+
import sys
11+
from unittest.mock import MagicMock
12+
13+
# -- Dummy environment variables (must be set before importing constants) ----
14+
_ENV = {
15+
"MODEL_BUCKET_ARN": "arn:aws:s3:::test-model-bucket",
16+
"MODEL_PACKAGE_GROUP_NAME": "test-model-group",
17+
"DEV_ACCOUNT_ID": "111111111111",
18+
"DEV_REGION": "us-east-1",
19+
"DEV_VPC_ID": "vpc-abc123",
20+
"DEV_SUBNET_IDS": json.dumps(["subnet-aaa"]),
21+
"DEV_SECURITY_GROUP_IDS": json.dumps(["sg-aaa"]),
22+
"PRE_PROD_ACCOUNT_ID": "222222222222",
23+
"PRE_PROD_REGION": "us-east-1",
24+
"PRE_PROD_VPC_ID": "vpc-def456",
25+
"PRE_PROD_SUBNET_IDS": json.dumps(["subnet-bbb"]),
26+
"PRE_PROD_SECURITY_GROUP_IDS": json.dumps(["sg-bbb"]),
27+
"PROD_ACCOUNT_ID": "333333333333",
28+
"PROD_REGION": "us-east-1",
29+
"PROD_VPC_ID": "vpc-ghi789",
30+
"PROD_SUBNET_IDS": json.dumps(["subnet-ccc"]),
31+
"PROD_SECURITY_GROUP_IDS": json.dumps(["sg-ccc"]),
32+
"PROJECT_NAME": "test-project",
33+
"PROJECT_ID": "test-id",
34+
"ENABLE_NETWORK_ISOLATION": "true",
35+
"ENABLE_DATA_CAPTURE": "true",
36+
"ENABLE_MANUAL_APPROVAL": "false",
37+
"ENABLE_EVENTBRIDGE_TRIGGER": "false",
38+
}
39+
40+
# Set env vars before any imports that read them
41+
os.environ.update(_ENV)
42+
43+
# Stub out the get_approved_package module before it tries to create a boto3 client.
44+
# The module creates a boto3.client("sagemaker") at import time, so we must
45+
# intercept it before deploy_endpoint_stack imports it.
46+
_fake_get_approved = MagicMock(return_value="arn:aws:sagemaker:us-east-1:111111111111:model-package/test-group/1")
47+
_fake_module = MagicMock()
48+
_fake_module.get_approved_package = _fake_get_approved
49+
sys.modules["deploy_app.get_approved_package"] = _fake_module
50+
51+
import aws_cdk as cdk # noqa: E402
52+
from aws_cdk import assertions # noqa: E402
53+
from deploy_app.deploy_endpoint_stack import DeployEndpointStack # noqa: E402
54+
55+
56+
def _synth_deploy_endpoint_stack(stage_name: str = "dev") -> assertions.Template:
57+
"""Synthesize a standalone DeployEndpointStack and return its Template."""
58+
# Reset the mock so each test gets a fresh call count
59+
_fake_get_approved.reset_mock()
60+
61+
# Use a Stage with a concrete account (matches real usage in pipeline_stack.py)
62+
app = cdk.App()
63+
stage = cdk.Stage(
64+
app,
65+
"TestStage",
66+
env=cdk.Environment(account="111111111111", region="us-east-1"),
67+
)
68+
stack = DeployEndpointStack(
69+
stage,
70+
"test-project-endpoint",
71+
stage_name=stage_name,
72+
vpc_id="vpc-abc123",
73+
subnet_ids=["subnet-aaa"],
74+
security_group_ids=["sg-aaa"],
75+
)
76+
77+
return assertions.Template.from_stack(stack)
78+
79+
80+
def _get_endpoint_name(template: assertions.Template) -> str:
81+
"""Extract the SageMaker endpoint name from a synthesized template."""
82+
endpoints = template.find_resources("AWS::SageMaker::Endpoint")
83+
assert len(endpoints) == 1, f"Expected 1 endpoint, found {len(endpoints)}"
84+
props = next(iter(endpoints.values()))["Properties"]
85+
return props["EndpointName"]
86+
87+
88+
def test_stage_name_produces_unique_endpoint_names():
89+
"""Different stage_name values must produce different SageMaker endpoint names."""
90+
dev_template = _synth_deploy_endpoint_stack(stage_name="dev")
91+
prod_template = _synth_deploy_endpoint_stack(stage_name="prod")
92+
93+
dev_ep = _get_endpoint_name(dev_template)
94+
prod_ep = _get_endpoint_name(prod_template)
95+
96+
assert dev_ep != prod_ep, f"dev and prod endpoint names should differ, both are '{dev_ep}'"
97+
assert "-dev-" in dev_ep, f"Expected '-dev-' in endpoint name, got '{dev_ep}'"
98+
assert "-prod-" in prod_ep, f"Expected '-prod-' in endpoint name, got '{prod_ep}'"
99+
100+
101+
def test_role_references_managed_policy():
102+
"""The ModelExecutionRole must reference the ManagedPolicy (dependency chain)."""
103+
template = _synth_deploy_endpoint_stack()
104+
105+
template.has_resource_properties(
106+
"AWS::IAM::Role",
107+
assertions.Match.object_like(
108+
{
109+
"AssumeRolePolicyDocument": {
110+
"Statement": [
111+
{
112+
"Action": "sts:AssumeRole",
113+
"Effect": "Allow",
114+
"Principal": {"Service": "sagemaker.amazonaws.com"},
115+
}
116+
],
117+
},
118+
"ManagedPolicyArns": assertions.Match.array_with(
119+
[{"Ref": assertions.Match.string_like_regexp("ModelExecutionPolicy.*")}]
120+
),
121+
}
122+
),
123+
)

0 commit comments

Comments
 (0)