Skip to content

Commit c376aab

Browse files
authored
fix: duplicate princpals for dev nonprod and prod (#402)
1 parent 9b3608d commit c376aab

7 files changed

Lines changed: 179 additions & 35 deletions

File tree

CHANGELOG.md

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

1212
### **Changed**
1313

14+
- fixed duplicate principal error in `sagemaker-templates` KMS key and S3 bucket policies when dev, pre-prod, and prod account IDs resolve to the same AWS account
15+
1416
## v3.2.0
1517

1618
### **Added**

modules/sagemaker/sagemaker-templates/settings.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -53,6 +53,7 @@ class ModuleSettings(CdkBaseSettings):
5353
sagemaker_project_name: str
5454
sagemaker_project_id: str
5555

56+
dev_account_id: Optional[str] = Field(default=None)
5657
dev_vpc_id: Optional[str] = Field(default=None)
5758
dev_subnet_ids: List[str] = Field(default=[])
5859
dev_security_group_ids: List[str] = Field(default=[])

modules/sagemaker/sagemaker-templates/stack.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -27,6 +27,7 @@ def __init__(
2727
sagemaker_domain_arn: str,
2828
sagemaker_project_name: str,
2929
sagemaker_project_id: str,
30+
dev_account_id: str,
3031
dev_vpc_id: str,
3132
dev_subnet_ids: List[str],
3233
dev_security_group_ids: List[str],
@@ -61,6 +62,7 @@ def __init__(
6162
self,
6263
"XGBoostAbaloneProject",
6364
build_app_asset=cast(s3_assets.Asset, build_app_asset),
65+
dev_account_id=dev_account_id,
6466
pre_prod_account_id=pre_prod_account_id,
6567
prod_account_id=prod_account_id,
6668
sagemaker_domain_id=sagemaker_domain_id,
@@ -100,6 +102,7 @@ def __init__(
100102
build_app_asset=cast(s3_assets.Asset, build_app_asset),
101103
sagemaker_project_name=sagemaker_project_name,
102104
sagemaker_project_id=sagemaker_project_id,
105+
dev_account_id=dev_account_id,
103106
pre_prod_account_id=pre_prod_account_id,
104107
prod_account_id=prod_account_id,
105108
sagemaker_domain_id=sagemaker_domain_id,
@@ -118,6 +121,7 @@ def __init__(
118121
sagemaker_domain_arn=sagemaker_domain_arn,
119122
sagemaker_project_name=sagemaker_project_name,
120123
sagemaker_project_id=sagemaker_project_id,
124+
dev_account_id=dev_account_id,
121125
pre_prod_account_id=pre_prod_account_id,
122126
prod_account_id=prod_account_id,
123127
hf_access_token_secret=hf_import_models_project_settings.hf_access_token_secret,

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

Lines changed: 6 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -32,6 +32,7 @@ def __init__(
3232
build_app_asset: s3_assets.Asset,
3333
sagemaker_project_name: str,
3434
sagemaker_project_id: str,
35+
dev_account_id: str,
3536
pre_prod_account_id: str,
3637
prod_account_id: str,
3738
sagemaker_domain_id: str,
@@ -44,12 +45,14 @@ def __init__(
4445
) -> None:
4546
super().__init__(scope, id)
4647

47-
dev_account_id = Aws.ACCOUNT_ID
48+
dev_account_id = Aws.ACCOUNT_ID if not dev_account_id else dev_account_id
4849
pre_prod_account_id = Aws.ACCOUNT_ID if not pre_prod_account_id else pre_prod_account_id
4950
prod_account_id = Aws.ACCOUNT_ID if not prod_account_id else prod_account_id
5051

5152
# Deduplicate account IDs to avoid "Duplicate principal" errors in single-account deployments
5253
unique_account_ids = list(dict.fromkeys([dev_account_id, pre_prod_account_id, prod_account_id]))
54+
# Deduplicate cross-account IDs (pre-prod and prod) for KMS and S3 policies
55+
unique_cross_account_ids = list(dict.fromkeys([pre_prod_account_id, prod_account_id]))
5356

5457
Tags.of(self).add("sagemaker:project-id", sagemaker_project_id)
5558
Tags.of(self).add("sagemaker:project-name", sagemaker_project_name)
@@ -83,10 +86,7 @@ def __init__(
8386
resources=[
8487
"*",
8588
],
86-
principals=[
87-
iam.AccountPrincipal(pre_prod_account_id),
88-
iam.AccountPrincipal(prod_account_id),
89-
],
89+
principals=[iam.AccountPrincipal(account_id) for account_id in unique_cross_account_ids],
9090
),
9191
]
9292
),
@@ -126,10 +126,7 @@ def __init__(
126126
model_bucket.arn_for_objects(key_pattern="*"),
127127
model_bucket.bucket_arn,
128128
],
129-
principals=[
130-
iam.AccountPrincipal(pre_prod_account_id),
131-
iam.AccountPrincipal(prod_account_id),
132-
],
129+
principals=[iam.AccountPrincipal(account_id) for account_id in unique_cross_account_ids],
133130
)
134131
)
135132

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

Lines changed: 8 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -30,6 +30,7 @@ def __init__(
3030
sagemaker_domain_arn: str,
3131
sagemaker_project_name: str,
3232
sagemaker_project_id: str,
33+
dev_account_id: str,
3334
pre_prod_account_id: str,
3435
prod_account_id: str,
3536
hf_access_token_secret: str,
@@ -42,11 +43,14 @@ def __init__(
4243
) -> None:
4344
super().__init__(scope, construct_id)
4445

46+
dev_account_id = Aws.ACCOUNT_ID if not dev_account_id else dev_account_id
4547
pre_prod_account_id = Aws.ACCOUNT_ID if not pre_prod_account_id else pre_prod_account_id
4648
prod_account_id = Aws.ACCOUNT_ID if not prod_account_id else prod_account_id
4749

4850
# Deduplicate account IDs to avoid "Duplicate principal" errors in single-account deployments
49-
unique_account_ids = list(dict.fromkeys([pre_prod_account_id, prod_account_id]))
51+
unique_account_ids = list(dict.fromkeys([dev_account_id, pre_prod_account_id, prod_account_id]))
52+
# Deduplicate cross-account IDs (pre-prod and prod) for KMS and S3 policies
53+
unique_cross_account_ids = list(dict.fromkeys([pre_prod_account_id, prod_account_id]))
5054

5155
Tags.of(self).add("sagemaker:project-id", sagemaker_project_id)
5256
Tags.of(self).add("sagemaker:project-name", sagemaker_project_name)
@@ -80,10 +84,7 @@ def __init__(
8084
resources=[
8185
"*",
8286
],
83-
principals=[
84-
iam.AccountPrincipal(pre_prod_account_id),
85-
iam.AccountPrincipal(prod_account_id),
86-
],
87+
principals=[iam.AccountPrincipal(account_id) for account_id in unique_cross_account_ids],
8788
),
8889
]
8990
),
@@ -117,8 +118,8 @@ def __init__(
117118
s3_artifact.grant_read_write(iam.AccountRootPrincipal())
118119

119120
# PROD account access to objects in the bucket
120-
s3_artifact.grant_read_write(iam.AccountPrincipal(pre_prod_account_id))
121-
s3_artifact.grant_read_write(iam.AccountPrincipal(prod_account_id))
121+
for account_id in unique_cross_account_ids:
122+
s3_artifact.grant_read_write(iam.AccountPrincipal(account_id))
122123

123124
# cross account model registry resource policy
124125
model_package_group_name = f"{sagemaker_project_name}-{sagemaker_project_id}"

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

Lines changed: 6 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -28,6 +28,7 @@ def __init__(
2828
scope: Construct,
2929
id: str,
3030
build_app_asset: s3_assets.Asset,
31+
dev_account_id: str,
3132
pre_prod_account_id: str,
3233
prod_account_id: str,
3334
sagemaker_domain_id: str,
@@ -46,12 +47,14 @@ def __init__(
4647
) -> None:
4748
super().__init__(scope, id)
4849

49-
dev_account_id = Aws.ACCOUNT_ID
50+
dev_account_id = Aws.ACCOUNT_ID if not dev_account_id else dev_account_id
5051
pre_prod_account_id = Aws.ACCOUNT_ID if not pre_prod_account_id else pre_prod_account_id
5152
prod_account_id = Aws.ACCOUNT_ID if not prod_account_id else prod_account_id
5253

5354
# Deduplicate account IDs to avoid "Duplicate principal" errors in single-account deployments
5455
unique_account_ids = list(dict.fromkeys([dev_account_id, pre_prod_account_id, prod_account_id]))
56+
# Deduplicate cross-account IDs (pre-prod and prod) for KMS and S3 policies
57+
unique_cross_account_ids = list(dict.fromkeys([pre_prod_account_id, prod_account_id]))
5558

5659
dev_vpc = None
5760
if dev_vpc_id:
@@ -89,10 +92,7 @@ def __init__(
8992
resources=[
9093
"*",
9194
],
92-
principals=[
93-
iam.AccountPrincipal(pre_prod_account_id),
94-
iam.AccountPrincipal(prod_account_id),
95-
],
95+
principals=[iam.AccountPrincipal(account_id) for account_id in unique_cross_account_ids],
9696
),
9797
]
9898
),
@@ -132,10 +132,7 @@ def __init__(
132132
model_bucket.arn_for_objects(key_pattern="*"),
133133
model_bucket.bucket_arn,
134134
],
135-
principals=[
136-
iam.AccountPrincipal(pre_prod_account_id),
137-
iam.AccountPrincipal(prod_account_id),
138-
],
135+
principals=[iam.AccountPrincipal(account_id) for account_id in unique_cross_account_ids],
139136
)
140137
)
141138

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

Lines changed: 152 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -50,6 +50,7 @@ def stack(stack_defaults, project_template_type) -> cdk.Stack:
5050
project_name = "test-project"
5151
dep_name = "test-deployment"
5252
mod_name = "test-module"
53+
dev_account_id = "dev_account_id"
5354
dev_vpc_id = "vpc"
5455
dev_subnet_ids = ["sub"]
5556
dev_security_group_ids = ["sg"]
@@ -108,6 +109,7 @@ def stack(stack_defaults, project_template_type) -> cdk.Stack:
108109
project_template_type=project_template_type,
109110
sagemaker_project_name=sagemaker_project_name,
110111
sagemaker_project_id=sagemaker_project_id,
112+
dev_account_id=dev_account_id,
111113
dev_vpc_id=dev_vpc_id,
112114
dev_subnet_ids=dev_subnet_ids,
113115
dev_security_group_ids=dev_security_group_ids,
@@ -185,6 +187,7 @@ def stack_single_account(stack_defaults, project_template_type) -> cdk.Stack:
185187
project_template_type=project_template_type,
186188
sagemaker_project_name="test-project",
187189
sagemaker_project_id="test-project-id",
190+
dev_account_id=same_account_id,
188191
dev_vpc_id="",
189192
dev_subnet_ids=[],
190193
dev_security_group_ids=[],
@@ -227,22 +230,161 @@ def stack_single_account(stack_defaults, project_template_type) -> cdk.Stack:
227230
def test_no_duplicate_principals_in_model_package_group_policy(
228231
stack_single_account: cdk.Stack, project_template_type
229232
) -> None:
233+
"""Test that single-account deployments don't create duplicate principals.
234+
235+
When all account IDs (dev, pre-prod, prod) are the same, the policy should
236+
deduplicate them to avoid SageMaker's "Duplicate principal" validation error.
237+
238+
This test checks each policy statement individually, as SageMaker rejects
239+
policies where the same principal appears multiple times within a single statement.
240+
"""
230241
import json
231242

232243
from aws_cdk.assertions import Template
233244

245+
def stringify_principal(principal) -> str:
246+
"""Convert a principal to a comparable string representation."""
247+
if isinstance(principal, str):
248+
return principal
249+
elif isinstance(principal, dict):
250+
# Handle Fn::Join and other intrinsic functions
251+
return json.dumps(principal, sort_keys=True)
252+
return str(principal)
253+
254+
def check_statement_for_duplicates(statement: dict, logical_id: str, statement_sid: str) -> None:
255+
"""Check a single policy statement for duplicate principals."""
256+
principal = statement.get("Principal", {})
257+
if isinstance(principal, dict):
258+
aws_principals = principal.get("AWS", [])
259+
if isinstance(aws_principals, list):
260+
principal_strs = [stringify_principal(p) for p in aws_principals]
261+
assert len(principal_strs) == len(
262+
set(principal_strs)
263+
), f"Duplicate principals in {logical_id} statement '{statement_sid}': {aws_principals}"
264+
234265
template = Template.from_stack(stack_single_account)
235266
model_package_groups = template.find_resources("AWS::SageMaker::ModelPackageGroup")
236267

268+
assert model_package_groups, "Expected at least one ModelPackageGroup resource"
269+
237270
for logical_id, resource in model_package_groups.items():
238271
policy = resource.get("Properties", {}).get("ModelPackageGroupPolicy")
239-
if policy and isinstance(policy, str):
240-
policy_doc = json.loads(policy)
241-
for statement in policy_doc.get("Statement", []):
242-
principal = statement.get("Principal", {})
243-
if isinstance(principal, dict):
244-
aws_principals = principal.get("AWS", [])
245-
if isinstance(aws_principals, list):
246-
assert len(aws_principals) == len(
247-
set(aws_principals)
248-
), f"Duplicate principals in {logical_id}: {aws_principals}"
272+
if policy:
273+
if isinstance(policy, str):
274+
policy = json.loads(policy)
275+
276+
if isinstance(policy, dict):
277+
for statement in policy.get("Statement", []):
278+
statement_sid = statement.get("Sid", "unknown")
279+
check_statement_for_duplicates(statement, logical_id, statement_sid)
280+
281+
282+
@pytest.mark.parametrize(
283+
"project_template_type",
284+
[
285+
ProjectTemplateType.XGBOOST_ABALONE,
286+
ProjectTemplateType.FINETUNE_LLM_EVALUATION,
287+
ProjectTemplateType.HF_IMPORT_MODELS,
288+
],
289+
indirect=True,
290+
)
291+
def test_no_duplicate_principals_in_kms_key_policy(stack_single_account: cdk.Stack, project_template_type) -> None:
292+
"""Test that single-account deployments don't create duplicate principals in KMS key policies.
293+
294+
When pre-prod and prod account IDs are the same, the KMS key policy should
295+
deduplicate them to avoid "Duplicate principal" validation errors.
296+
"""
297+
import json
298+
299+
from aws_cdk.assertions import Template
300+
301+
def stringify_principal(principal) -> str:
302+
"""Convert a principal to a comparable string representation."""
303+
if isinstance(principal, str):
304+
return principal
305+
elif isinstance(principal, dict):
306+
return json.dumps(principal, sort_keys=True)
307+
return str(principal)
308+
309+
def check_statement_for_duplicates(statement: dict, logical_id: str, statement_sid: str) -> None:
310+
"""Check a single policy statement for duplicate principals."""
311+
principal = statement.get("Principal", {})
312+
if isinstance(principal, dict):
313+
aws_principals = principal.get("AWS", [])
314+
if isinstance(aws_principals, list):
315+
principal_strs = [stringify_principal(p) for p in aws_principals]
316+
assert len(principal_strs) == len(
317+
set(principal_strs)
318+
), f"Duplicate principals in KMS key {logical_id} statement '{statement_sid}': {aws_principals}"
319+
320+
template = Template.from_stack(stack_single_account)
321+
kms_keys = template.find_resources("AWS::KMS::Key")
322+
323+
assert kms_keys, "Expected at least one KMS Key resource"
324+
325+
for logical_id, resource in kms_keys.items():
326+
policy = resource.get("Properties", {}).get("KeyPolicy")
327+
if policy:
328+
if isinstance(policy, str):
329+
policy = json.loads(policy)
330+
331+
if isinstance(policy, dict):
332+
for statement in policy.get("Statement", []):
333+
statement_sid = statement.get("Sid", "unknown")
334+
check_statement_for_duplicates(statement, logical_id, statement_sid)
335+
336+
337+
@pytest.mark.parametrize(
338+
"project_template_type",
339+
[
340+
ProjectTemplateType.XGBOOST_ABALONE,
341+
ProjectTemplateType.FINETUNE_LLM_EVALUATION,
342+
ProjectTemplateType.HF_IMPORT_MODELS,
343+
],
344+
indirect=True,
345+
)
346+
def test_no_duplicate_principals_in_s3_bucket_policy(stack_single_account: cdk.Stack, project_template_type) -> None:
347+
"""Test that single-account deployments don't create duplicate principals in S3 bucket policies.
348+
349+
When pre-prod and prod account IDs are the same, the S3 bucket policy should
350+
deduplicate them to avoid "Duplicate principal" validation errors.
351+
"""
352+
import json
353+
354+
from aws_cdk.assertions import Template
355+
356+
def stringify_principal(principal) -> str:
357+
"""Convert a principal to a comparable string representation."""
358+
if isinstance(principal, str):
359+
return principal
360+
elif isinstance(principal, dict):
361+
return json.dumps(principal, sort_keys=True)
362+
return str(principal)
363+
364+
def check_statement_for_duplicates(statement: dict, logical_id: str, statement_sid: str) -> None:
365+
"""Check a single policy statement for duplicate principals."""
366+
principal = statement.get("Principal", {})
367+
if isinstance(principal, dict):
368+
aws_principals = principal.get("AWS", [])
369+
if isinstance(aws_principals, list):
370+
principal_strs = [stringify_principal(p) for p in aws_principals]
371+
assert len(principal_strs) == len(set(principal_strs)), (
372+
f"Duplicate principals in S3 bucket policy {logical_id} "
373+
f"statement '{statement_sid}': {aws_principals}"
374+
)
375+
376+
template = Template.from_stack(stack_single_account)
377+
bucket_policies = template.find_resources("AWS::S3::BucketPolicy")
378+
379+
assert bucket_policies, "Expected at least one S3 BucketPolicy resource"
380+
381+
for logical_id, resource in bucket_policies.items():
382+
policy = resource.get("Properties", {}).get("PolicyDocument")
383+
if policy:
384+
if isinstance(policy, str):
385+
policy = json.loads(policy)
386+
387+
if isinstance(policy, dict):
388+
for statement in policy.get("Statement", []):
389+
statement_sid = statement.get("Sid", "unknown")
390+
check_statement_for_duplicates(statement, logical_id, statement_sid)

0 commit comments

Comments
 (0)