diff --git a/pathwaysutils/experimental/shared_pathways_service/deploy_pathways_service.py b/pathwaysutils/experimental/shared_pathways_service/deploy_pathways_service.py index 17d88d2..d6e296b 100644 --- a/pathwaysutils/experimental/shared_pathways_service/deploy_pathways_service.py +++ b/pathwaysutils/experimental/shared_pathways_service/deploy_pathways_service.py @@ -2,6 +2,7 @@ from collections.abc import Callable, Sequence import dataclasses +import datetime import logging import math from typing import Any @@ -10,6 +11,7 @@ from kubernetes import client from kubernetes import config from pathwaysutils.experimental.gke import jobset +from pathwaysutils.experimental.shared_pathways_service import gke_utils import yaml _logger = logging.getLogger(__name__) @@ -207,36 +209,65 @@ def run_deployment( if container.name == "pathways-rm": container.image = server_image # Mutate worker job. - for container in pw_jobset.worker_job_template.spec.template.spec.containers: + for ( + container + ) in pw_jobset.worker_job_template.spec.template.spec.containers: if container.name == "pathways-worker": container.image = server_image # Add colocated python sidecar. - pw_jobset.add_colocated_python(image=sidecar_image, shm_mount_path=_SIDECAR_SHM_DIR) + pw_jobset.add_colocated_python( + image=sidecar_image, shm_mount_path=_SIDECAR_SHM_DIR + ) # Mutate the sidecar configuration to match what HEAD expects. worker_spec = pw_jobset.worker_job_template.spec.template.spec # 1. Add extra logging env vars to sidecar. - for container in ((worker_spec.containers or []) + (worker_spec.init_containers or [])): + # These make the colocated Python sidecar logs as verbose as possible so that + # issues in the sidecar (e.g. array (de)serialization during checkpointing) + # can be debugged from the container logs. + for container in (worker_spec.containers or []) + ( + worker_spec.init_containers or [] + ): if container.name == "colocated-python-sidecar": container.env.extend([ + # Disable Python stdout/stderr buffering so logs are emitted + # immediately and are not lost if the container crashes. client.V1EnvVar(name="PYTHONUNBUFFERED", value="1"), + # Python logging level for the sidecar. client.V1EnvVar(name="LOGLEVEL", value="DEBUG"), + # glog (C++) logging: emit INFO and above (0 = INFO) and enable + # VLOG(n) messages for n <= 5. client.V1EnvVar(name="GLOG_minloglevel", value="0"), client.V1EnvVar(name="GLOG_v", value="5"), + # TSL/XLA (C++) logging used by JAX: emit INFO and above and enable + # VLOG(n) messages for n <= 5. client.V1EnvVar(name="TF_CPP_MIN_LOG_LEVEL", value="0"), client.V1EnvVar(name="TF_CPP_MIN_VLOG_LEVEL", value="5"), + # TPU runtime (libtpu) logging: emit INFO and above. client.V1EnvVar(name="TPU_MIN_LOG_LEVEL", value="0"), - client.V1EnvVar(name="GLOG_vmodule", value="jax_array_handlers=5,type_handlers=5,tensorstore_utils=5"), + # Per-module verbose logging (level 5) for the array serialization + # modules used when transferring/checkpointing arrays via the + # sidecar. + client.V1EnvVar( + name="GLOG_vmodule", + value=( + "jax_array_handlers=5,type_handlers=5,tensorstore_utils=5" + ), + ), ]) - # 2. Add arg to pathways-worker container (in addition to env var set by builder). + # 2. Add arg to pathways-worker container. for container in worker_spec.containers: if container.name == "pathways-worker": args = container.args or [] - if not any(a.startswith("--cloud_pathways_sidecar_shm_directory=") for a in args): - args.append(f"--cloud_pathways_sidecar_shm_directory={_SIDECAR_SHM_DIR}") + if not any( + a.startswith("--cloud_pathways_sidecar_shm_directory=") for a in args + ): + args.append( + f"--cloud_pathways_sidecar_shm_directory={_SIDECAR_SHM_DIR}" + ) container.args = args jobset_config = pw_jobset.to_dict() @@ -246,6 +277,31 @@ def run_deployment( if not dry_run: _logger.info("Deploying JobSet...") + cluster, project = gke_utils.get_current_cluster_and_project() + if not cluster or not project: + raise ValueError( + "Cluster or project could not be determined from kubeconfig. Run" + " 'gcloud container clusters get-credentials ... && kubectl config" + " set-context --current --namespace=default' OR 'kubectl config" + " set-context --current --user=... --cluster=...'" + " first." + ) + now = datetime.datetime.now(datetime.timezone.utc) + start_time = now.isoformat(timespec="milliseconds").replace("+00:00", "Z") + end_time = (now + gke_utils.LOG_LINK_WINDOW).isoformat( + timespec="milliseconds" + ).replace("+00:00", "Z") + cloud_logging_link = gke_utils.get_log_link( + cluster=cluster, + project=project, + job_name=jobset_name, + start_time=start_time, + end_time=end_time, + ) + _logger.info( + "View SPS deployment logs in Cloud Logging: %s", cloud_logging_link + ) + deploy_func(jobset_config) else: _logger.info("Dry run mode, not deploying.") diff --git a/pathwaysutils/experimental/shared_pathways_service/gke_utils.py b/pathwaysutils/experimental/shared_pathways_service/gke_utils.py index 184f8a2..144c999 100644 --- a/pathwaysutils/experimental/shared_pathways_service/gke_utils.py +++ b/pathwaysutils/experimental/shared_pathways_service/gke_utils.py @@ -1,7 +1,9 @@ """GKE utils for deploying and managing the Pathways proxy.""" +import datetime import functools import logging +import os import re import socket import subprocess @@ -15,6 +17,9 @@ _logger = logging.getLogger(__name__) +# Default size of the time window covered by Cloud Logging links. +LOG_LINK_WINDOW = datetime.timedelta(minutes=10) + # TODO(b/456189271): Evaluate and replace the subprocess calls with Kubernetes # Python API for kubectl calls. @@ -239,23 +244,167 @@ def check_pod_ready(pod_name: str, timeout: int = 30) -> str: return pod_name -def get_log_link(*, cluster: str, project: str, job_name: str) -> str: - """Returns a link to Cloud Logging for the given cluster and job name.""" +def _format_log_timestamp(timestamp: str | datetime.datetime) -> str: + """Formats a timestamp for use in a Cloud Logging query URL. + + Args: + timestamp: An ISO-8601 string or a datetime. Naive datetimes are assumed to + be UTC. + + Returns: + The timestamp as an ISO-8601 string, or the unmodified string input. + """ + if not isinstance(timestamp, datetime.datetime): + return str(timestamp) + + if timestamp.tzinfo is None: + timestamp = timestamp.replace(tzinfo=datetime.timezone.utc) + timestamp = timestamp.astimezone(datetime.timezone.utc) + return timestamp.isoformat(timespec="milliseconds").replace("+00:00", "Z") + + +def _parse_log_timestamp( + timestamp: str | datetime.datetime, +) -> datetime.datetime: + """Parses a timestamp into a timezone-aware datetime. + + Args: + timestamp: An ISO-8601 string or a datetime. Naive values are assumed to be + UTC. + + Returns: + A timezone-aware datetime. + + Raises: + ValueError: If the string is not a valid ISO-8601 timestamp. + """ + if isinstance(timestamp, datetime.datetime): + parsed = timestamp + else: + # `fromisoformat` does not accept a trailing "Z" before Python 3.11. + parsed = datetime.datetime.fromisoformat( + re.sub(r"[Zz]$", "+00:00", timestamp) + ) + if parsed.tzinfo is None: + parsed = parsed.replace(tzinfo=datetime.timezone.utc) + return parsed + + +def get_log_link( + *, + cluster: str, + project: str, + job_name: str, + namespace: str = "default", + start_time: str | datetime.datetime | None = None, + end_time: str | datetime.datetime | None = None, + duration: str | None = "PT1H", +) -> str: + """Returns a link to Cloud Logging for the given cluster and job name. + + Args: + cluster: The name of the GKE cluster. + project: The GCP project ID. + job_name: The name of the job or jobset. + namespace: The Kubernetes namespace. Defaults to "default". + start_time: The start time for the time window (ISO-8601 string or + datetime). If provided without end_time, the end time is set to + `LOG_LINK_WINDOW` after start_time. + end_time: The end time for the time window (ISO-8601 string or datetime). + If provided without start_time, the start time is set to + `LOG_LINK_WINDOW` before end_time. + duration: The duration string (e.g. "PT1H") used when neither start_time + nor end_time is provided. + + Returns: + The Cloud Logging query URL. + + Raises: + ValueError: If only one of start_time and end_time is provided and it is a + string that is not a valid ISO-8601 timestamp. + """ log_filter = ( 'resource.type="k8s_container"\n' f'resource.labels.cluster_name="{cluster}"\n' - 'resource.labels.namespace_name="default"\n' + f'resource.labels.namespace_name="{namespace}"\n' f'labels.k8s-pod/job-name:"{job_name}"' ) + + if start_time is not None and end_time is None: + end_time = _parse_log_timestamp(start_time) + LOG_LINK_WINDOW + elif end_time is not None and start_time is None: + start_time = _parse_log_timestamp(end_time) - LOG_LINK_WINDOW + + if start_time is not None and end_time is not None: + start_time_str = _format_log_timestamp(start_time) + end_time_str = _format_log_timestamp(end_time) + time_param = f"startTime={start_time_str};endTime={end_time_str}" + elif duration is not None: + time_param = f"duration={duration}" + else: + time_param = "" + encoded_filter = urllib.parse.quote(log_filter, safe="") + time_part = f";{time_param}" if time_param else "" return ( "https://console.cloud.google.com/logs/query;" - f"query={encoded_filter};duration=PT1H" + f"query={encoded_filter}{time_part}" f"?project={project}" ) +def get_current_kube_context() -> tuple[str | None, str | None, str | None]: + """Reads the cluster targeted by the active kube config context. + + Only the kube config is consulted; no environment variable fallbacks are + applied. Contexts written by `gcloud container clusters get-credentials` are + named `gke___`. + + Returns: + A (cluster, project, location) tuple. All three are populated for a GKE + context. For any other context, only the context name is returned as the + cluster and the project and location are None. All three are None if the + active context cannot be read or is a malformed GKE context. + """ + try: + _, active_context = k8s_config.list_kube_config_contexts() + except Exception as e: # pylint: disable=broad-except + _logger.debug("Could not read the current kube config context: %s", e) + return None, None, None + + if not active_context: + return None, None, None + + context_data = active_context.get("context", {}) + cluster_context = ( + context_data.get("cluster") or active_context.get("name", "") + ) + + if cluster_context.startswith("gke_"): + parts = cluster_context.split("_", 3) + if len(parts) != 4: + return None, None, None + _, project, location, cluster = parts + return cluster, project, location + + return cluster_context or None, None, None + + +def get_current_cluster_and_project() -> tuple[str | None, str | None]: + """Extracts cluster name and project ID from current kubeconfig or environment.""" + cluster, project, _ = get_current_kube_context() + + if not cluster: + cluster = os.environ.get("GKE_CLUSTER") or os.environ.get("CLUSTER") + if not project: + project = os.environ.get("PROJECT") or os.environ.get( + "GOOGLE_CLOUD_PROJECT" + ) + + return cluster, project + + def wait_for_pod(job_name: str) -> str: """Waits for the given job's pod to be ready. diff --git a/pathwaysutils/test/experimental/shared_pathways_service/deploy_pathways_service_test.py b/pathwaysutils/test/experimental/shared_pathways_service/deploy_pathways_service_test.py index 7ca426a..62853cd 100644 --- a/pathwaysutils/test/experimental/shared_pathways_service/deploy_pathways_service_test.py +++ b/pathwaysutils/test/experimental/shared_pathways_service/deploy_pathways_service_test.py @@ -1,10 +1,12 @@ """Unit tests for the deploy_pathways_service script.""" +import datetime from unittest import mock from absl import flags from absl.testing import absltest from absl.testing import parameterized from pathwaysutils.experimental.shared_pathways_service import deploy_pathways_service +from pathwaysutils.experimental.shared_pathways_service import gke_utils class DeployPathwaysServiceTest(parameterized.TestCase): @@ -69,8 +71,10 @@ def test_calculate_vms_per_slice_not_divisible(self): with self.assertRaises(ValueError): deploy_pathways_service.calculate_vms_per_slice("4x8", 5) + @mock.patch.object(gke_utils, "get_current_cluster_and_project") @mock.patch("pathwaysutils.experimental.shared_pathways_service.deploy_pathways_service.jobset.PathwaysJobSet") - def test_run_deployment(self, mock_jobset_cls): + def test_run_deployment(self, mock_jobset_cls, mock_detect): + mock_detect.return_value = ("test-cluster", "test-project") mock_jobset = mock_jobset_cls.return_value mock_jobset.to_dict.return_value = {"metadata": {"name": "test-jobset"}} @@ -146,7 +150,9 @@ def test_run_deployment(self, mock_jobset_cls): # Verify deploy_func was called with the dict mock_deploy.assert_called_once_with({"metadata": {"name": "test-jobset"}}) - def test_run_deployment_worker_backoff_limit(self): + @mock.patch.object(gke_utils, "get_current_cluster_and_project") + def test_run_deployment_worker_backoff_limit(self, mock_detect): + mock_detect.return_value = ("test-cluster", "test-project") captured_config = {} def capture_deploy(config): @@ -178,7 +184,9 @@ def capture_deploy(config): # Verify worker backoff limit is set to a large value self.assertGreaterEqual(worker_backoff, 1000000) - def test_run_deployment_max_restarts(self): + @mock.patch.object(gke_utils, "get_current_cluster_and_project") + def test_run_deployment_max_restarts(self, mock_detect): + mock_detect.return_value = ("test-cluster", "test-project") captured_config = {} def capture_deploy(config): @@ -207,6 +215,92 @@ def capture_deploy(config): {"restartStrategy": "Recreate", "maxRestarts": 3}, ) + @mock.patch.object(gke_utils, "get_log_link") + @mock.patch.object(gke_utils, "get_current_cluster_and_project") + def test_run_deployment_cloud_logging_link( + self, mock_detect, mock_get_log_link + ): + mock_detect.return_value = ("test-cluster", "test-project") + mock_get_log_link.return_value = ( + "https://console.cloud.google.com/logs/query;dummy" + ) + mock_deploy = mock.MagicMock() + + deploy_pathways_service.run_deployment( + tpu_type="v5e", + topology="4x8", + num_slices=2, + jobset_name="test-jobset", + gcs_bucket="gs://test-bucket", + server_image="server-image", + sidecar_image="sidecar-image", + dry_run=False, + deploy_func=mock_deploy, + ) + + mock_detect.assert_called_once() + mock_get_log_link.assert_called_once() + _, kwargs = mock_get_log_link.call_args + self.assertEqual(kwargs["cluster"], "test-cluster") + self.assertEqual(kwargs["project"], "test-project") + self.assertEqual(kwargs["job_name"], "test-jobset") + start_time = datetime.datetime.fromisoformat( + kwargs["start_time"].replace("Z", "+00:00") + ) + end_time = datetime.datetime.fromisoformat( + kwargs["end_time"].replace("Z", "+00:00") + ) + self.assertEqual(end_time - start_time, gke_utils.LOG_LINK_WINDOW) + + @parameterized.named_parameters( + dict(testcase_name="neither", cluster=None, project=None), + dict(testcase_name="no_cluster", cluster=None, project="test-project"), + dict(testcase_name="no_project", cluster="test-cluster", project=None), + ) + @mock.patch.object(gke_utils, "get_log_link") + @mock.patch.object(gke_utils, "get_current_cluster_and_project") + def test_run_deployment_missing_cluster_or_project_raises( + self, mock_detect, mock_get_log_link, cluster, project + ): + mock_detect.return_value = (cluster, project) + mock_deploy = mock.MagicMock() + + with self.assertRaises(ValueError): + deploy_pathways_service.run_deployment( + tpu_type="v5e", + topology="4x8", + num_slices=2, + jobset_name="test-jobset", + gcs_bucket="gs://test-bucket", + server_image="server-image", + sidecar_image="sidecar-image", + dry_run=False, + deploy_func=mock_deploy, + ) + + mock_detect.assert_called_once() + mock_get_log_link.assert_not_called() + mock_deploy.assert_not_called() + + @mock.patch.object(gke_utils, "get_log_link") + def test_run_deployment_dry_run_no_log_link(self, mock_get_log_link): + mock_deploy = mock.MagicMock() + + deploy_pathways_service.run_deployment( + tpu_type="v5e", + topology="4x8", + num_slices=2, + jobset_name="test-jobset", + gcs_bucket="gs://test-bucket", + server_image="server-image", + sidecar_image="sidecar-image", + dry_run=True, + deploy_func=mock_deploy, + ) + + mock_get_log_link.assert_not_called() + mock_deploy.assert_not_called() + if __name__ == "__main__": FLAGS = flags.FLAGS diff --git a/pathwaysutils/test/experimental/shared_pathways_service/gke_utils_test.py b/pathwaysutils/test/experimental/shared_pathways_service/gke_utils_test.py index fde1830..dbc9340 100644 --- a/pathwaysutils/test/experimental/shared_pathways_service/gke_utils_test.py +++ b/pathwaysutils/test/experimental/shared_pathways_service/gke_utils_test.py @@ -1,11 +1,13 @@ """Tests for gke_utils.py. """ +import datetime import io import socket import subprocess from typing import Any from unittest import mock +import urllib.parse from absl.testing import absltest from kubernetes import client @@ -337,22 +339,383 @@ def test_check_pod_ready_failure(self): ): gke_utils.check_pod_ready("test-pod-123") + def _assert_log_link( + self, + log_link: str, + *, + cluster: str = "test-cluster", + project: str = "test-project", + job_name: str = "test-job", + namespace: str = "default", + time_params: dict[str, str], + ) -> None: + """Asserts each component of a Cloud Logging link, independent of order. + + Args: + log_link: The Cloud Logging URL to check. + cluster: The expected GKE cluster name. + project: The expected GCP project ID. + job_name: The expected job name. + namespace: The expected Kubernetes namespace. + time_params: The expected time-range parameters (e.g. `startTime`, + `endTime`, `duration`). Empty if no time range is expected. + """ + parsed = urllib.parse.urlparse(log_link) + self.assertEqual(parsed.scheme, "https") + self.assertEqual(parsed.netloc, "console.cloud.google.com") + self.assertEqual(parsed.path, "/logs/query") + self.assertEqual( + dict(urllib.parse.parse_qsl(parsed.query)), {"project": project} + ) + + link_params = dict( + param.split("=", 1) for param in parsed.params.split(";") + ) + log_filter = urllib.parse.unquote(link_params.pop("query")) + self.assertCountEqual( + log_filter.splitlines(), + [ + 'resource.type="k8s_container"', + f'resource.labels.cluster_name="{cluster}"', + f'resource.labels.namespace_name="{namespace}"', + f'labels.k8s-pod/job-name:"{job_name}"', + ], + ) + self.assertEqual(link_params, time_params) + def test_get_log_link(self): - cluster = "test-cluster" - project = "test-project" - job_name = "test-job" log_link = gke_utils.get_log_link( - cluster=cluster, project=project, job_name=job_name + cluster="test-cluster", project="test-project", job_name="test-job" ) - self.assertEqual( + self._assert_log_link(log_link, time_params={"duration": "PT1H"}) + + def test_get_log_link_with_time_window_str(self): + log_link = gke_utils.get_log_link( + cluster="test-cluster", + project="test-project", + job_name="test-job", + start_time="2026-09-10T10:00:00.000Z", + end_time="2026-09-10T10:10:00.000Z", + ) + self._assert_log_link( + log_link, + time_params={ + "startTime": "2026-09-10T10:00:00.000Z", + "endTime": "2026-09-10T10:10:00.000Z", + }, + ) + + def test_get_log_link_with_time_window_datetime(self): + log_link = gke_utils.get_log_link( + cluster="test-cluster", + project="test-project", + job_name="test-job", + start_time=datetime.datetime( + 2026, 9, 10, 10, 0, 0, tzinfo=datetime.timezone.utc + ), + end_time=datetime.datetime( + 2026, 9, 10, 10, 10, 0, tzinfo=datetime.timezone.utc + ), + ) + self._assert_log_link( log_link, - r"https://console.cloud.google.com/logs/query;query=resource.type%3D" - r"%22k8s_container%22%0Aresource.labels.cluster_name%3D" - "%22test-cluster%22%0Aresource.labels.namespace_name%3D" - "%22default%22%0Alabels.k8s-pod%2Fjob-name%3A%22test-job%22;" - "duration=PT1H?project=test-project", + time_params={ + "startTime": "2026-09-10T10:00:00.000Z", + "endTime": "2026-09-10T10:10:00.000Z", + }, ) + def test_get_log_link_with_naive_datetime_assumes_utc(self): + log_link = gke_utils.get_log_link( + cluster="test-cluster", + project="test-project", + job_name="test-job", + start_time=datetime.datetime(2026, 9, 10, 10, 0, 0), + end_time=datetime.datetime(2026, 9, 10, 10, 10, 0), + ) + self._assert_log_link( + log_link, + time_params={ + "startTime": "2026-09-10T10:00:00.000Z", + "endTime": "2026-09-10T10:10:00.000Z", + }, + ) + + def test_get_log_link_with_aware_non_utc_datetime_converts_to_utc(self): + tz = datetime.timezone(datetime.timedelta(hours=2)) + log_link = gke_utils.get_log_link( + cluster="test-cluster", + project="test-project", + job_name="test-job", + start_time=datetime.datetime(2026, 9, 10, 12, 0, 0, tzinfo=tz), + end_time=datetime.datetime(2026, 9, 10, 12, 10, 0, tzinfo=tz), + ) + self._assert_log_link( + log_link, + time_params={ + "startTime": "2026-09-10T10:00:00.000Z", + "endTime": "2026-09-10T10:10:00.000Z", + }, + ) + + def test_get_log_link_without_time_window_omits_time_param(self): + log_link = gke_utils.get_log_link( + cluster="test-cluster", + project="test-project", + job_name="test-job", + duration=None, + ) + self._assert_log_link(log_link, time_params={}) + + def test_get_log_link_with_only_start_time_uses_window_after_start(self): + log_link = gke_utils.get_log_link( + cluster="test-cluster", + project="test-project", + job_name="test-job", + start_time=datetime.datetime( + 2026, 9, 10, 10, 0, 0, tzinfo=datetime.timezone.utc + ), + ) + self._assert_log_link( + log_link, + time_params={ + "startTime": "2026-09-10T10:00:00.000Z", + "endTime": "2026-09-10T10:10:00.000Z", + }, + ) + + def test_get_log_link_with_only_start_time_str_uses_window_after_start( + self, + ): + log_link = gke_utils.get_log_link( + cluster="test-cluster", + project="test-project", + job_name="test-job", + start_time="2026-09-10T10:00:00.000Z", + ) + self._assert_log_link( + log_link, + time_params={ + "startTime": "2026-09-10T10:00:00.000Z", + "endTime": "2026-09-10T10:10:00.000Z", + }, + ) + + def test_get_log_link_with_only_end_time_uses_window_before_end(self): + log_link = gke_utils.get_log_link( + cluster="test-cluster", + project="test-project", + job_name="test-job", + end_time="2026-09-10T10:10:00.000Z", + ) + self._assert_log_link( + log_link, + time_params={ + "startTime": "2026-09-10T10:00:00.000Z", + "endTime": "2026-09-10T10:10:00.000Z", + }, + ) + + def test_get_log_link_with_only_end_time_naive_datetime_assumes_utc(self): + log_link = gke_utils.get_log_link( + cluster="test-cluster", + project="test-project", + job_name="test-job", + end_time=datetime.datetime(2026, 9, 10, 10, 10, 0), + ) + self._assert_log_link( + log_link, + time_params={ + "startTime": "2026-09-10T10:00:00.000Z", + "endTime": "2026-09-10T10:10:00.000Z", + }, + ) + + def test_get_log_link_with_only_invalid_start_time_raises(self): + with self.assertRaises(ValueError): + gke_utils.get_log_link( + cluster="test-cluster", + project="test-project", + job_name="test-job", + start_time="not-a-timestamp", + ) + + def test_get_log_link_with_custom_namespace(self): + log_link = gke_utils.get_log_link( + cluster="test-cluster", + project="test-project", + job_name="test-job", + namespace="custom-ns", + ) + self._assert_log_link( + log_link, namespace="custom-ns", time_params={"duration": "PT1H"} + ) + + def test_get_current_kube_context_from_gke_context(self): + mock_active_context = { + "name": "gke_test-proj_us-central1_test-cl", + "context": {"cluster": "gke_test-proj_us-central1_test-cl"}, + } + with mock.patch.object( + k8s_config, + "list_kube_config_contexts", + return_value=([mock_active_context], mock_active_context), + ): + self.assertEqual( + gke_utils.get_current_kube_context(), + ("test-cl", "test-proj", "us-central1"), + ) + + def test_get_current_kube_context_non_gke_context(self): + mock_active_context = { + "name": "minikube", + "context": {"cluster": "minikube"}, + } + with mock.patch.object( + k8s_config, + "list_kube_config_contexts", + return_value=([mock_active_context], mock_active_context), + ): + self.assertEqual( + gke_utils.get_current_kube_context(), ("minikube", None, None) + ) + + def test_get_current_kube_context_malformed_gke_context(self): + mock_active_context = { + "name": "gke_test-proj_us-central1", + "context": {"cluster": "gke_test-proj_us-central1"}, + } + with mock.patch.object( + k8s_config, + "list_kube_config_contexts", + return_value=([mock_active_context], mock_active_context), + ): + self.assertEqual( + gke_utils.get_current_kube_context(), (None, None, None) + ) + + def test_get_current_kube_context_no_active_context(self): + with mock.patch.object( + k8s_config, + "list_kube_config_contexts", + return_value=([], None), + ): + self.assertEqual( + gke_utils.get_current_kube_context(), (None, None, None) + ) + + def test_get_current_kube_context_error(self): + with mock.patch.object( + k8s_config, + "list_kube_config_contexts", + side_effect=k8s_config.ConfigException("no kubeconfig"), + ): + self.assertEqual( + gke_utils.get_current_kube_context(), (None, None, None) + ) + + def test_get_current_kube_context_ignores_environment(self): + """Tests that the strict reader does not fall back to the environment.""" + with mock.patch.object( + k8s_config, + "list_kube_config_contexts", + return_value=([], None), + ): + with mock.patch.dict( + "os.environ", + {"GKE_CLUSTER": "env-cl", "PROJECT": "env-proj"}, + clear=False, + ): + self.assertEqual( + gke_utils.get_current_kube_context(), (None, None, None) + ) + + def test_get_current_cluster_and_project_from_kubeconfig(self): + mock_active_context = { + "name": "gke_test-proj_us-central1_test-cl", + "context": {"cluster": "gke_test-proj_us-central1_test-cl"}, + } + with mock.patch.object( + k8s_config, + "list_kube_config_contexts", + return_value=([mock_active_context], mock_active_context), + ): + cluster, project = gke_utils.get_current_cluster_and_project() + self.assertEqual(cluster, "test-cl") + self.assertEqual(project, "test-proj") + + def test_get_current_cluster_and_project_fallback_env(self): + with mock.patch.object( + k8s_config, + "list_kube_config_contexts", + return_value=([], None), + ): + with mock.patch.dict( + "os.environ", + {"GKE_CLUSTER": "env-cl", "PROJECT": "env-proj"}, + clear=False, + ): + cluster, project = gke_utils.get_current_cluster_and_project() + self.assertEqual(cluster, "env-cl") + self.assertEqual(project, "env-proj") + + def test_get_current_cluster_and_project_non_gke_context(self): + mock_active_context = { + "name": "minikube", + "context": {"cluster": "minikube"}, + } + with mock.patch.object( + k8s_config, + "list_kube_config_contexts", + return_value=([mock_active_context], mock_active_context), + ): + with mock.patch.dict("os.environ", {}, clear=True): + cluster, project = gke_utils.get_current_cluster_and_project() + self.assertEqual(cluster, "minikube") + self.assertIsNone(project) + + def test_get_current_cluster_and_project_empty_context_name(self): + mock_active_context = {"name": "", "context": {}} + with mock.patch.object( + k8s_config, + "list_kube_config_contexts", + return_value=([mock_active_context], mock_active_context), + ): + with mock.patch.dict("os.environ", {}, clear=True): + cluster, project = gke_utils.get_current_cluster_and_project() + self.assertIsNone(cluster) + self.assertIsNone(project) + + def test_get_current_cluster_and_project_malformed_gke_context(self): + mock_active_context = { + "name": "gke_test-proj_us-central1", + "context": {"cluster": "gke_test-proj_us-central1"}, + } + with mock.patch.object( + k8s_config, + "list_kube_config_contexts", + return_value=([mock_active_context], mock_active_context), + ): + with mock.patch.dict("os.environ", {}, clear=True): + cluster, project = gke_utils.get_current_cluster_and_project() + self.assertIsNone(cluster) + self.assertIsNone(project) + + def test_get_current_cluster_and_project_kubeconfig_error_falls_back(self): + with mock.patch.object( + k8s_config, + "list_kube_config_contexts", + side_effect=k8s_config.ConfigException("no kubeconfig"), + ): + with mock.patch.dict( + "os.environ", + {"CLUSTER": "env-cl", "GOOGLE_CLOUD_PROJECT": "env-proj"}, + clear=True, + ): + cluster, project = gke_utils.get_current_cluster_and_project() + self.assertEqual(cluster, "env-cl") + self.assertEqual(project, "env-proj") + def test_wait_for_pod_success(self): """Tests that wait_for_pod returns the pod name on success.""" mock_run = self.enter_context( @@ -575,7 +938,9 @@ def test_get_pod_from_job_invalid_format_empty(self): returncode=0, stdout="\n", ) - with self.assertRaisesRegex(RuntimeError, "Failed to get pod name. Expected format:"): + with self.assertRaisesRegex( + RuntimeError, "Failed to get pod name. Expected format:" + ): gke_utils.get_pod_from_job("test-job") def test_get_pod_from_job_invalid_format_no_prefix(self): @@ -587,7 +952,9 @@ def test_get_pod_from_job_invalid_format_no_prefix(self): returncode=0, stdout="test-pod-123\n", ) - with self.assertRaisesRegex(RuntimeError, "Failed to get pod name. Expected format:"): + with self.assertRaisesRegex( + RuntimeError, "Failed to get pod name. Expected format:" + ): gke_utils.get_pod_from_job("test-job") def test_get_pod_from_job_invalid_format_too_many_slashes(self): @@ -599,7 +966,9 @@ def test_get_pod_from_job_invalid_format_too_many_slashes(self): returncode=0, stdout="pod/test-pod/extra\n", ) - with self.assertRaisesRegex(RuntimeError, "Failed to get pod name. Expected format:"): + with self.assertRaisesRegex( + RuntimeError, "Failed to get pod name. Expected format:" + ): gke_utils.get_pod_from_job("test-job") def test_test_remote_connection_success(self): @@ -627,7 +996,11 @@ def test_test_remote_connection_refused(self): def test_enable_port_forwarding_pick_port_fails(self): self.enter_context( - mock.patch.object(portpicker, "pick_unused_port", side_effect=ValueError("pick failed")) + mock.patch.object( + portpicker, + "pick_unused_port", + side_effect=ValueError("pick failed"), + ) ) with self.assertRaisesRegex(ValueError, "pick failed"): gke_utils.enable_port_forwarding("test-pod", 8080) @@ -637,7 +1010,9 @@ def test_enable_port_forwarding_popen_fails(self): mock.patch.object(portpicker, "pick_unused_port", return_value=12345) ) self.enter_context( - mock.patch.object(subprocess, "Popen", side_effect=OSError("Popen failed")) + mock.patch.object( + subprocess, "Popen", side_effect=OSError("Popen failed") + ) ) with self.assertRaisesRegex(OSError, "Popen failed"): gke_utils.enable_port_forwarding("test-pod", 8080) @@ -652,8 +1027,12 @@ def test_enable_port_forwarding_stdout_none(self): mock_process = mock_popen.return_value mock_process.stdout = None mock_process.communicate.return_value = ("stdout", "stderr_out") - - with self.assertRaisesRegex(RuntimeError, "Failed to start port forwarding: stdout not available.\nSTDERR: stderr_out"): + + with self.assertRaisesRegex( + RuntimeError, + "Failed to start port forwarding: stdout not available.\nSTDERR: " + "stderr_out", + ): gke_utils.enable_port_forwarding("test-pod", 8080) mock_process.terminate.assert_called_once() mock_process.communicate.assert_called_once() @@ -685,7 +1064,9 @@ def test_wait_for_deployment_failure(self): mock_run.side_effect = subprocess.CalledProcessError( returncode=1, cmd="kubectl rollout", stderr="rollout failed" ) - with self.assertRaisesRegex(RuntimeError, "Deployment did not become ready: rollout failed"): + with self.assertRaisesRegex( + RuntimeError, "Deployment did not become ready: rollout failed" + ): gke_utils.wait_for_deployment("my-deploy", "my-ns") def test_wait_for_service_ip_success_first_try(self):