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
345 changes: 335 additions & 10 deletions pathwaysutils/experimental/shared_pathways_service/gke_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -150,7 +150,6 @@ def delete_gke_resource(
raise



def get_pod_from_job(job_name: str) -> str:
"""Returns the pod name for the given job.

Expand Down Expand Up @@ -672,7 +671,6 @@ def stream_pod_logs(pod_name: str) -> subprocess.Popen[str]:
raise



def wait_for_deployment(
name: str, namespace: str = "default", timeout: int = 300
) -> None:
Expand Down Expand Up @@ -764,21 +762,36 @@ def _get_k8s_custom_objects_api() -> client.CustomObjectsApi:
return client.CustomObjectsApi()


def get_pathways_service_images(
pathways_service: str, namespace: str = "default"
) -> tuple[str, str | None]:
"""Gets the server image and optional worker sidecar image from the JobSet."""
pathways_head_hostname = pathways_service.split(":")[0]
_validate_k8s_name(namespace)
def extract_jobset_name(pathways_service: str) -> str:
"""Extracts the JobSet name from the Pathways service address.

# Try to extract the jobset name from the Pathways service hostname.
Args:
pathways_service: The Pathways service address (e.g.
"my-jobset-pathways-head-0-0.my-jobset:29001" or
"my-jobset-pathways-head:29001").

Returns:
The JobSet name.

Raises:
ValueError: If the JobSet name cannot be extracted from the address.
"""
pathways_head_hostname = pathways_service.split(":")[0]
if "-pathways-head" not in pathways_head_hostname:
raise ValueError(
"Failed to extract jobset name from Pathways service hostname:"
f" {pathways_head_hostname}. Expected prefix format:"
" <jobset_name>-pathways-head"
)
jobset_name = pathways_head_hostname.split("-pathways-head")[0]
return pathways_head_hostname.split("-pathways-head")[0]


def get_pathways_service_images(
pathways_service: str, namespace: str = "default"
) -> tuple[str, str | None]:
"""Gets the server image and optional worker sidecar image from the JobSet."""
_validate_k8s_name(namespace)
jobset_name = extract_jobset_name(pathways_service)

try:
custom_api = _get_k8s_custom_objects_api()
Expand Down Expand Up @@ -850,4 +863,316 @@ def get_compatible_proxy_server_image(server_image: str) -> str:
return new_repo


def _get_pod_metadata(pod: Any) -> Any:
if isinstance(pod, dict):
return pod.get("metadata", {})
return getattr(pod, "metadata", None)


def _get_pod_labels(pod: Any) -> dict[str, str]:
meta = _get_pod_metadata(pod)
if isinstance(meta, dict):
return meta.get("labels", {}) or {}
return getattr(meta, "labels", None) or {}


def _get_pod_name(pod: Any) -> str:
meta = _get_pod_metadata(pod)
if isinstance(meta, dict):
return meta.get("name", "") or ""
return getattr(meta, "name", "") or ""


def _is_head_pod(pod: Any, jobset_name: str) -> bool:
labels = _get_pod_labels(pod)
rep_job = labels.get("jobset.sigs.k8s.io/replicatedjob-name", "")
if rep_job in ("pathways-head", "head"):
return True
name = _get_pod_name(pod)
return f"{jobset_name}-pathways-head" in name or f"{jobset_name}-head" in name


def _is_worker_pod(pod: Any, jobset_name: str) -> bool:
labels = _get_pod_labels(pod)
rep_job = labels.get("jobset.sigs.k8s.io/replicatedjob-name", "")
if rep_job in ("pathways-worker", "worker"):
return True
name = _get_pod_name(pod)
return (
f"{jobset_name}-pathways-worker" in name
or f"{jobset_name}-worker" in name
)


def _is_pod_ready(pod: Any) -> bool:
"""Returns True if the pod phase is Running and Ready condition is True."""
if isinstance(pod, dict):
status = pod.get("status", {})
if status.get("phase") != "Running":
return False
conditions = status.get("conditions", []) or []
for cond in conditions:
if cond.get("type") == "Ready" and cond.get("status") == "True":
return True
return False
else:
status = getattr(pod, "status", None)
if not status or getattr(status, "phase", None) != "Running":
return False
conditions = getattr(status, "conditions", None) or []
for cond in conditions:
if (
getattr(cond, "type", None) == "Ready"
and getattr(cond, "status", None) == "True"
):
return True
return False


def _get_pod_status_details(pod: Any) -> str:
"""Returns a human-readable summary of pod phase and container states."""
reasons: list[str] = []
if isinstance(pod, dict):
status = pod.get("status", {})
phase = status.get("phase", "Unknown")
container_statuses = (status.get("containerStatuses") or []) + (
status.get("initContainerStatuses") or []
)
for cs in container_statuses:
c_name = cs.get("name", "unknown")
state = cs.get("state", {})
if "waiting" in state and state["waiting"]:
reason = state["waiting"].get("reason", "Waiting")
msg = state["waiting"].get("message", "")
reasons.append(
f"container {c_name} waiting: {reason} ({msg})"
if msg
else f"container {c_name} waiting: {reason}"
)
elif "terminated" in state and state["terminated"]:
reason = state["terminated"].get("reason", "Terminated")
exit_code = state["terminated"].get("exitCode", "")
reasons.append(
f"container {c_name} terminated: {reason} (exit code {exit_code})"
)
else:
status = getattr(pod, "status", None)
phase = getattr(status, "phase", "Unknown") if status else "Unknown"
container_statuses = []
if status:
container_statuses = (
getattr(status, "container_statuses", None) or []
) + (getattr(status, "init_container_statuses", None) or [])
for cs in container_statuses:
c_name = getattr(cs, "name", "unknown")
state = getattr(cs, "state", None)
if state:
waiting = getattr(state, "waiting", None)
terminated = getattr(state, "terminated", None)
if waiting:
reason = getattr(waiting, "reason", "Waiting")
msg = getattr(waiting, "message", "")
reasons.append(
f"container {c_name} waiting: {reason} ({msg})"
if msg
else f"container {c_name} waiting: {reason}"
)
elif terminated:
reason = getattr(terminated, "reason", "Terminated")
exit_code = getattr(terminated, "exit_code", "")
reasons.append(
f"container {c_name} terminated: {reason} (exit code {exit_code})"
)

details = f"phase={phase}"
if reasons:
details += f", {', '.join(reasons)}"
return details


def verify_pathways_service_is_up(
*,
cluster: str,
project: str,
region: str,
pathways_service: str,
tpu_count: int | None = None,
namespace: str = "default",
) -> None:
"""Verifies that the Shared Pathways Service JobSet and pods are up and ready.

Args:
cluster: The name of the GKE cluster.
project: The GCP project ID.
region: The GCP region.
pathways_service: The Pathways service address.
tpu_count: Optional expected number of TPU slices.
namespace: The Kubernetes namespace.

Raises:
ValueError: If the service address or namespace is invalid.
RuntimeError: If the JobSet is not found, suspended, failed, or if head or
worker pods are not ready.
"""
_validate_k8s_name(namespace)
jobset_name = extract_jobset_name(pathways_service)
_logger.info(
"Verifying Shared Pathways Service '%s' in namespace '%s' on cluster"
" '%s'...",
jobset_name,
namespace,
cluster,
)
fetch_cluster_credentials(
cluster_name=cluster, project_id=project, location=region
)

custom_api = _get_k8s_custom_objects_api()
try:
jobset = custom_api.get_namespaced_custom_object(
group="jobset.x-k8s.io",
version="v1alpha2",
namespace=namespace,
plural="jobsets",
name=jobset_name,
)
except Exception as e:
status_code = getattr(e, "status", None)
if status_code == 404:
raise RuntimeError(
f"Shared Pathways Service JobSet '{jobset_name}' was not found in"
f" namespace '{namespace}' on cluster '{cluster}'."
) from e
_logger.exception("Failed to get JobSet '%s': %r", jobset_name, e)
raise RuntimeError(
f"Failed to get Shared Pathways Service JobSet '{jobset_name}': {e}"
) from e

if jobset.get("spec", {}).get("suspend", False):
raise RuntimeError(
f"Shared Pathways Service JobSet '{jobset_name}' is suspended."
)

status = jobset.get("status", {})
terminal_state = status.get("terminalState")
if terminal_state:
raise RuntimeError(
f"Shared Pathways Service JobSet '{jobset_name}' has terminated with"
f" state '{terminal_state}'."
)

for condition in status.get("conditions", []):
cond_type = condition.get("type")
cond_status = condition.get("status")
if cond_status == "True" and cond_type in ("Failed", "Suspended"):
reason = condition.get("reason", "")
message = condition.get("message", "")
detail = (
f" (reason: {reason}, message: {message})"
if (reason or message)
else ""
)
raise RuntimeError(
f"Shared Pathways Service JobSet '{jobset_name}' is in condition"
f" '{cond_type}'{detail}."
)

core_api = _get_k8s_core_api()
try:
pod_list = core_api.list_namespaced_pod(
namespace=namespace,
label_selector=f"jobset.sigs.k8s.io/jobset-name={jobset_name}",
)
except Exception as e:
_logger.exception("Failed to list pods for JobSet '%s': %r", jobset_name, e)
raise RuntimeError(
"Failed to list pods for Shared Pathways Service JobSet"
f" '{jobset_name}': {e}"
) from e

if isinstance(pod_list, dict):
pods = pod_list.get("items", []) or []
else:
pods = getattr(pod_list, "items", []) or []

if not pods:
raise RuntimeError(
f"No pods found for Shared Pathways Service JobSet '{jobset_name}' in"
f" namespace '{namespace}'."
)

head_pods = [p for p in pods if _is_head_pod(p, jobset_name)]
worker_pods = [p for p in pods if _is_worker_pod(p, jobset_name)]

if not head_pods:
raise RuntimeError(
f"Shared Pathways Service head pod not found for JobSet '{jobset_name}'"
f" in namespace '{namespace}'."
)

ready_head_pods = [p for p in head_pods if _is_pod_ready(p)]
if not ready_head_pods:
head_pod = head_pods[0]
head_name = _get_pod_name(head_pod)
details = _get_pod_status_details(head_pod)
raise RuntimeError(
f"Shared Pathways Service head pod '{head_name}' is not ready"
f" ({details}). Please ensure the service is deployed and healthy"
" before running workloads."
)

if not worker_pods:
raise RuntimeError(
"Shared Pathways Service worker pods not found for JobSet"
f" '{jobset_name}' in namespace '{namespace}'."
)

ready_worker_pods = [p for p in worker_pods if _is_pod_ready(p)]
if not ready_worker_pods:
pod_summaries = ", ".join(
f"{_get_pod_name(p)}: {_get_pod_status_details(p)}"
for p in worker_pods[:3]
)
if len(worker_pods) > 3:
pod_summaries += f", ... ({len(worker_pods) - 3} more)"
raise RuntimeError(
"No ready worker pods found for Shared Pathways Service JobSet"
f" '{jobset_name}'. Found {len(worker_pods)} worker pods, but 0 are"
f" ready. Status: {pod_summaries}."
)

configured_slices = 1
vms_per_slice = 1
for job in jobset.get("spec", {}).get("replicatedJobs", []):
if job.get("name") in ("pathways-worker", "worker"):
configured_slices = job.get("replicas", 1)
job_spec = job.get("template", {}).get("spec", {})
vms_per_slice = (
job_spec.get("parallelism")
or job_spec.get("completions")
or 1
)
break

if tpu_count is not None and tpu_count > 0:
if tpu_count > configured_slices:
raise RuntimeError(
f"Requested {tpu_count} TPU slice(s), but Shared Pathways Service"
f" JobSet '{jobset_name}' only has {configured_slices} slice(s)"
" configured."
)
required_worker_pods = tpu_count * vms_per_slice
if len(ready_worker_pods) < required_worker_pods:
raise RuntimeError(
f"Shared Pathways Service JobSet '{jobset_name}' has insufficient"
f" ready worker pods: requires at least {required_worker_pods} ready"
f" pods ({tpu_count} slice(s) x {vms_per_slice} VMs per slice), but"
f" only {len(ready_worker_pods)} of {len(worker_pods)} worker pods"
" are ready."
)

_logger.info(
"Shared Pathways Service '%s' is up and ready (%d ready worker pods).",
jobset_name,
len(ready_worker_pods),
)
Loading
Loading