diff --git a/rock/sandbox/base_manager.py b/rock/sandbox/base_manager.py index 5cd705587a..829468fb44 100644 --- a/rock/sandbox/base_manager.py +++ b/rock/sandbox/base_manager.py @@ -83,9 +83,15 @@ def _setup_job_check_scheduler(self): ) logger.info("auto_transition and reconcile jobs registered (primary pod)") else: - logger.info("auto_transition and reconcile jobs skipped (non-primary pod)") + self.scheduler.add_job( + func=self._auto_stop_expired, + trigger=IntervalTrigger(seconds=self._auto_transition_interval), + id="auto_stop_expired", + name="Sandbox Auto Stop Expired", + ) + logger.info("auto_stop_expired job registered (non-primary pod); other lifecycle jobs skipped") self.scheduler.start() - logger.info("APScheduler started for auto_transition and reconcile") + logger.info("APScheduler started for lifecycle jobs") async def _collect_and_report_metrics(self): start_time = time.time() @@ -169,6 +175,9 @@ async def _collect_sandbox_meta(self) -> tuple[int, dict[str, dict[str, str]]]: @abstractmethod async def _auto_transition(self): ... + @abstractmethod + async def _auto_stop_expired(self): ... + @abstractmethod async def _reconcile(self): ... diff --git a/rock/sandbox/operator/ray.py b/rock/sandbox/operator/ray.py index 2fe91e4766..098f23f843 100644 --- a/rock/sandbox/operator/ray.py +++ b/rock/sandbox/operator/ray.py @@ -1,5 +1,4 @@ import json -import asyncio import ray @@ -19,7 +18,6 @@ from rock.utils.service import build_sandbox_from_redis logger = init_logger(__name__) -RAY_KILL_TIMEOUT_SECONDS = 10 class RayOperator(AbstractOperator): @@ -28,17 +26,6 @@ def __init__(self, ray_service: RayService, runtime_config: RuntimeConfig): self._runtime_config = runtime_config self._disk_scheduling_enabled = ray.is_initialized() and "disk" in ray.cluster_resources() - async def _kill_actor(self, actor, *, sandbox_id: str, no_restart: bool | None = None) -> None: - kill_kwargs = {} - if no_restart is not None: - kill_kwargs["no_restart"] = no_restart - try: - await asyncio.wait_for(asyncio.to_thread(ray.kill, actor, **kill_kwargs), timeout=RAY_KILL_TIMEOUT_SECONDS) - except asyncio.TimeoutError: - logger.warning("ray.kill timed out while terminating sandbox %s", sandbox_id) - except Exception as e: - logger.warning("failed to kill sandbox %s actor: %s", sandbox_id, e) - def _get_actor_name(self, sandbox_id: str) -> str: return f"sandbox-{sandbox_id}" @@ -100,7 +87,11 @@ async def submit(self, config: DockerDeploymentConfig, user_info: dict = {}) -> try: sandbox_info: SandboxInfo = await self._ray_service.async_ray_get(sandbox_actor.sandbox_info.remote()) except Exception: - await self._kill_actor(sandbox_actor, sandbox_id=sandbox_id, no_restart=True) + try: + ray.kill(sandbox_actor, no_restart=True) + logger.info("[%s] force-killed actor after sandbox info failure", sandbox_id) + except Exception: + logger.exception("[%s] failed to force-kill actor after sandbox info failure", sandbox_id) raise sandbox_info["user_id"] = user_id sandbox_info["experiment_id"] = experiment_id @@ -144,7 +135,7 @@ async def stop(self, sandbox_id: str, reason: StopReason = StopReason.MANUAL) -> actor: SandboxActor = await self._ray_service.async_ray_get_actor(self._get_actor_name(sandbox_id)) await self._ray_service.async_ray_get(actor.stop.remote(reason)) logger.info(f"run time stop over {sandbox_id}") - await self._kill_actor(actor, sandbox_id=sandbox_id) + ray.kill(actor) return True async def delete(self, config: DockerDeploymentConfig, host_ip: str | None = None) -> bool: @@ -154,7 +145,7 @@ async def delete(self, config: DockerDeploymentConfig, host_ip: str | None = Non try: existing_actor = await self._ray_service.async_ray_get_actor(actor_name) - await self._kill_actor(existing_actor, sandbox_id=sandbox_id) + ray.kill(existing_actor) except Exception: logger.info(f"Actor {actor_name} already gone, proceeding with delete") @@ -176,7 +167,7 @@ async def delete(self, config: DockerDeploymentConfig, host_ip: str | None = Non logger.info(f"sandbox {sandbox_id} deleted on host_ip={host_ip}") return True finally: - await self._kill_actor(sandbox_actor, sandbox_id=sandbox_id) + ray.kill(sandbox_actor) async def restart(self, config: DockerDeploymentConfig, host_ip: str | None = None) -> SandboxInfo: """Restart an existing sandbox using docker start (container is preserved). diff --git a/tests/unit/sandbox/operator/test_ray_operator_delete.py b/tests/unit/sandbox/operator/test_ray_operator_delete.py index 0ea7c4e81e..10768afed8 100644 --- a/tests/unit/sandbox/operator/test_ray_operator_delete.py +++ b/tests/unit/sandbox/operator/test_ray_operator_delete.py @@ -1,13 +1,9 @@ from unittest.mock import AsyncMock, MagicMock, patch -import asyncio -import time - import pytest from rock.admin.core.ray_service import RayService from rock.config import RayConfig, RuntimeConfig -from rock.common.constants import StopReason from rock.deployments.config import DockerDeploymentConfig from rock.sandbox.operator.ray import RayOperator @@ -111,22 +107,3 @@ async def test_archive_actor_drops_sandbox_resources(): image_storage_config, archive_params, ) - - -@pytest.mark.asyncio -async def test_stop_with_hanging_kill_times_out_and_returns(): - operator, ray_service = _make_operator() - actor = MagicMock() - actor.stop.remote.return_value = object() - ray_service.async_ray_get_actor = AsyncMock(return_value=actor) - ray_service.async_ray_get = AsyncMock(return_value=None) - - def hanging_kill(*_args, **_kwargs): - time.sleep(5) - - with patch("rock.sandbox.operator.ray.ray.kill", side_effect=hanging_kill), patch( - "rock.sandbox.operator.ray.RAY_KILL_TIMEOUT_SECONDS", 0.2 - ): - result = await asyncio.wait_for(operator.stop("sb-1", reason=StopReason.MANUAL), timeout=1) - - assert result is True diff --git a/tests/unit/sandbox/test_base_manager_scheduler.py b/tests/unit/sandbox/test_base_manager_scheduler.py new file mode 100644 index 0000000000..e3969cf1f0 --- /dev/null +++ b/tests/unit/sandbox/test_base_manager_scheduler.py @@ -0,0 +1,50 @@ +from unittest.mock import MagicMock, patch + +from rock.sandbox.base_manager import BaseManager + + +class _ConcreteManager(BaseManager): + async def _auto_transition(self): ... + + async def _auto_stop_expired(self): ... + + async def _reconcile(self): ... + + +def _manager() -> _ConcreteManager: + manager = _ConcreteManager.__new__(_ConcreteManager) + manager._auto_transition_interval = 180 + manager._reconcile_interval = 30 + return manager + + +@patch("rock.sandbox.base_manager.AsyncIOScheduler") +@patch("rock.sandbox.base_manager.is_primary_pod", return_value=True) +def test_primary_registers_full_lifecycle_and_reconcile(mock_is_primary, mock_scheduler_cls): + scheduler = MagicMock() + mock_scheduler_cls.return_value = scheduler + + manager = _manager() + manager._setup_job_check_scheduler() + + assert [call.kwargs["id"] for call in scheduler.add_job.call_args_list] == ["auto_transition", "reconcile"] + assert scheduler.add_job.call_args_list[0].kwargs["func"] == manager._auto_transition + scheduler.start.assert_called_once_with() + mock_is_primary.assert_called_once_with() + + +@patch("rock.sandbox.base_manager.AsyncIOScheduler") +@patch("rock.sandbox.base_manager.is_primary_pod", return_value=False) +def test_non_primary_registers_only_auto_stop_expired(mock_is_primary, mock_scheduler_cls): + scheduler = MagicMock() + mock_scheduler_cls.return_value = scheduler + + manager = _manager() + manager._setup_job_check_scheduler() + + scheduler.add_job.assert_called_once() + job = scheduler.add_job.call_args.kwargs + assert job["id"] == "auto_stop_expired" + assert job["func"] == manager._auto_stop_expired + scheduler.start.assert_called_once_with() + mock_is_primary.assert_called_once_with() diff --git a/tests/unit/sandbox/test_collect_sandbox_meta.py b/tests/unit/sandbox/test_collect_sandbox_meta.py index 74f3c1b5d1..b3594c75fc 100644 --- a/tests/unit/sandbox/test_collect_sandbox_meta.py +++ b/tests/unit/sandbox/test_collect_sandbox_meta.py @@ -37,11 +37,11 @@ def base_manager(meta_store): from rock.sandbox.base_manager import BaseManager class _ConcreteManager(BaseManager): - async def _auto_transition(self): - ... + async def _auto_transition(self): ... - async def _reconcile(self): - ... + async def _auto_stop_expired(self): ... + + async def _reconcile(self): ... with patch.object(BaseManager, "__init__", lambda self, *a, **kw: None): mgr = _ConcreteManager.__new__(_ConcreteManager)