Skip to content
Open
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
Original file line number Diff line number Diff line change
Expand Up @@ -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__)
Expand Down Expand Up @@ -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.")

Expand All @@ -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
Expand Down
201 changes: 201 additions & 0 deletions pathwaysutils/experimental/shared_pathways_service/gke_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__)
Expand Down Expand Up @@ -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.
"<jobset_name>-pathways-head-0-0.<jobset_name>: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),
)




43 changes: 41 additions & 2 deletions pathwaysutils/experimental/shared_pathways_service/isc_pathways.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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)

Loading
Loading