From 444b016734efa1b4fea4f78d4e6577f0ea266f05 Mon Sep 17 00:00:00 2001 From: Akanksha Gupta Date: Fri, 26 Jun 2026 10:56:29 -0700 Subject: [PATCH] Add sidecar image version validation to ISC Pathways connection PiperOrigin-RevId: 938672500 --- .../shared_pathways_service/gke_utils.py | 56 +++++++++++ .../shared_pathways_service/isc_pathways.py | 23 +++-- .../shared_pathways_service/run_workload.py | 6 +- .../shared_pathways_service/validators.py | 99 +++++++++++++++++++ 4 files changed, 172 insertions(+), 12 deletions(-) diff --git a/pathwaysutils/experimental/shared_pathways_service/gke_utils.py b/pathwaysutils/experimental/shared_pathways_service/gke_utils.py index 2e08a49..ac17acf 100644 --- a/pathwaysutils/experimental/shared_pathways_service/gke_utils.py +++ b/pathwaysutils/experimental/shared_pathways_service/gke_utils.py @@ -1,5 +1,6 @@ """GKE utils for deploying and managing the Pathways proxy.""" +import json import logging import re import socket @@ -475,3 +476,58 @@ def is_local_port_free(port: int) -> bool: """Checks if a local port is free.""" return portpicker.is_port_free(port) + +def get_worker_sidecar_image( + pathways_service: str, namespace: str = "default" +) -> str | None: + """Gets the image of the sidecar container used by the workers.""" + pathways_head_hostname = pathways_service.split(":")[0] + _validate_k8s_name(namespace) + + # Try to extract the jobset name from the Pathways service hostname. + jobset_name = None + if "-pathways-head" in pathways_head_hostname: + jobset_name = pathways_head_hostname.split("-pathways-head")[0] + + command = ["kubectl", "get", "pods", "-n", namespace, "-o", "json"] + try: + result = subprocess.run( + command, + check=True, + capture_output=True, + text=True, + ) + except subprocess.CalledProcessError as e: + _logger.exception("Failed to get pods. kubectl output:\n%r", e.stderr) + return None + + try: + pods_data = json.loads(result.stdout) + except json.JSONDecodeError as e: + _logger.exception("Failed to parse kubectl get pods output: %r", e) + return None + + items = pods_data.get("items", []) + + # Look for pods belonging to the jobset and having the sidecar + # container/initContainer. + if jobset_name: + for pod in items: + metadata = pod.get("metadata", {}) + labels = metadata.get("labels", {}) + pod_jobset_name = labels.get("jobset.sigs.k8s.io/jobset-name") + pod_name = metadata.get("name", "") + + if pod_jobset_name == jobset_name or pod_name.startswith(jobset_name): + spec = pod.get("spec", {}) + for container in spec.get("initContainers", []) + spec.get( + "containers", [] + ): + if container.get("name") == "colocated-python-sidecar": + image = container.get("image") + if image: + return image + + return None + + diff --git a/pathwaysutils/experimental/shared_pathways_service/isc_pathways.py b/pathwaysutils/experimental/shared_pathways_service/isc_pathways.py index f34fcad..339b1f3 100644 --- a/pathwaysutils/experimental/shared_pathways_service/isc_pathways.py +++ b/pathwaysutils/experimental/shared_pathways_service/isc_pathways.py @@ -1,6 +1,6 @@ """Module for connecting to a Pathways server for interactive supercomputing.""" -from collections.abc import Iterable, Iterator, Mapping +from collections.abc import Iterable, Iterator, Mapping, Sequence import contextlib import dataclasses import gc @@ -435,7 +435,7 @@ def connect( expected_tpu_instances: Mapping[str, int], proxy_job_name: str | None = None, proxy_server_image: str = DEFAULT_PROXY_IMAGE, - proxy_options: ProxyOptions | None = None, + proxy_options: Sequence[str] | None = None, collect_service_metrics: bool = False, ) -> Iterator["_ISCPathways"]: """Connects to a Pathways server if the cluster exists. If not, creates it. @@ -465,17 +465,24 @@ def connect( validators.validate_tpu_instances(expected_tpu_instances) validators.validate_proxy_server_image(proxy_server_image) validators.validate_proxy_options(proxy_options) - _logger.info("Validation complete.") gke_utils.fetch_cluster_credentials( cluster_name=cluster, project_id=project, location=region ) - proxy_job_name = ( - proxy_job_name or f"isc-proxy-{os.environ.get('USER', 'user')}-{''.join( - random.choices(string.ascii_lowercase + string.digits, k=5) - )}" - ) proxy_options_obj = ProxyOptions.from_list(proxy_options) + if proxy_options_obj.sidecar: + sidecar_image = gke_utils.get_worker_sidecar_image( + pathways_service=pathways_service + ) + if sidecar_image: + validators.validate_sidecar_image_versions(sidecar_image) + _logger.info("Validation complete.") + + proxy_job_name = ( + proxy_job_name + or f"isc-proxy-{os.environ.get('USER', 'user')}-" + f"{''.join(random.choices(string.ascii_lowercase + string.digits, k=5))}" + ) _logger.info("Starting ISCPathways context.") with _ISCPathways( diff --git a/pathwaysutils/experimental/shared_pathways_service/run_workload.py b/pathwaysutils/experimental/shared_pathways_service/run_workload.py index be2e4bf..59c0d97 100644 --- a/pathwaysutils/experimental/shared_pathways_service/run_workload.py +++ b/pathwaysutils/experimental/shared_pathways_service/run_workload.py @@ -77,8 +77,6 @@ ) - - def run_command( *, cluster: str, @@ -107,8 +105,8 @@ def run_command( command: The command to run on TPUs. proxy_server_image: The proxy server image to use. proxy_options: Configuration options for the Pathways proxy. - collect_service_metrics: Whether to collect usage metrics for Shared Pathways - Service. Defaults to False. + collect_service_metrics: Whether to collect usage metrics for Shared + Pathways Service. Defaults to False. connect_fn: The function to use for establishing the connection context, expected to be a callable that returns a context manager. diff --git a/pathwaysutils/experimental/shared_pathways_service/validators.py b/pathwaysutils/experimental/shared_pathways_service/validators.py index 7b28440..3a13b25 100644 --- a/pathwaysutils/experimental/shared_pathways_service/validators.py +++ b/pathwaysutils/experimental/shared_pathways_service/validators.py @@ -3,11 +3,16 @@ from collections.abc import Iterable, Mapping import logging import re +import sys from typing import Any from absl import flags +import jax _logger = logging.getLogger(__name__) +_PYTHON_VERSION_REGEX = r"python[-_]?(\d+\.\d+(?:\.\d+)*)" +_JAX_VERSION_REGEX = r"jax[-_]?(\d+\.\d+(?:\.\d+)*)" + def validate_proxy_options(proxy_options: Iterable[str] | None) -> None: """Validates that proxy options are in the format 'key:value'.""" @@ -133,3 +138,97 @@ def validate_xla_flags(xla_flags: Iterable[str] | None) -> None: raise flags.ValidationError( f"XLA flag '{flag}' must start with '--xla_'." ) + + +def validate_sidecar_image_versions(sidecar_image: str) -> None: + """Checks compatibility of sidecar image versions with user environment. + + Compares the Python and JAX versions in the sidecar image tag with the user + environment's Python and JAX versions. + + Args: + sidecar_image: The sidecar image string, e.g., + "us-docker.pkg.dev/.../sidecar:20260423-python_3.12-jax_0.10.0". + + Raises: + ValueError: If the sidecar image Python or JAX versions do not match the + user environment. + """ + _logger.info( + "Checking sidecar image version compatibility: %s", sidecar_image + ) + + parts = sidecar_image.rsplit(":", 1) + if len(parts) < 2: + _logger.warning( + "No tag found in sidecar image: %s. Skipping version validation.", + sidecar_image, + ) + return + tag = parts[1] + + sidecar_python_match = re.search( + _PYTHON_VERSION_REGEX, tag, re.IGNORECASE + ) + sidecar_jax_match = re.search( + _JAX_VERSION_REGEX, tag, re.IGNORECASE + ) + if not sidecar_python_match and not sidecar_jax_match: + _logger.warning( + "No Python or JAX versions found in sidecar image tag: %s. Skipping " + "version validation.", + tag, + ) + return + + def clean_version(version_str: str) -> str: + match = re.match(r"^(\d+(?:\.\d+)*)", version_str) + return match.group(1) if match else version_str + + def versions_match(sidecar_ver: str, env_ver: str) -> bool: + sidecar_parts = sidecar_ver.split(".") + env_parts = env_ver.split(".") + compare_len = min(len(sidecar_parts), len(env_parts)) + if compare_len == 0: + return False + return sidecar_parts[:compare_len] == env_parts[:compare_len] + + if sidecar_python_match: + sidecar_python = clean_version(sidecar_python_match.group(1)) + env_python = ( + f"{sys.version_info.major}.{sys.version_info.minor}." + f"{sys.version_info.micro}" + ) + if not versions_match(sidecar_python, env_python): + raise ValueError( + f"Python version mismatch: sidecar image matches Python version " + f"{sidecar_python}, but the user environment is running Python " + f"{env_python}. Either rebuild the sidecar image with a matching " + "Python version or update the user environment to match the sidecar" + " image." + ) + _logger.info( + "Python version match: sidecar image matches Python version %s, and the" + " user environment is running Python %s.", + sidecar_python, + env_python, + ) + + if sidecar_jax_match: + sidecar_jax = clean_version(sidecar_jax_match.group(1)) + env_jax = clean_version(jax.__version__) + if not versions_match(sidecar_jax, env_jax): + raise ValueError( + f"JAX version mismatch: sidecar image matches JAX version " + f"{sidecar_jax}, but the user environment is running JAX " + f"{env_jax}. Either rebuild the sidecar image with a matching " + "JAX version or update the user environment to match the sidecar " + "image." + ) + _logger.info( + "JAX version match: sidecar image matches JAX version %s, and the user" + " environment is running JAX %s.", + sidecar_jax, + env_jax, + ) +