Skip to content

Commit a9cda72

Browse files
guptaakacopybara-github
authored andcommitted
Skip fetching GKE cluster credentials if the active kube context matches for the proxy job
PiperOrigin-RevId: 980079122
1 parent c572299 commit a9cda72

2 files changed

Lines changed: 174 additions & 3 deletions

File tree

‎pathwaysutils/experimental/shared_pathways_service/isc_pathways.py‎

Lines changed: 39 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -534,6 +534,43 @@ def _get_username() -> str:
534534
return username or "user"
535535

536536

537+
def _is_current_kube_context(
538+
*, cluster: str, project: str, location: str
539+
) -> bool:
540+
"""Checks whether the active kube config context points at the cluster.
541+
542+
Args:
543+
cluster: The name of the GKE cluster.
544+
project: The GCP project ID.
545+
location: The GCP region or zone of the cluster.
546+
547+
Returns:
548+
True if the current kube config context already targets the given cluster.
549+
"""
550+
return gke_utils.get_current_kube_context() == (cluster, project, location)
551+
552+
553+
def _ensure_cluster_credentials(
554+
*, cluster: str, project: str, location: str
555+
) -> None:
556+
"""Fetches the GKE cluster credentials unless kube config already has them."""
557+
if _is_current_kube_context(
558+
cluster=cluster, project=project, location=location
559+
):
560+
_logger.info(
561+
"The current kube config context already points to cluster '%s' in"
562+
" project '%s' and location '%s'. Skipping credential fetch.",
563+
cluster,
564+
project,
565+
location,
566+
)
567+
return
568+
569+
gke_utils.fetch_cluster_credentials(
570+
cluster_name=cluster, project_id=project, location=location
571+
)
572+
573+
537574
@contextlib.contextmanager
538575
def connect(
539576
*,
@@ -584,8 +621,8 @@ def connect(
584621
validators.validate_pathways_service(pathways_service)
585622
validators.validate_tpu_instances(expected_tpu_instances)
586623
validators.validate_proxy_options(proxy_options)
587-
gke_utils.fetch_cluster_credentials(
588-
cluster_name=cluster, project_id=project, location=region
624+
_ensure_cluster_credentials(
625+
cluster=cluster, project=project, location=region
589626
)
590627

591628
server_image, sidecar_image = gke_utils.get_pathways_service_images(

‎pathwaysutils/test/experimental/shared_pathways_service/isc_pathways_test.py‎

Lines changed: 135 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -12,7 +12,6 @@
1212
from pathwaysutils.experimental.shared_pathways_service import isc_pathways
1313

1414

15-
1615
class ISCPathwaysTest(parameterized.TestCase):
1716
"""Tests for the ISCPathways class."""
1817

@@ -23,6 +22,14 @@ def setUp(self):
2322
isc_pathways.gke_utils, "stream_pod_logs", autospec=True
2423
)
2524
)
25+
# By default, pretend the kube config does not already point at the
26+
# cluster so that credentials are fetched.
27+
self.mock_is_current_kube_context = self.enter_context(
28+
mock.patch.object(
29+
isc_pathways, "_is_current_kube_context", autospec=True
30+
)
31+
)
32+
self.mock_is_current_kube_context.return_value = False
2633

2734
def test_wait_for_placement_success(self):
2835
"""Tests that _wait_for_placement correctly processes logs."""
@@ -581,6 +588,49 @@ def test_isc_pathways(self):
581588
text=True,
582589
)
583590

591+
def test_connect_skips_credentials_when_kube_config_matches(self):
592+
"""Tests that connect skips fetching credentials for the active context."""
593+
self.mock_is_current_kube_context.return_value = True
594+
mock_fetch_creds = self.enter_context(
595+
mock.patch.object(
596+
isc_pathways.gke_utils, "fetch_cluster_credentials", autospec=True
597+
)
598+
)
599+
mock_get_images = self.enter_context(
600+
mock.patch.object(
601+
isc_pathways.gke_utils, "get_pathways_service_images", autospec=True
602+
)
603+
)
604+
mock_get_images.return_value = (
605+
"us-docker.pkg.dev/cloud-tpu-v2-images/pathways/server:latest",
606+
None,
607+
)
608+
mock_isc_pathways = self.enter_context(
609+
mock.patch.object(isc_pathways, "_ISCPathways", autospec=True)
610+
)
611+
self.enter_context(mock.patch("threading.Thread", autospec=True))
612+
613+
mock_manager_instance = (
614+
mock_isc_pathways.return_value.__enter__.return_value
615+
)
616+
mock_manager_instance.proxy_pod_name = "test-pod-123"
617+
mock_manager_instance.expected_tpu_instances = {"tpuv6e:2x2": 1}
618+
619+
with isc_pathways.connect(
620+
cluster="test-cluster",
621+
project="test-project",
622+
region="test-region",
623+
gcs_bucket="test-bucket",
624+
pathways_service="test-service-pathways-head:1234",
625+
expected_tpu_instances={"tpuv6e:2x2": 1},
626+
):
627+
pass
628+
629+
self.mock_is_current_kube_context.assert_called_once_with(
630+
cluster="test-cluster", project="test-project", location="test-region"
631+
)
632+
mock_fetch_creds.assert_not_called()
633+
584634
def test_connect_success(self):
585635
"""Tests that connect calls the dependencies and yields the manager."""
586636
# Arrange
@@ -1369,5 +1419,89 @@ def test_connect_proxy_server_image_deprecation_warning(self):
13691419
pass
13701420

13711421

1422+
class KubeConfigCredentialsTest(parameterized.TestCase):
1423+
"""Tests for the kube config context helpers."""
1424+
1425+
def setUp(self):
1426+
super().setUp()
1427+
self.mock_get_current_kube_context = self.enter_context(
1428+
mock.patch.object(
1429+
isc_pathways.gke_utils, "get_current_kube_context", autospec=True
1430+
)
1431+
)
1432+
self.mock_fetch_creds = self.enter_context(
1433+
mock.patch.object(
1434+
isc_pathways.gke_utils, "fetch_cluster_credentials", autospec=True
1435+
)
1436+
)
1437+
1438+
def test_is_current_kube_context_matching(self):
1439+
"""Tests that a matching active context is detected."""
1440+
self.mock_get_current_kube_context.return_value = (
1441+
"test-cluster",
1442+
"test-project",
1443+
"test-region",
1444+
)
1445+
1446+
self.assertTrue(
1447+
isc_pathways._is_current_kube_context(
1448+
cluster="test-cluster",
1449+
project="test-project",
1450+
location="test-region",
1451+
)
1452+
)
1453+
1454+
@parameterized.named_parameters(
1455+
("different_cluster", ("other-cluster", "test-project", "test-region")),
1456+
("different_project", ("test-cluster", "other-project", "test-region")),
1457+
("different_location", ("test-cluster", "test-project", "other-region")),
1458+
("non_gke_context", ("minikube", None, None)),
1459+
("no_context", (None, None, None)),
1460+
)
1461+
def test_is_current_kube_context_not_matching(self, current_context):
1462+
"""Tests that a non-matching active context is detected."""
1463+
self.mock_get_current_kube_context.return_value = current_context
1464+
1465+
self.assertFalse(
1466+
isc_pathways._is_current_kube_context(
1467+
cluster="test-cluster",
1468+
project="test-project",
1469+
location="test-region",
1470+
)
1471+
)
1472+
1473+
def test_ensure_cluster_credentials_skips_fetch_when_matching(self):
1474+
"""Tests that credentials are not fetched for the active context."""
1475+
self.mock_get_current_kube_context.return_value = (
1476+
"test-cluster",
1477+
"test-project",
1478+
"test-region",
1479+
)
1480+
1481+
isc_pathways._ensure_cluster_credentials(
1482+
cluster="test-cluster", project="test-project", location="test-region"
1483+
)
1484+
1485+
self.mock_fetch_creds.assert_not_called()
1486+
1487+
def test_ensure_cluster_credentials_fetches_when_not_matching(self):
1488+
"""Tests that credentials are fetched for a different context."""
1489+
self.mock_get_current_kube_context.return_value = (
1490+
"test-cluster",
1491+
"other-project",
1492+
"test-region",
1493+
)
1494+
1495+
isc_pathways._ensure_cluster_credentials(
1496+
cluster="test-cluster", project="test-project", location="test-region"
1497+
)
1498+
1499+
self.mock_fetch_creds.assert_called_once_with(
1500+
cluster_name="test-cluster",
1501+
project_id="test-project",
1502+
location="test-region",
1503+
)
1504+
1505+
13721506
if __name__ == "__main__":
13731507
absltest.main()

0 commit comments

Comments
 (0)