Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
56 changes: 56 additions & 0 deletions pathwaysutils/experimental/shared_pathways_service/gke_utils.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
"""GKE utils for deploying and managing the Pathways proxy."""

import json
import logging
import re
import socket
Expand Down Expand Up @@ -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


23 changes: 15 additions & 8 deletions pathwaysutils/experimental/shared_pathways_service/isc_pathways.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -77,8 +77,6 @@
)




def run_command(
*,
cluster: str,
Expand Down Expand Up @@ -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.

Expand Down
99 changes: 99 additions & 0 deletions pathwaysutils/experimental/shared_pathways_service/validators.py
Original file line number Diff line number Diff line change
Expand Up @@ -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'."""
Expand Down Expand Up @@ -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,
)

Loading