diff --git a/pathwaysutils/experimental/shared_pathways_service/gke_utils.py b/pathwaysutils/experimental/shared_pathways_service/gke_utils.py index 144c999..02e8682 100644 --- a/pathwaysutils/experimental/shared_pathways_service/gke_utils.py +++ b/pathwaysutils/experimental/shared_pathways_service/gke_utils.py @@ -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. @@ -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: @@ -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:" " -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() @@ -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), + ) diff --git a/pathwaysutils/experimental/shared_pathways_service/run_workload.py b/pathwaysutils/experimental/shared_pathways_service/run_workload.py index 3ba940a..2af949c 100644 --- a/pathwaysutils/experimental/shared_pathways_service/run_workload.py +++ b/pathwaysutils/experimental/shared_pathways_service/run_workload.py @@ -80,6 +80,38 @@ ) +def verify_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 is running and ready. + + Args: + cluster: The name of the GKE cluster. + project: The GCP project ID. + region: The GCP region. + pathways_service: The address and port of the Pathways Resource Manager. + tpu_count: Optional expected number of TPU slices. + namespace: The Kubernetes namespace. + + Raises: + RuntimeError: If the Shared Pathways Service is not up or not ready. + """ + gke_utils.verify_pathways_service_is_up( + cluster=cluster, + project=project, + region=region, + pathways_service=pathways_service, + tpu_count=tpu_count, + namespace=namespace, + ) + + def run_command( *, cluster: str, @@ -94,6 +126,7 @@ def run_command( proxy_options: Sequence[str] | None = None, collect_service_metrics: bool = False, connect_fn: Callable[..., ContextManager[Any]] = isc_pathways.connect, + verify_service_fn: Callable[..., None] | None = None, ) -> None: """Run the TPU workload within a Shared Pathways connection. @@ -113,8 +146,11 @@ def run_command( 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. + verify_service_fn: The function to use for verifying that the Shared + Pathways Service is up and ready before connecting. Raises: + RuntimeError: If the Shared Pathways Service is not up or ready. subprocess.CalledProcessError: If the workload command fails. """ if proxy_server_image: @@ -125,6 +161,16 @@ def run_command( DeprecationWarning, stacklevel=2, ) + if verify_service_fn is None: + verify_service_fn = verify_service_is_up + logging.info("Verifying Shared Pathways Service is up and ready...") + verify_service_fn( + cluster=cluster, + project=project, + region=region, + pathways_service=pathways_service, + tpu_count=tpu_count, + ) logging.info("Connecting to Shared Pathways Service...") with connect_fn( cluster=cluster, 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..0c9b1ad 100644 --- a/pathwaysutils/test/experimental/shared_pathways_service/gke_utils_test.py +++ b/pathwaysutils/test/experimental/shared_pathways_service/gke_utils_test.py @@ -1593,6 +1593,825 @@ def test_terminate_process_custom_timeout(self): mock_proc.terminate.assert_called_once() mock_proc.wait.assert_called_once_with(timeout=10) + def _make_service_jobset( + self, + name: str = "my-jobset", + namespace: str = "default", + num_slices: int = 2, + vms_per_slice: int = 2, + suspend: bool = False, + terminal_state: str | None = None, + conditions: list[dict[str, Any]] | None = None, + ) -> dict[str, Any]: + jobset: dict[str, Any] = { + "apiVersion": "jobset.x-k8s.io/v1alpha2", + "kind": "JobSet", + "metadata": { + "name": name, + "namespace": namespace, + }, + "spec": { + "suspend": suspend, + "replicatedJobs": [ + { + "name": "pathways-head", + "replicas": 1, + "template": { + "spec": { + "parallelism": 1, + "completions": 1, + } + }, + }, + { + "name": "pathways-worker", + "replicas": num_slices, + "template": { + "spec": { + "parallelism": vms_per_slice, + "completions": vms_per_slice, + } + }, + }, + ], + }, + "status": { + "conditions": conditions or [], + }, + } + if terminal_state: + jobset["status"]["terminalState"] = terminal_state + return jobset + + def _make_pod_dict( + self, + name: str, + jobset_name: str, + replicated_job_name: str, + phase: str = "Running", + ready: bool = True, + waiting_reason: str | None = None, + waiting_message: str | None = None, + terminated_reason: str | None = None, + exit_code: int = 0, + ) -> dict[str, Any]: + container_status: dict[str, Any] = {"name": replicated_job_name} + if waiting_reason: + container_status["state"] = { + "waiting": { + "reason": waiting_reason, + "message": waiting_message or "", + } + } + elif terminated_reason: + container_status["state"] = { + "terminated": { + "reason": terminated_reason, + "exitCode": exit_code, + } + } + else: + container_status["state"] = {"running": {}} + + return { + "metadata": { + "name": name, + "labels": { + "jobset.sigs.k8s.io/jobset-name": jobset_name, + "jobset.sigs.k8s.io/replicatedjob-name": replicated_job_name, + }, + }, + "status": { + "phase": phase, + "conditions": [ + { + "type": "Ready", + "status": "True" if ready else "False", + } + ], + "containerStatuses": [container_status], + }, + } + + def test_extract_jobset_name_valid(self): + self.assertEqual( + gke_utils.extract_jobset_name( + "my-jobset-pathways-head-0-0.my-jobset:29001" + ), + "my-jobset", + ) + self.assertEqual( + gke_utils.extract_jobset_name("my-jobset-pathways-head:29001"), + "my-jobset", + ) + self.assertEqual( + gke_utils.extract_jobset_name("sps-cluster-pathways-head-0-0:8000"), + "sps-cluster", + ) + + def test_extract_jobset_name_invalid_raises(self): + with self.assertRaises(ValueError): + gke_utils.extract_jobset_name("my-jobset:29001") + + def test_verify_pathways_service_is_up_success(self): + mock_creds = self.enter_context( + mock.patch.object( + gke_utils, "fetch_cluster_credentials", autospec=True + ) + ) + mock_custom_api = mock.MagicMock(spec=client.CustomObjectsApi) + mock_core_api = mock.MagicMock(spec=client.CoreV1Api) + self.enter_context( + mock.patch.object( + gke_utils, + "_get_k8s_custom_objects_api", + return_value=mock_custom_api, + ) + ) + self.enter_context( + mock.patch.object( + gke_utils, "_get_k8s_core_api", return_value=mock_core_api + ) + ) + mock_custom_api.get_namespaced_custom_object.return_value = ( + self._make_service_jobset(num_slices=2, vms_per_slice=2) + ) + head_pod = self._make_pod_dict( + name="my-jobset-pathways-head-0-0", + jobset_name="my-jobset", + replicated_job_name="pathways-head", + ready=True, + ) + worker_pods = [] + for slice_idx in range(2): + for vm_idx in range(2): + worker_pods.append( + self._make_pod_dict( + name=f"my-jobset-pathways-worker-{slice_idx}-{vm_idx}", + jobset_name="my-jobset", + replicated_job_name="pathways-worker", + ready=True, + ) + ) + mock_core_api.list_namespaced_pod.return_value = { + "items": [head_pod] + worker_pods + } + + gke_utils.verify_pathways_service_is_up( + cluster="my-cluster", + project="my-project", + region="us-central1", + pathways_service="my-jobset-pathways-head-0-0:29001", + tpu_count=2, + ) + + mock_creds.assert_called_once_with( + cluster_name="my-cluster", + project_id="my-project", + location="us-central1", + ) + mock_custom_api.get_namespaced_custom_object.assert_called_once_with( + group="jobset.x-k8s.io", + version="v1alpha2", + namespace="default", + plural="jobsets", + name="my-jobset", + ) + mock_core_api.list_namespaced_pod.assert_called_once_with( + namespace="default", + label_selector="jobset.sigs.k8s.io/jobset-name=my-jobset", + ) + + def test_verify_pathways_service_is_up_jobset_not_found_404(self): + self.enter_context( + mock.patch.object( + gke_utils, "fetch_cluster_credentials", autospec=True + ) + ) + mock_custom_api = mock.MagicMock(spec=client.CustomObjectsApi) + self.enter_context( + mock.patch.object( + gke_utils, + "_get_k8s_custom_objects_api", + return_value=mock_custom_api, + ) + ) + mock_custom_api.get_namespaced_custom_object.side_effect = ( + client.rest.ApiException(status=404, reason="Not Found") + ) + with self.assertRaisesRegex( + RuntimeError, "JobSet 'my-jobset' was not found in namespace 'default'" + ): + gke_utils.verify_pathways_service_is_up( + cluster="my-cluster", + project="my-project", + region="us-central1", + pathways_service="my-jobset-pathways-head-0-0:29001", + ) + + def test_verify_pathways_service_is_up_jobset_api_error(self): + self.enter_context( + mock.patch.object( + gke_utils, "fetch_cluster_credentials", autospec=True + ) + ) + mock_custom_api = mock.MagicMock(spec=client.CustomObjectsApi) + self.enter_context( + mock.patch.object( + gke_utils, + "_get_k8s_custom_objects_api", + return_value=mock_custom_api, + ) + ) + mock_custom_api.get_namespaced_custom_object.side_effect = ( + client.rest.ApiException(status=500, reason="Internal Server Error") + ) + with self.assertRaisesRegex( + RuntimeError, + "Failed to get Shared Pathways Service JobSet 'my-jobset'", + ): + gke_utils.verify_pathways_service_is_up( + cluster="my-cluster", + project="my-project", + region="us-central1", + pathways_service="my-jobset-pathways-head-0-0:29001", + ) + + def test_verify_pathways_service_is_up_jobset_suspended(self): + self.enter_context( + mock.patch.object( + gke_utils, "fetch_cluster_credentials", autospec=True + ) + ) + mock_custom_api = mock.MagicMock(spec=client.CustomObjectsApi) + self.enter_context( + mock.patch.object( + gke_utils, + "_get_k8s_custom_objects_api", + return_value=mock_custom_api, + ) + ) + mock_custom_api.get_namespaced_custom_object.return_value = ( + self._make_service_jobset(suspend=True) + ) + with self.assertRaisesRegex( + RuntimeError, "JobSet 'my-jobset' is suspended" + ): + gke_utils.verify_pathways_service_is_up( + cluster="my-cluster", + project="my-project", + region="us-central1", + pathways_service="my-jobset-pathways-head-0-0:29001", + ) + + def test_verify_pathways_service_is_up_jobset_terminal_state(self): + self.enter_context( + mock.patch.object( + gke_utils, "fetch_cluster_credentials", autospec=True + ) + ) + mock_custom_api = mock.MagicMock(spec=client.CustomObjectsApi) + self.enter_context( + mock.patch.object( + gke_utils, + "_get_k8s_custom_objects_api", + return_value=mock_custom_api, + ) + ) + mock_custom_api.get_namespaced_custom_object.return_value = ( + self._make_service_jobset(terminal_state="Failed") + ) + with self.assertRaisesRegex( + RuntimeError, "has terminated with state 'Failed'" + ): + gke_utils.verify_pathways_service_is_up( + cluster="my-cluster", + project="my-project", + region="us-central1", + pathways_service="my-jobset-pathways-head-0-0:29001", + ) + + def test_verify_pathways_service_is_up_jobset_failed_condition(self): + self.enter_context( + mock.patch.object( + gke_utils, "fetch_cluster_credentials", autospec=True + ) + ) + mock_custom_api = mock.MagicMock(spec=client.CustomObjectsApi) + self.enter_context( + mock.patch.object( + gke_utils, + "_get_k8s_custom_objects_api", + return_value=mock_custom_api, + ) + ) + mock_custom_api.get_namespaced_custom_object.return_value = ( + self._make_service_jobset( + conditions=[ + { + "type": "Failed", + "status": "True", + "reason": "DeadlineExceeded", + "message": "Job timed out", + } + ] + ) + ) + with self.assertRaisesRegex( + RuntimeError, "is in condition 'Failed'" + ): + gke_utils.verify_pathways_service_is_up( + cluster="my-cluster", + project="my-project", + region="us-central1", + pathways_service="my-jobset-pathways-head-0-0:29001", + ) + + def test_verify_pathways_service_is_up_list_pods_error(self): + self.enter_context( + mock.patch.object( + gke_utils, "fetch_cluster_credentials", autospec=True + ) + ) + mock_custom_api = mock.MagicMock(spec=client.CustomObjectsApi) + mock_core_api = mock.MagicMock(spec=client.CoreV1Api) + self.enter_context( + mock.patch.object( + gke_utils, + "_get_k8s_custom_objects_api", + return_value=mock_custom_api, + ) + ) + self.enter_context( + mock.patch.object( + gke_utils, "_get_k8s_core_api", return_value=mock_core_api + ) + ) + mock_custom_api.get_namespaced_custom_object.return_value = ( + self._make_service_jobset() + ) + mock_core_api.list_namespaced_pod.side_effect = client.rest.ApiException( + status=500, reason="Error" + ) + with self.assertRaisesRegex( + RuntimeError, + "Failed to list pods for Shared Pathways Service JobSet 'my-jobset'", + ): + gke_utils.verify_pathways_service_is_up( + cluster="my-cluster", + project="my-project", + region="us-central1", + pathways_service="my-jobset-pathways-head-0-0:29001", + ) + + def test_verify_pathways_service_is_up_no_pods_found(self): + self.enter_context( + mock.patch.object( + gke_utils, "fetch_cluster_credentials", autospec=True + ) + ) + mock_custom_api = mock.MagicMock(spec=client.CustomObjectsApi) + mock_core_api = mock.MagicMock(spec=client.CoreV1Api) + self.enter_context( + mock.patch.object( + gke_utils, + "_get_k8s_custom_objects_api", + return_value=mock_custom_api, + ) + ) + self.enter_context( + mock.patch.object( + gke_utils, "_get_k8s_core_api", return_value=mock_core_api + ) + ) + mock_custom_api.get_namespaced_custom_object.return_value = ( + self._make_service_jobset() + ) + mock_core_api.list_namespaced_pod.return_value = {"items": []} + with self.assertRaisesRegex( + RuntimeError, + "No pods found for Shared Pathways Service JobSet 'my-jobset'", + ): + gke_utils.verify_pathways_service_is_up( + cluster="my-cluster", + project="my-project", + region="us-central1", + pathways_service="my-jobset-pathways-head-0-0:29001", + ) + + def test_verify_pathways_service_is_up_no_head_pod(self): + self.enter_context( + mock.patch.object( + gke_utils, "fetch_cluster_credentials", autospec=True + ) + ) + mock_custom_api = mock.MagicMock(spec=client.CustomObjectsApi) + mock_core_api = mock.MagicMock(spec=client.CoreV1Api) + self.enter_context( + mock.patch.object( + gke_utils, + "_get_k8s_custom_objects_api", + return_value=mock_custom_api, + ) + ) + self.enter_context( + mock.patch.object( + gke_utils, "_get_k8s_core_api", return_value=mock_core_api + ) + ) + mock_custom_api.get_namespaced_custom_object.return_value = ( + self._make_service_jobset() + ) + worker_pod = self._make_pod_dict( + name="my-jobset-pathways-worker-0-0", + jobset_name="my-jobset", + replicated_job_name="pathways-worker", + ready=True, + ) + mock_core_api.list_namespaced_pod.return_value = {"items": [worker_pod]} + with self.assertRaisesRegex( + RuntimeError, "head pod not found for JobSet 'my-jobset'" + ): + gke_utils.verify_pathways_service_is_up( + cluster="my-cluster", + project="my-project", + region="us-central1", + pathways_service="my-jobset-pathways-head-0-0:29001", + ) + + def test_verify_pathways_service_is_up_head_pod_not_ready(self): + self.enter_context( + mock.patch.object( + gke_utils, "fetch_cluster_credentials", autospec=True + ) + ) + mock_custom_api = mock.MagicMock(spec=client.CustomObjectsApi) + mock_core_api = mock.MagicMock(spec=client.CoreV1Api) + self.enter_context( + mock.patch.object( + gke_utils, + "_get_k8s_custom_objects_api", + return_value=mock_custom_api, + ) + ) + self.enter_context( + mock.patch.object( + gke_utils, "_get_k8s_core_api", return_value=mock_core_api + ) + ) + mock_custom_api.get_namespaced_custom_object.return_value = ( + self._make_service_jobset() + ) + head_pod = self._make_pod_dict( + name="my-jobset-pathways-head-0-0", + jobset_name="my-jobset", + replicated_job_name="pathways-head", + phase="Running", + ready=False, + waiting_reason="CrashLoopBackOff", + waiting_message="back-off 5m0s restarting failed container=pathways-rm", + ) + worker_pod = self._make_pod_dict( + name="my-jobset-pathways-worker-0-0", + jobset_name="my-jobset", + replicated_job_name="pathways-worker", + ready=True, + ) + mock_core_api.list_namespaced_pod.return_value = { + "items": [head_pod, worker_pod] + } + with self.assertRaisesRegex( + RuntimeError, + "head pod 'my-jobset-pathways-head-0-0' is not ready.*CrashLoopBackOff", + ): + gke_utils.verify_pathways_service_is_up( + cluster="my-cluster", + project="my-project", + region="us-central1", + pathways_service="my-jobset-pathways-head-0-0:29001", + ) + + def test_verify_pathways_service_is_up_no_worker_pods(self): + self.enter_context( + mock.patch.object( + gke_utils, "fetch_cluster_credentials", autospec=True + ) + ) + mock_custom_api = mock.MagicMock(spec=client.CustomObjectsApi) + mock_core_api = mock.MagicMock(spec=client.CoreV1Api) + self.enter_context( + mock.patch.object( + gke_utils, + "_get_k8s_custom_objects_api", + return_value=mock_custom_api, + ) + ) + self.enter_context( + mock.patch.object( + gke_utils, "_get_k8s_core_api", return_value=mock_core_api + ) + ) + mock_custom_api.get_namespaced_custom_object.return_value = ( + self._make_service_jobset() + ) + head_pod = self._make_pod_dict( + name="my-jobset-pathways-head-0-0", + jobset_name="my-jobset", + replicated_job_name="pathways-head", + ready=True, + ) + mock_core_api.list_namespaced_pod.return_value = {"items": [head_pod]} + with self.assertRaisesRegex( + RuntimeError, "worker pods not found for JobSet 'my-jobset'" + ): + gke_utils.verify_pathways_service_is_up( + cluster="my-cluster", + project="my-project", + region="us-central1", + pathways_service="my-jobset-pathways-head-0-0:29001", + ) + + def test_verify_pathways_service_is_up_zero_ready_worker_pods(self): + self.enter_context( + mock.patch.object( + gke_utils, "fetch_cluster_credentials", autospec=True + ) + ) + mock_custom_api = mock.MagicMock(spec=client.CustomObjectsApi) + mock_core_api = mock.MagicMock(spec=client.CoreV1Api) + self.enter_context( + mock.patch.object( + gke_utils, + "_get_k8s_custom_objects_api", + return_value=mock_custom_api, + ) + ) + self.enter_context( + mock.patch.object( + gke_utils, "_get_k8s_core_api", return_value=mock_core_api + ) + ) + mock_custom_api.get_namespaced_custom_object.return_value = ( + self._make_service_jobset() + ) + head_pod = self._make_pod_dict( + name="my-jobset-pathways-head-0-0", + jobset_name="my-jobset", + replicated_job_name="pathways-head", + ready=True, + ) + worker_pod = self._make_pod_dict( + name="my-jobset-pathways-worker-0-0", + jobset_name="my-jobset", + replicated_job_name="pathways-worker", + phase="Pending", + ready=False, + waiting_reason="ImagePullBackOff", + ) + mock_core_api.list_namespaced_pod.return_value = { + "items": [head_pod, worker_pod] + } + with self.assertRaisesRegex( + RuntimeError, + "No ready worker pods found for Shared Pathways Service JobSet" + " 'my-jobset'", + ): + gke_utils.verify_pathways_service_is_up( + cluster="my-cluster", + project="my-project", + region="us-central1", + pathways_service="my-jobset-pathways-head-0-0:29001", + ) + + def test_verify_pathways_service_is_up_insufficient_ready_worker_pods(self): + self.enter_context( + mock.patch.object( + gke_utils, "fetch_cluster_credentials", autospec=True + ) + ) + mock_custom_api = mock.MagicMock(spec=client.CustomObjectsApi) + mock_core_api = mock.MagicMock(spec=client.CoreV1Api) + self.enter_context( + mock.patch.object( + gke_utils, + "_get_k8s_custom_objects_api", + return_value=mock_custom_api, + ) + ) + self.enter_context( + mock.patch.object( + gke_utils, "_get_k8s_core_api", return_value=mock_core_api + ) + ) + mock_custom_api.get_namespaced_custom_object.return_value = ( + self._make_service_jobset(num_slices=2, vms_per_slice=2) + ) + head_pod = self._make_pod_dict( + name="my-jobset-pathways-head-0-0", + jobset_name="my-jobset", + replicated_job_name="pathways-head", + ready=True, + ) + worker_pod_0 = self._make_pod_dict( + name="my-jobset-pathways-worker-0-0", + jobset_name="my-jobset", + replicated_job_name="pathways-worker", + ready=True, + ) + worker_pod_1 = self._make_pod_dict( + name="my-jobset-pathways-worker-0-1", + jobset_name="my-jobset", + replicated_job_name="pathways-worker", + ready=True, + ) + worker_pod_2 = self._make_pod_dict( + name="my-jobset-pathways-worker-1-0", + jobset_name="my-jobset", + replicated_job_name="pathways-worker", + phase="Failed", + ready=False, + terminated_reason="Error", + exit_code=1, + ) + worker_pod_3 = self._make_pod_dict( + name="my-jobset-pathways-worker-1-1", + jobset_name="my-jobset", + replicated_job_name="pathways-worker", + phase="Failed", + ready=False, + terminated_reason="Error", + exit_code=1, + ) + mock_core_api.list_namespaced_pod.return_value = { + "items": [ + head_pod, + worker_pod_0, + worker_pod_1, + worker_pod_2, + worker_pod_3, + ] + } + with self.assertRaisesRegex( + RuntimeError, + "requires at least 4 ready pods.*but only 2 of 4 worker pods are" + " ready", + ): + gke_utils.verify_pathways_service_is_up( + cluster="my-cluster", + project="my-project", + region="us-central1", + pathways_service="my-jobset-pathways-head-0-0:29001", + tpu_count=2, + ) + + def test_verify_pathways_service_is_up_tpu_count_exceeds_configured_slices( + self, + ): + self.enter_context( + mock.patch.object( + gke_utils, "fetch_cluster_credentials", autospec=True + ) + ) + mock_custom_api = mock.MagicMock(spec=client.CustomObjectsApi) + mock_core_api = mock.MagicMock(spec=client.CoreV1Api) + self.enter_context( + mock.patch.object( + gke_utils, + "_get_k8s_custom_objects_api", + return_value=mock_custom_api, + ) + ) + self.enter_context( + mock.patch.object( + gke_utils, "_get_k8s_core_api", return_value=mock_core_api + ) + ) + mock_custom_api.get_namespaced_custom_object.return_value = ( + self._make_service_jobset(num_slices=2, vms_per_slice=2) + ) + head_pod = self._make_pod_dict( + name="my-jobset-pathways-head-0-0", + jobset_name="my-jobset", + replicated_job_name="pathways-head", + ready=True, + ) + worker_pod = self._make_pod_dict( + name="my-jobset-pathways-worker-0-0", + jobset_name="my-jobset", + replicated_job_name="pathways-worker", + ready=True, + ) + mock_core_api.list_namespaced_pod.return_value = { + "items": [head_pod, worker_pod] + } + with self.assertRaisesRegex( + RuntimeError, + r"Requested 4 TPU slice\(s\), but Shared Pathways Service JobSet" + r" 'my-jobset' only has 2 slice\(s\) configured", + ): + gke_utils.verify_pathways_service_is_up( + cluster="my-cluster", + project="my-project", + region="us-central1", + pathways_service="my-jobset-pathways-head-0-0:29001", + tpu_count=4, + ) + + def test_verify_pathways_service_is_up_with_v1_pod_objects(self): + self.enter_context( + mock.patch.object( + gke_utils, "fetch_cluster_credentials", autospec=True + ) + ) + mock_custom_api = mock.MagicMock(spec=client.CustomObjectsApi) + mock_core_api = mock.MagicMock(spec=client.CoreV1Api) + self.enter_context( + mock.patch.object( + gke_utils, + "_get_k8s_custom_objects_api", + return_value=mock_custom_api, + ) + ) + self.enter_context( + mock.patch.object( + gke_utils, "_get_k8s_core_api", return_value=mock_core_api + ) + ) + mock_custom_api.get_namespaced_custom_object.return_value = ( + self._make_service_jobset(num_slices=1, vms_per_slice=1) + ) + head_pod = client.V1Pod( + metadata=client.V1ObjectMeta( + name="my-jobset-pathways-head-0-0", + labels={ + "jobset.sigs.k8s.io/jobset-name": "my-jobset", + "jobset.sigs.k8s.io/replicatedjob-name": "pathways-head", + }, + ), + status=client.V1PodStatus( + phase="Running", + conditions=[client.V1PodCondition(type="Ready", status="True")], + container_statuses=[ + client.V1ContainerStatus( + name="pathways-rm", + ready=True, + image="img", + image_id="img_id", + restart_count=0, + state=client.V1ContainerState( + running=client.V1ContainerStateRunning() + ), + ) + ], + ), + ) + worker_pod = client.V1Pod( + metadata=client.V1ObjectMeta( + name="my-jobset-pathways-worker-0-0", + labels={ + "jobset.sigs.k8s.io/jobset-name": "my-jobset", + "jobset.sigs.k8s.io/replicatedjob-name": "pathways-worker", + }, + ), + status=client.V1PodStatus( + phase="Running", + conditions=[client.V1PodCondition(type="Ready", status="True")], + container_statuses=[ + client.V1ContainerStatus( + name="pathways-worker", + ready=True, + image="img", + image_id="img_id", + restart_count=0, + state=client.V1ContainerState( + running=client.V1ContainerStateRunning() + ), + ) + ], + ), + ) + mock_core_api.list_namespaced_pod.return_value = client.V1PodList( + items=[head_pod, worker_pod] + ) + + gke_utils.verify_pathways_service_is_up( + cluster="my-cluster", + project="my-project", + region="us-central1", + pathways_service="my-jobset-pathways-head-0-0:29001", + tpu_count=1, + ) + + def test_verify_pathways_service_is_up_invalid_namespace(self): + with self.assertRaises(ValueError): + gke_utils.verify_pathways_service_is_up( + cluster="my-cluster", + project="my-project", + region="us-central1", + pathways_service="my-jobset-pathways-head-0-0:29001", + namespace="invalid namespace!", + ) + if __name__ == "__main__": absltest.main() diff --git a/pathwaysutils/test/experimental/shared_pathways_service/run_workload_test.py b/pathwaysutils/test/experimental/shared_pathways_service/run_workload_test.py index 7323000..1d6104b 100644 --- a/pathwaysutils/test/experimental/shared_pathways_service/run_workload_test.py +++ b/pathwaysutils/test/experimental/shared_pathways_service/run_workload_test.py @@ -7,6 +7,7 @@ from absl.testing import absltest from absl.testing import flagsaver +from pathwaysutils.experimental.shared_pathways_service import gke_utils from pathwaysutils.experimental.shared_pathways_service import run_workload @@ -31,8 +32,40 @@ def __exit__( self.exited = True +class VerifyServiceIsUpTest(absltest.TestCase): + + def test_verify_service_is_up_delegates_to_gke_utils(self): + mock_verify = self.enter_context( + mock.patch.object( + gke_utils, "verify_pathways_service_is_up", autospec=True + ) + ) + run_workload.verify_service_is_up( + cluster="test-cluster", + project="test-project", + region="test-region", + pathways_service="test-service-pathways-head-0-0:1234", + tpu_count=2, + namespace="custom-ns", + ) + mock_verify.assert_called_once_with( + cluster="test-cluster", + project="test-project", + region="test-region", + pathways_service="test-service-pathways-head-0-0:1234", + tpu_count=2, + namespace="custom-ns", + ) + + class RunTpuWorkloadTest(absltest.TestCase): + def setUp(self): + super().setUp() + self.mock_verify_service = self.enter_context( + mock.patch.object(run_workload, "verify_service_is_up", autospec=True) + ) + def test_run_workload_success(self): fake_instances = [] @@ -78,6 +111,78 @@ def fake_connect_fn(**kwargs: Any) -> Generator[FakeConnect, None, None]: ) mock_proc.wait.assert_called_once() + with self.subTest("Service verified before connecting"): + self.mock_verify_service.assert_called_once_with( + cluster="test-cluster", + project="test-project", + region="test-region", + pathways_service="test-service:1234", + tpu_count=1, + ) + + def test_run_command_fails_when_service_not_up(self): + self.mock_verify_service.side_effect = RuntimeError( + "Shared Pathways Service JobSet 'test-service' was not found" + ) + mock_connect_fn = mock.MagicMock() + mock_popen = self.enter_context( + mock.patch.object(subprocess, "Popen", autospec=True) + ) + + with self.assertRaisesRegex( + RuntimeError, + "Shared Pathways Service JobSet 'test-service' was not found", + ): + run_workload.run_command( + cluster="test-cluster", + project="test-project", + region="test-region", + gcs_bucket="test-bucket", + pathways_service="test-service-pathways-head-0-0:1234", + tpu_type="tpuv6e:4x8", + tpu_count=1, + command="echo hello", + connect_fn=mock_connect_fn, + ) + + mock_connect_fn.assert_not_called() + mock_popen.assert_not_called() + + def test_run_command_custom_verify_service_fn(self): + mock_custom_verify = mock.MagicMock() + mock_popen = self.enter_context( + mock.patch.object(subprocess, "Popen", autospec=True) + ) + mock_proc = mock_popen.return_value + mock_proc.wait.return_value = 0 + + @contextlib.contextmanager + def fake_connect_fn(**kwargs: Any) -> Generator[None, None, None]: + del kwargs + yield None + + run_workload.run_command( + cluster="test-cluster", + project="test-project", + region="test-region", + gcs_bucket="test-bucket", + pathways_service="test-service:1234", + tpu_type="tpuv6e:4x8", + tpu_count=1, + command="echo hello", + connect_fn=fake_connect_fn, + verify_service_fn=mock_custom_verify, + ) + + mock_custom_verify.assert_called_once_with( + cluster="test-cluster", + project="test-project", + region="test-region", + pathways_service="test-service:1234", + tpu_count=1, + ) + self.mock_verify_service.assert_not_called() + def test_run_command_runs_command_inside_context(self): """Verifies that the command is executed while the connection is active.""" connection_active_during_run = False