From ced1c77abdb9c3ae913f98e4532d3fb1068fee6c Mon Sep 17 00:00:00 2001 From: Akanksha Gupta Date: Wed, 30 Sep 2026 15:42:43 -0700 Subject: [PATCH] Skip fetching GKE cluster credentials if the active kube context matches for the proxy job PiperOrigin-RevId: 991261962 --- .../shared_pathways_service/isc_pathways.py | 41 +++++- .../isc_pathways_test.py | 136 +++++++++++++++++- 2 files changed, 174 insertions(+), 3 deletions(-) diff --git a/pathwaysutils/experimental/shared_pathways_service/isc_pathways.py b/pathwaysutils/experimental/shared_pathways_service/isc_pathways.py index 90f10dc..6c8411d 100644 --- a/pathwaysutils/experimental/shared_pathways_service/isc_pathways.py +++ b/pathwaysutils/experimental/shared_pathways_service/isc_pathways.py @@ -534,6 +534,43 @@ def _get_username() -> str: return username or "user" +def _is_current_kube_context( + *, cluster: str, project: str, location: str +) -> bool: + """Checks whether the active kube config context points at the cluster. + + Args: + cluster: The name of the GKE cluster. + project: The GCP project ID. + location: The GCP region or zone of the cluster. + + Returns: + True if the current kube config context already targets the given cluster. + """ + return gke_utils.get_current_kube_context() == (cluster, project, location) + + +def _ensure_cluster_credentials( + *, cluster: str, project: str, location: str +) -> None: + """Fetches the GKE cluster credentials unless kube config already has them.""" + if _is_current_kube_context( + cluster=cluster, project=project, location=location + ): + _logger.info( + "The current kube config context already points to cluster '%s' in" + " project '%s' and location '%s'. Skipping credential fetch.", + cluster, + project, + location, + ) + return + + gke_utils.fetch_cluster_credentials( + cluster_name=cluster, project_id=project, location=location + ) + + @contextlib.contextmanager def connect( *, @@ -584,8 +621,8 @@ def connect( validators.validate_pathways_service(pathways_service) validators.validate_tpu_instances(expected_tpu_instances) validators.validate_proxy_options(proxy_options) - gke_utils.fetch_cluster_credentials( - cluster_name=cluster, project_id=project, location=region + _ensure_cluster_credentials( + cluster=cluster, project=project, location=region ) server_image, sidecar_image = gke_utils.get_pathways_service_images( diff --git a/pathwaysutils/test/experimental/shared_pathways_service/isc_pathways_test.py b/pathwaysutils/test/experimental/shared_pathways_service/isc_pathways_test.py index 42d98d7..9eb7620 100644 --- a/pathwaysutils/test/experimental/shared_pathways_service/isc_pathways_test.py +++ b/pathwaysutils/test/experimental/shared_pathways_service/isc_pathways_test.py @@ -12,7 +12,6 @@ from pathwaysutils.experimental.shared_pathways_service import isc_pathways - class ISCPathwaysTest(parameterized.TestCase): """Tests for the ISCPathways class.""" @@ -23,6 +22,14 @@ def setUp(self): isc_pathways.gke_utils, "stream_pod_logs", autospec=True ) ) + # By default, pretend the kube config does not already point at the + # cluster so that credentials are fetched. + self.mock_is_current_kube_context = self.enter_context( + mock.patch.object( + isc_pathways, "_is_current_kube_context", autospec=True + ) + ) + self.mock_is_current_kube_context.return_value = False def test_wait_for_placement_success(self): """Tests that _wait_for_placement correctly processes logs.""" @@ -581,6 +588,49 @@ def test_isc_pathways(self): text=True, ) + def test_connect_skips_credentials_when_kube_config_matches(self): + """Tests that connect skips fetching credentials for the active context.""" + self.mock_is_current_kube_context.return_value = True + mock_fetch_creds = self.enter_context( + mock.patch.object( + isc_pathways.gke_utils, "fetch_cluster_credentials", autospec=True + ) + ) + mock_get_images = self.enter_context( + mock.patch.object( + isc_pathways.gke_utils, "get_pathways_service_images", autospec=True + ) + ) + mock_get_images.return_value = ( + "us-docker.pkg.dev/cloud-tpu-v2-images/pathways/server:latest", + None, + ) + mock_isc_pathways = self.enter_context( + mock.patch.object(isc_pathways, "_ISCPathways", autospec=True) + ) + self.enter_context(mock.patch("threading.Thread", autospec=True)) + + mock_manager_instance = ( + mock_isc_pathways.return_value.__enter__.return_value + ) + mock_manager_instance.proxy_pod_name = "test-pod-123" + mock_manager_instance.expected_tpu_instances = {"tpuv6e:2x2": 1} + + with isc_pathways.connect( + cluster="test-cluster", + project="test-project", + region="test-region", + gcs_bucket="test-bucket", + pathways_service="test-service-pathways-head:1234", + expected_tpu_instances={"tpuv6e:2x2": 1}, + ): + pass + + self.mock_is_current_kube_context.assert_called_once_with( + cluster="test-cluster", project="test-project", location="test-region" + ) + mock_fetch_creds.assert_not_called() + def test_connect_success(self): """Tests that connect calls the dependencies and yields the manager.""" # Arrange @@ -1369,5 +1419,89 @@ def test_connect_proxy_server_image_deprecation_warning(self): pass +class KubeConfigCredentialsTest(parameterized.TestCase): + """Tests for the kube config context helpers.""" + + def setUp(self): + super().setUp() + self.mock_get_current_kube_context = self.enter_context( + mock.patch.object( + isc_pathways.gke_utils, "get_current_kube_context", autospec=True + ) + ) + self.mock_fetch_creds = self.enter_context( + mock.patch.object( + isc_pathways.gke_utils, "fetch_cluster_credentials", autospec=True + ) + ) + + def test_is_current_kube_context_matching(self): + """Tests that a matching active context is detected.""" + self.mock_get_current_kube_context.return_value = ( + "test-cluster", + "test-project", + "test-region", + ) + + self.assertTrue( + isc_pathways._is_current_kube_context( + cluster="test-cluster", + project="test-project", + location="test-region", + ) + ) + + @parameterized.named_parameters( + ("different_cluster", ("other-cluster", "test-project", "test-region")), + ("different_project", ("test-cluster", "other-project", "test-region")), + ("different_location", ("test-cluster", "test-project", "other-region")), + ("non_gke_context", ("minikube", None, None)), + ("no_context", (None, None, None)), + ) + def test_is_current_kube_context_not_matching(self, current_context): + """Tests that a non-matching active context is detected.""" + self.mock_get_current_kube_context.return_value = current_context + + self.assertFalse( + isc_pathways._is_current_kube_context( + cluster="test-cluster", + project="test-project", + location="test-region", + ) + ) + + def test_ensure_cluster_credentials_skips_fetch_when_matching(self): + """Tests that credentials are not fetched for the active context.""" + self.mock_get_current_kube_context.return_value = ( + "test-cluster", + "test-project", + "test-region", + ) + + isc_pathways._ensure_cluster_credentials( + cluster="test-cluster", project="test-project", location="test-region" + ) + + self.mock_fetch_creds.assert_not_called() + + def test_ensure_cluster_credentials_fetches_when_not_matching(self): + """Tests that credentials are fetched for a different context.""" + self.mock_get_current_kube_context.return_value = ( + "test-cluster", + "other-project", + "test-region", + ) + + isc_pathways._ensure_cluster_credentials( + cluster="test-cluster", project="test-project", location="test-region" + ) + + self.mock_fetch_creds.assert_called_once_with( + cluster_name="test-cluster", + project_id="test-project", + location="test-region", + ) + + if __name__ == "__main__": absltest.main()