@@ -654,3 +654,38 @@ def test_build_leaderworkerset(
654654 self .assertIn ("spec" , lws_spec )
655655 self .assertIn ("replicas" , lws_spec ["spec" ])
656656 self .assertIn ("leaderWorkerTemplate" , lws_spec ["spec" ])
657+
658+ @parameterized .product (
659+ bundler_cls = [ArtifactRegistryBundler , CloudBuildBundler ],
660+ labels = [None , {"env" : "tpu-test" }, {"team" : "research" , "experiment" : "training" }],
661+ )
662+ def test_build_leaderworkerset_labels (
663+ self ,
664+ bundler_cls ,
665+ labels : Optional [dict ] = None ,
666+ ):
667+ """Test that labels are properly set in LeaderWorkerSet metadata."""
668+ cfg , bundler_cfg = self ._job_config (
669+ command = "test-command" ,
670+ bundler_cls = bundler_cls ,
671+ )
672+ # Set labels on the config
673+ cfg = cfg .set (labels = labels )
674+ gke_job : job .GKELeaderWorkerSet = cfg .instantiate (bundler = bundler_cfg .instantiate ())
675+ # pylint: disable-next=protected-access
676+ lws_spec = gke_job ._build_leaderworkerset ()
677+ lws_metadata = lws_spec ["metadata" ]
678+ lws_labels = lws_metadata .get ("labels" , {})
679+
680+ # Test basic metadata
681+ self .assertEqual (lws_metadata ["name" ], cfg .name )
682+
683+ # Test labels
684+ if labels is None :
685+ # When labels is None, labels dict should be empty
686+ self .assertEqual (lws_labels , {})
687+ else :
688+ # When labels is provided, they should be present in metadata
689+ for key , value in labels .items ():
690+ self .assertIn (key , lws_labels )
691+ self .assertEqual (lws_labels [key ], value )
0 commit comments