diff --git a/src/codeflare_sdk/ray/cluster/cluster.py b/src/codeflare_sdk/ray/cluster/cluster.py index 4a3b0f13..5431fe44 100644 --- a/src/codeflare_sdk/ray/cluster/cluster.py +++ b/src/codeflare_sdk/ray/cluster/cluster.py @@ -1199,10 +1199,22 @@ def _map_to_ray_cluster(rc) -> Optional[RayCluster]: ) +_CODEFLARE_TO_RAY_STATUS = { + CodeFlareClusterStatus.READY: RayClusterStatus.READY, + CodeFlareClusterStatus.FAILED: RayClusterStatus.FAILED, + CodeFlareClusterStatus.SUSPENDED: RayClusterStatus.SUSPENDED, + CodeFlareClusterStatus.UNKNOWN: RayClusterStatus.UNKNOWN, + CodeFlareClusterStatus.STARTING: RayClusterStatus.UNKNOWN, + CodeFlareClusterStatus.QUEUED: RayClusterStatus.UNKNOWN, + CodeFlareClusterStatus.QUEUEING: RayClusterStatus.UNKNOWN, +} + + def _copy_to_ray(cluster: Cluster) -> RayCluster: - ray = RayCluster( + cf_status = cluster.status(print_to_console=False)[0] + return RayCluster( name=cluster.config.name, - status=cluster.status(print_to_console=False)[0], + status=_CODEFLARE_TO_RAY_STATUS.get(cf_status, RayClusterStatus.UNKNOWN), num_workers=cluster.config.num_workers, worker_mem_requests=cluster.config.worker_memory_requests, worker_mem_limits=cluster.config.worker_memory_limits, @@ -1217,9 +1229,6 @@ def _copy_to_ray(cluster: Cluster) -> RayCluster: head_cpu_limits=cluster.config.head_cpu_limits, head_extended_resources=cluster.config.head_extended_resource_requests, ) - if ray.status == CodeFlareClusterStatus.READY: - ray.status = RayClusterStatus.READY - return ray # Check if the routes api exists diff --git a/src/codeflare_sdk/ray/cluster/test_cluster.py b/src/codeflare_sdk/ray/cluster/test_cluster.py index b4e53ac0..ce12c36d 100644 --- a/src/codeflare_sdk/ray/cluster/test_cluster.py +++ b/src/codeflare_sdk/ray/cluster/test_cluster.py @@ -32,6 +32,7 @@ route_list_retrieval, ) from codeflare_sdk.ray.cluster.cluster import _is_openshift_cluster +from codeflare_sdk.ray.cluster.status import CodeFlareClusterStatus, RayClusterStatus from pathlib import Path from unittest.mock import MagicMock from kubernetes import client @@ -47,6 +48,31 @@ cluster_dir = os.path.expanduser("~/.codeflare/resources/") +@pytest.mark.parametrize( + "cf_status,expected_ray_status", + [ + (CodeFlareClusterStatus.READY, RayClusterStatus.READY), + (CodeFlareClusterStatus.FAILED, RayClusterStatus.FAILED), + (CodeFlareClusterStatus.SUSPENDED, RayClusterStatus.SUSPENDED), + (CodeFlareClusterStatus.UNKNOWN, RayClusterStatus.UNKNOWN), + (CodeFlareClusterStatus.STARTING, RayClusterStatus.UNKNOWN), + (CodeFlareClusterStatus.QUEUED, RayClusterStatus.UNKNOWN), + (CodeFlareClusterStatus.QUEUEING, RayClusterStatus.UNKNOWN), + ], +) +def test_details_maps_codeflare_status_to_ray_status( + mocker, cf_status, expected_ray_status +): + cluster = create_cluster(mocker) + mocker.patch.object(cluster, "status", return_value=(cf_status, "")) + mocker.patch.object(cluster, "cluster_dashboard_uri", return_value="http://fake") + + ray_cluster = cluster.details(print_to_console=False) + + assert ray_cluster.status == expected_ray_status + assert isinstance(ray_cluster.status, RayClusterStatus) + + def test_cluster_apply_down(mocker): mocker.patch("kubernetes.client.ApisApi.get_api_versions") mocker.patch("kubernetes.config.load_kube_config", return_value="ignore")