@@ -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:
227230def 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