diff --git a/pathwaysutils/experimental/shared_pathways_service/deploy_pathways_service.py b/pathwaysutils/experimental/shared_pathways_service/deploy_pathways_service.py index d6e296b..93b4783 100644 --- a/pathwaysutils/experimental/shared_pathways_service/deploy_pathways_service.py +++ b/pathwaysutils/experimental/shared_pathways_service/deploy_pathways_service.py @@ -12,6 +12,7 @@ from kubernetes import config from pathwaysutils.experimental.gke import jobset from pathwaysutils.experimental.shared_pathways_service import gke_utils +from pathwaysutils.experimental.shared_pathways_service import validators import yaml _logger = logging.getLogger(__name__) @@ -303,6 +304,19 @@ def run_deployment( ) deploy_func(jobset_config) + if sidecar_image: + pathways_service = ( + f"{jobset_name}-pathways-head-0-0.{jobset_name}:29001" + ) + _, sidecar_versions = gke_utils.get_sidecar_versions( + pathways_service, sidecar_image=sidecar_image + ) + _logger.info( + "\n%s", + validators.format_sidecar_versions( + pathways_service, sidecar_image, sidecar_versions + ), + ) else: _logger.info("Dry run mode, not deploying.") @@ -311,6 +325,8 @@ def main(argv: Sequence[str]) -> None: if len(argv) > 1: raise app.UsageError("Too many command-line arguments.") + logging.getLogger().setLevel(logging.INFO) + try: if ( flags.FLAGS["jax_version"].present diff --git a/pathwaysutils/experimental/shared_pathways_service/gke_utils.py b/pathwaysutils/experimental/shared_pathways_service/gke_utils.py index 144c999..1b45638 100644 --- a/pathwaysutils/experimental/shared_pathways_service/gke_utils.py +++ b/pathwaysutils/experimental/shared_pathways_service/gke_utils.py @@ -10,9 +10,11 @@ import time from typing import Any import urllib.parse +import uuid from kubernetes import client from kubernetes import config as k8s_config +from pathwaysutils.experimental.shared_pathways_service import validators import portpicker _logger = logging.getLogger(__name__) @@ -850,4 +852,203 @@ def get_compatible_proxy_server_image(server_image: str) -> str: return new_repo +_PYTHON_VERSION_SNIPPET = ( + "import jax, jaxlib, sys; " + "print('SPS_VERSIONS:' + " + "f'{sys.version_info.major}.{sys.version_info.minor}," + "{jax.__version__},{jaxlib.__version__}')" +) + + +def _parse_sidecar_versions_output( + stdout: str, +) -> validators.SidecarVersions | None: + """Parses Python, JAX, and JAXLib versions from command stdout.""" + for line in stdout.strip().splitlines(): + line = line.strip() + if line.startswith("SPS_VERSIONS:"): + line = line[len("SPS_VERSIONS:") :].strip() + parts = line.split(",") + if len(parts) == 3 and all(p.strip() for p in parts): + return validators.SidecarVersions( + python_version=parts[0].strip(), + jax_version=parts[1].strip(), + jaxlib_version=parts[2].strip(), + ) + return None + + +def query_sidecar_image_versions( + sidecar_image: str, namespace: str = "default" +) -> validators.SidecarVersions | None: + """Queries Python, JAX, and JAXLib versions directly from a sidecar image. + + Launches a short-lived pod in the cluster running the sidecar image to + inspect the installed Python, JAX, and JAXLib versions when no running + worker pod is available. + + Args: + sidecar_image: The container image for the colocated python sidecar. + namespace: The Kubernetes namespace to launch the temporary pod in. + + Returns: + A SidecarVersions object if the image could be inspected, or None if + inspection failed. + """ + _validate_k8s_name(namespace) + if not sidecar_image or sidecar_image.startswith("-"): + raise ValueError(f"Invalid sidecar image: '{sidecar_image}'") + + pod_name = f"sps-version-check-{uuid.uuid4().hex[:8]}" + run_cmd = [ + "kubectl", + "run", + pod_name, + "--rm", + "-i", + "--restart=Never", + "-n", + namespace, + f"--image={sidecar_image}", + "--quiet", + "--command", + "--", + "python3", + "-c", + _PYTHON_VERSION_SNIPPET, + ] + try: + result = subprocess.run( + run_cmd, capture_output=True, text=True, check=False, timeout=120 + ) + if result.returncode == 0 and result.stdout: + return _parse_sidecar_versions_output(result.stdout) + except subprocess.TimeoutExpired: + _logger.debug( + "Timed out querying sidecar image %s for versions; cleaning up pod %s.", + sidecar_image, + pod_name, + ) + try: + subprocess.run( + [ + "kubectl", + "delete", + "pod", + pod_name, + "-n", + namespace, + "--ignore-not-found=true", + ], + capture_output=True, + text=True, + check=False, + timeout=10, + ) + except Exception as cleanup_err: # pylint: disable=broad-exception-caught + _logger.debug("Failed to clean up pod %s: %r", pod_name, cleanup_err) + except Exception as e: # pylint: disable=broad-exception-caught + _logger.debug( + "Could not query sidecar image %s for versions: %r", sidecar_image, e + ) + return None + + +def get_sidecar_versions( + pathways_service: str, + namespace: str = "default", + sidecar_image: str | None = None, +) -> tuple[str | None, validators.SidecarVersions]: + """Gets sidecar image and JAX/JAXLib/Python versions from the SPS instance. + + Attempts to query a running `colocated-python-sidecar` container in the + JobSet first. If no running worker pod is available, launches a temporary + pod from `sidecar_image` to inspect the versions directly from the image, + and finally falls back to parsing the image tag. + + Args: + pathways_service: The Pathways service address (e.g. + "-pathways-head-0-0.:29001"). + namespace: The Kubernetes namespace of the JobSet. + sidecar_image: Optional pre-fetched sidecar image string. If None, fetched + from the JobSet via `get_pathways_service_images`. + + Returns: + A tuple of (sidecar_image, SidecarVersions). If colocated python sidecar + is not enabled, sidecar_image is None and SidecarVersions is empty. + """ + _validate_k8s_name(namespace) + if sidecar_image is None: + _, sidecar_image = get_pathways_service_images( + pathways_service, namespace=namespace + ) + if not sidecar_image: + return (None, validators.SidecarVersions()) + + # 1. Attempt to query live sidecar pod for exact runtime versions. + pathways_head_hostname = pathways_service.split(":")[0] + if "-pathways-head" in pathways_head_hostname: + jobset_name = pathways_head_hostname.split("-pathways-head")[0] + try: + _validate_k8s_name(jobset_name) + cmd = [ + "kubectl", + "get", + "pods", + "-n", + namespace, + "-l", + f"jobset.sigs.k8s.io/jobset-name={jobset_name}", + "--field-selector=status.phase=Running", + "-o", + "jsonpath={.items[*].metadata.name}", + ] + pod_result = subprocess.run( + cmd, capture_output=True, text=True, check=True, timeout=10 + ) + pod_names = [ + p + for p in pod_result.stdout.strip().split() + if "-pathways-head-" not in p + ] + for pod_name in pod_names: + _validate_k8s_name(pod_name) + exec_cmd = [ + "kubectl", + "exec", + "-n", + namespace, + f"pod/{pod_name}", + "-c", + "colocated-python-sidecar", + "--", + "python3", + "-c", + _PYTHON_VERSION_SNIPPET, + ] + exec_result = subprocess.run( + exec_cmd, capture_output=True, text=True, check=False, timeout=10 + ) + if exec_result.returncode == 0 and exec_result.stdout: + parsed = _parse_sidecar_versions_output(exec_result.stdout) + if parsed is not None: + return (sidecar_image, parsed) + except Exception as e: # pylint: disable=broad-exception-caught + _logger.debug("Could not query live sidecar pod for versions: %r", e) + + # 2. Inspect the sidecar image directly via a short-lived pod. + image_versions = query_sidecar_image_versions( + sidecar_image, namespace=namespace + ) + if image_versions is not None: + return (sidecar_image, image_versions) + + # 3. Fall back to extracting versions from the sidecar image tag. + return ( + sidecar_image, + validators.extract_sidecar_image_versions(sidecar_image), + ) + + + diff --git a/pathwaysutils/experimental/shared_pathways_service/isc_pathways.py b/pathwaysutils/experimental/shared_pathways_service/isc_pathways.py index 6c8411d..aa9e1c7 100644 --- a/pathwaysutils/experimental/shared_pathways_service/isc_pathways.py +++ b/pathwaysutils/experimental/shared_pathways_service/isc_pathways.py @@ -645,8 +645,20 @@ def connect( proxy_server_image = compatible_proxy_image proxy_options_obj = ProxyOptions.from_list(proxy_options) - if proxy_options_obj.sidecar and sidecar_image: - validators.validate_sidecar_image_versions(sidecar_image) + if sidecar_image: + _, sidecar_versions = gke_utils.get_sidecar_versions( + pathways_service, sidecar_image=sidecar_image + ) + _logger.info( + "\n%s", + validators.format_sidecar_versions( + pathways_service, sidecar_image, sidecar_versions + ), + ) + if proxy_options_obj.sidecar: + validators.validate_sidecar_image_versions( + sidecar_image, sidecar_versions=sidecar_versions + ) _logger.info("Validation complete.") if not proxy_job_name: @@ -690,3 +702,30 @@ def connect( " _wait_for_placement." ) yield t + + +def get_sidecar_versions( + pathways_service: str, + cluster: str | None = None, + project: str | None = None, + region: str | None = None, + namespace: str = "default", +) -> tuple[str | None, validators.SidecarVersions]: + """Gets the sidecar image and versions for the given Pathways service. + + Args: + pathways_service: The service name and port of the Pathways head pod. + cluster: The name of the GKE cluster (optional). + project: The GCP project ID (optional). + region: The GCP region (optional). + namespace: The Kubernetes namespace (defaults to 'default'). + + Returns: + A tuple of (sidecar_image, SidecarVersions). + """ + if cluster and project and region: + _ensure_cluster_credentials( + cluster=cluster, project=project, location=region + ) + return gke_utils.get_sidecar_versions(pathways_service, namespace=namespace) + diff --git a/pathwaysutils/experimental/shared_pathways_service/validators.py b/pathwaysutils/experimental/shared_pathways_service/validators.py index bcba4b3..e9d1b6a 100644 --- a/pathwaysutils/experimental/shared_pathways_service/validators.py +++ b/pathwaysutils/experimental/shared_pathways_service/validators.py @@ -1,6 +1,8 @@ """Validation functions for Shared Pathways Service.""" from collections.abc import Iterable, Mapping +import dataclasses +import importlib import logging import re import sys @@ -12,6 +14,16 @@ _PYTHON_VERSION_REGEX = r"python[-_]?(\d+\.\d+(?:\.\d+)*)" _JAX_VERSION_REGEX = r"jax[-_]?(\d+\.\d+(?:\.\d+)*)" +_JAXLIB_VERSION_REGEX = r"jaxlib[-_]?(\d+\.\d+(?:\.\d+)*)" + + +@dataclasses.dataclass(frozen=True) +class SidecarVersions: + """Holds Python, JAX, and JAXLib versions for a colocated python sidecar.""" + + python_version: str | None = None + jax_version: str | None = None + jaxlib_version: str | None = None def validate_proxy_options(proxy_options: Iterable[str] | None) -> None: @@ -124,40 +136,141 @@ def validate_xla_flags(xla_flags: Iterable[str] | None) -> None: ) -def validate_sidecar_image_versions(sidecar_image: str) -> None: +def _extract_image_tag(sidecar_image: str) -> str | None: + """Extracts the tag from a container image string, ignoring digests.""" + image_without_digest = sidecar_image.split("@", 1)[0] + last_slash = image_without_digest.rfind("/") + if ":" not in image_without_digest[last_slash + 1 :]: + return None + return image_without_digest.rsplit(":", 1)[1] + + +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 extract_sidecar_image_versions(sidecar_image: str) -> SidecarVersions: + """Extracts Python, JAX, and JAXLib versions from the sidecar image tag. + + Args: + sidecar_image: The sidecar image string, e.g., + "us-docker.pkg.dev/.../sidecar:20260423-python_3.12-jax_0.10.0". + + Returns: + A SidecarVersions object with the extracted versions, or None for versions + that could not be determined. + """ + tag = _extract_image_tag(sidecar_image) + if not tag: + return SidecarVersions() + + python_version = None + py_match = re.search(_PYTHON_VERSION_REGEX, tag, re.IGNORECASE) + if py_match: + python_version = _clean_version(py_match.group(1)) + + jax_version = None + jax_match = re.search(_JAX_VERSION_REGEX, tag, re.IGNORECASE) + if jax_match: + jax_version = _clean_version(jax_match.group(1)) + + jaxlib_version = None + jaxlib_match = re.search(_JAXLIB_VERSION_REGEX, tag, re.IGNORECASE) + if jaxlib_match: + jaxlib_version = _clean_version(jaxlib_match.group(1)) + elif jax_version: + # JAX and JAXLib release versions correspond to each other by default. + jaxlib_version = jax_version + + return SidecarVersions( + python_version=python_version, + jax_version=jax_version, + jaxlib_version=jaxlib_version, + ) + + +def format_sidecar_versions( + pathways_service: str, + sidecar_image: str, + sidecar_versions: SidecarVersions, +) -> str: + """Formats the sidecar versions into a human-readable summary string.""" + lines = [ + ( + "Colocated Python sidecar found for Pathways service" + f" '{pathways_service}':" + ), + f" Sidecar Image: {sidecar_image}", + ] + if sidecar_versions.python_version: + lines.append(f" Python: {sidecar_versions.python_version}") + if sidecar_versions.jax_version: + lines.append(f" JAX: {sidecar_versions.jax_version}") + if sidecar_versions.jaxlib_version: + lines.append(f" JAXLib: {sidecar_versions.jaxlib_version}") + + if sidecar_versions.jax_version: + jaxlib = sidecar_versions.jaxlib_version or sidecar_versions.jax_version + lines.append( + "To install the matching JAX and JAXLib versions locally, run:" + ) + lines.append( + f" pip install jax=={sidecar_versions.jax_version} jaxlib=={jaxlib}" + ) + else: + lines.append( + "Could not determine JAX version from colocated python sidecar" + f" image: {sidecar_image}" + ) + return "\n".join(lines) + + +def validate_sidecar_image_versions( + sidecar_image: str, sidecar_versions: SidecarVersions | None = None +) -> 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. + Compares the Python, JAX, and JAXLib versions in the sidecar image or + container with the user environment's Python, JAX, and JAXLib versions. Args: sidecar_image: The sidecar image string, e.g., "us-docker.pkg.dev/.../sidecar:20260423-python_3.12-jax_0.10.0". + sidecar_versions: Optional pre-resolved SidecarVersions (e.g., queried from + the sidecar container or image). If omitted, versions are extracted from + the sidecar image tag. Raises: - ValueError: If the sidecar image Python or JAX versions do not match the - user environment. + ValueError: If the sidecar image Python, JAX, or JAXLib 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: + tag = _extract_image_tag(sidecar_image) + explicit_versions_provided = sidecar_versions is not None + + if sidecar_versions is None or ( + not sidecar_versions.python_version + and not sidecar_versions.jax_version + and not sidecar_versions.jaxlib_version + ): + if not tag: + _logger.warning( + "No tag found in sidecar image: %s. Skipping version validation.", + sidecar_image, + ) + return + sidecar_versions = extract_sidecar_image_versions(sidecar_image) + explicit_versions_provided = False + + if ( + not sidecar_versions.python_version + and not sidecar_versions.jax_version + and not sidecar_versions.jaxlib_version + ): _logger.warning( "No Python or JAX versions found in sidecar image tag: %s. Skipping " "version validation.", @@ -165,10 +278,6 @@ def validate_sidecar_image_versions(sidecar_image: str) -> None: ) 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(".") @@ -177,8 +286,15 @@ def versions_match(sidecar_ver: str, env_ver: str) -> bool: return False return sidecar_parts[:compare_len] == env_parts[:compare_len] - if sidecar_python_match: - sidecar_python = clean_version(sidecar_python_match.group(1)) + install_hint = "" + if sidecar_versions.jax_version and sidecar_versions.jaxlib_version: + install_hint = ( + f" by running: pip install jax=={sidecar_versions.jax_version}" + f" jaxlib=={sidecar_versions.jaxlib_version}" + ) + + if sidecar_versions.python_version: + sidecar_python = _clean_version(sidecar_versions.python_version) env_python = ( f"{sys.version_info.major}.{sys.version_info.minor}." f"{sys.version_info.micro}" @@ -198,16 +314,16 @@ def versions_match(sidecar_ver: str, env_ver: str) -> bool: env_python, ) - if sidecar_jax_match: - sidecar_jax = clean_version(sidecar_jax_match.group(1)) - env_jax = clean_version(jax.__version__) + if sidecar_versions.jax_version: + sidecar_jax = _clean_version(sidecar_versions.jax_version) + 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." + f"image{install_hint}." ) _logger.info( "JAX version match: sidecar image matches JAX version %s, and the user" @@ -216,3 +332,35 @@ def versions_match(sidecar_ver: str, env_ver: str) -> bool: env_jax, ) + should_check_jaxlib = explicit_versions_provided or bool( + tag and re.search(_JAXLIB_VERSION_REGEX, tag, re.IGNORECASE) + ) + if should_check_jaxlib and sidecar_versions.jaxlib_version: + sidecar_jaxlib = _clean_version(sidecar_versions.jaxlib_version) + env_jaxlib = None + try: + jaxlib_mod = sys.modules.get("jaxlib") + if jaxlib_mod is None: + jaxlib_mod = importlib.import_module("jaxlib") + if hasattr(jaxlib_mod, "__version__"): + env_jaxlib = _clean_version(jaxlib_mod.__version__) + except (ImportError, AttributeError): + env_jaxlib = None + + if env_jaxlib is not None: + if not versions_match(sidecar_jaxlib, env_jaxlib): + raise ValueError( + f"JAXLib version mismatch: sidecar image matches JAXLib version " + f"{sidecar_jaxlib}, but the user environment is running JAXLib " + f"{env_jaxlib}. Either rebuild the sidecar image with a matching " + "JAXLib version or update the user environment to match the" + f" sidecar image{install_hint}." + ) + _logger.info( + "JAXLib version match: sidecar image matches JAXLib version %s, and" + " the user environment is running JAXLib %s.", + sidecar_jaxlib, + env_jaxlib, + ) + + 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 62853cd..34579ea 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 @@ -7,10 +7,29 @@ from absl.testing import parameterized from pathwaysutils.experimental.shared_pathways_service import deploy_pathways_service from pathwaysutils.experimental.shared_pathways_service import gke_utils +from pathwaysutils.experimental.shared_pathways_service import validators class DeployPathwaysServiceTest(parameterized.TestCase): + def setUp(self): + super().setUp() + self.mock_get_sidecar_versions = self.enter_context( + mock.patch.object( + gke_utils, + "get_sidecar_versions", + autospec=True, + return_value=( + "custom-sidecar-image", + validators.SidecarVersions( + python_version="3.12", + jax_version="0.11.1", + jaxlib_version="0.11.1", + ), + ), + ) + ) + @parameterized.named_parameters( dict( testcase_name="v5p", 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 dbc9340..92669df 100644 --- a/pathwaysutils/test/experimental/shared_pathways_service/gke_utils_test.py +++ b/pathwaysutils/test/experimental/shared_pathways_service/gke_utils_test.py @@ -13,6 +13,7 @@ from kubernetes import client from kubernetes import config as k8s_config from pathwaysutils.experimental.shared_pathways_service import gke_utils +from pathwaysutils.experimental.shared_pathways_service import validators import portpicker @@ -1593,7 +1594,145 @@ def test_terminate_process_custom_timeout(self): mock_proc.terminate.assert_called_once() mock_proc.wait.assert_called_once_with(timeout=10) + def test_query_sidecar_image_versions_success(self): + with mock.patch.object(subprocess, "run") as mock_run: + mock_run.return_value = mock.Mock( + stdout=( + "SPS_VERSIONS:3.12,0.11.1,0.11.1\n" + 'pod "sps-version-check-12345678" deleted\n' + ), + returncode=0, + ) + versions = gke_utils.query_sidecar_image_versions( + "us-docker.pkg.dev/repo/sidecar@sha256:80b6718" + ) + self.assertEqual( + versions, + validators.SidecarVersions( + python_version="3.12", + jax_version="0.11.1", + jaxlib_version="0.11.1", + ), + ) + mock_run.assert_called_once() + cmd = mock_run.call_args[0][0] + self.assertEqual(cmd[:2], ["kubectl", "run"]) + self.assertIn( + "--image=us-docker.pkg.dev/repo/sidecar@sha256:80b6718", cmd + ) + + def test_query_sidecar_image_versions_failure(self): + with mock.patch.object( + subprocess, "run", side_effect=Exception("kubectl run failed") + ): + versions = gke_utils.query_sidecar_image_versions( + "us-docker.pkg.dev/repo/sidecar@sha256:80b6718" + ) + self.assertIsNone(versions) + + def test_query_sidecar_image_versions_timeout_cleans_up_pod(self): + with mock.patch.object(subprocess, "run") as mock_run: + mock_run.side_effect = [ + subprocess.TimeoutExpired(cmd=["kubectl", "run"], timeout=120), + mock.Mock(returncode=0), + ] + versions = gke_utils.query_sidecar_image_versions( + "us-docker.pkg.dev/repo/sidecar@sha256:80b6718", + namespace="custom-ns", + ) + self.assertIsNone(versions) + self.assertEqual(mock_run.call_count, 2) + delete_cmd = mock_run.call_args_list[1][0][0] + self.assertEqual(delete_cmd[:3], ["kubectl", "delete", "pod"]) + self.assertIn("-n", delete_cmd) + self.assertIn("custom-ns", delete_cmd) + + def test_get_sidecar_versions_no_sidecar(self): + with mock.patch.object( + gke_utils, + "get_pathways_service_images", + return_value=("server_img", None), + ): + sidecar_img, versions = gke_utils.get_sidecar_versions( + "my-jobset-pathways-head-0-0.my-jobset:8000" + ) + self.assertIsNone(sidecar_img) + self.assertIsNone(versions.python_version) + self.assertIsNone(versions.jax_version) + self.assertIsNone(versions.jaxlib_version) + + def test_get_sidecar_versions_live_query_success(self): + with mock.patch.object( + gke_utils, + "get_pathways_service_images", + return_value=( + "server_img", + "us-docker.pkg.dev/repo/sidecar:20260423-python_3.12-jax_0.10.0", + ), + ), mock.patch.object(subprocess, "run") as mock_run: + mock_run.side_effect = [ + mock.Mock(stdout="pod-0\n", returncode=0), + mock.Mock(stdout="SPS_VERSIONS:3.12,0.10.0,0.10.0\n", returncode=0), + ] + sidecar_img, versions = gke_utils.get_sidecar_versions( + "my-jobset-pathways-head-0-0.my-jobset:8000" + ) + self.assertEqual( + sidecar_img, + "us-docker.pkg.dev/repo/sidecar:20260423-python_3.12-jax_0.10.0", + ) + self.assertEqual(versions.python_version, "3.12") + self.assertEqual(versions.jax_version, "0.10.0") + self.assertEqual(versions.jaxlib_version, "0.10.0") + + def test_get_sidecar_versions_fallback_to_image_query(self): + digest_img = "us-docker.pkg.dev/repo/sidecar@sha256:80b671827c0e6995d9aa615635981c21" + with mock.patch.object( + gke_utils, + "get_pathways_service_images", + return_value=("server_img", digest_img), + ), mock.patch.object(subprocess, "run") as mock_run: + # 1st call: kubectl get pods returns no running pods + # 2nd call: kubectl run inspects the sidecar image directly + mock_run.side_effect = [ + mock.Mock(stdout="", returncode=0), + mock.Mock( + stdout="SPS_VERSIONS:3.12,0.10.0,0.10.0\npod deleted\n", + returncode=0, + ), + ] + sidecar_img, versions = gke_utils.get_sidecar_versions( + "my-jobset-pathways-head-0-0.my-jobset:8000" + ) + self.assertEqual(sidecar_img, digest_img) + self.assertEqual(versions.python_version, "3.12") + self.assertEqual(versions.jax_version, "0.10.0") + self.assertEqual(versions.jaxlib_version, "0.10.0") + + def test_get_sidecar_versions_fallback_to_tag(self): + with mock.patch.object( + gke_utils, + "get_pathways_service_images", + return_value=( + "server_img", + "us-docker.pkg.dev/repo/sidecar:20260423-python_3.12-jax_0.10.0", + ), + ), mock.patch.object( + subprocess, "run", side_effect=Exception("kubectl failed") + ): + sidecar_img, versions = gke_utils.get_sidecar_versions( + "my-jobset-pathways-head-0-0.my-jobset:8000" + ) + self.assertEqual( + sidecar_img, + "us-docker.pkg.dev/repo/sidecar:20260423-python_3.12-jax_0.10.0", + ) + self.assertEqual(versions.python_version, "3.12") + self.assertEqual(versions.jax_version, "0.10.0") + self.assertEqual(versions.jaxlib_version, "0.10.0") + if __name__ == "__main__": absltest.main() + 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 9eb7620..c3fceda 100644 --- a/pathwaysutils/test/experimental/shared_pathways_service/isc_pathways_test.py +++ b/pathwaysutils/test/experimental/shared_pathways_service/isc_pathways_test.py @@ -10,6 +10,7 @@ from absl.testing import absltest from absl.testing import parameterized from pathwaysutils.experimental.shared_pathways_service import isc_pathways +from pathwaysutils.experimental.shared_pathways_service import validators class ISCPathwaysTest(parameterized.TestCase): @@ -825,6 +826,7 @@ def test_connect_with_sidecar_validation_success(self): isc_pathways.gke_utils, "fetch_cluster_credentials", autospec=True ) ) + sidecar_img = "us-docker.pkg.dev/.../sidecar:20260423-python_3.12-jax_0.10.0" mock_get_images = self.enter_context( mock.patch.object( isc_pathways.gke_utils, "get_pathways_service_images", autospec=True @@ -832,7 +834,18 @@ def test_connect_with_sidecar_validation_success(self): ) mock_get_images.return_value = ( "us-docker.pkg.dev/cloud-tpu-v2-images/pathways/server:latest", - "us-docker.pkg.dev/.../sidecar:20260423-python_3.12-jax_0.10.0", + sidecar_img, + ) + expected_versions = validators.SidecarVersions( + python_version="3.12", jax_version="0.10.0", jaxlib_version="0.10.0" + ) + mock_get_sidecar_versions = self.enter_context( + mock.patch.object( + isc_pathways.gke_utils, + "get_sidecar_versions", + autospec=True, + return_value=(sidecar_img, expected_versions), + ) ) mock_validate_versions = self.enter_context( mock.patch.object( @@ -852,23 +865,108 @@ def test_connect_with_sidecar_validation_success(self): 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}, - proxy_options=["sidecar:true"], - ): - pass + with self.assertLogs(isc_pathways._logger, level="INFO") as log_cm: + 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}, + proxy_options=["sidecar:true"], + ): + pass + self.assertTrue( + any( + "Colocated Python sidecar found for Pathways service" in msg + and "pip install jax==0.10.0 jaxlib==0.10.0" in msg + for msg in log_cm.output + ) + ) mock_get_images.assert_called_once_with( "test-service-pathways-head:1234" ) + mock_get_sidecar_versions.assert_called_once_with( + "test-service-pathways-head:1234", sidecar_image=sidecar_img + ) mock_validate_versions.assert_called_once_with( - "us-docker.pkg.dev/.../sidecar:20260423-python_3.12-jax_0.10.0" + sidecar_img, sidecar_versions=expected_versions + ) + + def test_connect_logs_sidecar_versions_by_default_when_sidecar_image_present( + self, + ): + self.enter_context(mock.patch.dict(os.environ, {"USER": "testuser"})) + self.enter_context( + mock.patch.object( + isc_pathways.gke_utils, "fetch_cluster_credentials", autospec=True + ) + ) + sidecar_img = "us-docker.pkg.dev/.../sidecar:20260423-python_3.12-jax_0.10.0" + 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", + sidecar_img, ) + expected_versions = validators.SidecarVersions( + python_version="3.12", jax_version="0.10.0", jaxlib_version="0.10.0" + ) + mock_get_sidecar_versions = self.enter_context( + mock.patch.object( + isc_pathways.gke_utils, + "get_sidecar_versions", + autospec=True, + return_value=(sidecar_img, expected_versions), + ) + ) + mock_validate_versions = self.enter_context( + mock.patch.object( + isc_pathways.validators, + "validate_sidecar_image_versions", + autospec=True, + ) + ) + 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 self.assertLogs(isc_pathways._logger, level="INFO") as log_cm: + 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.assertTrue( + any( + "Colocated Python sidecar found for Pathways service" in msg + and "pip install jax==0.10.0 jaxlib==0.10.0" in msg + for msg in log_cm.output + ) + ) + mock_get_images.assert_called_once_with( + "test-service-pathways-head:1234" + ) + mock_get_sidecar_versions.assert_called_once_with( + "test-service-pathways-head:1234", sidecar_image=sidecar_img + ) + mock_validate_versions.assert_not_called() def test_connect_with_sidecar_validation_mismatch_raises_error(self): self.enter_context(mock.patch.dict(os.environ, {"USER": "testuser"})) @@ -877,6 +975,7 @@ def test_connect_with_sidecar_validation_mismatch_raises_error(self): isc_pathways.gke_utils, "fetch_cluster_credentials", autospec=True ) ) + sidecar_img = "us-docker.pkg.dev/.../sidecar:20260423-python_3.12-jax_0.10.0" mock_get_images = self.enter_context( mock.patch.object( isc_pathways.gke_utils, "get_pathways_service_images", autospec=True @@ -884,7 +983,18 @@ def test_connect_with_sidecar_validation_mismatch_raises_error(self): ) mock_get_images.return_value = ( "us-docker.pkg.dev/cloud-tpu-v2-images/pathways/server:latest", - "us-docker.pkg.dev/.../sidecar:20260423-python_3.12-jax_0.10.0", + sidecar_img, + ) + expected_versions = validators.SidecarVersions( + python_version="3.12", jax_version="0.10.0", jaxlib_version="0.10.0" + ) + mock_get_sidecar_versions = self.enter_context( + mock.patch.object( + isc_pathways.gke_utils, + "get_sidecar_versions", + autospec=True, + return_value=(sidecar_img, expected_versions), + ) ) mock_validate_versions = self.enter_context( mock.patch.object( @@ -910,8 +1020,11 @@ def test_connect_with_sidecar_validation_mismatch_raises_error(self): mock_get_images.assert_called_once_with( "test-service-pathways-head:1234" ) + mock_get_sidecar_versions.assert_called_once_with( + "test-service-pathways-head:1234", sidecar_image=sidecar_img + ) mock_validate_versions.assert_called_once_with( - "us-docker.pkg.dev/.../sidecar:20260423-python_3.12-jax_0.10.0" + sidecar_img, sidecar_versions=expected_versions ) @parameterized.named_parameters( @@ -1418,6 +1531,61 @@ def test_connect_proxy_server_image_deprecation_warning(self): ): pass + def test_get_sidecar_versions_fetches_credentials_when_cluster_provided(self): + mock_fetch = self.enter_context( + mock.patch.object( + isc_pathways.gke_utils, "fetch_cluster_credentials", autospec=True + ) + ) + mock_get = self.enter_context( + mock.patch.object( + isc_pathways.gke_utils, "get_sidecar_versions", autospec=True + ) + ) + mock_versions = mock.Mock() + mock_get.return_value = ("sidecar_img", mock_versions) + + img, versions = isc_pathways.get_sidecar_versions( + pathways_service="test-service-pathways-head:1234", + cluster="test-cluster", + project="test-project", + region="test-region", + ) + mock_fetch.assert_called_once_with( + cluster_name="test-cluster", + project_id="test-project", + location="test-region", + ) + mock_get.assert_called_once_with( + "test-service-pathways-head:1234", namespace="default" + ) + self.assertEqual(img, "sidecar_img") + self.assertIs(versions, mock_versions) + + def test_get_sidecar_versions_without_cluster_skips_fetch_credentials(self): + mock_fetch = self.enter_context( + mock.patch.object( + isc_pathways.gke_utils, "fetch_cluster_credentials", autospec=True + ) + ) + mock_get = self.enter_context( + mock.patch.object( + isc_pathways.gke_utils, "get_sidecar_versions", autospec=True + ) + ) + mock_versions = mock.Mock() + mock_get.return_value = ("sidecar_img", mock_versions) + + img, versions = isc_pathways.get_sidecar_versions( + pathways_service="test-service-pathways-head:1234" + ) + mock_fetch.assert_not_called() + mock_get.assert_called_once_with( + "test-service-pathways-head:1234", namespace="default" + ) + self.assertEqual(img, "sidecar_img") + self.assertIs(versions, mock_versions) + class KubeConfigCredentialsTest(parameterized.TestCase): """Tests for the kube config context helpers.""" diff --git a/pathwaysutils/test/experimental/shared_pathways_service/validators_test.py b/pathwaysutils/test/experimental/shared_pathways_service/validators_test.py index 484be59..d16ef79 100644 --- a/pathwaysutils/test/experimental/shared_pathways_service/validators_test.py +++ b/pathwaysutils/test/experimental/shared_pathways_service/validators_test.py @@ -1,5 +1,6 @@ """Tests for validation functions for the Shared Pathways service.""" +import sys from unittest import mock from absl import flags @@ -309,6 +310,126 @@ def test_validate_sidecar_image_versions_jax_mismatch(self): "us-docker.pkg.dev/.../sidecar:20260423-python_3.12-jax_0.10.0" ) + def test_extract_sidecar_image_versions(self): + versions = validators.extract_sidecar_image_versions( + "us-docker.pkg.dev/proj/repo/sidecar:20260423-python_3.12-jax_0.10.0" + ) + self.assertEqual(versions.python_version, "3.12") + self.assertEqual(versions.jax_version, "0.10.0") + self.assertEqual(versions.jaxlib_version, "0.10.0") + + versions2 = validators.extract_sidecar_image_versions( + "us-docker.pkg.dev/proj/repo/sidecar:" + "20260831-maxtext-v0.2.4-jax0.11.1-pwutils" + ) + self.assertIsNone(versions2.python_version) + self.assertEqual(versions2.jax_version, "0.11.1") + self.assertEqual(versions2.jaxlib_version, "0.11.1") + + versions3 = validators.extract_sidecar_image_versions( + "us-docker.pkg.dev/proj/repo/sidecar:python_3.11-jax_0.10.0-jaxlib_0.9.0" + ) + self.assertEqual(versions3.python_version, "3.11") + self.assertEqual(versions3.jax_version, "0.10.0") + self.assertEqual(versions3.jaxlib_version, "0.9.0") + + versions_digest = validators.extract_sidecar_image_versions( + "us-docker.pkg.dev/proj/repo/sidecar@sha256:80b671827c0e6995d9aa615635981c21" + ) + self.assertIsNone(versions_digest.python_version) + self.assertIsNone(versions_digest.jax_version) + self.assertIsNone(versions_digest.jaxlib_version) + + versions_latest = validators.extract_sidecar_image_versions( + "us-docker.pkg.dev/proj/repo/sidecar:latest" + ) + self.assertIsNone(versions_latest.python_version) + self.assertIsNone(versions_latest.jax_version) + self.assertIsNone(versions_latest.jaxlib_version) + + versions_no_tag = validators.extract_sidecar_image_versions("sidecar") + self.assertIsNone(versions_no_tag.python_version) + self.assertIsNone(versions_no_tag.jax_version) + self.assertIsNone(versions_no_tag.jaxlib_version) + + def test_validate_sidecar_image_versions_jaxlib_mismatch(self): + mock_sys_info = mock.Mock() + mock_sys_info.major = 3 + mock_sys_info.minor = 12 + mock_sys_info.micro = 8 + mock_jaxlib = mock.Mock(__version__="0.9.0") + with mock.patch("sys.version_info", mock_sys_info), mock.patch( + "jax.__version__", "0.10.0" + ), mock.patch.dict(sys.modules, {"jaxlib": mock_jaxlib}): + with self.assertRaisesRegex(ValueError, "JAXLib version mismatch"): + validators.validate_sidecar_image_versions( + "us-docker.pkg.dev/.../sidecar:python_3.12-jax_0.10.0-jaxlib_0.10.0" + ) + + def test_validate_sidecar_image_versions_with_explicit_sidecar_versions(self): + mock_sys_info = mock.Mock() + mock_sys_info.major = 3 + mock_sys_info.minor = 12 + mock_sys_info.micro = 8 + mock_jaxlib = mock.Mock(__version__="0.11.1") + digest_img = ( + "us-docker.pkg.dev/proj/repo/sidecar@sha256:80b671827c0e6995d9aa615635981c21" + ) + queried_versions = validators.SidecarVersions( + python_version="3.12", + jax_version="0.11.1", + jaxlib_version="0.11.1", + ) + with mock.patch("sys.version_info", mock_sys_info), mock.patch( + "jax.__version__", "0.11.1" + ), mock.patch.dict(sys.modules, {"jaxlib": mock_jaxlib}): + validators.validate_sidecar_image_versions( + digest_img, sidecar_versions=queried_versions + ) + + with mock.patch("sys.version_info", mock_sys_info), mock.patch( + "jax.__version__", "0.10.0" + ), mock.patch.dict(sys.modules, {"jaxlib": mock_jaxlib}): + with self.assertRaisesRegex( + ValueError, + r"JAX version mismatch.*pip install jax==0\.11\.1 jaxlib==0\.11\.1", + ): + validators.validate_sidecar_image_versions( + digest_img, sidecar_versions=queried_versions + ) + + def test_format_sidecar_versions(self): + versions = validators.SidecarVersions( + python_version="3.12", + jax_version="0.11.1", + jaxlib_version="0.11.1", + ) + text_out = validators.format_sidecar_versions( + pathways_service="test-service:29001", + sidecar_image="repo/sidecar:tag", + sidecar_versions=versions, + ) + self.assertEqual( + text_out, + "Colocated Python sidecar found for Pathways service" + " 'test-service:29001':\n" + " Sidecar Image: repo/sidecar:tag\n" + " Python: 3.12\n" + " JAX: 0.11.1\n" + " JAXLib: 0.11.1\n" + "To install the matching JAX and JAXLib versions locally, run:\n" + " pip install jax==0.11.1 jaxlib==0.11.1", + ) + + def test_format_sidecar_versions_no_jax_version(self): + versions = validators.SidecarVersions() + text_out = validators.format_sidecar_versions( + pathways_service="test-service:29001", + sidecar_image="repo/sidecar:latest", + sidecar_versions=versions, + ) + self.assertIn("Could not determine JAX version", text_out) + if __name__ == "__main__": absltest.main()