Skip to content
Merged
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
13 changes: 11 additions & 2 deletions rock/sandbox/base_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down Expand Up @@ -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): ...

Expand Down
25 changes: 8 additions & 17 deletions rock/sandbox/operator/ray.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,4 @@
import json
import asyncio

import ray

Expand All @@ -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):
Expand All @@ -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}"

Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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:
Expand All @@ -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")

Expand All @@ -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).
Expand Down
23 changes: 0 additions & 23 deletions tests/unit/sandbox/operator/test_ray_operator_delete.py
Original file line number Diff line number Diff line change
@@ -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

Expand Down Expand Up @@ -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
50 changes: 50 additions & 0 deletions tests/unit/sandbox/test_base_manager_scheduler.py
Original file line number Diff line number Diff line change
@@ -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()
8 changes: 4 additions & 4 deletions tests/unit/sandbox/test_collect_sandbox_meta.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
Loading